db: store API keys in the credentials table

Via the APIKey projection, leaving callers unchanged.
This commit is contained in:
Kristoffer Dalby 2026-06-27 12:19:57 +00:00
parent 30a0ddde4b
commit e276a94f25
3 changed files with 61 additions and 38 deletions

View file

@ -29,63 +29,70 @@ var (
func (hsdb *HSDatabase) CreateAPIKey(
expiration *time.Time,
) (string, *types.APIKey, error) {
keyStr, prefix, hash, err := generateSecret(apiKeyPrefix)
keyStr, identifier, hash, err := generateSecret(apiKeyPrefix)
if err != nil {
return "", nil, err
}
key := types.APIKey{
Prefix: prefix,
cred := types.Credential{
Kind: types.CredentialAPIKey,
Identifier: identifier,
Hash: hash,
Expiration: expiration,
}
if err := hsdb.DB.Save(&key).Error; err != nil { //nolint:noinlineerr
if err := hsdb.DB.Save(&cred).Error; err != nil { //nolint:noinlineerr
return "", nil, fmt.Errorf("saving API key to database: %w", err)
}
return keyStr, &key, nil
return keyStr, credentialToAPIKey(&cred), nil
}
// ListAPIKeys returns the list of [types.APIKey] values for a user.
func (hsdb *HSDatabase) ListAPIKeys() ([]types.APIKey, error) {
keys := []types.APIKey{}
var creds []types.Credential
err := hsdb.DB.Find(&keys).Error
err := hsdb.DB.Where("kind = ?", types.CredentialAPIKey).Find(&creds).Error
if err != nil {
return nil, err
}
keys := make([]types.APIKey, 0, len(creds))
for i := range creds {
keys = append(keys, *credentialToAPIKey(&creds[i]))
}
return keys, nil
}
// GetAPIKey returns a [types.APIKey] for a given key.
func (hsdb *HSDatabase) GetAPIKey(prefix string) (*types.APIKey, error) {
key := types.APIKey{}
if result := hsdb.DB.First(&key, "prefix = ?", prefix); result.Error != nil {
var cred types.Credential
if result := hsdb.DB.First(&cred, "kind = ? AND identifier = ?", types.CredentialAPIKey, prefix); result.Error != nil {
return nil, result.Error
}
return &key, nil
return credentialToAPIKey(&cred), nil
}
// GetAPIKeyByID returns a [types.APIKey] for a given id.
func (hsdb *HSDatabase) GetAPIKeyByID(id uint64) (*types.APIKey, error) {
key := types.APIKey{}
var cred types.Credential
// Query on an explicit primary-key clause: a struct condition would drop a
// zero-valued ID, making the lookup unconditional and returning the first
// row instead of not-found.
if result := hsdb.DB.First(&key, "id = ?", id); result.Error != nil {
if result := hsdb.DB.First(&cred, "kind = ? AND id = ?", types.CredentialAPIKey, id); result.Error != nil {
return nil, result.Error
}
return &key, nil
return credentialToAPIKey(&cred), nil
}
// DestroyAPIKey destroys a [types.APIKey]. Returns error if the [types.APIKey]
// does not exist.
func (hsdb *HSDatabase) DestroyAPIKey(key types.APIKey) error {
if result := hsdb.DB.Unscoped().Delete(key); result.Error != nil {
if result := hsdb.DB.Unscoped().
Delete(&types.Credential{}, "kind = ? AND id = ?", types.CredentialAPIKey, key.ID); result.Error != nil {
return result.Error
}
@ -94,12 +101,9 @@ func (hsdb *HSDatabase) DestroyAPIKey(key types.APIKey) error {
// ExpireAPIKey marks a [types.APIKey] as expired.
func (hsdb *HSDatabase) ExpireAPIKey(key *types.APIKey) error {
err := hsdb.DB.Model(&key).Update("Expiration", time.Now()).Error
if err != nil {
return err
}
return nil
return hsdb.DB.Model(&types.Credential{}).
Where("kind = ? AND id = ?", types.CredentialAPIKey, key.ID).
Update("expiration", time.Now()).Error
}
func (hsdb *HSDatabase) ValidateAPIKey(keyStr string) (bool, error) {
@ -135,8 +139,8 @@ func (hsdb *HSDatabase) AuthenticateAPIKey(keyStr string) (*types.APIKey, error)
// SetAPIKeyUser sets the owning user of an API key. Used when an admin mints a
// key on behalf of a user (headscale apikeys create --user).
func (hsdb *HSDatabase) SetAPIKeyUser(keyID uint64, userID types.UserID) error {
return hsdb.DB.Model(&types.APIKey{}).
Where("id = ?", keyID).
return hsdb.DB.Model(&types.Credential{}).
Where("kind = ? AND id = ?", types.CredentialAPIKey, keyID).
Update("user_id", uint(userID)).Error
}
@ -203,24 +207,24 @@ func validateAPIKey(db *gorm.DB, keyStr string) (*types.APIKey, error) {
return nil, err
}
// Look up by prefix (indexed)
var key types.APIKey
// Look up by identifier (indexed)
var cred types.Credential
err = db.First(&key, "prefix = ?", prefix).Error
err = db.First(&cred, "kind = ? AND identifier = ?", types.CredentialAPIKey, prefix).Error
if err != nil {
return nil, fmt.Errorf("API key not found: %w", err)
}
needsRehash, err := verifySecret(key.Hash, secret)
needsRehash, err := verifySecret(cred.Hash, secret)
if err != nil {
return nil, fmt.Errorf("invalid API key: %w", err)
}
if needsRehash {
rehashToArgon2id(db, &key, secret)
rehashToArgon2id(db, &cred, secret)
}
return &key, nil
return credentialToAPIKey(&cred), nil
}
// validateLegacyAPIKey validates a legacy format API key (prefix.secret).
@ -236,21 +240,21 @@ func validateLegacyAPIKey(db *gorm.DB, keyStr string) (*types.APIKey, error) {
return nil, fmt.Errorf("%w: legacy prefix length mismatch", ErrAPIKeyFailedToParse)
}
var key types.APIKey
var cred types.Credential
err := db.First(&key, "prefix = ?", prefix).Error
err := db.First(&cred, "kind = ? AND identifier = ?", types.CredentialAPIKey, prefix).Error
if err != nil {
return nil, fmt.Errorf("API key not found: %w", err)
}
needsRehash, err := verifySecret(key.Hash, secret)
needsRehash, err := verifySecret(cred.Hash, secret)
if err != nil {
return nil, fmt.Errorf("invalid API key: %w", err)
}
if needsRehash {
rehashToArgon2id(db, &key, secret)
rehashToArgon2id(db, &cred, secret)
}
return &key, nil
return credentialToAPIKey(&cred), nil
}

View file

@ -190,9 +190,9 @@ func TestAPIKeyWithPrefix(t *testing.T) {
now := time.Now()
err = db.DB.Exec(`
INSERT INTO api_keys (prefix, hash, created_at)
VALUES (?, ?, ?)
`, legacyPrefix, hash, now).Error
INSERT INTO credentials (kind, identifier, hash, created_at)
VALUES (?, ?, ?, ?)
`, types.CredentialAPIKey, legacyPrefix, hash, now).Error
require.NoError(t, err)
// Validate legacy key
@ -289,8 +289,8 @@ func TestAPIKeyLazyRehashesBcrypt(t *testing.T) {
require.NoError(t, err)
err = db.DB.Exec(
`INSERT INTO api_keys (prefix, hash, created_at) VALUES (?, ?, ?)`,
prefix, hash, time.Now(),
`INSERT INTO credentials (kind, identifier, hash, created_at) VALUES (?, ?, ?, ?)`,
types.CredentialAPIKey, prefix, hash, time.Now(),
).Error
require.NoError(t, err)

View file

@ -0,0 +1,19 @@
package db
import (
"github.com/juanfont/headscale/hscontrol/types"
)
// credentialToAPIKey projects a unified credentials row onto the [types.APIKey]
// shape the API, state, and CLI layers still consume.
func credentialToAPIKey(c *types.Credential) *types.APIKey {
return &types.APIKey{
ID: c.ID,
Prefix: c.Identifier,
Hash: c.Hash,
UserID: c.UserID,
CreatedAt: c.CreatedAt,
Expiration: c.Expiration,
LastSeen: c.LastSeen,
}
}