mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-06 13:26:28 +00:00
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:
@@ -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 := §ypes.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 §ypes.LoginResponse{
|
||||
Token: fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix()),
|
||||
User: §ypes.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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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(), §ypes.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(), §ypes.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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user