mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 04:21:58 +00:00
Verify bcrypt passwords in Direct mode and the shipped procedures, hash on register/reset, ignore client-supplied roles and level at registration and drop the password from the jwt_login payload. Legacy cleartext upgrade is opt-in. Row security templates now bind the user as a parameter, validate identifiers, attach via common.SelectQuery and fail the request if the filter cannot be attached. ApplyColumnSecurity and GetRowSecurityTemplate convert panics to errors and the hooks fail closed. Update audit status.
1316 lines
41 KiB
Go
1316 lines
41 KiB
Go
package security
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
|
)
|
|
|
|
// Production-Ready Authenticators
|
|
// =================================
|
|
|
|
// HeaderAuthenticator provides simple header-based authentication
|
|
// Expects: X-User-ID, X-User-Name, X-User-Level, X-Session-ID, X-Remote-ID, X-User-Roles, X-User-Email
|
|
type HeaderAuthenticator struct{}
|
|
|
|
func NewHeaderAuthenticator() *HeaderAuthenticator {
|
|
return &HeaderAuthenticator{}
|
|
}
|
|
|
|
func (a *HeaderAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
|
return nil, fmt.Errorf("header authentication does not support login")
|
|
}
|
|
|
|
func (a *HeaderAuthenticator) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) {
|
|
return a.Login(ctx, req)
|
|
}
|
|
|
|
func (a *HeaderAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
|
|
return nil
|
|
}
|
|
|
|
func (a *HeaderAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
|
return a.Logout(ctx, req)
|
|
}
|
|
|
|
func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
|
userIDStr := r.Header.Get("X-User-ID")
|
|
if userIDStr == "" {
|
|
return nil, fmt.Errorf("X-User-ID header required")
|
|
}
|
|
|
|
userID, err := strconv.Atoi(userIDStr)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("invalid user ID: %w", err)
|
|
}
|
|
|
|
return &UserContext{
|
|
UserID: userID,
|
|
UserName: r.Header.Get("X-User-Name"),
|
|
UserLevel: parseIntHeader(r, "X-User-Level", 0),
|
|
SessionID: r.Header.Get("X-Session-ID"),
|
|
RemoteID: r.Header.Get("X-Remote-ID"),
|
|
Email: r.Header.Get("X-User-Email"),
|
|
Roles: parseRoles(r.Header.Get("X-User-Roles")),
|
|
}, nil
|
|
}
|
|
|
|
// maxAuthTokens caps the comma-separated credentials tried per request so one
|
|
// request cannot drive unbounded session lookups.
|
|
const maxAuthTokens = 4
|
|
|
|
// sessionActivityTimeout bounds the detached last-activity update.
|
|
const sessionActivityTimeout = 5 * time.Second
|
|
|
|
// DatabaseAuthenticator provides session-based authentication with database storage
|
|
// All database operations go through stored procedures for security and consistency
|
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
|
// See database_schema.sql for procedure definitions
|
|
// Also supports multiple OAuth2 providers configured with WithOAuth2()
|
|
// Also supports passkey authentication configured with WithPasskey()
|
|
type DatabaseAuthenticator struct {
|
|
db *sql.DB
|
|
dbMu sync.RWMutex
|
|
dbFactory func() (*sql.DB, error)
|
|
cache *cache.Cache
|
|
cacheTTL time.Duration
|
|
sqlNames *SQLNames
|
|
tableNames *TableNames
|
|
queryMode QueryMode
|
|
capability *dbCapability
|
|
|
|
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
|
|
// on successful login (opt-in, see DatabaseAuthenticatorOptions).
|
|
upgradePasswordHash bool
|
|
|
|
// activityWG tracks in-flight asynchronous session activity updates
|
|
activityWG sync.WaitGroup
|
|
|
|
// Cookie session support (optional, gated by enableCookieSession)
|
|
enableCookieSession bool
|
|
cookieOptions SessionCookieOptions
|
|
|
|
// OAuth2 providers registry (multiple providers supported)
|
|
oauth2Providers map[string]*OAuth2Provider
|
|
oauth2ProvidersMutex sync.RWMutex
|
|
|
|
// Passkey provider (optional)
|
|
passkeyProvider PasskeyProvider
|
|
|
|
// Optional fallback called when primary authentication fails
|
|
authenticateCallback func(r *http.Request) (*UserContext, error)
|
|
}
|
|
|
|
// DatabaseAuthenticatorOptions configures the database authenticator
|
|
type DatabaseAuthenticatorOptions struct {
|
|
// CacheTTL is the duration to cache user contexts
|
|
// Default: 5 minutes
|
|
CacheTTL time.Duration
|
|
// Cache is an optional cache instance. If nil, uses the default cache
|
|
Cache *cache.Cache
|
|
// PasskeyProvider is an optional passkey provider for WebAuthn/FIDO2 authentication
|
|
PasskeyProvider PasskeyProvider
|
|
// SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames().
|
|
// Partial overrides are supported: only set the fields you want to change.
|
|
SQLNames *SQLNames
|
|
// TableNames provides custom table names for Direct mode. If nil, uses DefaultTableNames().
|
|
TableNames *TableNames
|
|
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
|
|
QueryMode QueryMode
|
|
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
|
// If nil, reconnection is disabled.
|
|
DBFactory func() (*sql.DB, error)
|
|
// EnableCookieSession enables cookie-based session management.
|
|
// When true, Authenticate reads the session token from the cookie named by
|
|
// CookieOptions.Name (default "session_token") in addition to the Authorization header,
|
|
// and LoginWithCookie / LogoutWithCookie automatically set / clear the cookie.
|
|
EnableCookieSession bool
|
|
// UpgradePasswordHash, when true, rewrites a legacy cleartext password as a
|
|
// bcrypt hash after a successful login. It is off by default and is never
|
|
// enabled automatically: legacy cleartext values are still accepted at login,
|
|
// but stored rows are left untouched unless this is set.
|
|
UpgradePasswordHash bool
|
|
// CookieOptions configures the session cookie written by LoginWithCookie.
|
|
// Only used when EnableCookieSession is true.
|
|
CookieOptions SessionCookieOptions
|
|
// AuthenticateCallback is a fallback called when the primary authentication (database
|
|
// session lookup) fails. If non-nil and the callback returns a non-nil UserContext,
|
|
// that result is used in place of the failure.
|
|
AuthenticateCallback func(r *http.Request) (*UserContext, error)
|
|
}
|
|
|
|
func NewDatabaseAuthenticator(db *sql.DB) *DatabaseAuthenticator {
|
|
return NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{
|
|
CacheTTL: 5 * time.Minute,
|
|
})
|
|
}
|
|
|
|
func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorOptions) *DatabaseAuthenticator {
|
|
if opts.CacheTTL == 0 {
|
|
opts.CacheTTL = 5 * time.Minute
|
|
}
|
|
|
|
cacheInstance := opts.Cache
|
|
if cacheInstance == nil {
|
|
cacheInstance = cache.GetDefaultCache()
|
|
}
|
|
|
|
sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames)
|
|
tableNames := resolveTableNames(opts.TableNames)
|
|
|
|
return &DatabaseAuthenticator{
|
|
db: db,
|
|
dbFactory: opts.DBFactory,
|
|
cache: cacheInstance,
|
|
cacheTTL: opts.CacheTTL,
|
|
sqlNames: sqlNames,
|
|
tableNames: tableNames,
|
|
queryMode: opts.QueryMode,
|
|
capability: newDBCapability(),
|
|
passkeyProvider: opts.PasskeyProvider,
|
|
enableCookieSession: opts.EnableCookieSession,
|
|
upgradePasswordHash: opts.UpgradePasswordHash,
|
|
cookieOptions: opts.CookieOptions,
|
|
authenticateCallback: opts.AuthenticateCallback,
|
|
}
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) getDB() *sql.DB {
|
|
a.dbMu.RLock()
|
|
defer a.dbMu.RUnlock()
|
|
return a.db
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) reconnectDB() error {
|
|
if a.dbFactory == nil {
|
|
return fmt.Errorf("no db factory configured for reconnect")
|
|
}
|
|
newDB, err := a.dbFactory()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
a.dbMu.Lock()
|
|
a.db = newDB
|
|
a.dbMu.Unlock()
|
|
if a.capability != nil {
|
|
a.capability.reset()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) runDBOpWithReconnect(run func(*sql.DB) error) error {
|
|
db := a.getDB()
|
|
if db == nil {
|
|
return fmt.Errorf("database connection is nil")
|
|
}
|
|
|
|
err := run(db)
|
|
if isDBClosed(err) {
|
|
if reconnErr := a.reconnectDB(); reconnErr == nil {
|
|
err = run(a.getDB())
|
|
}
|
|
}
|
|
|
|
return err
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error)) {
|
|
a.authenticateCallback = fn
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Login) {
|
|
return a.loginDirect(ctx, req)
|
|
}
|
|
// Convert LoginRequest to JSON
|
|
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)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var dataJSON sql.NullString
|
|
|
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.sqlNames.Login)
|
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
|
})
|
|
if err != nil {
|
|
return nil, fmt.Errorf("login query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("login failed")
|
|
}
|
|
|
|
// Parse response
|
|
var response LoginResponse
|
|
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
|
return nil, fmt.Errorf("failed to parse login response: %w", err)
|
|
}
|
|
|
|
return &response, nil
|
|
}
|
|
|
|
// Register implements Registrable interface
|
|
func (a *DatabaseAuthenticator) Register(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Register) {
|
|
return a.registerDirect(ctx, req)
|
|
}
|
|
// Convert RegisterRequest to JSON
|
|
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 sql.NullString
|
|
var dataJSON sql.NullString
|
|
|
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.sqlNames.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 {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("registration failed")
|
|
}
|
|
|
|
// Parse response
|
|
var response 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
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Logout) {
|
|
return a.logoutDirect(ctx, req)
|
|
}
|
|
// Convert LogoutRequest to JSON
|
|
reqJSON, err := json.Marshal(req)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to marshal logout request: %w", err)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var dataJSON sql.NullString
|
|
|
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.sqlNames.Logout)
|
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
|
})
|
|
if err != nil {
|
|
return fmt.Errorf("logout query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return fmt.Errorf("logout failed")
|
|
}
|
|
|
|
// Clear cache for this token
|
|
if req.Token != "" {
|
|
cacheKey := fmt.Sprintf("auth:session:%s", req.Token)
|
|
_ = a.cache.Delete(ctx, cacheKey)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// LoginWithCookie performs a login and, when EnableCookieSession is true, writes the
|
|
// session cookie to w using the configured CookieOptions. The LoginResponse is returned
|
|
// regardless of whether cookie sessions are enabled.
|
|
func (a *DatabaseAuthenticator) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) {
|
|
resp, err := a.Login(ctx, req)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if a.enableCookieSession {
|
|
SetSessionCookie(w, resp, a.cookieOptions)
|
|
}
|
|
return resp, nil
|
|
}
|
|
|
|
// LogoutWithCookie performs a logout and, when EnableCookieSession is true, clears the
|
|
// session cookie on w. The logout itself is performed regardless of the cookie flag.
|
|
func (a *DatabaseAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
|
err := a.Logout(ctx, req)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if a.enableCookieSession {
|
|
ClearSessionCookie(w, a.cookieOptions)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
|
// Extract session token from header or cookie
|
|
sessionToken := r.Header.Get("Authorization")
|
|
reference := "authenticate"
|
|
var tokens []string
|
|
|
|
if sessionToken == "" {
|
|
if a.enableCookieSession {
|
|
if token := GetSessionCookie(r, a.cookieOptions); token != "" {
|
|
tokens = []string{token}
|
|
reference = "cookie"
|
|
}
|
|
}
|
|
} else {
|
|
// Parse Authorization header which may contain multiple comma-separated tokens
|
|
// Format: "Token abc, Token def" or "Bearer abc" or just "abc"
|
|
rawTokens := strings.SplitN(sessionToken, ",", maxAuthTokens+2)
|
|
if len(rawTokens) > maxAuthTokens {
|
|
return nil, fmt.Errorf("too many authorization tokens")
|
|
}
|
|
for _, token := range rawTokens {
|
|
token = strings.TrimSpace(token)
|
|
// Remove "Bearer " prefix if present
|
|
token = strings.TrimPrefix(token, "Bearer ")
|
|
// Remove "Token " prefix if present
|
|
token = strings.TrimPrefix(token, "Token ")
|
|
token = strings.TrimSpace(token)
|
|
if token != "" {
|
|
tokens = append(tokens, token)
|
|
}
|
|
}
|
|
}
|
|
|
|
if len(tokens) == 0 {
|
|
if a.authenticateCallback != nil {
|
|
return a.authenticateCallback(r)
|
|
}
|
|
return nil, fmt.Errorf("session token required")
|
|
}
|
|
|
|
// Log warning if multiple tokens are provided
|
|
if len(tokens) > 1 {
|
|
logger.Warn("Multiple authentication tokens provided in Authorization header (%d tokens). This is unusual and may indicate a misconfigured client. Header: %s", len(tokens), sessionToken)
|
|
}
|
|
|
|
// Try each token until one succeeds
|
|
var lastErr error
|
|
for _, token := range tokens {
|
|
// Build cache key
|
|
cacheKey := fmt.Sprintf("auth:session:%s", token)
|
|
|
|
// Use cache.GetOrSet to get from cache or load from database
|
|
var userCtx UserContext
|
|
err := a.cache.GetOrSet(r.Context(), cacheKey, &userCtx, a.cacheTTL, func() (any, error) {
|
|
// This function is called only if cache miss
|
|
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
|
|
return a.sessionDirect(r.Context(), token)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var userJSON sql.NullString
|
|
|
|
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
|
|
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
|
})
|
|
if err != nil {
|
|
return nil, fmt.Errorf("session query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("invalid or expired session")
|
|
}
|
|
|
|
if !userJSON.Valid {
|
|
return nil, fmt.Errorf("no user data in session")
|
|
}
|
|
|
|
// Parse UserContext
|
|
var user 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
|
|
})
|
|
|
|
if err != nil {
|
|
lastErr = err
|
|
continue // Try next token
|
|
}
|
|
|
|
// Authentication succeeded with this token
|
|
// Update last activity timestamp asynchronously
|
|
activityCtx := userCtx
|
|
// Detach from the request (it is cancelled when the handler returns) but
|
|
// keep a deadline, and never let a panic here take the process down.
|
|
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
|
|
a.activityWG.Add(1)
|
|
go func(ctx context.Context, token string) {
|
|
defer a.activityWG.Done()
|
|
defer cancel()
|
|
defer logger.CatchPanic("updateSessionActivity")()
|
|
a.updateSessionActivity(ctx, token, &activityCtx)
|
|
}(detached, token)
|
|
|
|
return &userCtx, nil
|
|
}
|
|
|
|
// All tokens failed — try callback before returning error
|
|
if a.authenticateCallback != nil {
|
|
return a.authenticateCallback(r)
|
|
}
|
|
if lastErr != nil {
|
|
return nil, lastErr
|
|
}
|
|
return nil, fmt.Errorf("authentication failed for all provided tokens")
|
|
}
|
|
|
|
// ClearCache removes a specific token from the cache or clears all cache if token is empty
|
|
func (a *DatabaseAuthenticator) ClearCache(token string) error {
|
|
ctx := context.Background()
|
|
if token != "" {
|
|
cacheKey := fmt.Sprintf("auth:session:%s", token)
|
|
return a.cache.Delete(ctx, cacheKey)
|
|
}
|
|
// Clear all auth cache entries
|
|
return a.cache.DeleteByPattern(ctx, "auth:session:*")
|
|
}
|
|
|
|
// ClearUserCache removes all cache entries for a specific user ID
|
|
func (a *DatabaseAuthenticator) ClearUserCache(userID int) error {
|
|
ctx := context.Background()
|
|
// Clear all sessions for this user
|
|
pattern := "auth:session:*"
|
|
return a.cache.DeleteByPattern(ctx, pattern)
|
|
}
|
|
|
|
// updateSessionActivity updates the last activity timestamp for the session
|
|
func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.SessionUpdate) {
|
|
_ = a.updateSessionActivityDirect(ctx, sessionToken)
|
|
return
|
|
}
|
|
|
|
// Convert UserContext to JSON
|
|
userJSON, err := json.Marshal(userCtx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var updatedUserJSON sql.NullString
|
|
|
|
_ = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.sqlNames.SessionUpdate)
|
|
return db.QueryRowContext(ctx, query, sessionToken, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
|
})
|
|
}
|
|
|
|
// RefreshToken implements Refreshable interface
|
|
func (a *DatabaseAuthenticator) RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.RefreshToken) {
|
|
return a.refreshTokenDirect(ctx, refreshToken)
|
|
}
|
|
// First, we need to get the current user context for the refresh token
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var userJSON sql.NullString
|
|
// Get current session to pass to refresh
|
|
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.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 {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("invalid refresh token")
|
|
}
|
|
|
|
var newSuccess bool
|
|
var newErrorMsg sql.NullString
|
|
var newUserJSON sql.NullString
|
|
|
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
refreshQuery := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.sqlNames.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 {
|
|
if newErrorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", newErrorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("failed to refresh token")
|
|
}
|
|
|
|
// Parse refreshed user context
|
|
var userCtx 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 := &LoginResponse{
|
|
Token: userCtx.SessionID, // New session token from stored procedure
|
|
User: &userCtx,
|
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
|
}
|
|
if refreshToken, ok := userCtx.Claims["refresh_token"].(string); ok && refreshToken != "" {
|
|
resp.RefreshToken = refreshToken
|
|
}
|
|
if expiresIn, ok := userCtx.Claims["expires_in"].(float64); ok && expiresIn > 0 {
|
|
resp.ExpiresIn = int64(expiresIn)
|
|
}
|
|
return resp, nil
|
|
}
|
|
|
|
// JWTAuthenticator provides JWT token-based authentication
|
|
// All database operations go through stored procedures
|
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
|
// NOTE: JWT signing/verification requires github.com/golang-jwt/jwt/v5 to be installed and imported
|
|
type JWTAuthenticator struct {
|
|
secretKey []byte
|
|
db *sql.DB
|
|
dbMu sync.RWMutex
|
|
dbFactory func() (*sql.DB, error)
|
|
sqlNames *SQLNames
|
|
tableNames *TableNames
|
|
queryMode QueryMode
|
|
capability *dbCapability
|
|
|
|
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
|
|
// on successful login. Off by default; enable with WithPasswordHashUpgrade.
|
|
upgradePasswordHash bool
|
|
}
|
|
|
|
// WithPasswordHashUpgrade explicitly enables (or disables) upgrading legacy
|
|
// cleartext passwords to bcrypt after a successful login. Off by default.
|
|
func (a *JWTAuthenticator) WithPasswordHashUpgrade(enabled bool) *JWTAuthenticator {
|
|
a.upgradePasswordHash = enabled
|
|
return a
|
|
}
|
|
|
|
func NewJWTAuthenticator(secretKey string, db *sql.DB, names ...*SQLNames) *JWTAuthenticator {
|
|
return &JWTAuthenticator{
|
|
secretKey: []byte(secretKey),
|
|
db: db,
|
|
sqlNames: resolveSQLNames(names...),
|
|
tableNames: DefaultTableNames(),
|
|
capability: newDBCapability(),
|
|
}
|
|
}
|
|
|
|
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
|
func (a *JWTAuthenticator) WithDBFactory(factory func() (*sql.DB, error)) *JWTAuthenticator {
|
|
a.dbFactory = factory
|
|
return a
|
|
}
|
|
|
|
// WithTableNames configures Direct-mode table names. If names is nil, defaults are used.
|
|
func (a *JWTAuthenticator) WithTableNames(names *TableNames) *JWTAuthenticator {
|
|
a.tableNames = resolveTableNames(names)
|
|
return a
|
|
}
|
|
|
|
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
|
|
func (a *JWTAuthenticator) WithQueryMode(mode QueryMode) *JWTAuthenticator {
|
|
a.queryMode = mode
|
|
return a
|
|
}
|
|
|
|
func (a *JWTAuthenticator) getDB() *sql.DB {
|
|
a.dbMu.RLock()
|
|
defer a.dbMu.RUnlock()
|
|
return a.db
|
|
}
|
|
|
|
func (a *JWTAuthenticator) reconnectDB() error {
|
|
if a.dbFactory == nil {
|
|
return fmt.Errorf("no db factory configured for reconnect")
|
|
}
|
|
newDB, err := a.dbFactory()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
a.dbMu.Lock()
|
|
a.db = newDB
|
|
a.dbMu.Unlock()
|
|
if a.capability != nil {
|
|
a.capability.reset()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.JWTLogin) {
|
|
return a.jwtLoginDirect(ctx, req)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var userJSON []byte
|
|
|
|
runLoginQuery := func() error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user FROM %s($1, $2)`, a.sqlNames.JWTLogin)
|
|
return a.getDB().QueryRowContext(ctx, query, req.Username, req.Password).Scan(&success, &errorMsg, &userJSON)
|
|
}
|
|
err := runLoginQuery()
|
|
if isDBClosed(err) {
|
|
if reconnErr := a.reconnectDB(); reconnErr == nil {
|
|
err = runLoginQuery()
|
|
}
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("login query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("invalid credentials")
|
|
}
|
|
|
|
// Parse user data
|
|
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)
|
|
}
|
|
|
|
// The password is verified inside resolvespec_jwt_login; the hash is never
|
|
// returned to Go.
|
|
|
|
// Generate token (placeholder - implement JWT signing when library is available)
|
|
expiresAt := time.Now().Add(24 * time.Hour)
|
|
tokenString := fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix())
|
|
|
|
return &LoginResponse{
|
|
Token: tokenString,
|
|
User: &UserContext{
|
|
UserID: user.ID,
|
|
UserName: user.Username,
|
|
Email: user.Email,
|
|
UserLevel: user.UserLevel,
|
|
Roles: parseRoles(user.Roles),
|
|
},
|
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
|
}, nil
|
|
}
|
|
|
|
func (a *JWTAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.JWTLogout) {
|
|
return a.jwtLogoutDirect(ctx, req)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
|
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, a.sqlNames.JWTLogout)
|
|
err := a.getDB().QueryRowContext(ctx, query, req.Token, req.UserID).Scan(&success, &errorMsg)
|
|
if err != nil {
|
|
return fmt.Errorf("logout query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return fmt.Errorf("logout failed")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (a *JWTAuthenticator) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) {
|
|
return a.Login(ctx, req)
|
|
}
|
|
|
|
func (a *JWTAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
|
return a.Logout(ctx, req)
|
|
}
|
|
|
|
func (a *JWTAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
|
authHeader := r.Header.Get("Authorization")
|
|
if authHeader == "" {
|
|
return nil, fmt.Errorf("authorization header required")
|
|
}
|
|
|
|
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
|
if tokenString == authHeader {
|
|
return nil, fmt.Errorf("bearer token required")
|
|
}
|
|
|
|
// TODO: Implement JWT parsing when library is available
|
|
return nil, fmt.Errorf("JWT parsing not implemented - install github.com/golang-jwt/jwt/v5")
|
|
}
|
|
|
|
// Production-Ready Security Providers
|
|
// ====================================
|
|
|
|
// DatabaseColumnSecurityProvider loads column security from database
|
|
// All database operations go through stored procedures
|
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
|
type DatabaseColumnSecurityProvider struct {
|
|
db *sql.DB
|
|
dbMu sync.RWMutex
|
|
dbFactory func() (*sql.DB, error)
|
|
sqlNames *SQLNames
|
|
queryMode QueryMode
|
|
capability *dbCapability
|
|
}
|
|
|
|
func NewDatabaseColumnSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseColumnSecurityProvider {
|
|
return &DatabaseColumnSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()}
|
|
}
|
|
|
|
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
|
|
// Direct mode is unsupported for column security (see ErrDirectModeUnsupported).
|
|
func (p *DatabaseColumnSecurityProvider) WithQueryMode(mode QueryMode) *DatabaseColumnSecurityProvider {
|
|
p.queryMode = mode
|
|
return p
|
|
}
|
|
|
|
func (p *DatabaseColumnSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseColumnSecurityProvider {
|
|
p.dbFactory = factory
|
|
return p
|
|
}
|
|
|
|
func (p *DatabaseColumnSecurityProvider) getDB() *sql.DB {
|
|
p.dbMu.RLock()
|
|
defer p.dbMu.RUnlock()
|
|
return p.db
|
|
}
|
|
|
|
func (p *DatabaseColumnSecurityProvider) reconnectDB() error {
|
|
if p.dbFactory == nil {
|
|
return fmt.Errorf("no db factory configured for reconnect")
|
|
}
|
|
newDB, err := p.dbFactory()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
p.dbMu.Lock()
|
|
p.db = newDB
|
|
p.dbMu.Unlock()
|
|
if p.capability != nil {
|
|
p.capability.reset()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
|
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
|
|
return nil, ErrDirectModeUnsupported
|
|
}
|
|
|
|
var rules []ColumnSecurity
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var rulesJSON []byte
|
|
|
|
runQuery := func() error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_rules FROM %s($1, $2, $3)`, p.sqlNames.ColumnSecurity)
|
|
return p.getDB().QueryRowContext(ctx, query, userID, schema, table).Scan(&success, &errorMsg, &rulesJSON)
|
|
}
|
|
err := runQuery()
|
|
if isDBClosed(err) {
|
|
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
|
err = runQuery()
|
|
}
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to load column security: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("failed to load column security")
|
|
}
|
|
|
|
// Parse the JSON array of security records
|
|
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)
|
|
}
|
|
|
|
// Convert records to ColumnSecurity rules
|
|
for _, rec := range records {
|
|
parts := strings.Split(rec.Control, ".")
|
|
if len(parts) < 3 {
|
|
continue
|
|
}
|
|
|
|
rule := ColumnSecurity{
|
|
Schema: schema,
|
|
Tablename: table,
|
|
Path: parts[2:],
|
|
Accesstype: rec.Accesstype,
|
|
UserID: userID,
|
|
}
|
|
|
|
rules = append(rules, rule)
|
|
}
|
|
|
|
return rules, nil
|
|
}
|
|
|
|
// DatabaseRowSecurityProvider loads row security from database
|
|
// All database operations go through stored procedures
|
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
|
type DatabaseRowSecurityProvider struct {
|
|
db *sql.DB
|
|
dbMu sync.RWMutex
|
|
dbFactory func() (*sql.DB, error)
|
|
sqlNames *SQLNames
|
|
queryMode QueryMode
|
|
capability *dbCapability
|
|
}
|
|
|
|
func NewDatabaseRowSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseRowSecurityProvider {
|
|
return &DatabaseRowSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()}
|
|
}
|
|
|
|
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
|
|
// Direct mode is unsupported for row security (see ErrDirectModeUnsupported).
|
|
func (p *DatabaseRowSecurityProvider) WithQueryMode(mode QueryMode) *DatabaseRowSecurityProvider {
|
|
p.queryMode = mode
|
|
return p
|
|
}
|
|
|
|
func (p *DatabaseRowSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseRowSecurityProvider {
|
|
p.dbFactory = factory
|
|
return p
|
|
}
|
|
|
|
func (p *DatabaseRowSecurityProvider) getDB() *sql.DB {
|
|
p.dbMu.RLock()
|
|
defer p.dbMu.RUnlock()
|
|
return p.db
|
|
}
|
|
|
|
func (p *DatabaseRowSecurityProvider) reconnectDB() error {
|
|
if p.dbFactory == nil {
|
|
return fmt.Errorf("no db factory configured for reconnect")
|
|
}
|
|
newDB, err := p.dbFactory()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
p.dbMu.Lock()
|
|
p.db = newDB
|
|
p.dbMu.Unlock()
|
|
if p.capability != nil {
|
|
p.capability.reset()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
|
|
return RowSecurity{}, ErrDirectModeUnsupported
|
|
}
|
|
|
|
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
|
|
// hand back the full *UserContext so non-DB providers can inspect claims;
|
|
// unwrap it here before it reaches the SQL args.
|
|
switch v := userRef.(type) {
|
|
case *UserContext:
|
|
if v != nil {
|
|
userRef = v.UserID
|
|
}
|
|
case UserContext:
|
|
userRef = v.UserID
|
|
}
|
|
|
|
var template sql.NullString
|
|
var hasBlock sql.NullBool
|
|
|
|
runQuery := func() error {
|
|
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.sqlNames.RowSecurity)
|
|
return p.getDB().QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
|
|
}
|
|
err := runQuery()
|
|
if isDBClosed(err) {
|
|
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
|
err = runQuery()
|
|
}
|
|
}
|
|
if err != nil {
|
|
return RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
|
|
}
|
|
|
|
return RowSecurity{
|
|
Schema: schema,
|
|
Tablename: table,
|
|
UserID: userRef,
|
|
Template: template.String,
|
|
HasBlock: hasBlock.Bool,
|
|
}, nil
|
|
}
|
|
|
|
// ConfigColumnSecurityProvider provides static column security configuration
|
|
type ConfigColumnSecurityProvider struct {
|
|
rules map[string][]ColumnSecurity
|
|
}
|
|
|
|
func NewConfigColumnSecurityProvider(rules map[string][]ColumnSecurity) *ConfigColumnSecurityProvider {
|
|
return &ConfigColumnSecurityProvider{rules: rules}
|
|
}
|
|
|
|
func (p *ConfigColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
|
|
key := fmt.Sprintf("%s.%s", schema, table)
|
|
rules, ok := p.rules[key]
|
|
if !ok {
|
|
return []ColumnSecurity{}, nil
|
|
}
|
|
return rules, nil
|
|
}
|
|
|
|
// ConfigRowSecurityProvider provides static row security configuration
|
|
type ConfigRowSecurityProvider struct {
|
|
templates map[string]string
|
|
blocked map[string]bool
|
|
}
|
|
|
|
func NewConfigRowSecurityProvider(templates map[string]string, blocked map[string]bool) *ConfigRowSecurityProvider {
|
|
return &ConfigRowSecurityProvider{
|
|
templates: templates,
|
|
blocked: blocked,
|
|
}
|
|
}
|
|
|
|
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
|
key := fmt.Sprintf("%s.%s", schema, table)
|
|
|
|
if p.blocked[key] {
|
|
return RowSecurity{
|
|
Schema: schema,
|
|
Tablename: table,
|
|
UserID: userRef,
|
|
HasBlock: true,
|
|
}, nil
|
|
}
|
|
|
|
template := p.templates[key]
|
|
return RowSecurity{
|
|
Schema: schema,
|
|
Tablename: table,
|
|
UserID: userRef,
|
|
Template: template,
|
|
HasBlock: false,
|
|
}, nil
|
|
}
|
|
|
|
// Helper functions
|
|
// ================
|
|
|
|
// isDBClosed reports whether err indicates the *sql.DB has been closed.
|
|
func isDBClosed(err error) bool {
|
|
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
|
|
}
|
|
|
|
func parseRoles(rolesStr string) []string {
|
|
if rolesStr == "" {
|
|
return []string{}
|
|
}
|
|
return strings.Split(rolesStr, ",")
|
|
}
|
|
|
|
func parseIntHeader(r *http.Request, key string, defaultVal int) int {
|
|
val := r.Header.Get(key)
|
|
if val == "" {
|
|
return defaultVal
|
|
}
|
|
intVal, err := strconv.Atoi(val)
|
|
if err != nil {
|
|
return defaultVal
|
|
}
|
|
return intVal
|
|
}
|
|
|
|
func generateRandomString(length int) string {
|
|
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
|
|
b := make([]byte, length)
|
|
for i := range b {
|
|
b[i] = charset[time.Now().UnixNano()%int64(len(charset))]
|
|
}
|
|
return string(b)
|
|
}
|
|
|
|
// func getClaimString(claims map[string]any, key string) string {
|
|
// if claims == nil {
|
|
// return ""
|
|
// }
|
|
// if val, ok := claims[key]; ok {
|
|
// if str, ok := val.(string); ok {
|
|
// return str
|
|
// }
|
|
// }
|
|
// return ""
|
|
// }
|
|
|
|
// Password reset methods
|
|
// ======================
|
|
|
|
// RequestPasswordReset implements PasswordResettable. It calls the stored procedure
|
|
// resolvespec_password_reset_request and returns the reset token and expiry.
|
|
func (a *DatabaseAuthenticator) RequestPasswordReset(ctx context.Context, req PasswordResetRequest) (*PasswordResetResponse, error) {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.PasswordResetRequest) {
|
|
return a.requestPasswordResetDirect(ctx, req)
|
|
}
|
|
reqJSON, err := json.Marshal(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to marshal password reset request: %w", err)
|
|
}
|
|
|
|
var success bool
|
|
var errorMsg sql.NullString
|
|
var dataJSON sql.NullString
|
|
|
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.sqlNames.PasswordResetRequest)
|
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
|
})
|
|
if err != nil {
|
|
return nil, fmt.Errorf("password reset request query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("password reset request failed")
|
|
}
|
|
|
|
var response PasswordResetResponse
|
|
if dataJSON.Valid && dataJSON.String != "" {
|
|
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
|
return nil, fmt.Errorf("failed to parse password reset response: %w", err)
|
|
}
|
|
}
|
|
|
|
return &response, nil
|
|
}
|
|
|
|
// CompletePasswordReset implements PasswordResettable. It validates the token and
|
|
// updates the user's password via resolvespec_password_reset.
|
|
func (a *DatabaseAuthenticator) CompletePasswordReset(ctx context.Context, req PasswordResetCompleteRequest) error {
|
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.PasswordResetComplete) {
|
|
return a.completePasswordResetDirect(ctx, req)
|
|
}
|
|
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.runDBOpWithReconnect(func(db *sql.DB) error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1::jsonb)`, a.sqlNames.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 {
|
|
if errorMsg.Valid {
|
|
return fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return fmt.Errorf("password reset failed")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// Passkey authentication methods
|
|
// ==============================
|
|
|
|
// WithPasskey configures the DatabaseAuthenticator with a passkey provider
|
|
func (a *DatabaseAuthenticator) WithPasskey(provider PasskeyProvider) *DatabaseAuthenticator {
|
|
a.passkeyProvider = provider
|
|
return a
|
|
}
|
|
|
|
// BeginPasskeyRegistration initiates passkey registration for a user
|
|
func (a *DatabaseAuthenticator) BeginPasskeyRegistration(ctx context.Context, req PasskeyBeginRegistrationRequest) (*PasskeyRegistrationOptions, error) {
|
|
if a.passkeyProvider == nil {
|
|
return nil, fmt.Errorf("passkey provider not configured")
|
|
}
|
|
return a.passkeyProvider.BeginRegistration(ctx, req.UserID, req.Username, req.DisplayName)
|
|
}
|
|
|
|
// CompletePasskeyRegistration completes passkey registration
|
|
func (a *DatabaseAuthenticator) CompletePasskeyRegistration(ctx context.Context, req PasskeyRegisterRequest) (*PasskeyCredential, error) {
|
|
if a.passkeyProvider == nil {
|
|
return nil, fmt.Errorf("passkey provider not configured")
|
|
}
|
|
|
|
cred, err := a.passkeyProvider.CompleteRegistration(ctx, req.UserID, req.Response, req.ExpectedChallenge)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Update credential name if provided
|
|
if req.CredentialName != "" && cred.ID != "" {
|
|
_ = a.passkeyProvider.UpdateCredentialName(ctx, req.UserID, cred.ID, req.CredentialName)
|
|
}
|
|
|
|
return cred, nil
|
|
}
|
|
|
|
// BeginPasskeyAuthentication initiates passkey authentication
|
|
func (a *DatabaseAuthenticator) BeginPasskeyAuthentication(ctx context.Context, req PasskeyBeginAuthenticationRequest) (*PasskeyAuthenticationOptions, error) {
|
|
if a.passkeyProvider == nil {
|
|
return nil, fmt.Errorf("passkey provider not configured")
|
|
}
|
|
return a.passkeyProvider.BeginAuthentication(ctx, req.Username)
|
|
}
|
|
|
|
// LoginWithPasskey authenticates a user using a passkey and creates a session
|
|
func (a *DatabaseAuthenticator) LoginWithPasskey(ctx context.Context, req PasskeyLoginRequest) (*LoginResponse, error) {
|
|
if a.passkeyProvider == nil {
|
|
return nil, fmt.Errorf("passkey provider not configured")
|
|
}
|
|
|
|
// Verify passkey assertion
|
|
userID, err := a.passkeyProvider.CompleteAuthentication(ctx, req.Response, req.ExpectedChallenge)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("passkey authentication failed: %w", err)
|
|
}
|
|
|
|
// Build request JSON for passkey login stored procedure
|
|
reqData := map[string]any{
|
|
"user_id": userID,
|
|
}
|
|
if req.Claims != nil {
|
|
if ip, ok := req.Claims["ip_address"].(string); ok {
|
|
reqData["ip_address"] = ip
|
|
}
|
|
if ua, ok := req.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 sql.NullString
|
|
var dataJSON sql.NullString
|
|
|
|
runPasskeyQuery := func() error {
|
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.sqlNames.PasskeyLogin)
|
|
return a.getDB().QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
|
}
|
|
err = runPasskeyQuery()
|
|
if isDBClosed(err) {
|
|
if reconnErr := a.reconnectDB(); reconnErr == nil {
|
|
err = runPasskeyQuery()
|
|
}
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("passkey login query failed: %w", err)
|
|
}
|
|
|
|
if !success {
|
|
if errorMsg.Valid {
|
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
|
}
|
|
return nil, fmt.Errorf("passkey login failed")
|
|
}
|
|
|
|
var response 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
|
|
}
|
|
|
|
// GetPasskeyCredentials returns all passkey credentials for a user
|
|
func (a *DatabaseAuthenticator) GetPasskeyCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error) {
|
|
if a.passkeyProvider == nil {
|
|
return nil, fmt.Errorf("passkey provider not configured")
|
|
}
|
|
return a.passkeyProvider.GetCredentials(ctx, userID)
|
|
}
|
|
|
|
// DeletePasskeyCredential removes a passkey credential
|
|
func (a *DatabaseAuthenticator) DeletePasskeyCredential(ctx context.Context, userID int, credentialID string) error {
|
|
if a.passkeyProvider == nil {
|
|
return fmt.Errorf("passkey provider not configured")
|
|
}
|
|
return a.passkeyProvider.DeleteCredential(ctx, userID, credentialID)
|
|
}
|
|
|
|
// UpdatePasskeyCredentialName updates the friendly name of a credential
|
|
func (a *DatabaseAuthenticator) UpdatePasskeyCredentialName(ctx context.Context, userID int, credentialID string, name string) error {
|
|
if a.passkeyProvider == nil {
|
|
return fmt.Errorf("passkey provider not configured")
|
|
}
|
|
return a.passkeyProvider.UpdateCredentialName(ctx, userID, credentialID, name)
|
|
}
|