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
This commit is contained in:
Hein
2026-10-01 13:19:44 +02:00
parent 60bd0a6dd3
commit c9fa8c60f2
118 changed files with 11218 additions and 5565 deletions
+318
View File
@@ -0,0 +1,318 @@
package procedure
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Auth implements lookup.AuthStore with stored procedures.
type Auth struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.AuthStore = (*Auth)(nil)
// NewAuth creates the procedure-backed AuthStore.
func NewAuth(run Runner, procs lookup.ProcNames) *Auth { return &Auth{run: run, procs: procs} }
// callData runs "SELECT p_success, p_error, p_data::text FROM proc($1::jsonb)".
func (a *Auth) callData(ctx context.Context, proc, queryErrOp string, arg any) (sql.NullString, error) {
var success bool
var errorMsg, dataJSON sql.NullString
err := a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, proc) //nolint:gosec // G201: identifier comes from validated config
return db.QueryRowContext(ctx, query, arg).Scan(&success, &errorMsg, &dataJSON)
})
if err != nil {
return sql.NullString{}, fmt.Errorf("%s query failed: %w", queryErrOp, err)
}
if !success {
return sql.NullString{}, failure(errorMsg, queryErrOp+" failed")
}
return dataJSON, nil
}
// Login implements lookup.AuthStore.
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return nil, fmt.Errorf("failed to marshal login request: %w", err)
}
data, err := a.callData(ctx, a.procs.Login, "login", string(reqJSON))
if err != nil {
return nil, err
}
var response sectypes.LoginResponse
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
return nil, fmt.Errorf("failed to parse login response: %w", err)
}
return &response, nil
}
// Register implements lookup.AuthStore.
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return nil, fmt.Errorf("failed to marshal register request: %w", err)
}
var success bool
var errorMsg, dataJSON sql.NullString
err = a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.Register)
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
})
if err != nil {
return nil, fmt.Errorf("register query failed: %w", err)
}
if !success {
return nil, failure(errorMsg, "registration failed")
}
var response sectypes.LoginResponse
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
return nil, fmt.Errorf("failed to parse register response: %w", err)
}
return &response, nil
}
// Logout implements lookup.AuthStore.
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
reqJSON, err := json.Marshal(req)
if err != nil {
return fmt.Errorf("failed to marshal logout request: %w", err)
}
_, err = a.callData(ctx, a.procs.Logout, "logout", string(reqJSON))
return err
}
// Session implements lookup.AuthStore.
func (a *Auth) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) {
var success bool
var errorMsg, userJSON sql.NullString
err := a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
return db.QueryRowContext(ctx, query, token, reference).Scan(&success, &errorMsg, &userJSON)
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
}
if !success {
return nil, failure(errorMsg, "invalid or expired session")
}
if !userJSON.Valid {
return nil, fmt.Errorf("no user data in session")
}
var user sectypes.UserContext
if err := json.Unmarshal([]byte(userJSON.String), &user); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
return &user, nil
}
// TouchSession implements lookup.AuthStore.
func (a *Auth) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error {
userJSON, err := json.Marshal(user)
if err != nil {
return err
}
var success bool
var errorMsg, updatedUserJSON sql.NullString
return a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.SessionUpdate)
return db.QueryRowContext(ctx, query, token, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
}
// Refresh implements lookup.AuthStore.
func (a *Auth) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) {
// Get the current session to pass to refresh.
var success bool
var errorMsg, userJSON sql.NullString
err := a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session)
return db.QueryRowContext(ctx, query, refreshToken, "refresh").Scan(&success, &errorMsg, &userJSON)
})
if err != nil {
return nil, fmt.Errorf("refresh token query failed: %w", err)
}
if !success {
return nil, failure(errorMsg, "invalid refresh token")
}
var newSuccess bool
var newErrorMsg, newUserJSON sql.NullString
err = a.run.Run(func(db *sql.DB) error {
refreshQuery := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.RefreshToken)
return db.QueryRowContext(ctx, refreshQuery, refreshToken, userJSON).Scan(&newSuccess, &newErrorMsg, &newUserJSON)
})
if err != nil {
return nil, fmt.Errorf("refresh token generation failed: %w", err)
}
if !newSuccess {
return nil, failure(newErrorMsg, "failed to refresh token")
}
var userCtx sectypes.UserContext
if err := json.Unmarshal([]byte(newUserJSON.String), &userCtx); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
// A resolvespec_refresh_token implementation that issues its own rotating
// refresh token (independent of the access/session token) returns it
// under claims.refresh_token, since UserContext has no dedicated field
// for it. Surface that into LoginResponse.RefreshToken so callers don't
// need to reach into User.Claims themselves. claims.expires_in
// (seconds) similarly overrides the default access-token ExpiresIn when
// the procedure provides a real value. Implementations that don't set
// these claims keep today's behavior unchanged (empty RefreshToken,
// 24h ExpiresIn default).
resp := &sectypes.LoginResponse{
Token: userCtx.SessionID, // New session token from stored procedure
User: &userCtx,
ExpiresIn: int64(24 * time.Hour.Seconds()),
}
if rt, ok := userCtx.Claims["refresh_token"].(string); ok && rt != "" {
resp.RefreshToken = rt
}
if expiresIn, ok := userCtx.Claims["expires_in"].(float64); ok && expiresIn > 0 {
resp.ExpiresIn = int64(expiresIn)
}
return resp, nil
}
// LoginAPIKey implements lookup.AuthStore. Unknown, expired and inactive keys all return
// lookup.ErrInvalidAPIKey; the raw key is never logged.
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
if rawKey == "" {
return nil, lookup.ErrInvalidAPIKey
}
reqJSON, err := json.Marshal(map[string]any{"api_key": rawKey, "claims": claims})
if err != nil {
return nil, fmt.Errorf("failed to marshal api key login request: %w", err)
}
var success bool
var errorMsg, dataJSON sql.NullString
err = a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.LoginAPIKey)
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
})
if err != nil {
return nil, fmt.Errorf("api key login query failed: %w", err)
}
if !success {
return nil, lookup.ErrInvalidAPIKey
}
var response sectypes.LoginResponse
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
return nil, fmt.Errorf("failed to parse api key login response: %w", err)
}
return &response, nil
}
// JWTLogin implements lookup.AuthStore. The password is verified inside the procedure;
// the hash is never returned. The token is a placeholder until JWT signing is wired in.
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
var success bool
var errorMsg sql.NullString
var userJSON []byte
err := a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user FROM %s($1, $2)`, a.procs.JWTLogin)
return db.QueryRowContext(ctx, query, req.Username, req.Password).Scan(&success, &errorMsg, &userJSON)
})
if err != nil {
return nil, fmt.Errorf("login query failed: %w", err)
}
if !success {
return nil, failure(errorMsg, "invalid credentials")
}
var user struct {
ID int `json:"id"`
Username string `json:"username"`
Email string `json:"email"`
UserLevel int `json:"user_level"`
Roles string `json:"roles"`
}
if err := json.Unmarshal(userJSON, &user); err != nil {
return nil, fmt.Errorf("failed to parse user data: %w", err)
}
roles := []string{}
if user.Roles != "" {
roles = strings.Split(user.Roles, ",")
}
expiresAt := time.Now().Add(24 * time.Hour)
return &sectypes.LoginResponse{
Token: fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix()),
User: &sectypes.UserContext{
UserID: user.ID,
UserName: user.Username,
Email: user.Email,
UserLevel: user.UserLevel,
Roles: roles,
},
ExpiresIn: int64(24 * time.Hour.Seconds()),
}, nil
}
// JWTLogout implements lookup.AuthStore.
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
var success bool
var errorMsg sql.NullString
err := a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, a.procs.JWTLogout)
return db.QueryRowContext(ctx, query, req.Token, req.UserID).Scan(&success, &errorMsg)
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
if !success {
return failure(errorMsg, "logout failed")
}
return nil
}
// ResetRequest implements lookup.AuthStore.
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
reqJSON, err := json.Marshal(req)
if err != nil {
return nil, fmt.Errorf("failed to marshal password reset request: %w", err)
}
data, err := a.callData(ctx, a.procs.PasswordResetRequest, "password reset request", string(reqJSON))
if err != nil {
return nil, err
}
var response sectypes.PasswordResetResponse
if data.Valid && data.String != "" {
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
return nil, fmt.Errorf("failed to parse password reset response: %w", err)
}
}
return &response, nil
}
// ResetComplete implements lookup.AuthStore.
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
reqJSON, err := json.Marshal(req)
if err != nil {
return fmt.Errorf("failed to marshal password reset complete request: %w", err)
}
var success bool
var errorMsg sql.NullString
err = a.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1::jsonb)`, a.procs.PasswordResetComplete)
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg)
})
if err != nil {
return fmt.Errorf("password reset complete query failed: %w", err)
}
if !success {
return failure(errorMsg, "password reset failed")
}
return nil
}
+54
View File
@@ -0,0 +1,54 @@
package procedure
import (
"encoding/json"
"strings"
"time"
)
// normalizeTimes rewrites the zone-less timestamps the procedures emit (Postgres `timestamp`
// columns serialise as "2026-01-02T03:04:05.123456") as UTC RFC 3339, so the standard
// time.Time decoder accepts them. It handles a JSON object or an array of objects and only
// touches string fields whose name ends in "_at" or is "expiry". Anything else, including
// input that is not valid JSON, is returned unchanged.
func normalizeTimes(raw []byte) []byte {
var v any
if err := json.Unmarshal(raw, &v); err != nil {
return raw
}
switch x := v.(type) {
case map[string]any:
fixTimes(x)
case []any:
for _, e := range x {
if m, ok := e.(map[string]any); ok {
fixTimes(m)
}
}
default:
return raw
}
out, err := json.Marshal(v)
if err != nil {
return raw
}
return out
}
func fixTimes(m map[string]any) {
for k, v := range m {
s, ok := v.(string)
if !ok || !(strings.HasSuffix(k, "_at") || k == "expiry") {
continue
}
if _, err := time.Parse(time.RFC3339Nano, s); err == nil {
continue
}
for _, layout := range []string{"2006-01-02T15:04:05.999999999", "2006-01-02 15:04:05.999999999"} {
if t, err := time.Parse(layout, s); err == nil {
m[k] = t.UTC().Format(time.RFC3339Nano)
break
}
}
}
}
+156
View File
@@ -0,0 +1,156 @@
package procedure
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Keys implements lookup.KeyStore with the resolvespec_keystore_* procedures.
type Keys struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.KeyStore = (*Keys)(nil)
// NewKeys creates the procedure-backed KeyStore.
func NewKeys(run Runner, procs lookup.ProcNames) *Keys { return &Keys{run: run, procs: procs} }
// orDefault returns the procedure's error message when it is non-empty, otherwise def.
func orDefault(s sql.NullString, def string) string {
if s.Valid && s.String != "" {
return s.String
}
return def
}
// Create implements lookup.KeyStore.
func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
type createRequest struct {
UserID int `json:"user_id"`
KeyType sectypes.KeyType `json:"key_type"`
KeyHash string `json:"key_hash"`
Name string `json:"name"`
Scopes []string `json:"scopes,omitempty"`
Meta map[string]any `json:"meta,omitempty"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
}
reqJSON, err := json.Marshal(createRequest{
UserID: req.UserID,
KeyType: req.KeyType,
KeyHash: keyHash,
Name: req.Name,
Scopes: req.Scopes,
Meta: req.Meta,
ExpiresAt: req.ExpiresAt,
})
if err != nil {
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
}
var success bool
var errorMsg, keyJSON sql.NullString
err = k.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, k.procs.KeystoreCreateKey)
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON)
})
if err != nil {
return nil, fmt.Errorf("create key procedure failed: %w", err)
}
if !success {
return nil, errors.New(orDefault(errorMsg, "create key failed"))
}
key, err := decodeKey([]byte(keyJSON.String))
if err != nil {
return nil, fmt.Errorf("failed to parse created key: %w", err)
}
return key, nil
}
// List implements lookup.KeyStore.
func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
var success bool
var errorMsg, keysJSON sql.NullString
err := k.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, k.procs.KeystoreGetUserKeys)
return db.QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON)
})
if err != nil {
return nil, fmt.Errorf("get user keys procedure failed: %w", err)
}
if !success {
return nil, errors.New(orDefault(errorMsg, "get user keys failed"))
}
var keys []sectypes.UserKey
if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" {
var raw []json.RawMessage
if err := json.Unmarshal([]byte(keysJSON.String), &raw); err != nil {
return nil, fmt.Errorf("failed to parse user keys: %w", err)
}
for _, r := range raw {
k, err := decodeKey(r)
if err != nil {
return nil, fmt.Errorf("failed to parse user keys: %w", err)
}
keys = append(keys, *k)
}
}
if keys == nil {
keys = []sectypes.UserKey{}
}
return keys, nil
}
// Delete implements lookup.KeyStore. The procedure returns the key hash.
func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
var success bool
var errorMsg, keyHash sql.NullString
err := k.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, k.procs.KeystoreDeleteKey)
return db.QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash)
})
if err != nil {
return "", fmt.Errorf("delete key procedure failed: %w", err)
}
if !success {
return "", errors.New(orDefault(errorMsg, "delete key failed"))
}
return keyHash.String, nil
}
// Validate implements lookup.KeyStore.
func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
var success bool
var errorMsg, keyJSON sql.NullString
err := k.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, k.procs.KeystoreValidateKey)
return db.QueryRowContext(ctx, query, keyHash, string(keyType)).Scan(&success, &errorMsg, &keyJSON)
})
if err != nil {
return nil, fmt.Errorf("validate key procedure failed: %w", err)
}
if !success {
return nil, errors.New(orDefault(errorMsg, "invalid or expired key"))
}
key, err := decodeKey([]byte(keyJSON.String))
if err != nil {
return nil, fmt.Errorf("failed to parse validated key: %w", err)
}
return key, nil
}
// decodeKey reads one key record from a key procedure.
func decodeKey(raw []byte) (*sectypes.UserKey, error) {
var k sectypes.UserKey
if err := json.Unmarshal(normalizeTimes(raw), &k); err != nil {
return nil, err
}
return &k, nil
}
+315
View File
@@ -0,0 +1,315 @@
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
}
+282
View File
@@ -0,0 +1,282 @@
package procedure
import (
"context"
"database/sql"
"encoding/base64"
"encoding/json"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Passkey implements lookup.PasskeyStore with the resolvespec_passkey_* procedures.
// Credential ids cross the lookup interface as base64 text; the procedures that take a
// bytea credential id receive the decoded bytes.
type Passkey struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.PasskeyStore = (*Passkey)(nil)
// NewPasskey creates the procedure-backed PasskeyStore.
func NewPasskey(run Runner, procs lookup.ProcNames) *Passkey {
return &Passkey{run: run, procs: procs}
}
func decodeCredentialID(b64 string) ([]byte, error) {
id, err := base64.StdEncoding.DecodeString(b64)
if err != nil {
return nil, fmt.Errorf("invalid credential ID: %w", err)
}
return id, nil
}
// Store implements lookup.PasskeyStore.
func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
credJSON, err := json.Marshal(map[string]any{
"user_id": rec.UserID,
"credential_id": rec.CredentialID,
"public_key": rec.PublicKey,
"attestation_type": rec.AttestationType,
"sign_count": rec.SignCount,
"transports": rec.Transports,
"backup_eligible": rec.BackupEligible,
"backup_state": rec.BackupState,
"name": rec.Name,
})
if err != nil {
return 0, fmt.Errorf("failed to marshal credential data: %w", err)
}
var success bool
var errorMsg sql.NullString
var credentialID sql.NullInt64
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential_id FROM %s($1::jsonb)`, p.procs.PasskeyStoreCredential)
return db.QueryRowContext(ctx, query, string(credJSON)).Scan(&success, &errorMsg, &credentialID)
})
if err != nil {
return 0, fmt.Errorf("failed to store credential: %w", err)
}
if !success {
return 0, failure(errorMsg, "failed to store credential")
}
return credentialID.Int64, nil
}
// Get implements lookup.PasskeyStore.
func (p *Passkey) Get(ctx context.Context, credentialID string) (int, uint32, error) {
raw, err := decodeCredentialID(credentialID)
if err != nil {
return 0, 0, err
}
var success bool
var errorMsg, credentialJSON sql.NullString
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential::text FROM %s($1)`, p.procs.PasskeyGetCredential)
return db.QueryRowContext(ctx, query, raw).Scan(&success, &errorMsg, &credentialJSON)
})
if err != nil {
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
}
if !success {
return 0, 0, failure(errorMsg, "credential not found")
}
var cred struct {
UserID int `json:"user_id"`
SignCount uint32 `json:"sign_count"`
}
if err := json.Unmarshal(normalizeTimes([]byte(credentialJSON.String)), &cred); err != nil {
return 0, 0, fmt.Errorf("failed to parse credential: %w", err)
}
return cred.UserID, cred.SignCount, nil
}
// UpdateCounter implements lookup.PasskeyStore. Like the code it replaces, it only reports
// an error when the query itself fails; the procedure's success flag is not checked.
func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
raw, err := decodeCredentialID(credentialID)
if err != nil {
return false, err
}
var success bool
var errorMsg sql.NullString
var cloneWarning sql.NullBool
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_clone_warning FROM %s($1, $2)`, p.procs.PasskeyUpdateCounter)
return db.QueryRowContext(ctx, query, raw, newCounter).Scan(&success, &errorMsg, &cloneWarning)
})
if err != nil {
return false, err
}
return cloneWarning.Valid && cloneWarning.Bool, nil
}
// List implements lookup.PasskeyStore.
func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
var success bool
var errorMsg, credentialsJSON sql.NullString
err := p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetUserCredentials)
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &credentialsJSON)
})
if err != nil {
return nil, fmt.Errorf("failed to get credentials: %w", err)
}
if !success {
return nil, failure(errorMsg, "failed to get credentials")
}
var rawCreds []struct {
ID int `json:"id"`
UserID int `json:"user_id"`
CredentialID string `json:"credential_id"`
PublicKey string `json:"public_key"`
AttestationType string `json:"attestation_type"`
AAGUID string `json:"aaguid"`
SignCount uint32 `json:"sign_count"`
CloneWarning bool `json:"clone_warning"`
Transports []string `json:"transports"`
BackupEligible bool `json:"backup_eligible"`
BackupState bool `json:"backup_state"`
Name string `json:"name"`
CreatedAt time.Time `json:"created_at"`
LastUsedAt time.Time `json:"last_used_at"`
}
if err := json.Unmarshal(normalizeTimes([]byte(credentialsJSON.String)), &rawCreds); err != nil {
return nil, fmt.Errorf("failed to parse credentials: %w", err)
}
credentials := make([]sectypes.PasskeyCredential, 0, len(rawCreds))
for i := range rawCreds {
raw := rawCreds[i]
credID, err := base64.StdEncoding.DecodeString(raw.CredentialID)
if err != nil {
continue
}
pubKey, err := base64.StdEncoding.DecodeString(raw.PublicKey)
if err != nil {
continue
}
aaguid, _ := base64.StdEncoding.DecodeString(raw.AAGUID)
credentials = append(credentials, sectypes.PasskeyCredential{
ID: fmt.Sprintf("%d", raw.ID),
UserID: raw.UserID,
CredentialID: credID,
PublicKey: pubKey,
AttestationType: raw.AttestationType,
AAGUID: aaguid,
SignCount: raw.SignCount,
CloneWarning: raw.CloneWarning,
Transports: raw.Transports,
BackupEligible: raw.BackupEligible,
BackupState: raw.BackupState,
Name: raw.Name,
CreatedAt: raw.CreatedAt,
LastUsedAt: raw.LastUsedAt,
})
}
return credentials, nil
}
// Delete implements lookup.PasskeyStore.
func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error {
raw, err := decodeCredentialID(credentialID)
if err != nil {
return err
}
var success bool
var errorMsg sql.NullString
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, p.procs.PasskeyDeleteCredential)
return db.QueryRowContext(ctx, query, userID, raw).Scan(&success, &errorMsg)
})
if err != nil {
return fmt.Errorf("failed to delete credential: %w", err)
}
if !success {
return failure(errorMsg, "failed to delete credential")
}
return nil
}
// Rename implements lookup.PasskeyStore.
func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error {
raw, err := decodeCredentialID(credentialID)
if err != nil {
return err
}
var success bool
var errorMsg sql.NullString
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3)`, p.procs.PasskeyUpdateName)
return db.QueryRowContext(ctx, query, userID, raw, name).Scan(&success, &errorMsg)
})
if err != nil {
return fmt.Errorf("failed to update credential name: %w", err)
}
if !success {
return failure(errorMsg, "failed to update credential name")
}
return nil
}
// ByUsername implements lookup.PasskeyStore.
func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
var success bool
var errorMsg, credentialsJSON sql.NullString
var userID sql.NullInt64
err := p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetCredsByUsername)
return db.QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
})
if err != nil {
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
}
if !success {
return 0, nil, failure(errorMsg, "failed to get credentials")
}
var creds []lookup.PasskeyCredentialRef
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
return 0, nil, fmt.Errorf("failed to parse credentials: %w", err)
}
return int(userID.Int64), creds, nil
}
// Login implements lookup.PasskeyStore: it creates the session for a user whose passkey
// assertion was already verified.
func (p *Passkey) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
reqData := map[string]any{"user_id": userID}
if claims != nil {
if ip, ok := claims["ip_address"].(string); ok {
reqData["ip_address"] = ip
}
if ua, ok := claims["user_agent"].(string); ok {
reqData["user_agent"] = ua
}
}
reqJSON, err := json.Marshal(reqData)
if err != nil {
return nil, fmt.Errorf("failed to marshal passkey login request: %w", err)
}
var success bool
var errorMsg, dataJSON sql.NullString
err = p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, p.procs.PasskeyLogin)
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
})
if err != nil {
return nil, fmt.Errorf("passkey login query failed: %w", err)
}
if !success {
return nil, failure(errorMsg, "passkey login failed")
}
var response sectypes.LoginResponse
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
return nil, fmt.Errorf("failed to parse passkey login response: %w", err)
}
return &response, nil
}
+96
View File
@@ -0,0 +1,96 @@
package procedure
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Policy implements lookup.PolicyStore with the column and row security procedures.
type Policy struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.PolicyStore = (*Policy)(nil)
// NewPolicy creates the procedure-backed PolicyStore.
func NewPolicy(run Runner, procs lookup.ProcNames) *Policy { return &Policy{run: run, procs: procs} }
// ColumnSecurity implements lookup.PolicyStore.
func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
var success bool
var errorMsg sql.NullString
var rulesJSON []byte
err := p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_rules FROM %s($1, $2, $3)`, p.procs.ColumnSecurity)
return db.QueryRowContext(ctx, query, userID, schema, table).Scan(&success, &errorMsg, &rulesJSON)
})
if err != nil {
return nil, fmt.Errorf("failed to load column security: %w", err)
}
if !success {
return nil, failure(errorMsg, "failed to load column security")
}
type securityRecord struct {
Control string `json:"control"`
Accesstype string `json:"accesstype"`
JSONValue string `json:"jsonvalue"`
}
var records []securityRecord
if err := json.Unmarshal(rulesJSON, &records); err != nil {
return nil, fmt.Errorf("failed to parse security rules: %w", err)
}
var rules []sectypes.ColumnSecurity
for _, rec := range records {
parts := strings.Split(rec.Control, ".")
if len(parts) < 3 {
continue
}
rules = append(rules, sectypes.ColumnSecurity{
Schema: schema,
Tablename: table,
Path: parts[2:],
Accesstype: rec.Accesstype,
UserID: userID,
})
}
return rules, nil
}
// RowSecurity implements lookup.PolicyStore. userRef is unwrapped to a scalar user id
// because the procedure's p_user_id is an integer.
func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
switch v := userRef.(type) {
case *sectypes.UserContext:
if v != nil {
userRef = v.UserID
}
case sectypes.UserContext:
userRef = v.UserID
}
var template sql.NullString
var hasBlock sql.NullBool
err := p.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.procs.RowSecurity)
return db.QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
})
if err != nil {
return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
}
return sectypes.RowSecurity{
Schema: schema,
Tablename: table,
UserID: userRef,
Template: template.String,
HasBlock: hasBlock.Bool,
}, nil
}
@@ -0,0 +1,92 @@
// Package procedure is the stored-procedure backend of lookup: it calls the
// resolvespec_* functions (names from lookup.ProcNames) and keeps their
// p_success / p_error / p_data contract. Error texts match the ones the security
// package returned before the extraction.
package procedure
import (
"database/sql"
"fmt"
"strings"
"sync"
)
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
type Runner interface {
Run(run func(*sql.DB) error) error
}
// RunFunc adapts a function to Runner. The security package passes its own
// reconnecting helper this way.
type RunFunc func(run func(*sql.DB) error) error
// Run implements Runner.
func (f RunFunc) Run(run func(*sql.DB) error) error { return f(run) }
// DB is a standalone Runner over a *sql.DB with an optional reconnect factory.
type DB struct {
mu sync.RWMutex
db *sql.DB
factory func() (*sql.DB, error)
onReconnect func()
}
// NewDB wraps db. factory (optional) is called to obtain a fresh handle when the current
// one is closed; onReconnect (optional) runs after a successful reconnect, e.g. to reset
// cached procedure probes.
func NewDB(db *sql.DB, factory func() (*sql.DB, error), onReconnect func()) *DB {
return &DB{db: db, factory: factory, onReconnect: onReconnect}
}
// Get returns the current handle.
func (d *DB) Get() *sql.DB {
d.mu.RLock()
defer d.mu.RUnlock()
return d.db
}
func (d *DB) reconnect() error {
if d.factory == nil {
return fmt.Errorf("no db factory configured for reconnect")
}
newDB, err := d.factory()
if err != nil {
return err
}
d.mu.Lock()
d.db = newDB
d.mu.Unlock()
if d.onReconnect != nil {
d.onReconnect()
}
return nil
}
// Run implements Runner.
func (d *DB) Run(run func(*sql.DB) error) error {
db := d.Get()
if db == nil {
return fmt.Errorf("database connection is nil")
}
err := run(db)
if IsClosed(err) {
if reconnErr := d.reconnect(); reconnErr == nil {
err = run(d.Get())
}
}
return err
}
// IsClosed reports whether err indicates the *sql.DB has been closed.
func IsClosed(err error) bool {
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
}
// failure builds the error for p_success = false: the procedure's own message when it
// returned one, otherwise def.
func failure(errMsg sql.NullString, def string) error {
if errMsg.Valid {
return fmt.Errorf("%s", errMsg.String)
}
return fmt.Errorf("%s", def)
}
@@ -0,0 +1,226 @@
package procedure
import (
"context"
"database/sql"
"encoding/base64"
"errors"
"regexp"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func newMock(t *testing.T) (*DB, sqlmock.Sqlmock) {
t.Helper()
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
return NewDB(db, nil, nil), mock
}
func q(s string) string { return regexp.QuoteMeta(s) }
func TestAuthLoginCallsProcedureWithJSON(t *testing.T) {
run, mock := newMock(t)
procs := lookup.DefaultProcNames()
procs.Login = "custom_login"
a := NewAuth(run, procs)
mock.ExpectQuery(q("SELECT p_success, p_error, p_data::text FROM custom_login($1::jsonb)")).
WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
AddRow(true, nil, `{"token":"sess_1","user":{"user_id":7,"user_name":"bob"}}`))
resp, err := a.Login(context.Background(), sectypes.LoginRequest{Username: "bob", Password: "x"})
if err != nil {
t.Fatal(err)
}
if resp.Token != "sess_1" || resp.User == nil || resp.User.UserID != 7 {
t.Fatalf("unexpected response: %+v", resp)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestAuthLoginFailureUsesProcedureMessage(t *testing.T) {
run, mock := newMock(t)
a := NewAuth(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, "bad credentials", nil))
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "bad credentials" {
t.Fatalf("got %v", err)
}
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() == "" {
t.Fatalf("expected default error, got %v", err)
}
}
func TestRunnerReconnectsOnClosedDB(t *testing.T) {
first, _, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
_ = first.Close()
second, mock, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
defer second.Close()
reconnected := false
run := NewDB(first, func() (*sql.DB, error) { return second, nil }, func() { reconnected = true })
mock.ExpectQuery("SELECT 1").WillReturnRows(sqlmock.NewRows([]string{"x"}).AddRow(1))
err = run.Run(func(db *sql.DB) error {
var x int
return db.QueryRow("SELECT 1").Scan(&x)
})
if err != nil {
t.Fatal(err)
}
if !reconnected || run.Get() != second {
t.Fatal("expected reconnect to the new handle")
}
}
func TestRunnerNoFactoryReturnsClosedError(t *testing.T) {
db, _, _ := sqlmock.New()
_ = db.Close()
err := NewDB(db, nil, nil).Run(func(db *sql.DB) error { return db.QueryRow("SELECT 1").Scan(new(int)) })
if !IsClosed(err) {
t.Fatalf("got %v", err)
}
}
func TestPasskeyGetDecodesCredentialID(t *testing.T) {
run, mock := newMock(t)
p := NewPasskey(run, lookup.DefaultProcNames())
raw := []byte{1, 2, 3, 4}
mock.ExpectQuery("resolvespec_passkey_get_credential").WithArgs(raw).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_credential"}).
AddRow(true, nil, `{"user_id":9,"sign_count":4}`))
uid, count, err := p.Get(context.Background(), base64.StdEncoding.EncodeToString(raw))
if err != nil || uid != 9 || count != 4 {
t.Fatalf("got %d %d %v", uid, count, err)
}
}
func TestPasskeyInvalidBase64(t *testing.T) {
run, _ := newMock(t)
p := NewPasskey(run, lookup.DefaultProcNames())
if _, _, err := p.Get(context.Background(), "***"); err == nil {
t.Fatal("expected error")
}
if err := p.Delete(context.Background(), 1, "***"); err == nil {
t.Fatal("expected error")
}
}
func TestPasskeyUpdateCounterReportsCloneWarning(t *testing.T) {
run, mock := newMock(t)
p := NewPasskey(run, lookup.DefaultProcNames())
id := base64.StdEncoding.EncodeToString([]byte("abc"))
mock.ExpectQuery("resolvespec_passkey_update_counter").WithArgs([]byte("abc"), uint32(5)).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_clone_warning"}).AddRow(true, nil, true))
warn, err := p.UpdateCounter(context.Background(), id, 5)
if err != nil || !warn {
t.Fatalf("got %v %v", warn, err)
}
mock.ExpectQuery("resolvespec_passkey_update_counter").WillReturnError(errors.New("boom"))
if _, err := p.UpdateCounter(context.Background(), id, 6); err == nil {
t.Fatal("expected error")
}
}
func TestPasskeyByUsername(t *testing.T) {
run, mock := newMock(t)
p := NewPasskey(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_passkey_get_credentials_by_username").WithArgs("bob").
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id", "p_credentials"}).
AddRow(true, nil, 3, `[{"credential_id":"YWJj","transports":["usb"]}]`))
uid, refs, err := p.ByUsername(context.Background(), "bob")
if err != nil || uid != 3 || len(refs) != 1 || refs[0].CredentialID != "YWJj" || refs[0].Transports[0] != "usb" {
t.Fatalf("got %d %+v %v", uid, refs, err)
}
}
func TestOAuthUsersGetOrCreateUser(t *testing.T) {
run, mock := newMock(t)
o := NewOAuthUsers(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, 11))
id, err := o.GetOrCreateUser(context.Background(), &sectypes.UserContext{UserName: "u", Email: "e"}, "github")
if err != nil || id != 11 {
t.Fatalf("got %d %v", id, err)
}
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, nil))
if _, err := o.GetOrCreateUser(context.Background(), &sectypes.UserContext{}, "github"); err == nil || err.Error() != "user ID not returned" {
t.Fatalf("got %v", err)
}
}
func TestOAuthUsersRefreshRoundTrip(t *testing.T) {
run, mock := newMock(t)
o := NewOAuthUsers(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_oauth_getrefreshtoken").WithArgs("r1").
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
AddRow(true, nil, `{"user_id":2,"access_token":"a","token_type":"Bearer","expiry":"2030-01-01T00:00:00Z"}`))
s, err := o.GetByRefreshToken(context.Background(), "r1")
if err != nil || s.UserID != 2 || s.AccessToken != "a" {
t.Fatalf("got %+v %v", s, err)
}
mock.ExpectQuery("resolvespec_oauth_updaterefreshtoken").WithArgs(sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(false, "session not found"))
err = o.UpdateRefreshToken(context.Background(), 2, "r1", "s2", "a2", "r2", time.Now())
if err == nil || err.Error() != "session not found" {
t.Fatalf("got %v", err)
}
}
func TestOAuthClientsExchangeCodeSetsCode(t *testing.T) {
run, mock := newMock(t)
c := NewOAuthClients(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("abc").
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, `{"client_id":"cid"}`))
code, err := c.ExchangeCode(context.Background(), "abc")
if err != nil || code.Code != "abc" || code.ClientID != "cid" {
t.Fatalf("got %+v %v", code, err)
}
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("zzz").
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
if _, err := c.ExchangeCode(context.Background(), "zzz"); err == nil || err.Error() != "invalid or expired code" {
t.Fatalf("got %v", err)
}
}
func TestOAuthClientsRevoke(t *testing.T) {
run, mock := newMock(t)
c := NewOAuthClients(run, lookup.DefaultProcNames())
mock.ExpectQuery("resolvespec_oauth_revoke").WithArgs("t").
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(true, nil))
if err := c.Revoke(context.Background(), "t"); err != nil {
t.Fatal(err)
}
}
+121
View File
@@ -0,0 +1,121 @@
package procedure
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// TOTP implements lookup.TOTPStore with the resolvespec_totp_* procedures.
type TOTP struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.TOTPStore = (*TOTP)(nil)
// NewTOTP creates the procedure-backed TOTPStore.
func NewTOTP(run Runner, procs lookup.ProcNames) *TOTP { return &TOTP{run: run, procs: procs} }
// exec runs a "p_success, p_error" procedure and maps failure to an error.
func (t *TOTP) exec(ctx context.Context, query, op, def string, args ...any) error {
var success bool
var errorMsg sql.NullString
err := t.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, query, args...).Scan(&success, &errorMsg)
})
if err != nil {
return fmt.Errorf("%s query failed: %w", op, err)
}
if !success {
return failure(errorMsg, def)
}
return nil
}
// Enable implements lookup.TOTPStore.
func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
codesJSON, err := json.Marshal(hashedCodes)
if err != nil {
return fmt.Errorf("failed to marshal backup codes: %w", err)
}
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, t.procs.TOTPEnable)
return t.exec(ctx, query, "enable 2FA", "failed to enable 2FA", userID, secret, string(codesJSON))
}
// Disable implements lookup.TOTPStore.
func (t *TOTP) Disable(ctx context.Context, userID int) error {
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, t.procs.TOTPDisable)
return t.exec(ctx, query, "disable 2FA", "failed to disable 2FA", userID)
}
// Status implements lookup.TOTPStore.
func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) {
var success, enabled bool
var errorMsg sql.NullString
err := t.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, t.procs.TOTPGetStatus)
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &enabled)
})
if err != nil {
return false, fmt.Errorf("get 2FA status query failed: %w", err)
}
if !success {
return false, failure(errorMsg, "failed to get 2FA status")
}
return enabled, nil
}
// Secret implements lookup.TOTPStore.
func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) {
var success bool
var errorMsg, secret sql.NullString
err := t.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, t.procs.TOTPGetSecret)
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &secret)
})
if err != nil {
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
}
if !success {
return "", failure(errorMsg, "failed to get 2FA secret")
}
if !secret.Valid {
return "", fmt.Errorf("2FA secret not found")
}
return secret.String, nil
}
// RegenerateBackupCodes implements lookup.TOTPStore.
func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
codesJSON, err := json.Marshal(hashedCodes)
if err != nil {
return fmt.Errorf("failed to marshal backup codes: %w", err)
}
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, t.procs.TOTPRegenerateBackup)
return t.exec(ctx, query, "regenerate backup codes", "failed to regenerate backup codes", userID, string(codesJSON))
}
// ValidateBackupCode implements lookup.TOTPStore. A failure without a message means
// "not valid", not an error.
func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
var success, valid bool
var errorMsg sql.NullString
err := t.run.Run(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, t.procs.TOTPValidateBackupCode)
return db.QueryRowContext(ctx, query, userID, codeHash).Scan(&success, &errorMsg, &valid)
})
if err != nil {
return false, fmt.Errorf("validate backup code query failed: %w", err)
}
if !success {
if errorMsg.Valid {
return false, fmt.Errorf("%s", errorMsg.String)
}
return false, nil
}
return valid, nil
}