db: store OAuth clients and tokens in the credentials table

Via the OAuthClient and OAuthAccessToken projections.
This commit is contained in:
Kristoffer Dalby 2026-06-27 12:20:08 +00:00
parent e276a94f25
commit 7acf557102
3 changed files with 74 additions and 34 deletions

View file

@ -17,3 +17,35 @@ func credentialToAPIKey(c *types.Credential) *types.APIKey {
LastSeen: c.LastSeen,
}
}
// credentialToOAuthClient projects a unified credentials row onto the
// [types.OAuthClient] shape. The client id is stored as the row's identifier.
func credentialToOAuthClient(c *types.Credential) *types.OAuthClient {
return &types.OAuthClient{
ID: c.ID,
ClientID: c.Identifier,
SecretHash: c.Hash,
Scopes: c.Scopes,
Tags: c.Tags,
Description: c.Description,
UserID: c.UserID,
CreatedAt: c.CreatedAt,
Revoked: c.Revoked,
}
}
// credentialToOAuthAccessToken projects a unified credentials row onto the
// [types.OAuthAccessToken] shape. The token's lookup prefix is stored as the
// row's identifier; ClientID links back to the issuing client's identifier.
func credentialToOAuthAccessToken(c *types.Credential) *types.OAuthAccessToken {
return &types.OAuthAccessToken{
ID: c.ID,
Prefix: c.Identifier,
Hash: c.Hash,
ClientID: c.ClientID,
Scopes: c.Scopes,
Tags: c.Tags,
Expiration: c.Expiration,
CreatedAt: c.CreatedAt,
}
}

View file

@ -60,9 +60,10 @@ func (hsdb *HSDatabase) CreateOAuthClient(
}
now := time.Now().UTC()
client := types.OAuthClient{
ClientID: clientID,
SecretHash: hash,
cred := types.Credential{
Kind: types.CredentialOAuthClient,
Identifier: clientID,
Hash: hash,
Scopes: scopes,
Tags: tags,
Description: description,
@ -71,13 +72,13 @@ func (hsdb *HSDatabase) CreateOAuthClient(
}
err = hsdb.Write(func(tx *gorm.DB) error {
return tx.Save(&client).Error
return tx.Save(&cred).Error
})
if err != nil {
return "", nil, fmt.Errorf("saving oauth client: %w", err)
}
return secretStr, &client, nil
return secretStr, credentialToOAuthClient(&cred), nil
}
// AuthenticateOAuthClient validates a presented client secret and returns the
@ -109,41 +110,46 @@ func (hsdb *HSDatabase) AuthenticateOAuthClient(secretStr string) (*types.OAuthC
return nil, err
}
var client types.OAuthClient
if err := hsdb.DB.First(&client, "client_id = ?", clientID).Error; err != nil { //nolint:noinlineerr
var cred types.Credential
if err := hsdb.DB.First(&cred, "kind = ? AND identifier = ?", types.CredentialOAuthClient, clientID).Error; err != nil { //nolint:noinlineerr
return nil, ErrOAuthClientNotFound
}
if _, err := verifySecret(client.SecretHash, secret); err != nil { //nolint:noinlineerr
if _, err := verifySecret(cred.Hash, secret); err != nil { //nolint:noinlineerr
return nil, fmt.Errorf("invalid oauth client secret: %w", err)
}
if client.Revoked != nil {
if cred.Revoked != nil {
return nil, ErrOAuthClientRevoked
}
return &client, nil
return credentialToOAuthClient(&cred), nil
}
// GetOAuthClientByClientID returns a [types.OAuthClient] by its public client id.
func (hsdb *HSDatabase) GetOAuthClientByClientID(clientID string) (*types.OAuthClient, error) {
var client types.OAuthClient
if result := hsdb.DB.First(&client, "client_id = ?", clientID); result.Error != nil {
var cred types.Credential
if result := hsdb.DB.First(&cred, "kind = ? AND identifier = ?", types.CredentialOAuthClient, clientID); result.Error != nil {
return nil, result.Error
}
return &client, nil
return credentialToOAuthClient(&cred), nil
}
// ListOAuthClients returns every [types.OAuthClient].
func (hsdb *HSDatabase) ListOAuthClients() ([]types.OAuthClient, error) {
clients := []types.OAuthClient{}
var creds []types.Credential
err := hsdb.DB.Find(&clients).Error
err := hsdb.DB.Where("kind = ?", types.CredentialOAuthClient).Find(&creds).Error
if err != nil {
return nil, err
}
clients := make([]types.OAuthClient, 0, len(creds))
for i := range creds {
clients = append(clients, *credentialToOAuthClient(&creds[i]))
}
return clients, nil
}
@ -153,13 +159,14 @@ func (hsdb *HSDatabase) ListOAuthClients() ([]types.OAuthClient, error) {
// OAuth client has no such history and is removed outright, matching Tailscale.
func (hsdb *HSDatabase) RevokeOAuthClient(clientID string) error {
return hsdb.Write(func(tx *gorm.DB) error {
err := tx.Where("client_id = ?", clientID).
Delete(&types.OAuthAccessToken{}).Error
err := tx.Where("kind = ? AND client_id = ?", types.CredentialOAuthToken, clientID).
Delete(&types.Credential{}).Error
if err != nil {
return fmt.Errorf("deleting oauth access tokens: %w", err)
}
res := tx.Where("client_id = ?", clientID).Delete(&types.OAuthClient{})
res := tx.Where("kind = ? AND identifier = ?", types.CredentialOAuthClient, clientID).
Delete(&types.Credential{})
if res.Error != nil {
return res.Error
}
@ -186,8 +193,9 @@ func (hsdb *HSDatabase) MintAccessToken(
}
now := time.Now().UTC()
token := types.OAuthAccessToken{
Prefix: prefix,
cred := types.Credential{
Kind: types.CredentialOAuthToken,
Identifier: prefix,
Hash: hash,
ClientID: clientID,
Scopes: scopes,
@ -199,9 +207,9 @@ func (hsdb *HSDatabase) MintAccessToken(
// Mint inside a transaction that re-checks the client still exists and is
// not revoked, so a mint cannot complete against a client being deleted.
err = hsdb.Write(func(tx *gorm.DB) error {
var client types.OAuthClient
var client types.Credential
err := tx.First(&client, "client_id = ?", clientID).Error
err := tx.First(&client, "kind = ? AND identifier = ?", types.CredentialOAuthClient, clientID).Error
if err != nil {
return ErrOAuthClientNotFound
}
@ -210,13 +218,13 @@ func (hsdb *HSDatabase) MintAccessToken(
return ErrOAuthClientRevoked
}
return tx.Save(&token).Error
return tx.Save(&cred).Error
})
if err != nil {
return "", nil, fmt.Errorf("saving oauth access token: %w", err)
}
return tokenStr, &token, nil
return tokenStr, credentialToOAuthAccessToken(&cred), nil
}
// AuthenticateAccessToken validates a presented bearer token and returns the
@ -242,8 +250,8 @@ func (hsdb *HSDatabase) AuthenticateAccessToken(tokenStr string) (*types.OAuthAc
return nil, err
}
var token types.OAuthAccessToken
if err := hsdb.DB.First(&token, "prefix = ?", prefix).Error; err != nil { //nolint:noinlineerr
var token types.Credential
if err := hsdb.DB.First(&token, "kind = ? AND identifier = ?", types.CredentialOAuthToken, prefix).Error; err != nil { //nolint:noinlineerr
return nil, ErrAccessTokenNotFound
}
@ -259,8 +267,8 @@ func (hsdb *HSDatabase) AuthenticateAccessToken(tokenStr string) (*types.OAuthAc
// revoked or deleted is rejected. This closes a mint/revoke race (where a
// token could be inserted after the client's tokens were purged) and any
// orphan left by manual deletion or a future soft-revoke path.
var client types.OAuthClient
if err := hsdb.DB.First(&client, "client_id = ?", token.ClientID).Error; err != nil { //nolint:noinlineerr
var client types.Credential
if err := hsdb.DB.First(&client, "kind = ? AND identifier = ?", types.CredentialOAuthClient, token.ClientID).Error; err != nil { //nolint:noinlineerr
return nil, ErrAccessTokenClientRevoked
}
@ -268,7 +276,7 @@ func (hsdb *HSDatabase) AuthenticateAccessToken(tokenStr string) (*types.OAuthAc
return nil, ErrAccessTokenClientRevoked
}
return &token, nil
return credentialToOAuthAccessToken(&token), nil
}
// DeleteExpiredAccessTokens hard-deletes every access token that expired before
@ -276,8 +284,8 @@ func (hsdb *HSDatabase) AuthenticateAccessToken(tokenStr string) (*types.OAuthAc
// expired tokens; the hourly reaper (see app.go) calls this only to keep the
// table from growing unbounded.
func (hsdb *HSDatabase) DeleteExpiredAccessTokens(cutoff time.Time) (int64, error) {
res := hsdb.DB.Where("expiration IS NOT NULL AND expiration < ?", cutoff).
Delete(&types.OAuthAccessToken{})
res := hsdb.DB.Where("kind = ? AND expiration IS NOT NULL AND expiration < ?", types.CredentialOAuthToken, cutoff).
Delete(&types.Credential{})
return res.RowsAffected, res.Error
}

View file

@ -188,7 +188,7 @@ func TestAccessTokenRejectedWhenClientGone(t *testing.T) {
// Delete only the client row, leaving the token orphaned (the state a
// mint/revoke race or manual deletion would produce).
require.NoError(t, db.DB.Where("client_id = ?", client.ClientID).Delete(&types.OAuthClient{}).Error)
require.NoError(t, db.DB.Where("kind = ? AND identifier = ?", types.CredentialOAuthClient, client.ClientID).Delete(&types.Credential{}).Error)
_, err = db.AuthenticateAccessToken(tokenStr)
require.ErrorIs(t, err, ErrAccessTokenClientRevoked)
@ -201,8 +201,8 @@ func TestAccessTokenRejectedWhenClientGone(t *testing.T) {
require.NoError(t, err)
now := time.Now()
require.NoError(t, db.DB.Model(&types.OAuthClient{}).
Where("client_id = ?", client2.ClientID).Update("revoked", now).Error)
require.NoError(t, db.DB.Model(&types.Credential{}).
Where("kind = ? AND identifier = ?", types.CredentialOAuthClient, client2.ClientID).Update("revoked", now).Error)
_, err = db.AuthenticateAccessToken(tokenStr2)
require.ErrorIs(t, err, ErrAccessTokenClientRevoked)