Files
ResolveSpec/pkg/security/lookup/procedure/oauth.go
T
Hein c9fa8c60f2 refactor(security): move all database access into pkg/security/lookup
pkg/security no longer contains SQL. Every provider calls a store interface
from lookup, implemented by a procedure backend (Postgres stored procedures,
the default there) and a direct backend (dialect-driven SQL for postgres,
sqlite, mysql and mssql with configurable table and column names).

- add sectypes, lookup, lookup/{dialect,procedure,direct,backends,ddl,conformance}
- split totp and providers sub packages out of the core package
- replace SQLNames/TableNames/QueryMode with lookup.Config (see breaking_changes.md)
- direct backend now covers column/row security and API-key login
- move txsettings SQL to lookup.ApplyTxSettings; remove password.go
- move schema scripts under lookup/, add reference DDL per dialect
- add a shared conformance suite; run it on sqlite, and on Postgres in a
  podman/docker container (RESOLVESPEC_TEST_CONTAINERS=1)
- fix procedure schema bugs found on real Postgres: duplicate p_data
  parameter, JSON null arrays, expires_at timezone casts, passkey list
  GROUP BY, missing resolvespec_passkey_login; accept zone-less timestamps
2026-10-01 13:19:44 +02:00

316 lines
9.9 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
}