Files
ResolveSpec/pkg/security/lookup/procedure/oauth.go
T
Hein 640faeeeaf feat(security): full OAuth 2.1 / OpenID Connect server and OIDC relying-party client
Authorization server: consent and scopes, OIDC (nonce, auth_time, acr, sid,
at_hash, signed userinfo, RP-initiated and back-channel logout), managed
refresh tokens with rotation and reuse detection, RFC 9068 JWT access tokens,
DPoP, PAR, device grant, token exchange, private_key_jwt, RFC 7591/7592
registration, RFC 9207 iss, signing keyring with rotation.

State is DB-backed through a new lookup.OAuthGrantStore (procedure and direct
backends, four dialect DDLs, conformance cases).

Client side: WithOIDC discovery, PKCE, nonce, id_token validation, OAuth2LogoutURL.

PeekRefresh now returns already rotated tokens so RotateRefresh can detect reuse.

Docs: OAUTH2_SERVER.md, oauth2_full_example.go, breaking_changes.md step 8.
2026-10-01 14:42:12 +02:00

351 lines
11 KiB
Go

package procedure
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// OAuthUsers implements lookup.OAuthUserStore with the resolvespec_oauth_* procedures.
type OAuthUsers struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.OAuthUserStore = (*OAuthUsers)(nil)
// NewOAuthUsers creates the procedure-backed OAuthUserStore.
func NewOAuthUsers(run Runner, procs lookup.ProcNames) *OAuthUsers {
return &OAuthUsers{run: run, procs: procs}
}
// GetOrCreateUser implements lookup.OAuthUserStore.
func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
userJSON, err := json.Marshal(map[string]any{
"username": user.UserName,
"email": user.Email,
"remote_id": user.RemoteID,
"user_level": user.UserLevel,
"roles": user.Roles,
"auth_provider": provider,
})
if err != nil {
return 0, fmt.Errorf("failed to marshal user data: %w", err)
}
var success bool
var errMsg sql.NullString
var userID sql.NullInt64
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_user_id
FROM %s($1::jsonb)
`, o.procs.OAuthGetOrCreateUser), userJSON).Scan(&success, &errMsg, &userID)
})
if err != nil {
return 0, fmt.Errorf("failed to get or create user: %w", err)
}
if !success {
return 0, failure(errMsg, "failed to get or create user")
}
if !userID.Valid {
return 0, fmt.Errorf("user ID not returned")
}
return int(userID.Int64), nil
}
// CreateSession implements lookup.OAuthUserStore.
func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error {
sessionJSON, err := json.Marshal(map[string]any{
"session_token": s.SessionToken,
"user_id": s.UserID,
"access_token": s.AccessToken,
"refresh_token": s.RefreshToken,
"token_type": s.TokenType,
"expires_at": s.ExpiresAt,
"auth_provider": s.Provider,
})
if err != nil {
return fmt.Errorf("failed to marshal session data: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, o.procs.OAuthCreateSession), sessionJSON).Scan(&success, &errMsg)
})
if err != nil {
return fmt.Errorf("failed to create session: %w", err)
}
if !success {
return failure(errMsg, "failed to create session")
}
return nil
}
// GetByRefreshToken implements lookup.OAuthUserStore.
func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
var success bool
var errMsg sql.NullString
var data []byte
err := o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, o.procs.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &data)
})
if err != nil {
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
}
if !success {
return nil, failure(errMsg, "invalid or expired refresh token")
}
var session lookup.OAuthRefreshSession
if err := json.Unmarshal(normalizeTimes(data), &session); err != nil {
return nil, fmt.Errorf("failed to parse session data: %w", err)
}
return &session, nil
}
// UpdateRefreshToken implements lookup.OAuthUserStore.
func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
updateJSON, err := json.Marshal(map[string]any{
"user_id": userID,
"old_refresh_token": oldRefreshToken,
"new_session_token": newSessionToken,
"new_access_token": newAccessToken,
"new_refresh_token": newRefreshToken,
"expires_at": expiresAt,
})
if err != nil {
return fmt.Errorf("failed to marshal update data: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, o.procs.OAuthUpdateRefreshToken), updateJSON).Scan(&success, &errMsg)
})
if err != nil {
return fmt.Errorf("failed to update session: %w", err)
}
if !success {
return failure(errMsg, "failed to update session")
}
return nil
}
// GetUser implements lookup.OAuthUserStore.
func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
var success bool
var errMsg sql.NullString
var data []byte
err := o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, o.procs.OAuthGetUser), userID).Scan(&success, &errMsg, &data)
})
if err != nil {
return nil, fmt.Errorf("failed to get user data: %w", err)
}
if !success {
return nil, failure(errMsg, "failed to get user data")
}
var userCtx sectypes.UserContext
if err := json.Unmarshal(data, &userCtx); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
return &userCtx, nil
}
// OAuthClients implements lookup.OAuthClientStore with the resolvespec_oauth_* server procedures.
type OAuthClients struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.OAuthClientStore = (*OAuthClients)(nil)
// NewOAuthClients creates the procedure-backed OAuthClientStore.
func NewOAuthClients(run Runner, procs lookup.ProcNames) *OAuthClients {
return &OAuthClients{run: run, procs: procs}
}
// callData runs a `(p_success, p_error, p_data)` procedure with one argument.
func (o *OAuthClients) callData(ctx context.Context, proc string, arg any) (data []byte, ok bool, errMsg sql.NullString, err error) {
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, proc), arg).Scan(&ok, &errMsg, &data)
})
return
}
// callNoData runs a `(p_success, p_error)` procedure with one argument.
func (o *OAuthClients) callNoData(ctx context.Context, proc string, arg any) (ok bool, errMsg sql.NullString, err error) {
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1)
`, proc), arg).Scan(&ok, &errMsg)
})
return
}
// RegisterClient implements lookup.OAuthClientStore.
func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
input, err := json.Marshal(client)
if err != nil {
return nil, fmt.Errorf("failed to marshal client: %w", err)
}
var success bool
var errMsg sql.NullString
var data []byte
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1::jsonb)
`, o.procs.OAuthRegisterClient), input).Scan(&success, &errMsg, &data)
})
if err != nil {
return nil, fmt.Errorf("failed to register client: %w", err)
}
if !success {
return nil, failure(errMsg, "failed to register client")
}
var result sectypes.OAuthServerClient
if err := json.Unmarshal(data, &result); err != nil {
return nil, fmt.Errorf("failed to parse registered client: %w", err)
}
return &result, nil
}
// GetClient implements lookup.OAuthClientStore.
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthGetClient, clientID)
if err != nil {
return nil, fmt.Errorf("failed to get client: %w", err)
}
if !ok {
return nil, failure(errMsg, "client not found")
}
var result sectypes.OAuthServerClient
if err := json.Unmarshal(data, &result); err != nil {
return nil, fmt.Errorf("failed to parse client: %w", err)
}
return &result, nil
}
// SaveCode implements lookup.OAuthClientStore.
func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
input, err := json.Marshal(code) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return fmt.Errorf("failed to marshal code: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, o.procs.OAuthSaveCode), input).Scan(&success, &errMsg)
})
if err != nil {
return fmt.Errorf("failed to save code: %w", err)
}
if !success {
return failure(errMsg, "failed to save code")
}
return nil
}
// ExchangeCode implements lookup.OAuthClientStore.
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthExchangeCode, code)
if err != nil {
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
if !ok {
return nil, failure(errMsg, "invalid or expired code")
}
var result sectypes.OAuthCode
if err := json.Unmarshal(normalizeTimes(data), &result); err != nil {
return nil, fmt.Errorf("failed to parse code data: %w", err)
}
result.Code = code
return &result, nil
}
// Introspect implements lookup.OAuthClientStore.
func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthIntrospect, token)
if err != nil {
return nil, fmt.Errorf("failed to introspect token: %w", err)
}
if !ok {
return nil, failure(errMsg, "introspection failed")
}
var result sectypes.OAuthTokenInfo
if err := json.Unmarshal(data, &result); err != nil {
return nil, fmt.Errorf("failed to parse token info: %w", err)
}
return &result, nil
}
// Revoke implements lookup.OAuthClientStore.
func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthRevoke, token)
if err != nil {
return fmt.Errorf("failed to revoke token: %w", err)
}
if !ok {
return failure(errMsg, "failed to revoke token")
}
return nil
}
// UpdateClient implements lookup.OAuthClientStore.
func (o *OAuthClients) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
input, err := json.Marshal(client)
if err != nil {
return fmt.Errorf("failed to marshal client: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, o.procs.OAuthUpdateClient), input).Scan(&success, &errMsg)
})
if err != nil {
return fmt.Errorf("failed to update client: %w", err)
}
if !success {
return failure(errMsg, "failed to update client")
}
return nil
}
// DeleteClient implements lookup.OAuthClientStore.
func (o *OAuthClients) DeleteClient(ctx context.Context, clientID string) error {
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthDeleteClient, clientID)
if err != nil {
return fmt.Errorf("failed to delete client: %w", err)
}
if !ok {
return failure(errMsg, "failed to delete client")
}
return nil
}