diff --git a/hscontrol/db/credential.go b/hscontrol/db/credential.go index be7deefdc..39e95d8ee 100644 --- a/hscontrol/db/credential.go +++ b/hscontrol/db/credential.go @@ -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, + } +} diff --git a/hscontrol/db/oauth.go b/hscontrol/db/oauth.go index ec205aa24..d957f1fa8 100644 --- a/hscontrol/db/oauth.go +++ b/hscontrol/db/oauth.go @@ -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 } diff --git a/hscontrol/db/oauth_test.go b/hscontrol/db/oauth_test.go index 93be4b7e5..20caab232 100644 --- a/hscontrol/db/oauth_test.go +++ b/hscontrol/db/oauth_test.go @@ -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)