mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +00:00
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
283 lines
9.6 KiB
Go
283 lines
9.6 KiB
Go
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
|
|
}
|