From c9fa8c60f2fb2c891e53372385f271dff5ee107d Mon Sep 17 00:00:00 2001 From: Hein Date: Thu, 1 Oct 2026 13:19:44 +0200 Subject: [PATCH] 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 --- README.md | 2 +- pkg/common/adapters/database/bun.go | 8 + pkg/common/adapters/database/gorm.go | 15 + pkg/common/adapters/database/pgsql.go | 5 + pkg/common/interfaces.go | 10 + pkg/security/KEYSTORE.md | 41 +- pkg/security/OAUTH2.md | 4 +- .../OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md | 6 +- pkg/security/PASSKEY_QUICK_REFERENCE.md | 2 +- pkg/security/QUICK_REFERENCE.md | 27 +- pkg/security/README.md | 195 ++-- pkg/security/breaking_changes.md | 126 +++ pkg/security/chain.go | 15 + pkg/security/composite.go | 8 + pkg/security/direct_mode_test.go | 64 +- pkg/security/interfaces.go | 82 +- pkg/security/keystore.go | 59 +- pkg/security/keystore_database.go | 218 +---- pkg/security/keystore_database_direct.go | 216 ----- pkg/security/keystore_sql_names.go | 61 -- pkg/security/keystore_table_names.go | 44 - pkg/security/login_api_key_test.go | 39 + pkg/security/lookup/backends/backends.go | 162 ++++ pkg/security/lookup/backends/backends_test.go | 119 +++ .../lookup/backends/conformance_test.go | 126 +++ .../lookup/backends/container_test.go | 141 +++ pkg/security/lookup/backends/routers.go | 396 ++++++++ .../lookup/conformance/conformance.go | 564 +++++++++++ pkg/security/lookup/database.go | 44 + pkg/security/{ => lookup}/database_schema.sql | 211 +++- pkg/security/lookup/database_test.go | 84 ++ pkg/security/lookup/ddl/ddl.go | 50 + pkg/security/lookup/ddl/ddl_test.go | 58 ++ pkg/security/lookup/ddl/mssql.sql | 259 +++++ pkg/security/lookup/ddl/mysql.sql | 223 +++++ pkg/security/lookup/ddl/postgres.sql | 239 +++++ .../ddl/sqlite.sql} | 133 ++- pkg/security/lookup/deps_test.go | 34 + pkg/security/lookup/dialect/builtin.go | 93 ++ pkg/security/lookup/dialect/dialect.go | 312 ++++++ pkg/security/lookup/dialect/dialect_test.go | 270 ++++++ pkg/security/lookup/direct/auth.go | 608 ++++++++++++ pkg/security/lookup/direct/auth_test.go | 268 ++++++ pkg/security/lookup/direct/base.go | 462 +++++++++ pkg/security/lookup/direct/builder_test.go | 75 ++ pkg/security/lookup/direct/helpers_test.go | 46 + pkg/security/lookup/direct/keys.go | 197 ++++ pkg/security/lookup/direct/keys_test.go | 71 ++ pkg/security/lookup/direct/oauth.go | 380 ++++++++ pkg/security/lookup/direct/oauth_test.go | 130 +++ pkg/security/lookup/direct/passkey.go | 294 ++++++ .../lookup/direct/passkey_totp_test.go | 162 ++++ pkg/security/{ => lookup/direct}/password.go | 22 +- pkg/security/lookup/direct/password_test.go | 19 + pkg/security/lookup/direct/policy.go | 200 ++++ pkg/security/lookup/direct/policy_test.go | 93 ++ pkg/security/lookup/direct/totp.go | 159 +++ pkg/security/{ => lookup}/keystore_schema.sql | 80 +- pkg/security/lookup/lookup.go | 208 ++++ pkg/security/lookup/lookup_test.go | 165 ++++ pkg/security/lookup/mode.go | 169 ++++ pkg/security/lookup/procedure/auth.go | 318 ++++++ pkg/security/lookup/procedure/flextime.go | 54 ++ pkg/security/lookup/procedure/keys.go | 156 +++ pkg/security/lookup/procedure/oauth.go | 315 ++++++ pkg/security/lookup/procedure/passkey.go | 282 ++++++ pkg/security/lookup/procedure/policy.go | 96 ++ pkg/security/lookup/procedure/procedure.go | 92 ++ .../lookup/procedure/procedure_test.go | 226 +++++ pkg/security/lookup/procedure/totp.go | 121 +++ pkg/security/lookup/procs.go | 143 +++ pkg/security/lookup/schema.go | 346 +++++++ pkg/security/lookup/txsettings.go | 47 + pkg/security/lookup/txsettings_test.go | 93 ++ pkg/security/lookup_bridge.go | 52 + pkg/security/nosql_test.go | 42 + pkg/security/oauth2_examples.go | 2 +- pkg/security/oauth2_methods.go | 99 +- pkg/security/oauth2_methods_direct.go | 242 ----- pkg/security/oauth_server_db.go | 207 +--- pkg/security/oauth_server_db_direct.go | 209 ---- pkg/security/oauth_server_test.go | 2 +- pkg/security/passkey.go | 110 --- pkg/security/passkey_provider.go | 393 +------- pkg/security/passkey_provider_direct.go | 256 ----- pkg/security/provider.go | 91 -- pkg/security/providers.go | 908 ++---------------- pkg/security/providers/header.go | 76 ++ .../{ => providers}/keystore_authenticator.go | 28 +- .../{ => providers}/keystore_config.go | 25 +- pkg/security/providers/policy_config.go | 61 ++ pkg/security/providers/providers_test.go | 180 ++++ pkg/security/providers_direct.go | 535 ----------- pkg/security/providers_test.go | 171 ---- pkg/security/query_mode.go | 172 ---- pkg/security/query_mode_test.go | 165 ---- pkg/security/sectypes/auth.go | 76 ++ pkg/security/sectypes/deps_test.go | 20 + pkg/security/sectypes/keys.go | 62 ++ pkg/security/sectypes/oauth.go | 40 + pkg/security/sectypes/passkey.go | 112 +++ pkg/security/sectypes/policy.go | 98 ++ pkg/security/sectypes/twofactor.go | 10 + pkg/security/sql_names.go | 267 ----- pkg/security/sql_names_test.go | 145 --- pkg/security/table_names.go | 101 -- pkg/security/table_names_test.go | 134 --- .../authenticator.go} | 45 +- .../integration_test.go} | 96 +- .../memory.go} | 37 +- pkg/security/{ => totp}/totp.go | 52 +- pkg/security/{ => totp}/totp_test.go | 42 +- pkg/security/totp_provider_database.go | 241 +---- pkg/security/totp_provider_database_direct.go | 151 --- pkg/security/totp_provider_database_test.go | 2 +- pkg/security/txsettings.go | 39 +- pkg/security/txsettings_test.go | 71 -- pkg/security/types.go | 54 ++ 118 files changed, 11218 insertions(+), 5565 deletions(-) create mode 100644 pkg/security/breaking_changes.md delete mode 100644 pkg/security/keystore_database_direct.go delete mode 100644 pkg/security/keystore_sql_names.go delete mode 100644 pkg/security/keystore_table_names.go create mode 100644 pkg/security/login_api_key_test.go create mode 100644 pkg/security/lookup/backends/backends.go create mode 100644 pkg/security/lookup/backends/backends_test.go create mode 100644 pkg/security/lookup/backends/conformance_test.go create mode 100644 pkg/security/lookup/backends/container_test.go create mode 100644 pkg/security/lookup/backends/routers.go create mode 100644 pkg/security/lookup/conformance/conformance.go create mode 100644 pkg/security/lookup/database.go rename pkg/security/{ => lookup}/database_schema.sql (89%) create mode 100644 pkg/security/lookup/database_test.go create mode 100644 pkg/security/lookup/ddl/ddl.go create mode 100644 pkg/security/lookup/ddl/ddl_test.go create mode 100644 pkg/security/lookup/ddl/mssql.sql create mode 100644 pkg/security/lookup/ddl/mysql.sql create mode 100644 pkg/security/lookup/ddl/postgres.sql rename pkg/security/{database_schema_sqlite.sql => lookup/ddl/sqlite.sql} (51%) create mode 100644 pkg/security/lookup/deps_test.go create mode 100644 pkg/security/lookup/dialect/builtin.go create mode 100644 pkg/security/lookup/dialect/dialect.go create mode 100644 pkg/security/lookup/dialect/dialect_test.go create mode 100644 pkg/security/lookup/direct/auth.go create mode 100644 pkg/security/lookup/direct/auth_test.go create mode 100644 pkg/security/lookup/direct/base.go create mode 100644 pkg/security/lookup/direct/builder_test.go create mode 100644 pkg/security/lookup/direct/helpers_test.go create mode 100644 pkg/security/lookup/direct/keys.go create mode 100644 pkg/security/lookup/direct/keys_test.go create mode 100644 pkg/security/lookup/direct/oauth.go create mode 100644 pkg/security/lookup/direct/oauth_test.go create mode 100644 pkg/security/lookup/direct/passkey.go create mode 100644 pkg/security/lookup/direct/passkey_totp_test.go rename pkg/security/{ => lookup/direct}/password.go (66%) create mode 100644 pkg/security/lookup/direct/password_test.go create mode 100644 pkg/security/lookup/direct/policy.go create mode 100644 pkg/security/lookup/direct/policy_test.go create mode 100644 pkg/security/lookup/direct/totp.go rename pkg/security/{ => lookup}/keystore_schema.sql (67%) create mode 100644 pkg/security/lookup/lookup.go create mode 100644 pkg/security/lookup/lookup_test.go create mode 100644 pkg/security/lookup/mode.go create mode 100644 pkg/security/lookup/procedure/auth.go create mode 100644 pkg/security/lookup/procedure/flextime.go create mode 100644 pkg/security/lookup/procedure/keys.go create mode 100644 pkg/security/lookup/procedure/oauth.go create mode 100644 pkg/security/lookup/procedure/passkey.go create mode 100644 pkg/security/lookup/procedure/policy.go create mode 100644 pkg/security/lookup/procedure/procedure.go create mode 100644 pkg/security/lookup/procedure/procedure_test.go create mode 100644 pkg/security/lookup/procedure/totp.go create mode 100644 pkg/security/lookup/procs.go create mode 100644 pkg/security/lookup/schema.go create mode 100644 pkg/security/lookup/txsettings.go create mode 100644 pkg/security/lookup/txsettings_test.go create mode 100644 pkg/security/lookup_bridge.go create mode 100644 pkg/security/nosql_test.go delete mode 100644 pkg/security/oauth2_methods_direct.go delete mode 100644 pkg/security/oauth_server_db_direct.go delete mode 100644 pkg/security/passkey_provider_direct.go create mode 100644 pkg/security/providers/header.go rename pkg/security/{ => providers}/keystore_authenticator.go (76%) rename pkg/security/{ => providers}/keystore_config.go (84%) create mode 100644 pkg/security/providers/policy_config.go create mode 100644 pkg/security/providers/providers_test.go delete mode 100644 pkg/security/providers_direct.go delete mode 100644 pkg/security/query_mode.go delete mode 100644 pkg/security/query_mode_test.go create mode 100644 pkg/security/sectypes/auth.go create mode 100644 pkg/security/sectypes/deps_test.go create mode 100644 pkg/security/sectypes/keys.go create mode 100644 pkg/security/sectypes/oauth.go create mode 100644 pkg/security/sectypes/passkey.go create mode 100644 pkg/security/sectypes/policy.go create mode 100644 pkg/security/sectypes/twofactor.go delete mode 100644 pkg/security/sql_names.go delete mode 100644 pkg/security/sql_names_test.go delete mode 100644 pkg/security/table_names.go delete mode 100644 pkg/security/table_names_test.go rename pkg/security/{totp_middleware.go => totp/authenticator.go} (62%) rename pkg/security/{totp_integration_test.go => totp/integration_test.go} (77%) rename pkg/security/{totp_provider_memory.go => totp/memory.go} (68%) rename pkg/security/{ => totp}/totp.go (73%) rename pkg/security/{ => totp}/totp_test.go (87%) delete mode 100644 pkg/security/totp_provider_database_direct.go create mode 100644 pkg/security/types.go diff --git a/README.md b/README.md index 4af3516..56daeb0 100644 --- a/README.md +++ b/README.md @@ -646,7 +646,7 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md). #### Security -Authentication and authorization framework with hooks integration. Database-backed providers use PostgreSQL stored procedures by default, with a portable Direct mode (plain Go/SQL) for SQLite, MySQL, or Postgres without the procedures installed. +Authentication and authorization framework with hooks integration. Database-backed providers use PostgreSQL stored procedures by default, with a direct SQL backend (SQLite, MySQL, SQL Server, or Postgres without the procedures) selected through `lookup.Config`. For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Direct Mode" for the SQLite/portable-SQL path). diff --git a/pkg/common/adapters/database/bun.go b/pkg/common/adapters/database/bun.go index 9fcf939..7d87399 100644 --- a/pkg/common/adapters/database/bun.go +++ b/pkg/common/adapters/database/bun.go @@ -275,6 +275,14 @@ func (b *BunAdapter) GetUnderlyingDB() interface{} { return b.getDB() } +// SQLDB implements common.SQLDBProvider. +func (b *BunAdapter) SQLDB() *sql.DB { + if db := b.getDB(); db != nil { + return db.DB + } + return nil +} + func (b *BunAdapter) DriverName() string { // Normalize Bun's dialect name to match the project's canonical vocabulary. // Bun returns "pg" for PostgreSQL; the rest of the project uses "postgres". diff --git a/pkg/common/adapters/database/gorm.go b/pkg/common/adapters/database/gorm.go index a681447..46bcb35 100644 --- a/pkg/common/adapters/database/gorm.go +++ b/pkg/common/adapters/database/gorm.go @@ -2,6 +2,7 @@ package database import ( "context" + "database/sql" "fmt" "reflect" "strings" @@ -227,6 +228,20 @@ func (g *GormAdapter) GetUnderlyingDB() interface{} { return g.getDB() } +// SQLDB implements common.SQLDBProvider. It returns nil when GORM has no *sql.DB +// (for example a ConnPool that is not database/sql). +func (g *GormAdapter) SQLDB() *sql.DB { + db := g.getDB() + if db == nil { + return nil + } + sqlDB, err := db.DB() + if err != nil { + return nil + } + return sqlDB +} + func (g *GormAdapter) DriverName() string { return normalizeGormDriverName(g.getDB()) } diff --git a/pkg/common/adapters/database/pgsql.go b/pkg/common/adapters/database/pgsql.go index 8d7ff7a..9616601 100644 --- a/pkg/common/adapters/database/pgsql.go +++ b/pkg/common/adapters/database/pgsql.go @@ -223,6 +223,11 @@ func (p *PgSQLAdapter) GetUnderlyingDB() interface{} { return p.db } +// SQLDB implements common.SQLDBProvider. +func (p *PgSQLAdapter) SQLDB() *sql.DB { + return p.db +} + func (p *PgSQLAdapter) DriverName() string { return p.driverName } diff --git a/pkg/common/interfaces.go b/pkg/common/interfaces.go index 83cd3dc..00e67e7 100644 --- a/pkg/common/interfaces.go +++ b/pkg/common/interfaces.go @@ -2,6 +2,7 @@ package common import ( "context" + "database/sql" "encoding/json" "io" "net/http" @@ -38,6 +39,15 @@ type Database interface { DriverName() string } +// SQLDBProvider is implemented by adapters that wrap a *sql.DB (directly or through an ORM). +// It lets packages that work on database/sql, such as pkg/security/lookup, reuse the +// connection an application already configured. Transaction adapters do not implement it. +type SQLDBProvider interface { + // SQLDB returns the current underlying *sql.DB. After an adapter reconnects it + // returns the new handle, so do not cache it across reconnects. + SQLDB() *sql.DB +} + // SelectQuery interface for building SELECT queries (compatible with both GORM and Bun) type SelectQuery interface { Model(model interface{}) SelectQuery diff --git a/pkg/security/KEYSTORE.md b/pkg/security/KEYSTORE.md index dab4d6e..a1fe312 100644 --- a/pkg/security/KEYSTORE.md +++ b/pkg/security/KEYSTORE.md @@ -19,7 +19,7 @@ In-memory store seeded from a static list. Suitable for a small, fixed set of se ```go // Pre-load keys from config (KeyHash = SHA-256 hex of the raw key) -store := security.NewConfigKeyStore([]security.UserKey{ +store := providers.NewConfigKeyStore([]security.UserKey{ { UserID: 1, KeyType: security.KeyTypeGenericAPI, @@ -33,7 +33,7 @@ store := security.NewConfigKeyStore([]security.UserKey{ ### DatabaseKeyStore -Backed by PostgreSQL stored procedures. Supports optional caching (default 2-minute TTL). Apply `keystore_schema.sql` before use. +Backed by PostgreSQL stored procedures by default, or by the `user_keys` table in direct mode (any supported dialect). Supports optional caching (default 2-minute TTL). Apply `lookup/keystore_schema.sql` before use. ```go db, _ := sql.Open("postgres", dsn) @@ -43,8 +43,8 @@ store := security.NewDatabaseKeyStore(db) // With options store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{ CacheTTL: 5 * time.Minute, - SQLNames: &security.KeyStoreSQLNames{ - ValidateKey: "myapp_keystore_validate", // override one procedure name + Lookup: lookup.Config{ + Procs: lookup.ProcNames{KeystoreValidateKey: "myapp_keystore_validate"}, // override one procedure name }, }) ``` @@ -85,9 +85,9 @@ Keys are extracted from the request in this order: 3. `X-API-Key: ` ```go -auth := security.NewKeyStoreAuthenticator(store, "") // "" = accept any key type +auth := providers.NewKeyStoreAuthenticator(store, "") // "" = accept any key type // Restrict to a specific type: -auth = security.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI) +auth = providers.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI) ``` Plug it into a handler: @@ -109,10 +109,10 @@ On successful validation the request context receives a `UserContext` where: ## Database setup -Apply `keystore_schema.sql` to your PostgreSQL database. It requires the `users` table from the main `database_schema.sql`. +Apply `lookup/keystore_schema.sql` to your PostgreSQL database. It requires the `users` table from the main `lookup/database_schema.sql`. ```sql -\i pkg/security/keystore_schema.sql +\i pkg/security/lookup/keystore_schema.sql ``` This creates: @@ -123,28 +123,23 @@ This creates: - `resolvespec_keystore_delete_key(p_user_id, p_key_id)` - `resolvespec_keystore_validate_key(p_key_hash, p_key_type)` -### Custom procedure names +### Custom names and modes ```go store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{ - SQLNames: &security.KeyStoreSQLNames{ - GetUserKeys: "myschema_get_keys", - CreateKey: "myschema_create_key", - DeleteKey: "myschema_delete_key", - ValidateKey: "myschema_validate_key", + Lookup: lookup.Config{ + Procs: lookup.ProcNames{ + KeystoreGetUserKeys: "myschema_get_keys", + KeystoreCreateKey: "myschema_create_key", + KeystoreDeleteKey: "myschema_delete_key", + KeystoreValidateKey: "myschema_validate_key", + }, }, }) - -// Validate names at startup -names := &security.KeyStoreSQLNames{ - GetUserKeys: "myschema_get_keys", - // ... -} -if err := security.ValidateKeyStoreSQLNames(names); err != nil { - log.Fatal(err) -} ``` +Names are validated when the store is first used. On Postgres the key store calls the procedures by default; on SQLite, MySQL and SQL Server (or with `Mode: lookup.ModeDirect`) it reads and writes the `user_keys` table directly (see `lookup/ddl`). Table and column names are configurable through `lookup.Config.Schema`. + ## Security notes - Raw keys are never stored. Only the SHA-256 hex digest is persisted. diff --git a/pkg/security/OAUTH2.md b/pkg/security/OAUTH2.md index cbfcbec..7dd5075 100644 --- a/pkg/security/OAUTH2.md +++ b/pkg/security/OAUTH2.md @@ -21,7 +21,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl ### 1. Database Setup ```sql --- Run the schema from database_schema.sql +-- Run the schema from lookup/database_schema.sql CREATE TABLE IF NOT EXISTS users ( id SERIAL PRIMARY KEY, username VARCHAR(255) NOT NULL UNIQUE, @@ -53,7 +53,7 @@ CREATE TABLE IF NOT EXISTS user_sessions ( ); -- OAuth2 stored procedures (7 functions) --- See database_schema.sql for full implementation +-- See lookup/database_schema.sql for full implementation ``` ### 2. Google OAuth2 diff --git a/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md b/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md index b261152..13a3f3e 100644 --- a/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md +++ b/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md @@ -43,16 +43,16 @@ CREATE TABLE IF NOT EXISTS user_sessions ( **`resolvespec_oauth_getrefreshtoken(p_refresh_token)`** - Gets OAuth2 session data by refresh token - Returns: `{user_id, access_token, token_type, expiry}` -- Location: `database_schema.sql:714` +- Location: `lookup/database_schema.sql:714` **`resolvespec_oauth_updaterefreshtoken(p_update_data)`** - Updates session with new tokens after refresh - Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}` -- Location: `database_schema.sql:752` +- Location: `lookup/database_schema.sql:752` **`resolvespec_oauth_getuser(p_user_id)`** - Gets user data by ID for building UserContext -- Location: `database_schema.sql:791` +- Location: `lookup/database_schema.sql:791` --- diff --git a/pkg/security/PASSKEY_QUICK_REFERENCE.md b/pkg/security/PASSKEY_QUICK_REFERENCE.md index 74f41fa..4bdbe84 100644 --- a/pkg/security/PASSKEY_QUICK_REFERENCE.md +++ b/pkg/security/PASSKEY_QUICK_REFERENCE.md @@ -6,7 +6,7 @@ Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthe ## Setup ### Database Schema -Run the passkey SQL schema (in database_schema.sql): +Run the passkey SQL schema (in lookup/database_schema.sql): - Creates `user_passkey_credentials` table - Adds stored procedures for passkey operations diff --git a/pkg/security/QUICK_REFERENCE.md b/pkg/security/QUICK_REFERENCE.md index ac9971d..36e54db 100644 --- a/pkg/security/QUICK_REFERENCE.md +++ b/pkg/security/QUICK_REFERENCE.md @@ -6,7 +6,7 @@ // Step 1: Create security providers auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended) // OR: auth := security.NewJWTAuthenticator("secret-key", db) -// OR: auth := security.NewHeaderAuthenticator() +// OR: auth := providers.NewHeaderAuthenticator() // OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2 colSec := security.NewDatabaseColumnSecurityProvider(db) @@ -55,7 +55,7 @@ All stored procedures return structured results: - Session/Login: `(p_success bool, p_error text, p_data jsonb)` - Security: `(p_success bool, p_error text, p_rules jsonb)` -See `database_schema.sql` for complete definitions. +See `lookup/database_schema.sql` for complete definitions. --- @@ -182,7 +182,7 @@ auth := security.NewDatabaseAuthenticator(db) // Requires these tables: // - users (id, username, email, password, user_level, roles, is_active) // - user_sessions (session_token, user_id, expires_at, created_at, last_activity_at) -// See database_schema.sql for full schema +// See lookup/database_schema.sql for full schema // Features: // - Login with username/password @@ -313,16 +313,15 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context, } query := ` - SELECT control, accesstype, jsonvalue - FROM core.secaccess - WHERE rid_hub IN ( - SELECT rid_hub_parent FROM core.hub_link - WHERE rid_hub_child = ? AND parent_hubtype = 'secgroup' - ) - AND control ILIKE ? + SELECT schema_name || '.' || table_name || '.' || column_path AS control, + access_type AS accesstype, COALESCE(extra_filters, '') AS jsonvalue + FROM sec_column_rules + WHERE is_active = true + AND lower(schema_name) = lower(?) AND lower(table_name) = lower(?) + AND (user_id = ? OR group_id IN (SELECT group_id FROM sec_group_members WHERE user_id = ?)) ` - err := p.db.WithContext(ctx).Raw(query, userID, fmt.Sprintf("%s.%s%%", schema, table)).Scan(&records).Error + err := p.db.WithContext(ctx).Raw(query, schema, table, userID, userID).Scan(&records).Error if err != nil { return nil, err } @@ -378,19 +377,19 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i ```go // Test Authenticator -auth := security.NewHeaderAuthenticator() +auth := providers.NewHeaderAuthenticator() req := httptest.NewRequest("GET", "/", nil) req.Header.Set("X-User-ID", "123") userCtx, err := auth.Authenticate(req) assert.Equal(t, 123, userCtx.UserID) // Test ColumnSecurityProvider -colSec := security.NewConfigColumnSecurityProvider(rules) +colSec := providers.NewConfigColumnSecurityProvider(rules) cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees") assert.Equal(t, "mask", cols[0].Accesstype) // Test RowSecurityProvider -rowSec := security.NewConfigRowSecurityProvider(templates, blocked) +rowSec := providers.NewConfigRowSecurityProvider(templates, blocked) row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders") assert.Equal(t, "user_id = {UserID}", row.Template) ``` diff --git a/pkg/security/README.md b/pkg/security/README.md index 1a51538..e6e8cbc 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -35,6 +35,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic | `resolvespec_login` | Session-based login | DatabaseAuthenticator | | `resolvespec_logout` | Session invalidation | DatabaseAuthenticator | | `resolvespec_session` | Session validation | DatabaseAuthenticator | +| `resolvespec_login_api_key` | Exchange a raw header/generic API key for a session (defined in `lookup/keystore_schema.sql`; direct mode reads `user_keys`) | DatabaseAuthenticator.LoginWithAPIKey | | `resolvespec_session_update` | Update session activity | DatabaseAuthenticator | | `resolvespec_refresh_token` | Token refresh | DatabaseAuthenticator | | `resolvespec_jwt_login` | JWT user validation | JWTAuthenticator | @@ -50,95 +51,96 @@ Type-safe, composable security system for ResolveSpec with support for authentic | `resolvespec_password_reset_request` | Create password reset token | DatabaseAuthenticator | | `resolvespec_password_reset` | Validate token and set new password | DatabaseAuthenticator | -See `database_schema.sql` for complete stored procedure definitions and examples. +See `lookup/database_schema.sql` for complete stored procedure definitions and examples. -**Not on Postgres, or don't have the procedures installed?** See [Direct Mode](#direct-mode-portable-sql-without-stored-procedures) below — every provider that calls a `resolvespec_*` procedure also has a portable Go/SQL implementation that works on SQLite, MySQL, or plain Postgres. +**Not on Postgres, or don't have the procedures installed?** See [Database access (lookup)](#database-access-lookup) below: every provider can also work directly on tables, on SQLite, MySQL, SQL Server or plain Postgres. -## Direct Mode (portable SQL without stored procedures) +## Database access (lookup) -Every database-backed provider (`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseTwoFactorProvider`, `DatabasePasskeyProvider`, the OAuth2 methods/server, `DatabaseKeyStore`) has two code paths: +`pkg/security` itself contains no SQL. Every database-backed provider (`DatabaseAuthenticator`, `JWTAuthenticator`, column/row security, `DatabaseTwoFactorProvider`, `DatabasePasskeyProvider`, the OAuth2 methods/server, `DatabaseKeyStore`) calls a store interface from `pkg/security/lookup`. Two backends implement each store: -- **Procedure mode** — calls the configured `resolvespec_*` stored procedure (original behavior, Postgres-only). -- **Direct mode** — reimplements the same logic in Go using plain parameterized SQL against configurable table names. Works on SQLite, MySQL, or a Postgres database where the procedures were never deployed. +- **procedure** (`lookup/procedure`): calls the `resolvespec_*` stored procedures (`p_success` / `p_error` / `p_data` contract). Postgres only. +- **direct** (`lookup/direct`): plain parameterized SQL on tables, rendered by a per-database dialect (`lookup/dialect`: postgres, sqlite, mysql, mssql). Table and column names are configurable. -### QueryMode +`lookup/backends.New(db, cfg, opts)` builds a `lookup.Provider` (all stores) and routes each operation to a backend. The security constructors do this for you from `lookup.Config`; pass a ready `*lookup.Provider` with `LookupProvider` / `WithLookupProvider` to share one between components. -Selection is controlled per-provider by a `QueryMode`: +### Choosing the mode ```go -type QueryMode int - -const ( - ModeAuto QueryMode = iota // default - ModeProcedure - ModeDirect -) -``` - -- **`ModeAuto`** (default, zero value) — auto-detects per connection: - - SQLite/MySQL drivers → Direct mode, no probing. - - Postgres drivers (`lib/pq`, `pgx`) → probes `pg_proc` for the configured procedure name and uses it **only if it actually exists**; otherwise falls back to Direct mode. The result is cached per procedure name and reset on reconnect. - - Any other/unrecognized driver (including `sqlmock` test doubles) → defaults to Procedure mode, preserving existing behavior for callers that don't expose an identifiable driver type. -- **`ModeProcedure`** — always calls the stored procedure, regardless of dialect. -- **`ModeDirect`** — always uses the portable Go/SQL path, never the stored procedure. - -Set it via the provider's `Options` struct or `With...` chain method: - -```go -auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{ - QueryMode: security.ModeDirect, // force Direct mode, e.g. for SQLite -}) - -tfaProvider := security.NewDatabaseTwoFactorProvider(sqliteDB, nil). - WithQueryMode(security.ModeDirect) -``` - -On a real SQLite/MySQL connection you can usually leave `QueryMode` unset — `ModeAuto` detects the dialect and uses Direct mode automatically. - -### TableNames / KeyStoreTableNames - -Direct mode reads/writes plain tables instead of calling procedures, so table names are configurable the same way procedure names are (`SQLNames`): - -```go -type TableNames struct { - Users string // default: "users" - UserSessions string // default: "user_sessions" - TokenBlacklist string // default: "token_blacklist" - UserTOTPBackupCodes string // default: "user_totp_backup_codes" - UserPasskeyCredentials string // default: "user_passkey_credentials" - UserPasswordResets string // default: "user_password_resets" - OAuthClients string // default: "oauth_clients" - OAuthCodes string // default: "oauth_codes" -} - -type KeyStoreTableNames struct { - UserKeys string // default: "user_keys" — used by DatabaseKeyStore +type Config struct { + Dialect string // "postgres", "sqlite", "mysql", "mssql", or one you registered; empty = detect from the driver + Mode lookup.Mode // default for every operation + Overrides map[lookup.Op]lookup.Mode // per-operation mode, e.g. lookup.OpSession: lookup.ModeDirect + Procs lookup.ProcNames // procedure names, empty fields keep the default + Schema lookup.Schema // table/column names, missing entries keep the default } ``` -`DefaultTableNames()` / `MergeTableNames()` / `ValidateTableNames()` mirror `DefaultSQLNames()` / `MergeSQLNames()` / `ValidateSQLNames()`. Set custom names via the same `Options`/`With...` surface as `QueryMode`: +| Mode | Behaviour | +|---|---| +| `ModeDefault` (zero value) | stored procedure on Postgres, direct SQL on every other dialect | +| `ModeProcedure` | always the procedure; an error on a non-Postgres dialect | +| `ModeDirect` | always direct SQL | +| `ModeAuto` | Postgres: probe `pg_proc` once per procedure (cached, reset on reconnect), use it if present, else direct. Other dialects: direct | + +An impossible combination (procedure on SQLite) fails when the provider is built, not on the first request. If the dialect is not set and cannot be detected from the driver, Postgres is assumed. A bad configuration makes every call return the error (fail closed). ```go -auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{ - TableNames: &security.TableNames{Users: "app_users"}, // only override what differs +// SQLite or MySQL: nothing to configure, direct SQL is the default. +auth := security.NewDatabaseAuthenticator(sqliteDB) + +// Postgres without the procedures installed: use tables only. +auth = security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{ + Lookup: lookup.Config{Mode: lookup.ModeDirect}, }) + +// Postgres, procedures for everything except session lookups. +auth = security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{ + Lookup: lookup.Config{Overrides: map[lookup.Op]lookup.Mode{lookup.OpSession: lookup.ModeDirect}}, +}) + +tfa := security.NewDatabaseTwoFactorProvider(db, nil).WithLookup(lookup.Config{Mode: lookup.ModeDirect}) ``` -`oauth2_methods.go` and `oauth_server_db.go` are methods on `*DatabaseAuthenticator` and reuse its `TableNames`/`QueryMode`; there's no separate config for them. +Other components take the same `Lookup` / `LookupProvider` options (or `WithLookup` / `WithLookupProvider` on the chain-style types). -### Schema +### Custom names -`database_schema_sqlite.sql` is the portable companion to `database_schema.sql` — plain `CREATE TABLE` statements only (no functions, no triggers, no `jsonb`/`bytea`/array types), covering every table Direct mode reads or writes. Use it to stand up a SQLite (or adapt for MySQL) database for Direct mode. +```go +cfg := lookup.Config{ + Procs: lookup.ProcNames{Login: "myapp_login"}, // only override what differs + Schema: lookup.Schema{ + lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"username": "login_name"}}, + }, +} +``` -### What's NOT covered +Procedure names, table names and column names are validated as identifiers at construction. `Schema` entries may also set `Schema` to qualify a table (`schema.table`). -`ColumnSecurityProvider`/`RowSecurityProvider` (`resolvespec_column_security` / `resolvespec_row_security`) query an external `core.secaccess`/`core.hub_link` schema this package doesn't own. Direct mode has no portable equivalent to fabricate for these and returns `security.ErrDirectModeUnsupported` — use `ConfigColumnSecurityProvider`/`ConfigRowSecurityProvider` instead when not running against Postgres with those procedures installed. +### Schemas + +| File | Purpose | +|---|---| +| `lookup/database_schema.sql` | Postgres: tables **and** stored procedures (procedure backend) | +| `lookup/keystore_schema.sql` | Postgres: `user_keys` table and key store procedures | +| `lookup/ddl/{postgres,sqlite,mysql,mssql}.sql` | tables only, for the direct backend, with the default names | + +Read them from Go with `ddl.SQL("sqlite")` or, for drivers that reject multi-statement execution (MySQL, SQL Server), `ddl.Statements("mysql")`. Do not mix `ddl/postgres.sql` with `database_schema.sql`: the procedure schema stores passkey credential ids as `bytea` and OAuth lists as `text[]`, the direct backend stores base64 / JSON text. The `ddl` files are starting points: adjust types and collations to your deployment, and set `lookup.Config.Schema` if you rename anything. + +### Column and row security + +`ColumnSecurityProvider` / `RowSecurityProvider` read `sec_group_members` (optional), `sec_column_rules` and `sec_row_rules`, in both backends. A rule belongs to one user or one group; rules apply to the exact schema and table (case-insensitive, never a prefix); a blocking row rule wins, otherwise row templates are combined with `AND`. A non-numeric user reference is an error, and no rule means no restriction from this provider. `WithNoGroupTables()` skips the membership table. ### Behavioral notes -- Direct mode matches Procedure mode's current behavior exactly, including its TODOs — e.g. passwords are compared as-is (the stored procedures don't verify bcrypt hashes yet either; see the TODO in `resolvespec_login`/`resolvespec_password_reset`). -- Session tokens generated by Direct mode use the same `sess__` shape as the plpgsql procedures. -- `bytea`/array/`jsonb` Postgres-only columns (passkey credentials, OAuth2 client scopes, keystore `meta`) are stored as base64/JSON-encoded `TEXT` in Direct mode — transparent to callers, since the Go-level API already deals in those same encodings. +- Direct login, register, refresh, API-key login, password reset and passkey login run in one transaction. +- Passwords are stored as bcrypt; legacy cleartext values are accepted at login and only rewritten when `UpgradePasswordHash` is enabled. +- Session tokens use the shape `sess__` in both backends. +- Direct mode stores `bytea` / array / `jsonb` values (passkey credentials, OAuth client lists, key meta) as base64 / JSON text; the Go API is unchanged. +- OAuth authorization codes are consumed atomically. +- Adding a database: implement `dialect.Dialect`, register it with `dialect.Register`, then set `Config.Dialect`. +- Backend conformance: `lookup/conformance` is one behavioural suite run against every backend (`go test ./pkg/security/lookup/backends -run TestConformance`). SQLite runs always; Postgres (procedure and direct), MySQL and SQL Server run when `RESOLVESPEC_TEST_PG_DSN`, `RESOLVESPEC_TEST_PG_DIRECT_DSN`, `RESOLVESPEC_TEST_MYSQL_DSN` or `RESOLVESPEC_TEST_MSSQL_DSN` is set (see the comment in `backends/conformance_test.go`). Rows are prefixed and removed afterwards. +- Migration from the old `SQLNames` / `TableNames` / `QueryMode` API: see `breaking_changes.md`. ## Quick Start @@ -297,7 +299,7 @@ type UserContext struct { **HeaderAuthenticator** - Simple header-based authentication: ```go -auth := security.NewHeaderAuthenticator() +auth := providers.NewHeaderAuthenticator() // Expects: X-User-ID, X-User-Name, X-User-Level, etc. ``` @@ -307,7 +309,7 @@ auth := security.NewDatabaseAuthenticator(db) // Supports: Login, Logout, Session management, Token refresh // All operations use stored procedures: resolvespec_login, resolvespec_logout, // resolvespec_session, resolvespec_session_update, resolvespec_refresh_token -// Requires: users and user_sessions tables + stored procedures (see database_schema.sql) +// Requires: users and user_sessions tables + stored procedures (see lookup/database_schema.sql) ``` **JWTAuthenticator** - JWT token authentication with login/logout: @@ -318,19 +320,19 @@ auth := security.NewJWTAuthenticator("secret-key", db) // Note: Requires JWT library installation for token signing/verification ``` -**TwoFactorAuthenticator** - Wraps any authenticator with TOTP 2FA: +**totp.Authenticator** - Wraps any authenticator with TOTP 2FA: ```go baseAuth := security.NewDatabaseAuthenticator(db) // Use in-memory provider (for testing) -tfaProvider := security.NewMemoryTwoFactorProvider(nil) +tfaProvider := totp.NewMemoryProvider(nil) // Or use database provider (for production) tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil) // Requires: users table with totp fields, user_totp_backup_codes table -// Requires: resolvespec_totp_* stored procedures (see totp_database_schema.sql) +// Requires: resolvespec_totp_* stored procedures (see lookup/database_schema.sql) -auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) +auth := totp.NewAuthenticator(baseAuth, tfaProvider, nil) // Supports: TOTP codes, backup codes, QR code generation // Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc. ``` @@ -341,7 +343,7 @@ auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) ```go colSec := security.NewDatabaseColumnSecurityProvider(db) // Uses stored procedure: resolvespec_column_security -// Queries core.secaccess and core.hub_link tables +// Reads sec_column_rules (user rules + rules of the user's sec_group_members groups) ``` **ConfigColumnSecurityProvider** - Static configuration: @@ -351,7 +353,7 @@ rules := map[string][]security.ColumnSecurity{ {Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5}, }, } -colSec := security.NewConfigColumnSecurityProvider(rules) +colSec := providers.NewConfigColumnSecurityProvider(rules) ``` ### Row Security Providers @@ -370,7 +372,7 @@ templates := map[string]string{ blocked := map[string]bool{ "public.admin_logs": true, } -rowSec := security.NewConfigRowSecurityProvider(templates, blocked) +rowSec := providers.NewConfigRowSecurityProvider(templates, blocked) ``` ## Usage Examples @@ -381,7 +383,7 @@ rowSec := security.NewConfigRowSecurityProvider(templates, blocked) func main() { db := setupDatabase() - // Run migrations (see database_schema.sql) + // Run migrations (see lookup/database_schema.sql) // db.Exec("CREATE TABLE users ...") // db.Exec("CREATE TABLE user_sessions ...") @@ -475,8 +477,8 @@ func handleRefresh(securityList *security.SecurityList) http.HandlerFunc { ```go // 1. Wrap existing authenticator with 2FA support baseAuth := security.NewDatabaseAuthenticator(db) -tfaProvider := security.NewMemoryTwoFactorProvider(nil) // Use custom DB implementation in production -tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) +tfaProvider := totp.NewMemoryProvider(nil) // Use custom DB implementation in production +tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil) // 2. Use as normal authenticator provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec) @@ -548,18 +550,18 @@ has2FA, err := tfaProvider.Get2FAStatus(userID) // Uses PostgreSQL stored procedures for all operations db := setupDatabase() -// Run migrations from totp_database_schema.sql +// Run migrations from lookup/database_schema.sql // - Add totp_secret, totp_enabled, totp_enabled_at to users table // - Create user_totp_backup_codes table // - Create resolvespec_totp_* stored procedures tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil) -tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) +tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil) ``` **Option 2: Implement Custom Provider** -Implement `TwoFactorAuthProvider` for custom storage: +Implement `totp.AuthProvider` for custom storage: ```go type DBTwoFactorProvider struct { @@ -585,15 +587,15 @@ func (p *DBTwoFactorProvider) Get2FASecret(userID int) (string, error) { ### Configuration ```go -config := &security.TwoFactorConfig{ +config := &totp.Config{ Algorithm: "SHA256", // SHA1, SHA256, SHA512 Digits: 8, // 6 or 8 Period: 30, // Seconds per code SkewWindow: 2, // Accept codes ±2 periods } -totp := security.NewTOTPGenerator(config) -tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, config) +totp := totp.NewGenerator(config) +tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, config) ``` ### API Response Structure @@ -662,9 +664,9 @@ func main() { } // Create providers - auth := security.NewHeaderAuthenticator() - colSec := security.NewConfigColumnSecurityProvider(columnRules) - rowSec := security.NewConfigRowSecurityProvider(rowTemplates, nil) + auth := providers.NewHeaderAuthenticator() + colSec := providers.NewConfigColumnSecurityProvider(columnRules) + rowSec := providers.NewConfigRowSecurityProvider(rowTemplates, nil) // Combine providers and register hooks provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec) @@ -1013,7 +1015,7 @@ restheadspec.RegisterSecurityHooks(handler, securityList) // or funcspec/resolve ### DB Requirements -Run the migrations in `database_schema.sql`: +Run the migrations in `lookup/database_schema.sql`: - `user_password_resets` table (`user_id`, `token_hash` SHA-256, `expires_at`, `used`, `used_at`) - `resolvespec_password_reset_request` stored procedure - `resolvespec_password_reset` stored procedure @@ -1050,13 +1052,14 @@ err = auth.CompletePasswordReset(ctx, security.PasswordResetCompleteRequest{ - `RequestPasswordReset` always returns success even when the email/username is not found, preventing user enumeration - Hash the new password with bcrypt before storing (pgcrypto `crypt`/`gen_salt`) — see the TODO comment in `resolvespec_password_reset` -### SQLNames +### Procedure names + +Set through `lookup.Config.Procs` (`lookup.ProcNames`): ```go -type SQLNames struct { - // ... - PasswordResetRequest string // default: "resolvespec_password_reset_request" - PasswordResetComplete string // default: "resolvespec_password_reset" +lookup.ProcNames{ + PasswordResetRequest: "resolvespec_password_reset_request", // default + PasswordResetComplete: "resolvespec_password_reset", // default } ``` @@ -1151,7 +1154,7 @@ http.ListenAndServe(":8080", mux) When `PersistClients: true` or `PersistCodes: true`, the server calls the corresponding `DatabaseAuthenticator` methods. Both flags default to `false` (in-memory maps). Enable both for multi-instance deployments. -Requires `oauth_clients` and `oauth_codes` tables + 6 stored procedures from `database_schema.sql`. +Requires `oauth_clients` and `oauth_codes` tables + 6 stored procedures from `lookup/database_schema.sql`. #### New DB Types @@ -1198,10 +1201,10 @@ auth.OAuthIntrospectToken(ctx, token) // RFC 7662 — returns OAuthTokenInfo auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session ``` -#### SQLNames Fields +#### Procedure names ```go -type SQLNames struct { +type ProcNames struct { // ... existing fields ... OAuthRegisterClient string // default: "resolvespec_oauth_register_client" OAuthGetClient string // default: "resolvespec_oauth_get_client" diff --git a/pkg/security/breaking_changes.md b/pkg/security/breaking_changes.md new file mode 100644 index 0000000..a564837 --- /dev/null +++ b/pkg/security/breaking_changes.md @@ -0,0 +1,126 @@ +# Breaking changes + +Appended as each step of `audit/sec_query_builder.plan.md` lands. + +## Step 0: shared types moved to `sectypes` (no action needed) + +The plain data types now live in `pkg/security/sectypes` and are aliased in `pkg/security` +(`types.go`), so `security.UserContext` and `sectypes.UserContext` are the same type. +Moved: `UserContext`, `LoginRequest/Response`, `RegisterRequest`, `LogoutRequest`, +`PasswordReset*`, `KeyType` (+ constants), `UserKey`, `CreateKey*`, `OAuthServerClient`, +`OAuthCode`, `OAuthTokenInfo`, `Passkey*` data structs, `TwoFactorSecret`, `ColumnSecurity`, +`RowSecurity` (incl. `GetTemplate`). + +## Step 0b (totp): moved to `pkg/security/totp` + +Import `github.com/bitechdev/ResolveSpec/pkg/security/totp`. No aliases (import cycle). + +| Old | New | +|---|---| +| `security.TwoFactorAuthProvider` | `totp.AuthProvider` | +| `security.TwoFactorConfig` / `DefaultTwoFactorConfig` | `totp.Config` / `totp.DefaultConfig` | +| `security.TOTPGenerator` / `NewTOTPGenerator` | `totp.Generator` / `totp.NewGenerator` | +| `security.GenerateBackupCodes` | `totp.GenerateBackupCodes` | +| `security.MemoryTwoFactorProvider` / `NewMemoryTwoFactorProvider` | `totp.MemoryProvider` / `totp.NewMemoryProvider` | +| `security.TwoFactorAuthenticator` / `NewTwoFactorAuthenticator` | `totp.Authenticator` / `totp.NewAuthenticator` | + +`totp.NewAuthenticator` takes a `totp.BaseAuthenticator` (Login, Logout, Authenticate) instead of +`security.Authenticator`; any `security.Authenticator` satisfies it. + +`DatabaseTwoFactorProvider` stays in `security` for now (it uses the core SQL internals) and moves +into `totp` once the lookup `TOTPStore` replaces them. Until then core imports `totp`, so `totp` +must not import `security`. + +## Step 0b (providers, first part): moved to `pkg/security/providers` + +Import `github.com/bitechdev/ResolveSpec/pkg/security/providers`. Names unchanged, no aliases (import cycle). + +| Old | New | +|---|---| +| `security.HeaderAuthenticator` / `NewHeaderAuthenticator` | `providers.HeaderAuthenticator` / `providers.NewHeaderAuthenticator` | +| `security.ConfigKeyStore` / `NewConfigKeyStore` | `providers.ConfigKeyStore` / `providers.NewConfigKeyStore` | +| `security.KeyStoreAuthenticator` / `NewKeyStoreAuthenticator` | `providers.KeyStoreAuthenticator` / `providers.NewKeyStoreAuthenticator` | +| `security.ConfigColumnSecurityProvider` / `NewConfigColumnSecurityProvider` | `providers.ConfigColumnSecurityProvider` / `providers.NewConfigColumnSecurityProvider` | +| `security.ConfigRowSecurityProvider` / `NewConfigRowSecurityProvider` | `providers.ConfigRowSecurityProvider` / `providers.NewConfigRowSecurityProvider` | + +The SHA-256 key hash helper is now `sectypes.HashKey`. The database-backed providers +(`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseKeyStore`, `DatabaseColumn/RowSecurityProvider`) +stay in `security` until the lookup stores replace their SQL. + +## Additions (no action needed) + +- `common.SQLDBProvider` (`SQLDB() *sql.DB`) is implemented by the bun, gorm and pgsql adapters + (not their transaction adapters). `lookup.FromDatabase(common.Database)` uses it, plus the + adapter's `DriverName()`, to get the `*sql.DB` and dialect name. + +## Steps 2–3: lookup dialects and procedure backend + +- New `lookup/dialect` (postgres, sqlite, mysql, mssql) and `lookup/procedure` packages. No + existing exported `security` API changed in these steps. +- Procedure-mode code paths in `DatabaseAuthenticator`, `JWTAuthenticator`, the policy providers, + `DatabaseKeyStore`, `DatabaseTwoFactorProvider` and `DatabasePasskeyProvider` now delegate to + `lookup/procedure`. Error texts are unchanged. +- Behaviour change (improvement): these procedure paths now reconnect once on a closed `*sql.DB` + (JWT logout, key create, TOTP, passkey, OAuth previously used the handle directly). + +## Step 4: lookup/direct backend + +- New `lookup/direct` package: table-backed stores for auth, keys, OAuth (client + user), passkey, + TOTP and policy, built from `lookup.Schema` and the dialect. Nothing in `pkg/security` calls it + yet (wiring is step 5), so no existing API changes here. +- Direct `LoginAPIKey` is new: `header_api` / `api` keys only; unknown, expired, inactive and + wrong-type keys (and inactive users) all return `lookup.ErrInvalidAPIKey`. +- Policy tables (`sec_group_members`, `sec_column_rules`, `sec_row_rules`) are required for the + direct policy store; `PolicyOptions.NoGroups` skips the membership table. +- Direct behaviour that changes when step 5 switches over: login, register, refresh, API-key login, + password reset and passkey login now write in one transaction; `Keys.Create` stores NULL (not the + text `null`) for empty scopes/meta; OAuth code exchange consumes the code atomically; a + non-numeric row-security user reference is an error instead of loading no rules. + +## Step 5: pkg/security uses lookup + +`pkg/security` no longer contains SQL (guarded by `TestCoreContainsNoSQL`). Every database call goes +through a `lookup.Provider` built by `lookup/backends.New`. + +Removed (replaced by `lookup.Config`: `Dialect`, `Mode`, `Overrides`, `Procs`, `Schema`): +- Types and functions `SQLNames`, `DefaultSQLNames`, `MergeSQLNames`, `ValidateSQLNames`, + `TableNames` (+ Default/Merge/Validate), `KeyStoreSQLNames`, `KeyStoreTableNames` (+ same), + `QueryMode`, `ModeAuto`/`ModeProcedure`/`ModeDirect`, `ErrDirectModeUnsupported`. +- Options fields `SQLNames`, `TableNames`, `QueryMode` on `DatabaseAuthenticatorOptions`, + `DatabaseKeyStoreOptions`, `DatabasePasskeyProviderOptions`; replaced by `Lookup lookup.Config` + and `LookupProvider *lookup.Provider`. +- Builders `WithQueryMode`, `WithTableNames` on `JWTAuthenticator`, the column/row providers and + `DatabaseTwoFactorProvider`; replaced by `WithLookup(cfg)` and `WithLookupProvider(p)`. +- The variadic `names ...*SQLNames` argument of `NewJWTAuthenticator`, + `NewDatabaseColumnSecurityProvider`, `NewDatabaseRowSecurityProvider` and + `NewDatabaseTwoFactorProvider`. + +Behaviour changes: +- Default mode is per dialect: stored procedures on Postgres, direct SQL elsewhere. `ModeAuto` + (probe `pg_proc` once per procedure) is now opt-in via `lookup.ModeAuto`; it used to be the + default everywhere. Procedure mode on a non-Postgres dialect is a configuration error. +- A bad lookup configuration no longer panics or is silently ignored: the component logs it and + every call returns the error. +- Dialect is detected from the driver; if detection fails the postgres dialect is assumed. +- Column and row security now work in direct mode (tables `sec_group_members`, `sec_column_rules`, + `sec_row_rules`); they used to return `ErrDirectModeUnsupported`. `WithNoGroupTables()` skips the + group membership table. +- `LoginWithAPIKey` works in direct mode; `DatabaseAuthenticator.Logout` now clears the session + cache in every mode (direct mode used to skip it). +- `Authenticate` now holds the session lookup in the configured backend only; the activity update no + longer silently falls back to a direct write when the procedure is missing. +- Direct-mode behaviour changes listed under step 4 take effect here. + +## Step 6: schemas and docs + +- SQL files moved from `pkg/security/` to `pkg/security/lookup/` (`database_schema.sql`, + `keystore_schema.sql`). `database_schema_sqlite.sql` is superseded by `lookup/ddl/sqlite.sql` + (now also includes `sec_group_members`, `sec_column_rules`, `sec_row_rules`). +- New `lookup/ddl` package: embedded reference table schemas for `postgres`, `sqlite`, `mysql`, + `mssql` (`ddl.SQL(dialect)`, `ddl.Statements(dialect)`). `ddl/postgres.sql` is tables only and uses + base64 / JSON text columns, so it cannot be combined with the procedure schema + (`database_schema.sql`, `bytea` / `text[]` columns) on the same tables. +- `security.ApplyTxSettings` is unchanged; its SQL moved to `lookup.ApplyTxSettings(ctx, tx, settings)`. +- Removed the unexported `password.go` from `pkg/security` (bcrypt helpers live in `lookup/direct`). +- `README.md`, `KEYSTORE.md` and the root README describe `lookup.Config` instead of `QueryMode`, + `SQLNames` and `TableNames`. diff --git a/pkg/security/chain.go b/pkg/security/chain.go index 8a3c158..3be5bbd 100644 --- a/pkg/security/chain.go +++ b/pkg/security/chain.go @@ -55,3 +55,18 @@ func (c *ChainAuthenticator) Logout(ctx context.Context, req LogoutRequest) erro func (c *ChainAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error { return c.authenticators[0].LogoutWithCookie(ctx, req, w) } + +// LoginWithAPIKey tries each authenticator that supports API key login and +// returns the first success. Failures collapse to one generic error. +func (c *ChainAuthenticator) LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) { + for _, a := range c.authenticators { + l, ok := a.(APIKeyLoginable) + if !ok { + continue + } + if resp, err := l.LoginWithAPIKey(ctx, rawKey, claims); err == nil { + return resp, nil + } + } + return nil, errInvalidAPIKey +} diff --git a/pkg/security/composite.go b/pkg/security/composite.go index fa5150e..a4ec4ec 100644 --- a/pkg/security/composite.go +++ b/pkg/security/composite.go @@ -88,6 +88,14 @@ func (c *CompositeSecurityProvider) RefreshToken(ctx context.Context, refreshTok return nil, fmt.Errorf("authenticator does not support token refresh") } +// LoginWithAPIKey implements APIKeyLoginable if the authenticator supports it +func (c *CompositeSecurityProvider) LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) { + if l, ok := c.auth.(APIKeyLoginable); ok { + return l.LoginWithAPIKey(ctx, rawKey, claims) + } + return nil, fmt.Errorf("authenticator does not support API key login") +} + // ValidateToken implements Validatable if the authenticator supports it func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) { if validatable, ok := c.auth.(Validatable); ok { diff --git a/pkg/security/direct_mode_test.go b/pkg/security/direct_mode_test.go index 3cdf196..70dff73 100644 --- a/pkg/security/direct_mode_test.go +++ b/pkg/security/direct_mode_test.go @@ -5,37 +5,41 @@ import ( "database/sql" "encoding/base64" "net/http" - "os" - "path/filepath" + "strings" "testing" "time" - _ "github.com/mattn/go-sqlite3" + _ "github.com/glebarez/go-sqlite" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" ) +// directConfig forces the direct (table) backend for every operation. +var directConfig = lookup.Config{Mode: lookup.ModeDirect} + func futureTime() time.Time { return time.Now().Add(1 * time.Hour) } // newDirectTestDB opens a fresh in-memory SQLite database and applies the -// portable Direct-mode schema (database_schema_sqlite.sql), giving every +// portable Direct-mode schema (lookup/ddl/sqlite.sql), giving every // Direct-mode test a real, isolated database to exercise end-to-end. func newDirectTestDB(t *testing.T) *sql.DB { t.Helper() - db, err := sql.Open("sqlite3", "file::memory:?cache=shared") + db, err := sql.Open("sqlite", ":memory:") if err != nil { t.Fatalf("failed to open sqlite db: %v", err) } db.SetMaxOpenConns(1) // keep the shared in-memory db single-connection so state isn't lost t.Cleanup(func() { _ = db.Close() }) - schemaPath := filepath.Join("database_schema_sqlite.sql") - schema, err := os.ReadFile(schemaPath) + schema, err := ddl.SQL("sqlite") if err != nil { t.Fatalf("failed to read schema: %v", err) } - if _, err := db.Exec(string(schema)); err != nil { + if _, err := db.Exec(schema); err != nil { t.Fatalf("failed to apply schema: %v", err) } return db @@ -49,7 +53,7 @@ func authenticatedRequest(token string) *http.Request { func TestDirectMode_RegisterThenLogin(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() regResp, err := auth.Register(ctx, RegisterRequest{ @@ -105,7 +109,7 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) { func TestDirectMode_SessionLifecycle(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() loginResp, err := auth.Register(ctx, RegisterRequest{Username: "bob", Password: "p", Email: "bob@example.com"}) @@ -140,7 +144,7 @@ func TestDirectMode_SessionLifecycle(t *testing.T) { func TestDirectMode_PasswordReset(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() if _, err := auth.Register(ctx, RegisterRequest{Username: "carol", Password: "old", Email: "carol@example.com"}); err != nil { @@ -167,8 +171,8 @@ func TestDirectMode_PasswordReset(t *testing.T) { func TestDirectMode_JWTLoginAndLogout(t *testing.T) { db := newDirectTestDB(t) - jwtAuth := NewJWTAuthenticator("secret", db).WithQueryMode(ModeDirect) - directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + jwtAuth := NewJWTAuthenticator("secret", db).WithLookup(directConfig) + directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() if _, err := directAuth.Register(ctx, RegisterRequest{Username: "dave", Password: "p", Email: "dave@example.com"}); err != nil { @@ -190,7 +194,7 @@ func TestDirectMode_JWTLoginAndLogout(t *testing.T) { func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() regResp, err := auth.Register(ctx, RegisterRequest{Username: "erin", Password: "p", Email: "erin@example.com"}) @@ -199,7 +203,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) { } userID := regResp.User.UserID - totp := NewDatabaseTwoFactorProvider(db, nil).WithQueryMode(ModeDirect) + totp := NewDatabaseTwoFactorProvider(db, nil).WithLookup(directConfig) if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil { t.Fatalf("Enable2FA() error = %v", err) @@ -248,7 +252,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) { func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() regResp, err := auth.Register(ctx, RegisterRequest{Username: "frank", Password: "p", Email: "frank@example.com"}) @@ -258,7 +262,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) { userID := regResp.User.UserID passkeys := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{ - RPID: "example.com", RPName: "Example", RPOrigin: "https://example.com", QueryMode: ModeDirect, + RPID: "example.com", RPName: "Example", RPOrigin: "https://example.com", Lookup: directConfig, }) cred, err := passkeys.CompleteRegistration(ctx, userID, PasskeyRegistrationResponse{ @@ -317,7 +321,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) { func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() userCtx := &UserContext{UserName: "gina", Email: "gina@example.com", Roles: []string{"user"}} @@ -341,7 +345,7 @@ func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) { func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() regResp, err := auth.Register(ctx, RegisterRequest{Username: "henry", Password: "p", Email: "henry@example.com"}) @@ -349,7 +353,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) { t.Fatalf("Register() error = %v", err) } - ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{QueryMode: ModeDirect}) + ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{Lookup: directConfig}) createResp, err := ks.CreateKey(ctx, CreateKeyRequest{ UserID: regResp.User.UserID, @@ -399,7 +403,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) { func TestDirectMode_OAuthServerClientAndCode(t *testing.T) { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) ctx := context.Background() client := &OAuthServerClient{ @@ -489,7 +493,7 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) { func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) { for _, enabled := range []bool{false, true} { db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig, UpgradePasswordHash: enabled}) ctx := context.Background() if _, err := db.Exec(`DELETE FROM users`); err != nil { @@ -518,18 +522,6 @@ func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) { } } -func TestVerifyPasswordEdgeCases(t *testing.T) { - h, _ := hashPassword("pw") - if ok, _ := verifyPassword(h, "pw"); !ok { - t.Error("bcrypt match failed") - } - if ok, _ := verifyPassword("", "pw"); ok { - t.Error("empty stored must not match") - } - if ok, _ := verifyPassword("pw", ""); ok { - t.Error("empty supplied must not match") - } - if _, err := hashPassword(string(make([]byte, 73))); err == nil { - t.Error("73-byte password must be rejected") - } +func isBcryptHash(s string) bool { + return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$") } diff --git a/pkg/security/interfaces.go b/pkg/security/interfaces.go index 1e2c18a..16442a9 100644 --- a/pkg/security/interfaces.go +++ b/pkg/security/interfaces.go @@ -5,81 +5,6 @@ import ( "net/http" ) -// UserContext holds authenticated user information -type UserContext struct { - UserID int `json:"user_id"` - UserName string `json:"user_name"` - UserLevel int `json:"user_level"` - SessionID string `json:"session_id"` - SessionRID int64 `json:"session_rid"` - RemoteID string `json:"remote_id"` - Roles []string `json:"roles"` - Email string `json:"email"` - Claims map[string]any `json:"claims"` - Meta map[string]any `json:"meta"` // Additional metadata that can hold any JSON-serializable values - TwoFactorEnabled bool `json:"two_factor_enabled"` // Indicates if 2FA is enabled for this user - ProgramUserID int `json:"program_user_id"` - ProgramUserTable string `json:"program_user_table"` -} - -// LoginRequest contains credentials for login -type LoginRequest struct { - Username string `json:"username"` - Password string `json:"password"` - TwoFactorCode string `json:"two_factor_code,omitempty"` // TOTP or backup code - Claims map[string]any `json:"claims"` // Additional login data - Meta map[string]any `json:"meta"` // Additional metadata to be set on user context -} - -// RegisterRequest contains information for new user registration -type RegisterRequest struct { - Username string `json:"username"` - Password string `json:"password"` - Email string `json:"email"` - UserLevel int `json:"user_level"` - Roles []string `json:"roles"` - Claims map[string]any `json:"claims"` // Additional registration data - Meta map[string]any `json:"meta"` // Additional metadata -} - -// LoginResponse contains the result of a login attempt -type LoginResponse struct { - Token string `json:"token"` - RefreshToken string `json:"refresh_token"` - User *UserContext `json:"user"` - ExpiresIn int64 `json:"expires_in"` // Token expiration in seconds - Requires2FA bool `json:"requires_2fa"` // True if 2FA code is required - TwoFactorSetupData *TwoFactorSecret `json:"two_factor_setup,omitempty"` // Present when setting up 2FA - Meta map[string]any `json:"meta"` // Additional metadata to be set on user context -} - -// LogoutRequest contains information for logout -type LogoutRequest struct { - Token string `json:"token"` - UserID int `json:"user_id"` -} - -// PasswordResetRequest initiates a password reset for a user -type PasswordResetRequest struct { - Email string `json:"email,omitempty"` - Username string `json:"username,omitempty"` -} - -// PasswordResetResponse is returned when a reset is initiated -type PasswordResetResponse struct { - // Token is the reset token to be delivered out-of-band (e.g. email). - // The stored procedure may return it for delivery or leave it empty - // if the delivery is handled entirely in the database. - Token string `json:"token"` - ExpiresIn int64 `json:"expires_in"` // seconds -} - -// PasswordResetCompleteRequest completes a password reset using the token -type PasswordResetCompleteRequest struct { - Token string `json:"token"` - NewPassword string `json:"new_password"` -} - // Authenticator handles user authentication operations type Authenticator interface { // Login authenticates credentials and returns a token @@ -144,6 +69,13 @@ type Refreshable interface { RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) } +// APIKeyLoginable allows providers to exchange a raw API key for a session. +type APIKeyLoginable interface { + // LoginWithAPIKey validates the raw API key and creates a session for its user. + // Unknown, expired and inactive keys all yield the same generic error. + LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) +} + // Validatable allows providers to validate tokens without full authentication type Validatable interface { // ValidateToken checks if a token is valid without extracting full user context diff --git a/pkg/security/keystore.go b/pkg/security/keystore.go index 1d442a1..be5fd8a 100644 --- a/pkg/security/keystore.go +++ b/pkg/security/keystore.go @@ -2,64 +2,11 @@ package security import ( "context" - "crypto/sha256" - "encoding/hex" - "time" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" ) -// hashSHA256Hex returns the lowercase hex SHA-256 digest of the given string. -// Used by all keystore implementations to hash raw keys before storage or lookup. -func hashSHA256Hex(raw string) string { - sum := sha256.Sum256([]byte(raw)) - return hex.EncodeToString(sum[:]) -} - -// KeyType identifies the category of an auth key. -type KeyType string - -const ( - // KeyTypeJWTSecret is a per-user JWT signing secret for token generation. - KeyTypeJWTSecret KeyType = "jwt_secret" - // KeyTypeHeaderAPI is a static API key sent via a request header. - KeyTypeHeaderAPI KeyType = "header_api" - // KeyTypeOAuth2 holds OAuth2 client credentials (client_id / client_secret). - KeyTypeOAuth2 KeyType = "oauth2" - // KeyTypeGenericAPI is a generic application API key. - KeyTypeGenericAPI KeyType = "api" -) - -// UserKey represents a single named auth key belonging to a user. -// KeyHash stores the SHA-256 hex digest of the raw key; the raw key is never persisted. -type UserKey struct { - ID int64 `json:"id"` - UserID int `json:"user_id"` - KeyType KeyType `json:"key_type"` - KeyHash string `json:"key_hash"` // SHA-256 hex; never the raw key - Name string `json:"name"` - Scopes []string `json:"scopes,omitempty"` - Meta map[string]any `json:"meta,omitempty"` - ExpiresAt *time.Time `json:"expires_at,omitempty"` - CreatedAt time.Time `json:"created_at"` - LastUsedAt *time.Time `json:"last_used_at,omitempty"` - IsActive bool `json:"is_active"` -} - -// CreateKeyRequest specifies the parameters for a new key. -type CreateKeyRequest struct { - UserID int - KeyType KeyType - Name string - Scopes []string - Meta map[string]any - ExpiresAt *time.Time -} - -// CreateKeyResponse is returned exactly once when a key is created. -// The caller is responsible for persisting RawKey; it is not stored anywhere. -type CreateKeyResponse struct { - Key UserKey - RawKey string // crypto/rand 32 bytes, base64url-encoded -} +// hashSHA256Hex is kept as a short alias for sectypes.HashKey inside this package. +func hashSHA256Hex(raw string) string { return sectypes.HashKey(raw) } // KeyStore manages per-user auth keys with pluggable storage backends. // Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures). diff --git a/pkg/security/keystore_database.go b/pkg/security/keystore_database.go index f0643e1..d9b0848 100644 --- a/pkg/security/keystore_database.go +++ b/pkg/security/keystore_database.go @@ -5,16 +5,16 @@ import ( "crypto/rand" "database/sql" "encoding/base64" - "encoding/json" "errors" "fmt" - "sync" "time" "golang.org/x/sync/singleflight" "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/dbtrace" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends" ) // DatabaseKeyStoreOptions configures DatabaseKeyStore. @@ -24,36 +24,28 @@ type DatabaseKeyStoreOptions struct { // CacheTTL is the duration to cache ValidateKey results. // Default: 2 minutes. CacheTTL time.Duration - // SQLNames provides custom procedure names. If nil, uses DefaultKeyStoreSQLNames(). - SQLNames *KeyStoreSQLNames - // TableNames provides custom table names for Direct mode. If nil, uses DefaultKeyStoreTableNames(). - TableNames *KeyStoreTableNames - // QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto. - QueryMode QueryMode + // Lookup selects dialect, query mode and procedure/table/column names. + // The zero value uses stored procedures on Postgres and direct SQL elsewhere. + Lookup lookup.Config + // LookupProvider, when set, is used instead of building one from Lookup and the db. + LookupProvider *lookup.Provider // 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) } -// DatabaseKeyStore is a KeyStore backed by PostgreSQL stored procedures. -// All DB operations go through configurable procedure names; the raw key is -// never passed to the database. +// DatabaseKeyStore is a KeyStore backed by the lookup package (stored procedures on +// Postgres by default, direct SQL elsewhere). The raw key is never passed to the database. // -// See keystore_schema.sql for the required table and procedure definitions. +// See lookup/keystore_schema.sql for the required table and procedure definitions. // // Note: DeleteKey invalidates the cache entry for the deleted key. Due to the // cache TTL, a deleted key may continue to authenticate for up to CacheTTL // (default 2 minutes) if the cache entry cannot be invalidated. type DatabaseKeyStore struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - sqlNames *KeyStoreSQLNames - tableNames *KeyStoreTableNames - queryMode QueryMode - capability *dbCapability - cache *cache.Cache - cacheTTL time.Duration + src *lookupSource + cache *cache.Cache + cacheTTL time.Duration // validateLoads collapses concurrent key lookups for the same key validateLoads singleflight.Group @@ -72,42 +64,14 @@ func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseK if c == nil { c = cache.GetDefaultCache() } - names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames) - tableNames := resolveKeyStoreTableNames(o.TableNames) - return &DatabaseKeyStore{ - db: db, - dbFactory: o.DBFactory, - sqlNames: names, - tableNames: tableNames, - queryMode: o.QueryMode, - capability: newDBCapability(), - cache: c, - cacheTTL: o.CacheTTL, - } + src := newLookupSource(db) + src.cfg = o.Lookup + src.provider = o.LookupProvider + src.opts = backends.Options{DBFactory: o.DBFactory} + return &DatabaseKeyStore{src: src, cache: c, cacheTTL: o.CacheTTL} } -func (ks *DatabaseKeyStore) getDB() *sql.DB { - ks.dbMu.RLock() - defer ks.dbMu.RUnlock() - return ks.db -} - -func (ks *DatabaseKeyStore) reconnectDB() error { - if ks.dbFactory == nil { - return fmt.Errorf("no db factory configured for reconnect") - } - newDB, err := ks.dbFactory() - if err != nil { - return err - } - ks.dbMu.Lock() - ks.db = newDB - ks.dbMu.Unlock() - if ks.capability != nil { - ks.capability.reset() - } - return nil -} +func (ks *DatabaseKeyStore) keys() lookup.KeyStore { return ks.src.get().Keys } // CreateKey generates a raw key, stores its SHA-256 hash via the create procedure, // and returns the raw key once. @@ -119,110 +83,29 @@ func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest) rawKey := base64.RawURLEncoding.EncodeToString(rawBytes) hash := hashSHA256Hex(rawKey) - if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.CreateKey) { - key, err := ks.createKeyDirect(ctx, req, hash) - if err != nil { - return nil, err - } - return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil - } - - type createRequest struct { - UserID int `json:"user_id"` - KeyType 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: hash, - Name: req.Name, - Scopes: req.Scopes, - Meta: req.Meta, - ExpiresAt: req.ExpiresAt, - }) + key, err := ks.keys().Create(ctx, req, hash) if err != nil { - return nil, fmt.Errorf("failed to marshal create key request: %w", err) + return nil, err } - - var success bool - var errorMsg sql.NullString - var keyJSON sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, ks.sqlNames.CreateKey) - if err = ks.getDB().QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON); err != nil { - return nil, fmt.Errorf("create key procedure failed: %w", err) - } - if !success { - return nil, errors.New(nullStringOr(errorMsg, "create key failed")) - } - - var key UserKey - if err = json.Unmarshal([]byte(keyJSON.String), &key); err != nil { - return nil, fmt.Errorf("failed to parse created key: %w", err) - } - - return &CreateKeyResponse{Key: key, RawKey: rawKey}, nil + return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil } // GetUserKeys returns all active, non-expired keys for the given user. // Pass an empty KeyType to return all types. func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) { - if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.GetUserKeys) { - return ks.getUserKeysDirect(ctx, userID, keyType) - } - - var success bool - var errorMsg sql.NullString - var keysJSON sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, ks.sqlNames.GetUserKeys) - if err := ks.getDB().QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON); err != nil { - return nil, fmt.Errorf("get user keys procedure failed: %w", err) - } - if !success { - return nil, errors.New(nullStringOr(errorMsg, "get user keys failed")) - } - - var keys []UserKey - if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" { - if err := json.Unmarshal([]byte(keysJSON.String), &keys); err != nil { - return nil, fmt.Errorf("failed to parse user keys: %w", err) - } - } - if keys == nil { - keys = []UserKey{} - } - return keys, nil + return ks.keys().List(ctx, userID, keyType) } // DeleteKey soft-deletes a key after verifying ownership and invalidates its cache entry. // The delete procedure returns the key_hash so no separate lookup is needed. // Note: cache invalidation is best-effort; a cached entry may persist for up to CacheTTL. func (ks *DatabaseKeyStore) DeleteKey(ctx context.Context, userID int, keyID int64) error { - if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.DeleteKey) { - return ks.deleteKeyDirect(ctx, userID, keyID) + keyHash, err := ks.keys().Delete(ctx, userID, keyID) + if err != nil { + return err } - - var success bool - var errorMsg sql.NullString - var keyHash sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, ks.sqlNames.DeleteKey) - if err := ks.getDB().QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash); err != nil { - return fmt.Errorf("delete key procedure failed: %w", err) - } - if !success { - return errors.New(nullStringOr(errorMsg, "delete key failed")) - } - - if keyHash.Valid && keyHash.String != "" && ks.cache != nil { - _ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash.String)) + if keyHash != "" && ks.cache != nil { + _ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash)) } return nil } @@ -261,51 +144,16 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT // validateKeyLoad validates against the database and fills the cache. func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) { dbtrace.Raw(ctx, "keystore.validate") - if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) { - key, err := ks.validateKeyDirect(ctx, hash, keyType) - if err != nil { - return nil, err - } - if ks.cache != nil { - _ = ks.cache.Set(ctx, cacheKey, *key, ks.cacheTTL) - } - return key, nil - } - - var success bool - var errorMsg sql.NullString - var keyJSON sql.NullString - - runQuery := func() error { - query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, ks.sqlNames.ValidateKey) - return ks.getDB().QueryRowContext(ctx, query, hash, string(keyType)).Scan(&success, &errorMsg, &keyJSON) - } - if err := runQuery(); err != nil { - if isDBClosed(err) { - if reconnErr := ks.reconnectDB(); reconnErr == nil { - err = runQuery() - } - if err != nil { - return nil, fmt.Errorf("validate key procedure failed: %w", err) - } - } else { - return nil, fmt.Errorf("validate key procedure failed: %w", err) - } - } - if !success { - return nil, errors.New(nullStringOr(errorMsg, "invalid or expired key")) - } - - var key UserKey - if err := json.Unmarshal([]byte(keyJSON.String), &key); err != nil { - return nil, fmt.Errorf("failed to parse validated key: %w", err) + key, err := ks.keys().Validate(ctx, hash, keyType) + if err != nil { + return nil, err } if ks.cache != nil { - _ = ks.cache.Set(ctx, cacheKey, key, ks.cacheTTL) + _ = ks.cache.Set(ctx, cacheKey, *key, ks.cacheTTL) } - return &key, nil + return key, nil } func keystoreCacheKey(hash string) string { diff --git a/pkg/security/keystore_database_direct.go b/pkg/security/keystore_database_direct.go deleted file mode 100644 index af444e7..0000000 --- a/pkg/security/keystore_database_direct.go +++ /dev/null @@ -1,216 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "encoding/json" - "errors" - "fmt" - "time" -) - -// Direct-mode implementations mirroring the resolvespec_keystore_* stored -// procedures in keystore_schema.sql using plain SQL against -// TableNames.UserKeys. meta/scopes are stored as JSON-encoded TEXT instead -// of Postgres JSONB. - -func (ks *DatabaseKeyStore) createKeyDirect(ctx context.Context, req CreateKeyRequest, keyHash string) (*UserKey, error) { - scopesJSON, err := json.Marshal(req.Scopes) - if err != nil { - return nil, fmt.Errorf("failed to marshal scopes: %w", err) - } - var metaJSON []byte - if req.Meta != nil { - metaJSON, err = json.Marshal(req.Meta) - if err != nil { - return nil, fmt.Errorf("failed to marshal meta: %w", err) - } - } - - now := time.Now() - var id int64 - err = ks.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (user_id, key_type, key_hash, name, scopes, meta, expires_at, created_at, is_active) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, - ks.tableNames.UserKeys)) - res, err := db.ExecContext(ctx, query, req.UserID, string(req.KeyType), keyHash, req.Name, string(scopesJSON), nullableString(metaJSON), req.ExpiresAt, now, true) - if err != nil { - return err - } - id, err = res.LastInsertId() - return err - }) - if err != nil { - return nil, fmt.Errorf("create key query failed: %w", err) - } - - return &UserKey{ - ID: id, - UserID: req.UserID, - KeyType: req.KeyType, - KeyHash: keyHash, - Name: req.Name, - Scopes: req.Scopes, - Meta: req.Meta, - ExpiresAt: req.ExpiresAt, - CreatedAt: now, - IsActive: true, - }, nil -} - -func (ks *DatabaseKeyStore) runDBOpWithReconnect(run func(*sql.DB) error) error { - db := ks.getDB() - if db == nil { - return fmt.Errorf("database connection is nil") - } - err := run(db) - if isDBClosed(err) { - if reconnErr := ks.reconnectDB(); reconnErr == nil { - err = run(ks.getDB()) - } - } - return err -} - -func (ks *DatabaseKeyStore) getUserKeysDirect(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) { - keys := []UserKey{} - err := ks.runDBOpWithReconnect(func(db *sql.DB) error { - var query string - var args []any - if keyType == "" { - query = rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, last_used_at, is_active - FROM %s WHERE user_id = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?)`, ks.tableNames.UserKeys)) - args = []any{userID, true, time.Now()} - } else { - query = rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, last_used_at, is_active - FROM %s WHERE user_id = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?) AND key_type = ?`, ks.tableNames.UserKeys)) - args = []any{userID, true, time.Now(), string(keyType)} - } - - rows, err := db.QueryContext(ctx, query, args...) - if err != nil { - return err - } - defer rows.Close() - - for rows.Next() { - var k UserKey - var kt string - var scopesJSON, metaJSON sql.NullString - var expiresAt, lastUsedAt sql.NullTime - if err := rows.Scan(&k.ID, &k.UserID, &kt, &k.Name, &scopesJSON, &metaJSON, &expiresAt, &k.CreatedAt, &lastUsedAt, &k.IsActive); err != nil { - return err - } - k.KeyType = KeyType(kt) - if scopesJSON.Valid && scopesJSON.String != "" { - _ = json.Unmarshal([]byte(scopesJSON.String), &k.Scopes) - } - if metaJSON.Valid && metaJSON.String != "" { - _ = json.Unmarshal([]byte(metaJSON.String), &k.Meta) - } - if expiresAt.Valid { - t := expiresAt.Time - k.ExpiresAt = &t - } - if lastUsedAt.Valid { - t := lastUsedAt.Time - k.LastUsedAt = &t - } - keys = append(keys, k) - } - return rows.Err() - }) - if err != nil { - return nil, fmt.Errorf("get user keys query failed: %w", err) - } - return keys, nil -} - -func (ks *DatabaseKeyStore) deleteKeyDirect(ctx context.Context, userID int, keyID int64) error { - var keyHash string - err := ks.runDBOpWithReconnect(func(db *sql.DB) error { - selQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT key_hash FROM %s WHERE id = ? AND user_id = ? AND is_active = ?`, ks.tableNames.UserKeys)) - if err := db.QueryRowContext(ctx, selQuery, keyID, userID, true).Scan(&keyHash); err != nil { - return err - } - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET is_active = ? WHERE id = ? AND user_id = ? AND is_active = ?`, ks.tableNames.UserKeys)) - _, err := db.ExecContext(ctx, updQuery, false, keyID, userID, true) - return err - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return errors.New("key not found or already deleted") - } - return fmt.Errorf("delete key query failed: %w", err) - } - - if keyHash != "" && ks.cache != nil { - _ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash)) - } - return nil -} - -func (ks *DatabaseKeyStore) validateKeyDirect(ctx context.Context, keyHash string, keyType KeyType) (*UserKey, error) { - var k UserKey - var kt string - var scopesJSON, metaJSON sql.NullString - var expiresAt, lastUsedAt sql.NullTime - - err := ks.runDBOpWithReconnect(func(db *sql.DB) error { - var query string - var args []any - if keyType == "" { - query = rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, is_active - FROM %s WHERE key_hash = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?)`, ks.tableNames.UserKeys)) - args = []any{keyHash, true, time.Now()} - } else { - query = rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, is_active - FROM %s WHERE key_hash = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?) AND key_type = ?`, ks.tableNames.UserKeys)) - args = []any{keyHash, true, time.Now(), string(keyType)} - } - - if err := db.QueryRowContext(ctx, query, args...).Scan(&k.ID, &k.UserID, &kt, &k.Name, &scopesJSON, &metaJSON, &expiresAt, &k.CreatedAt, &k.IsActive); err != nil { - return err - } - - now := time.Now() - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_used_at = ? WHERE id = ?`, ks.tableNames.UserKeys)) - _, err := db.ExecContext(ctx, updQuery, now, k.ID) - return err - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, errors.New("invalid or expired key") - } - return nil, fmt.Errorf("validate key query failed: %w", err) - } - - k.KeyType = KeyType(kt) - k.KeyHash = keyHash - if scopesJSON.Valid && scopesJSON.String != "" { - _ = json.Unmarshal([]byte(scopesJSON.String), &k.Scopes) - } - if metaJSON.Valid && metaJSON.String != "" { - _ = json.Unmarshal([]byte(metaJSON.String), &k.Meta) - } - if expiresAt.Valid { - t := expiresAt.Time - k.ExpiresAt = &t - } - _ = lastUsedAt - now := time.Now() - k.LastUsedAt = &now - - return &k, nil -} - -func nullableString(b []byte) any { - if b == nil { - return nil - } - return string(b) -} diff --git a/pkg/security/keystore_sql_names.go b/pkg/security/keystore_sql_names.go deleted file mode 100644 index ab182c6..0000000 --- a/pkg/security/keystore_sql_names.go +++ /dev/null @@ -1,61 +0,0 @@ -package security - -import "fmt" - -// KeyStoreSQLNames holds the configurable stored procedure names used by DatabaseKeyStore. -// Use DefaultKeyStoreSQLNames() for defaults and MergeKeyStoreSQLNames() for partial overrides. -type KeyStoreSQLNames struct { - GetUserKeys string // default: "resolvespec_keystore_get_user_keys" - CreateKey string // default: "resolvespec_keystore_create_key" - DeleteKey string // default: "resolvespec_keystore_delete_key" - ValidateKey string // default: "resolvespec_keystore_validate_key" -} - -// DefaultKeyStoreSQLNames returns a KeyStoreSQLNames with all default resolvespec_keystore_* values. -func DefaultKeyStoreSQLNames() *KeyStoreSQLNames { - return &KeyStoreSQLNames{ - GetUserKeys: "resolvespec_keystore_get_user_keys", - CreateKey: "resolvespec_keystore_create_key", - DeleteKey: "resolvespec_keystore_delete_key", - ValidateKey: "resolvespec_keystore_validate_key", - } -} - -// MergeKeyStoreSQLNames returns a copy of base with any non-empty fields from override applied. -// If override is nil, a copy of base is returned. -func MergeKeyStoreSQLNames(base, override *KeyStoreSQLNames) *KeyStoreSQLNames { - if override == nil { - copied := *base - return &copied - } - merged := *base - if override.GetUserKeys != "" { - merged.GetUserKeys = override.GetUserKeys - } - if override.CreateKey != "" { - merged.CreateKey = override.CreateKey - } - if override.DeleteKey != "" { - merged.DeleteKey = override.DeleteKey - } - if override.ValidateKey != "" { - merged.ValidateKey = override.ValidateKey - } - return &merged -} - -// ValidateKeyStoreSQLNames checks that all non-empty procedure names are valid SQL identifiers. -func ValidateKeyStoreSQLNames(names *KeyStoreSQLNames) error { - fields := map[string]string{ - "GetUserKeys": names.GetUserKeys, - "CreateKey": names.CreateKey, - "DeleteKey": names.DeleteKey, - "ValidateKey": names.ValidateKey, - } - for field, val := range fields { - if val != "" && !validSQLIdentifier.MatchString(val) { - return fmt.Errorf("KeyStoreSQLNames.%s contains invalid characters: %q", field, val) - } - } - return nil -} diff --git a/pkg/security/keystore_table_names.go b/pkg/security/keystore_table_names.go deleted file mode 100644 index 2ecc761..0000000 --- a/pkg/security/keystore_table_names.go +++ /dev/null @@ -1,44 +0,0 @@ -package security - -import "fmt" - -// KeyStoreTableNames holds the configurable table name used by DatabaseKeyStore -// in Direct mode. Use DefaultKeyStoreTableNames() for defaults and -// MergeKeyStoreTableNames() for partial overrides. -type KeyStoreTableNames struct { - UserKeys string // default: "user_keys" -} - -// DefaultKeyStoreTableNames returns a KeyStoreTableNames with default table names. -func DefaultKeyStoreTableNames() *KeyStoreTableNames { - return &KeyStoreTableNames{ - UserKeys: "user_keys", - } -} - -// MergeKeyStoreTableNames returns a copy of base with any non-empty fields from override applied. -// If override is nil, a copy of base is returned. -func MergeKeyStoreTableNames(base, override *KeyStoreTableNames) *KeyStoreTableNames { - if override == nil { - copied := *base - return &copied - } - merged := *base - if override.UserKeys != "" { - merged.UserKeys = override.UserKeys - } - return &merged -} - -// ValidateKeyStoreTableNames checks that all non-empty table names are valid SQL identifiers. -func ValidateKeyStoreTableNames(names *KeyStoreTableNames) error { - if names.UserKeys != "" && !validSQLIdentifier.MatchString(names.UserKeys) { - return fmt.Errorf("KeyStoreTableNames.UserKeys contains invalid characters: %q", names.UserKeys) - } - return nil -} - -// resolveKeyStoreTableNames merges an optional override with defaults. -func resolveKeyStoreTableNames(override *KeyStoreTableNames) *KeyStoreTableNames { - return MergeKeyStoreTableNames(DefaultKeyStoreTableNames(), override) -} diff --git a/pkg/security/login_api_key_test.go b/pkg/security/login_api_key_test.go new file mode 100644 index 0000000..294eff8 --- /dev/null +++ b/pkg/security/login_api_key_test.go @@ -0,0 +1,39 @@ +package security + +import ( + "context" + "errors" + "testing" +) + +type fakeAPIKeyAuth struct { + Authenticator + key string +} + +func (f *fakeAPIKeyAuth) LoginWithAPIKey(_ context.Context, rawKey string, _ map[string]any) (*LoginResponse, error) { + if rawKey != f.key { + return nil, errInvalidAPIKey + } + return &LoginResponse{Token: "tok"}, nil +} + +func TestChainLoginWithAPIKey(t *testing.T) { + ctx := context.Background() + chain := NewChainAuthenticator(&fakeAPIKeyAuth{key: "a"}, &fakeAPIKeyAuth{key: "b"}) + for _, k := range []string{"a", "b"} { + if resp, err := chain.LoginWithAPIKey(ctx, k, nil); err != nil || resp.Token != "tok" { + t.Errorf("chain LoginWithAPIKey(%q) = %v, %v", k, resp, err) + } + } + if _, err := chain.LoginWithAPIKey(ctx, "bad", nil); !errors.Is(err, errInvalidAPIKey) { + t.Errorf("chain bad key error = %v, want errInvalidAPIKey", err) + } +} + +func TestDatabaseAuthenticatorLoginWithAPIKey_EmptyKey(t *testing.T) { + auth := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{}) + if _, err := auth.LoginWithAPIKey(context.Background(), "", nil); !errors.Is(err, errInvalidAPIKey) { + t.Errorf("empty key error = %v, want errInvalidAPIKey", err) + } +} diff --git a/pkg/security/lookup/backends/backends.go b/pkg/security/lookup/backends/backends.go new file mode 100644 index 0000000..73b766a --- /dev/null +++ b/pkg/security/lookup/backends/backends.go @@ -0,0 +1,162 @@ +// Package backends assembles a lookup.Provider: it builds the procedure and direct stores for +// one database and routes every operation to one of them according to lookup.Config. +// It lives apart from package lookup because both backends import lookup. +package backends + +import ( + "context" + "database/sql" + "fmt" + "sync" + + "github.com/bitechdev/ResolveSpec/pkg/dbtrace" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/direct" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure" +) + +// Options are the settings that are not naming or mode. +type Options struct { + // DBFactory is called to obtain a fresh *sql.DB when the current one has been closed. + // Nil disables reconnecting. + DBFactory func() (*sql.DB, error) + // UpgradePasswordHash rewrites a legacy cleartext password as bcrypt after a successful + // direct-mode login. Off by default. + UpgradePasswordHash bool + // NoGroupTables skips the group membership table when loading direct-mode policy rules. + NoGroupTables bool +} + +// New builds a Provider for db. cfg is merged with the defaults and validated; the dialect +// is cfg.Dialect or detected from the driver. Every operation's mode is resolved up front so +// an impossible combination (procedure mode on a non-Postgres dialect) fails here, not on the +// first request. +func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error) { + if db == nil { + return nil, fmt.Errorf("backends: nil database") + } + res, err := cfg.Resolve() + if err != nil { + return nil, err + } + d, err := cfg.ResolveDialect(db) + if err != nil { + return nil, err + } + for _, op := range lookup.AllOps() { + if _, err := res.EffectiveMode(op, d.Name()); err != nil { + return nil, err + } + } + + c := &chooser{cfg: res, dialect: d.Name(), procs: res.Procs} + run := procedure.NewDB(db, opts.DBFactory, c.resetProbes) + c.db = run + + base, err := direct.NewBase(run, d, res.Schema) + if err != nil { + return nil, err + } + p := procedure.NewPasskey(run, res.Procs) + return &lookup.Provider{ + Auth: &authRouter{c: c, + proc: procedure.NewAuth(run, res.Procs), + direct: direct.NewAuth(base, direct.AuthOptions{UpgradePasswordHash: opts.UpgradePasswordHash})}, + Keys: &keysRouter{c: c, + proc: procedure.NewKeys(run, res.Procs), + direct: direct.NewKeys(base)}, + OAuthClient: &oauthClientRouter{c: c, + proc: procedure.NewOAuthClients(run, res.Procs), + direct: direct.NewOAuthClients(base)}, + OAuthUser: &oauthUserRouter{c: c, + proc: procedure.NewOAuthUsers(run, res.Procs), + direct: direct.NewOAuthUsers(base)}, + Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)}, + TOTP: &totpRouter{c: c, + proc: procedure.NewTOTP(run, res.Procs), + direct: direct.NewTOTP(base)}, + Policy: &policyRouter{c: c, + proc: procedure.NewPolicy(run, res.Procs), + direct: direct.NewPolicy(base, direct.PolicyOptions{NoGroups: opts.NoGroupTables})}, + }, nil +} + +// Failed returns a Provider whose every operation returns err. Constructors that cannot +// return an error use it so a bad configuration fails closed on first use. +func Failed(err error) *lookup.Provider { + c := &chooser{fail: err} + return &lookup.Provider{ + Auth: &authRouter{c: c}, + Keys: &keysRouter{c: c}, + OAuthClient: &oauthClientRouter{c: c}, + OAuthUser: &oauthUserRouter{c: c}, + Passkey: &passkeyRouter{c: c}, + TOTP: &totpRouter{c: c}, + Policy: &policyRouter{c: c}, + } +} + +// chooser decides per operation whether the procedure or the direct store runs. +type chooser struct { + cfg *lookup.Resolved + dialect string + procs lookup.ProcNames + db *procedure.DB + probes sync.Map // proc name -> bool + fail error // set by Failed: every operation returns it +} + +func (c *chooser) resetProbes() { + c.probes.Range(func(k, _ any) bool { c.probes.Delete(k); return true }) +} + +// useProc reports whether op should call the stored procedure proc. In auto mode on Postgres +// the catalog is probed once per procedure (cached until a reconnect). +func (c *chooser) useProc(ctx context.Context, op lookup.Op, proc string) (bool, error) { + if c.fail != nil { + return false, c.fail + } + m, err := c.cfg.EffectiveMode(op, c.dialect) + if err != nil { + return false, err + } + switch m { + case lookup.ModeProcedure: + return true, nil + case lookup.ModeAuto: + if v, ok := c.probes.Load(proc); ok { + return v.(bool), nil + } + exists := probeProcedure(ctx, c.db.Get(), proc) + c.probes.Store(proc, exists) + return exists, nil + } + return false, nil +} + +// probeProcedure asks the Postgres catalog whether a function exists. Any failure counts as +// "does not exist" so the probe can never block an operation. +func probeProcedure(ctx context.Context, db *sql.DB, proc string) (exists bool) { + if db == nil { + return false + } + defer func() { _ = recover() }() + dbtrace.Raw(ctx, "probe.pg_proc") + if err := db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_proc WHERE proname = $1 LIMIT 1)`, proc).Scan(&exists); err != nil { + return false + } + return exists +} + +// pick returns the store that should serve op. +func pick[T any](c *chooser, ctx context.Context, op lookup.Op, proc string, p, d T) (T, error) { + use, err := c.useProc(ctx, op, proc) + if err != nil { + var zero T + return zero, err + } + if use { + return p, nil + } + return d, nil +} diff --git a/pkg/security/lookup/backends/backends_test.go b/pkg/security/lookup/backends/backends_test.go new file mode 100644 index 0000000..0838284 --- /dev/null +++ b/pkg/security/lookup/backends/backends_test.go @@ -0,0 +1,119 @@ +package backends + +import ( + "context" + "database/sql" + "errors" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + _ "github.com/glebarez/go-sqlite" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +func sqliteDB(t *testing.T) *sql.DB { + t.Helper() + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatal(err) + } + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + ddl, err := ddl.SQL("sqlite") + if err != nil { + t.Fatal(err) + } + if _, err := db.Exec(ddl); err != nil { + t.Fatal(err) + } + return db +} + +func TestSQLiteDefaultsToDirect(t *testing.T) { + ctx := context.Background() + p, err := New(sqliteDB(t), lookup.Config{}, Options{}) + if err != nil { + t.Fatal(err) + } + reg, err := p.Auth.Register(ctx, sectypes.RegisterRequest{Username: "a", Email: "a@x.io", Password: "pw"}) + if err != nil { + t.Fatal(err) + } + if _, err := p.Auth.Session(ctx, reg.Token, "authenticate"); err != nil { + t.Fatal(err) + } + if on, err := p.TOTP.Status(ctx, reg.User.UserID); err != nil || on { + t.Fatalf("%v %v", on, err) + } +} + +func TestProcedureModeRejectedOnSQLite(t *testing.T) { + _, err := New(sqliteDB(t), lookup.Config{Overrides: map[lookup.Op]lookup.Mode{lookup.OpLogin: lookup.ModeProcedure}}, Options{}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestCustomSchemaAndUnknownDialect(t *testing.T) { + if _, err := New(sqliteDB(t), lookup.Config{Dialect: "nosuch"}, Options{}); err == nil { + t.Fatal("unknown dialect accepted") + } + bad := lookup.Config{Schema: lookup.Schema{lookup.EntityUsers: {Name: "x; drop"}}} + if _, err := New(sqliteDB(t), bad, Options{}); err == nil { + t.Fatal("unsafe schema accepted") + } + if _, err := New(nil, lookup.Config{}, Options{}); err == nil { + t.Fatal("nil db accepted") + } +} + +func TestPostgresDefaultsToProcedure(t *testing.T) { + db, mock, _ := sqlmock.New() + defer db.Close() + p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres}, Options{}) + if err != nil { + t.Fatal(err) + } + mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1). + WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true)) + if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on { + t.Fatalf("%v %v", on, err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestPostgresAutoProbesOnce(t *testing.T) { + db, mock, _ := sqlmock.New() + defer db.Close() + p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres, Mode: lookup.ModeAuto}, Options{}) + if err != nil { + t.Fatal(err) + } + mock.ExpectQuery("pg_proc").WithArgs("resolvespec_totp_get_status"). + WillReturnRows(sqlmock.NewRows([]string{"e"}).AddRow(true)) + for i := 0; i < 2; i++ { // second call must reuse the cached probe + mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1). + WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true)) + if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on { + t.Fatalf("%v %v", on, err) + } + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestFailed(t *testing.T) { + p := Failed(errors.New("boom")) + if _, err := p.Auth.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "boom" { + t.Fatalf("got %v", err) + } + if _, err := p.Policy.RowSecurity(context.Background(), 1, "s", "t"); err == nil { + t.Fatal("want error") + } +} diff --git a/pkg/security/lookup/backends/conformance_test.go b/pkg/security/lookup/backends/conformance_test.go new file mode 100644 index 0000000..70c27e5 --- /dev/null +++ b/pkg/security/lookup/backends/conformance_test.go @@ -0,0 +1,126 @@ +package backends + +import ( + "crypto/rand" + "database/sql" + "encoding/hex" + "fmt" + "os" + "slices" + "testing" + + _ "github.com/jackc/pgx/v5/stdlib" + _ "github.com/microsoft/go-mssqldb" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/conformance" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" +) + +// Every backend/dialect runs the same suite. SQLite runs always. The others run only when a +// DSN is set, and only against a database you are happy to add rows to (all names the suite +// creates carry a unique "cf_" prefix and are removed afterwards): +// +// RESOLVESPEC_TEST_PG_DSN Postgres with the procedure schema installed +// (lookup/database_schema.sql + keystore_schema.sql): procedure mode +// RESOLVESPEC_TEST_PG_DIRECT_DSN Postgres for direct mode; ddl/postgres.sql is applied if the +// tables are missing. Do not point it at the procedure schema: +// the column types differ. +// RESOLVESPEC_TEST_MYSQL_DSN MySQL (needs a "mysql" database/sql driver linked into the test binary) +// RESOLVESPEC_TEST_MSSQL_DSN SQL Server (driver "sqlserver") +func TestConformance(t *testing.T) { + t.Run("sqlite/direct", func(t *testing.T) { + runConformance(t, sqliteDB(t), "sqlite", lookup.Config{Mode: lookup.ModeDirect}, false) + }) + t.Run("sqlite/default", func(t *testing.T) { + runConformance(t, sqliteDB(t), "sqlite", lookup.Config{}, false) + }) + + real := []struct { + name, env, driver, dialect string + cfg lookup.Config + applyDDL bool + }{ + {"postgres/procedure", "RESOLVESPEC_TEST_PG_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeProcedure}, false}, + {"postgres/direct", "RESOLVESPEC_TEST_PG_DIRECT_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeDirect}, true}, + {"mysql/direct", "RESOLVESPEC_TEST_MYSQL_DSN", "mysql", "mysql", lookup.Config{}, true}, + {"mssql/direct", "RESOLVESPEC_TEST_MSSQL_DSN", "sqlserver", "mssql", lookup.Config{}, true}, + } + for _, r := range real { + t.Run(r.name, func(t *testing.T) { + dsn := os.Getenv(r.env) + if dsn == "" { + t.Skipf("%s not set", r.env) + } + runOnServer(t, r.driver, dsn, r.dialect, r.cfg, r.applyDDL) + }) + } +} + +// runOnServer opens dsn, optionally applies the reference DDL, and runs the suite. +func runOnServer(t *testing.T, driver, dsn, dialectName string, cfg lookup.Config, applyDDL bool) { + t.Helper() + if !slices.Contains(sql.Drivers(), driver) { + t.Skipf("database/sql driver %q is not linked into this test binary", driver) + } + db, err := sql.Open(driver, dsn) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + if err := db.Ping(); err != nil { + t.Fatalf("ping: %v", err) + } + if applyDDL { + stmts, err := ddl.Statements(dialectName) + if err != nil { + t.Fatal(err) + } + for _, s := range stmts { + if _, err := db.Exec(s); err != nil { + t.Fatalf("apply ddl: %v\n%s", err, s) + } + } + } + runConformance(t, db, dialectName, cfg, true) +} + +func runConformance(t *testing.T, db *sql.DB, dialectName string, cfg lookup.Config, shared bool) { + t.Helper() + d, err := dialect.Get(dialectName) + if err != nil { + t.Fatal(err) + } + cfg.Dialect = dialectName + p, err := New(db, cfg, Options{}) + if err != nil { + t.Fatal(err) + } + var b [3]byte + _, _ = rand.Read(b[:]) + env := conformance.Env{Provider: p, DB: db, Dialect: d, Prefix: "cf" + hex.EncodeToString(b[:]) + "_"} + if shared { + env.Cleanup = func(t *testing.T) { cleanup(t, db, d, env.Prefix) } + } + conformance.Run(t, env) +} + +// cleanup deletes the rows a conformance run created, by prefix. Child rows go with their user +// through the foreign keys; tables without one are cleaned explicitly. +func cleanup(t *testing.T, db *sql.DB, d dialect.Dialect, prefix string) { + t.Helper() + like := prefix + "%" + for _, q := range []struct{ table, col string }{ + {"oauth_codes", "code"}, + {"oauth_clients", "client_id"}, + {"token_blacklist", "token"}, + {"sec_column_rules", "schema_name"}, + {"sec_row_rules", "schema_name"}, + {"users", "username"}, + } { + if _, err := db.Exec(fmt.Sprintf("DELETE FROM %s WHERE %s LIKE %s", q.table, q.col, d.Placeholder(1)), like); err != nil { + t.Logf("cleanup %s: %v", q.table, err) + } + } +} diff --git a/pkg/security/lookup/backends/container_test.go b/pkg/security/lookup/backends/container_test.go new file mode 100644 index 0000000..a391caa --- /dev/null +++ b/pkg/security/lookup/backends/container_test.go @@ -0,0 +1,141 @@ +package backends + +import ( + "bytes" + "context" + "database/sql" + "fmt" + "net" + "os" + "os/exec" + "strings" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" +) + +// Container tests start a throwaway database server with podman or docker (whichever is +// installed, podman first) and run the conformance suite against it. They pull an image, so +// they only run when RESOLVESPEC_TEST_CONTAINERS=1 and not with -short. The container is +// removed when the test ends. + +const containerPassword = "Resolve_Spec_1" + +func containerRuntime(t *testing.T) string { + t.Helper() + if testing.Short() { + t.Skip("container tests are skipped with -short") + } + if os.Getenv("RESOLVESPEC_TEST_CONTAINERS") != "1" { + t.Skip("set RESOLVESPEC_TEST_CONTAINERS=1 to run tests that start a podman/docker container") + } + for _, rt := range []string{"podman", "docker"} { + if p, err := exec.LookPath(rt); err == nil { + return p + } + } + t.Skip("neither podman nor docker found in PATH") + return "" +} + +func run(t *testing.T, timeout time.Duration, name string, args ...string) string { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + var out, errb bytes.Buffer + cmd := exec.CommandContext(ctx, name, args...) + cmd.Stdout, cmd.Stderr = &out, &errb + if err := cmd.Run(); err != nil { + t.Fatalf("%s %s: %v\n%s", name, strings.Join(args, " "), err, errb.String()) + } + return strings.TrimSpace(out.String()) +} + +// startContainer runs image publishing containerPort on a random localhost port and returns +// the host port. The container is force-removed on cleanup. +func startContainer(t *testing.T, rt, image, containerPort string, env map[string]string) string { + t.Helper() + args := []string{"run", "-d", "--rm", "-p", "127.0.0.1::" + containerPort} + for k, v := range env { + args = append(args, "-e", k+"="+v) + } + args = append(args, image) + id := run(t, 10*time.Minute, rt, args...) // first run may pull the image + t.Cleanup(func() { _ = exec.Command(rt, "rm", "-f", id).Run() }) + + // "127.0.0.1:49153" (docker may print one line per address family) + out := run(t, 30*time.Second, rt, "port", id, containerPort) + line := strings.Fields(out)[len(strings.Fields(out))-1] + for _, l := range strings.Split(out, "\n") { + if strings.HasPrefix(strings.TrimSpace(l), "127.0.0.1") || strings.Contains(l, " 127.0.0.1:") { + line = l[strings.LastIndex(l, " ")+1:] + break + } + } + _, port, err := net.SplitHostPort(line) + if err != nil { + t.Fatalf("cannot parse published port %q: %v", out, err) + } + return port +} + +// waitReady retries until the server accepts queries or the deadline passes. +func waitReady(t *testing.T, driver, dsn string, d time.Duration) *sql.DB { + t.Helper() + deadline := time.Now().Add(d) + var last error + for time.Now().Before(deadline) { + db, err := sql.Open(driver, dsn) + if err == nil { + if last = db.Ping(); last == nil { + t.Cleanup(func() { _ = db.Close() }) + return db + } + _ = db.Close() + } else { + last = err + } + time.Sleep(time.Second) + } + t.Fatalf("database did not become ready within %s: %v", d, last) + return nil +} + +func TestConformancePostgresContainer(t *testing.T) { + rt := containerRuntime(t) + port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword}) + dsn := func(db string) string { + return fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/%s?sslmode=disable", containerPassword, port, db) + } + admin := waitReady(t, "pgx", dsn("postgres"), 90*time.Second) + // The official image restarts once during init: make sure the second start is the one we use. + time.Sleep(2 * time.Second) + admin = waitReady(t, "pgx", dsn("postgres"), 60*time.Second) + for _, name := range []string{"cf_proc", "cf_direct"} { + if _, err := admin.Exec("CREATE DATABASE " + name); err != nil { + t.Fatal(err) + } + } + + t.Run("procedure", func(t *testing.T) { + db, err := sql.Open("pgx", dsn("cf_proc")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + for _, f := range []string{"../database_schema.sql", "../keystore_schema.sql"} { + b, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + if _, err := db.Exec(string(b)); err != nil { + t.Fatalf("apply %s: %v", f, err) + } + } + runConformance(t, db, "postgres", lookup.Config{Mode: lookup.ModeProcedure}, true) + }) + t.Run("direct", func(t *testing.T) { + runOnServer(t, "pgx", dsn("cf_direct"), "postgres", lookup.Config{Mode: lookup.ModeDirect}, true) + }) +} diff --git a/pkg/security/lookup/backends/routers.go b/pkg/security/lookup/backends/routers.go new file mode 100644 index 0000000..bdfd480 --- /dev/null +++ b/pkg/security/lookup/backends/routers.go @@ -0,0 +1,396 @@ +package backends + +import ( + "context" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Routers: each store method asks the chooser which backend serves that operation. + +type authRouter struct { + c *chooser + proc, direct lookup.AuthStore +} + +var _ lookup.AuthStore = (*authRouter)(nil) + +func (r *authRouter) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogin, r.c.procs.Login, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Login(ctx, req) +} + +func (r *authRouter) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRegister, r.c.procs.Register, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Register(ctx, req) +} + +func (r *authRouter) Logout(ctx context.Context, req sectypes.LogoutRequest) error { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogout, r.c.procs.Logout, r.proc, r.direct) + if err != nil { + return err + } + return st.Logout(ctx, req) +} + +func (r *authRouter) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpSession, r.c.procs.Session, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Session(ctx, token, reference) +} + +func (r *authRouter) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpTouchSession, r.c.procs.SessionUpdate, r.proc, r.direct) + if err != nil { + return err + } + return st.TouchSession(ctx, token, user) +} + +func (r *authRouter) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRefresh, r.c.procs.RefreshToken, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Refresh(ctx, refreshToken) +} + +func (r *authRouter) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLoginAPIKey, r.c.procs.LoginAPIKey, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.LoginAPIKey(ctx, rawKey, claims) +} + +func (r *authRouter) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogin, r.c.procs.JWTLogin, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.JWTLogin(ctx, req) +} + +func (r *authRouter) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogout, r.c.procs.JWTLogout, r.proc, r.direct) + if err != nil { + return err + } + return st.JWTLogout(ctx, req) +} + +func (r *authRouter) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetRequest, r.c.procs.PasswordResetRequest, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.ResetRequest(ctx, req) +} + +func (r *authRouter) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error { + st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetComplete, r.c.procs.PasswordResetComplete, r.proc, r.direct) + if err != nil { + return err + } + return st.ResetComplete(ctx, req) +} + +type keysRouter struct { + c *chooser + proc, direct lookup.KeyStore +} + +var _ lookup.KeyStore = (*keysRouter)(nil) + +func (r *keysRouter) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) { + st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyCreate, r.c.procs.KeystoreCreateKey, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Create(ctx, req, keyHash) +} + +func (r *keysRouter) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) { + st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyList, r.c.procs.KeystoreGetUserKeys, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.List(ctx, userID, keyType) +} + +func (r *keysRouter) Delete(ctx context.Context, userID int, keyID int64) (string, error) { + st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyDelete, r.c.procs.KeystoreDeleteKey, r.proc, r.direct) + if err != nil { + return "", err + } + return st.Delete(ctx, userID, keyID) +} + +func (r *keysRouter) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) { + st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyValidate, r.c.procs.KeystoreValidateKey, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Validate(ctx, keyHash, keyType) +} + +type oauthClientRouter struct { + c *chooser + proc, direct lookup.OAuthClientStore +} + +var _ lookup.OAuthClientStore = (*oauthClientRouter)(nil) + +func (r *oauthClientRouter) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRegisterClient, r.c.procs.OAuthRegisterClient, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.RegisterClient(ctx, client) +} + +func (r *oauthClientRouter) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthGetClient, r.c.procs.OAuthGetClient, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.GetClient(ctx, clientID) +} + +func (r *oauthClientRouter) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthSaveCode, r.c.procs.OAuthSaveCode, r.proc, r.direct) + if err != nil { + return err + } + return st.SaveCode(ctx, code) +} + +func (r *oauthClientRouter) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthExchangeCode, r.c.procs.OAuthExchangeCode, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.ExchangeCode(ctx, code) +} + +func (r *oauthClientRouter) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthIntrospect, r.c.procs.OAuthIntrospect, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Introspect(ctx, token) +} + +func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error { + st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRevoke, r.c.procs.OAuthRevoke, r.proc, r.direct) + if err != nil { + return err + } + return st.Revoke(ctx, token) +} + +type oauthUserRouter struct { + c *chooser + proc, direct lookup.OAuthUserStore +} + +var _ lookup.OAuthUserStore = (*oauthUserRouter)(nil) + +func (r *oauthUserRouter) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) { + st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetOrCreateUser, r.c.procs.OAuthGetOrCreateUser, r.proc, r.direct) + if err != nil { + return 0, err + } + return st.GetOrCreateUser(ctx, user, provider) +} + +func (r *oauthUserRouter) CreateSession(ctx context.Context, session lookup.OAuthSession) error { + st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthCreateSession, r.c.procs.OAuthCreateSession, r.proc, r.direct) + if err != nil { + return err + } + return st.CreateSession(ctx, session) +} + +func (r *oauthUserRouter) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) { + st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetRefreshToken, r.c.procs.OAuthGetRefreshToken, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.GetByRefreshToken(ctx, refreshToken) +} + +func (r *oauthUserRouter) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error { + st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthUpdateRefreshToken, r.c.procs.OAuthUpdateRefreshToken, r.proc, r.direct) + if err != nil { + return err + } + return st.UpdateRefreshToken(ctx, userID, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken, expiresAt) +} + +func (r *oauthUserRouter) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) { + st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetUser, r.c.procs.OAuthGetUser, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.GetUser(ctx, userID) +} + +type passkeyRouter struct { + c *chooser + proc, direct lookup.PasskeyStore +} + +var _ lookup.PasskeyStore = (*passkeyRouter)(nil) + +func (r *passkeyRouter) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyStore, r.c.procs.PasskeyStoreCredential, r.proc, r.direct) + if err != nil { + return 0, err + } + return st.Store(ctx, rec) +} + +func (r *passkeyRouter) Get(ctx context.Context, credentialID string) (int, uint32, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyGet, r.c.procs.PasskeyGetCredential, r.proc, r.direct) + if err != nil { + return 0, 0, err + } + return st.Get(ctx, credentialID) +} + +func (r *passkeyRouter) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyUpdateCounter, r.c.procs.PasskeyUpdateCounter, r.proc, r.direct) + if err != nil { + return false, err + } + return st.UpdateCounter(ctx, credentialID, newCounter) +} + +func (r *passkeyRouter) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyList, r.c.procs.PasskeyGetUserCredentials, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.List(ctx, userID) +} + +func (r *passkeyRouter) Delete(ctx context.Context, userID int, credentialID string) error { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyDelete, r.c.procs.PasskeyDeleteCredential, r.proc, r.direct) + if err != nil { + return err + } + return st.Delete(ctx, userID, credentialID) +} + +func (r *passkeyRouter) Rename(ctx context.Context, userID int, credentialID, name string) error { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyRename, r.c.procs.PasskeyUpdateName, r.proc, r.direct) + if err != nil { + return err + } + return st.Rename(ctx, userID, credentialID, name) +} + +func (r *passkeyRouter) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyByUsername, r.c.procs.PasskeyGetCredsByUsername, r.proc, r.direct) + if err != nil { + return 0, nil, err + } + return st.ByUsername(ctx, username) +} + +func (r *passkeyRouter) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) { + st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyLogin, r.c.procs.PasskeyLogin, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.Login(ctx, userID, claims) +} + +type totpRouter struct { + c *chooser + proc, direct lookup.TOTPStore +} + +var _ lookup.TOTPStore = (*totpRouter)(nil) + +func (r *totpRouter) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPEnable, r.c.procs.TOTPEnable, r.proc, r.direct) + if err != nil { + return err + } + return st.Enable(ctx, userID, secret, hashedCodes) +} + +func (r *totpRouter) Disable(ctx context.Context, userID int) error { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPDisable, r.c.procs.TOTPDisable, r.proc, r.direct) + if err != nil { + return err + } + return st.Disable(ctx, userID) +} + +func (r *totpRouter) Status(ctx context.Context, userID int) (bool, error) { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPStatus, r.c.procs.TOTPGetStatus, r.proc, r.direct) + if err != nil { + return false, err + } + return st.Status(ctx, userID) +} + +func (r *totpRouter) Secret(ctx context.Context, userID int) (string, error) { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPSecret, r.c.procs.TOTPGetSecret, r.proc, r.direct) + if err != nil { + return "", err + } + return st.Secret(ctx, userID) +} + +func (r *totpRouter) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPRegenerateBackup, r.c.procs.TOTPRegenerateBackup, r.proc, r.direct) + if err != nil { + return err + } + return st.RegenerateBackupCodes(ctx, userID, hashedCodes) +} + +func (r *totpRouter) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) { + st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPValidateBackupCode, r.c.procs.TOTPValidateBackupCode, r.proc, r.direct) + if err != nil { + return false, err + } + return st.ValidateBackupCode(ctx, userID, codeHash) +} + +type policyRouter struct { + c *chooser + proc, direct lookup.PolicyStore +} + +var _ lookup.PolicyStore = (*policyRouter)(nil) + +func (r *policyRouter) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) { + st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpColumnSecurity, r.c.procs.ColumnSecurity, r.proc, r.direct) + if err != nil { + return nil, err + } + return st.ColumnSecurity(ctx, userID, schema, table) +} + +func (r *policyRouter) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) { + st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpRowSecurity, r.c.procs.RowSecurity, r.proc, r.direct) + if err != nil { + return sectypes.RowSecurity{}, err + } + return st.RowSecurity(ctx, userRef, schema, table) +} diff --git a/pkg/security/lookup/conformance/conformance.go b/pkg/security/lookup/conformance/conformance.go new file mode 100644 index 0000000..1e11815 --- /dev/null +++ b/pkg/security/lookup/conformance/conformance.go @@ -0,0 +1,564 @@ +// Package conformance is the shared behavioural suite every lookup backend must pass. +// It only uses the store interfaces, so the same cases run against the direct backend on +// every dialect and against the procedure backend on Postgres. Error messages are not +// asserted (backends word them differently), only whether an operation succeeds or fails +// and the values it returns. +// +// The suite names everything it creates with Env.Prefix and never assumes empty tables, so +// it can run against a shared database. Env.Cleanup, when set, removes the prefixed rows. +package conformance + +import ( + "context" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "strings" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Env is one backend under test. +type Env struct { + Provider *lookup.Provider + // DB and Dialect are used only to seed policy rules, which have no store method. + DB *sql.DB + Dialect dialect.Dialect + // Prefix makes every created name unique to this run. + Prefix string + // Cleanup removes rows whose names start with Prefix. Optional. + Cleanup func(t *testing.T) +} + +// Run executes the suite. +func Run(t *testing.T, env Env) { + if env.Cleanup != nil { + t.Cleanup(func() { env.Cleanup(t) }) + } + s := &suite{Env: env} + t.Run("AuthSessionLifecycle", s.authSessionLifecycle) + t.Run("AuthRejectsBadCredentials", s.authRejectsBadCredentials) + t.Run("RegisterIgnoresPrivileges", s.registerIgnoresPrivileges) + t.Run("RegisterRejectsDuplicates", s.registerRejectsDuplicates) + t.Run("PasswordReset", s.passwordReset) + t.Run("JWT", s.jwt) + t.Run("Keys", s.keys) + t.Run("LoginAPIKey", s.loginAPIKey) + t.Run("OAuthClientAndCodes", s.oauthClientAndCodes) + t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke) + t.Run("OAuthUsers", s.oauthUsers) + t.Run("Passkey", s.passkey) + t.Run("TOTP", s.totp) + t.Run("Policy", s.policy) +} + +type suite struct{ Env } + +var ctx = context.Background() + +func (s *suite) name(n string) string { return s.Prefix + n } + +func (s *suite) register(t *testing.T, n string) *sectypes.LoginResponse { + t.Helper() + resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{ + Username: s.name(n), Email: s.name(n) + "@example.test", Password: "pw-" + n, + }) + if err != nil { + t.Fatalf("register %s: %v", n, err) + } + if resp == nil || resp.User == nil || resp.Token == "" || resp.User.UserID == 0 { + t.Fatalf("register %s: incomplete response %+v", n, resp) + } + return resp +} + +func rejected(t *testing.T, what string, err error) { + t.Helper() + if err == nil { + t.Fatalf("%s: expected an error", what) + } +} + +// notOK asserts an operation did not validate: it either failed or returned false. +func notOK(t *testing.T, what string, ok bool, err error) { + t.Helper() + if err == nil && ok { + t.Fatalf("%s: accepted", what) + } +} + +func (s *suite) authSessionLifecycle(t *testing.T) { + a := s.Provider.Auth + reg := s.register(t, "life") + + login, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("life"), Password: "pw-life", + Claims: map[string]any{"ip_address": "10.0.0.1", "user_agent": "conformance"}}) + if err != nil || login.Token == "" || login.User.UserName != s.name("life") { + t.Fatalf("login: %+v %v", login, err) + } + if login.Token == reg.Token { + t.Fatal("login reused the registration session") + } + + u, err := a.Session(ctx, login.Token, "authenticate") + if err != nil || u.UserName != s.name("life") || u.UserID != reg.User.UserID { + t.Fatalf("session: %+v %v", u, err) + } + if err := a.TouchSession(ctx, login.Token, u); err != nil { + t.Fatalf("touch: %v", err) + } + _, err = a.Session(ctx, s.name("no-such-token"), "authenticate") + rejected(t, "unknown session", err) + + ref, err := a.Refresh(ctx, login.Token) + if err != nil || ref.Token == "" || ref.Token == login.Token { + t.Fatalf("refresh: %+v %v", ref, err) + } + _, err = a.Session(ctx, login.Token, "") + rejected(t, "session after refresh", err) + _, err = a.Refresh(ctx, login.Token) + rejected(t, "second refresh of the same token", err) + + if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err != nil { + t.Fatalf("logout: %v", err) + } + _, err = a.Session(ctx, ref.Token, "") + rejected(t, "session after logout", err) +} + +func (s *suite) authRejectsBadCredentials(t *testing.T) { + a := s.Provider.Auth + s.register(t, "creds") + _, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds"), Password: "wrong"}) + rejected(t, "wrong password", err) + _, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds")}) + rejected(t, "empty password", err) + _, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("nobody"), Password: "pw"}) + rejected(t, "unknown user", err) +} + +func (s *suite) registerIgnoresPrivileges(t *testing.T) { + resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{ + Username: s.name("priv"), Email: s.name("priv") + "@example.test", Password: "x", + UserLevel: 99, Roles: []string{"admin"}, + }) + if err != nil { + t.Fatal(err) + } + if resp.User.UserLevel != 0 || len(resp.User.Roles) != 0 { + t.Fatalf("client-supplied privileges honoured: %+v", resp.User) + } +} + +func (s *suite) registerRejectsDuplicates(t *testing.T) { + s.register(t, "dup") + _, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup"), Email: s.name("dup2") + "@example.test", Password: "x"}) + rejected(t, "duplicate username", err) + _, err = s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup2"), Email: s.name("dup") + "@example.test", Password: "x"}) + rejected(t, "duplicate email", err) +} + +func (s *suite) passwordReset(t *testing.T) { + a := s.Provider.Auth + reg := s.register(t, "reset") + + if r, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("nobody") + "@example.test"}); err != nil || (r != nil && r.Token != "") { + t.Fatalf("unknown email must succeed without a token (user enumeration): %+v %v", r, err) + } + req, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("reset") + "@example.test"}) + if err != nil || req == nil || req.Token == "" { + t.Fatalf("reset request: %+v %v", req, err) + } + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: "bogus", NewPassword: "x"}); err == nil { + t.Fatal("bogus reset token accepted") + } + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "new-pw"}); err != nil { + t.Fatalf("reset complete: %v", err) + } + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "again"}); err == nil { + t.Fatal("reset token reused") + } + _, err = a.Session(ctx, reg.Token, "") + rejected(t, "session surviving a password reset", err) + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "new-pw"}); err != nil { + t.Fatalf("login with new password: %v", err) + } + _, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "pw-reset"}) + rejected(t, "old password after reset", err) +} + +func (s *suite) jwt(t *testing.T) { + a := s.Provider.Auth + reg := s.register(t, "jwt") + resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "pw-jwt"}) + if err != nil || resp.Token == "" || resp.User.UserID != reg.User.UserID { + t.Fatalf("jwt login: %+v %v", resp, err) + } + _, err = a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "bad"}) + rejected(t, "jwt login with wrong password", err) + if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: s.name("jwt-tok"), UserID: reg.User.UserID}); err != nil { + t.Fatalf("jwt logout: %v", err) + } +} + +func (s *suite) createKey(t *testing.T, uid int, typ sectypes.KeyType, raw string, exp *time.Time) *sectypes.UserKey { + t.Helper() + k, err := s.Provider.Keys.Create(ctx, sectypes.CreateKeyRequest{UserID: uid, KeyType: typ, Name: s.name("key"), + Scopes: []string{"read"}, ExpiresAt: exp}, sectypes.HashKey(raw)) + if err != nil || k == nil || k.ID == 0 { + t.Fatalf("create key: %+v %v", k, err) + } + return k +} + +func (s *suite) keys(t *testing.T) { + k := s.Provider.Keys + uid := s.register(t, "keys").User.UserID + raw := s.name("raw-keys") + created := s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, raw, nil) + s.createKey(t, uid, sectypes.KeyTypeJWTSecret, s.name("raw-keys-jwt"), nil) + s.createKey(t, uid, sectypes.KeyTypeGenericAPI, s.name("raw-keys-old"), ptr(time.Now().Add(-time.Hour))) + + all, err := k.List(ctx, uid, "") + if err != nil || len(all) != 2 { + t.Fatalf("list must hide expired keys: %d %v", len(all), err) + } + one, err := k.List(ctx, uid, sectypes.KeyTypeHeaderAPI) + if err != nil || len(one) != 1 || one[0].ID != created.ID || len(one[0].Scopes) != 1 { + t.Fatalf("typed list: %+v %v", one, err) + } + + got, err := k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeHeaderAPI) + if err != nil || got.UserID != uid { + t.Fatalf("validate: %+v %v", got, err) + } + _, err = k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeGenericAPI) + rejected(t, "wrong key type", err) + _, err = k.Validate(ctx, sectypes.HashKey(s.name("raw-keys-old")), "") + rejected(t, "expired key", err) + _, err = k.Validate(ctx, sectypes.HashKey(s.name("unknown")), "") + rejected(t, "unknown key", err) + + _, err = k.Delete(ctx, uid+1_000_000, created.ID) + rejected(t, "deleting another user's key", err) + if _, err := k.Delete(ctx, uid, created.ID); err != nil { + t.Fatalf("delete: %v", err) + } + _, err = k.Delete(ctx, uid, created.ID) + rejected(t, "deleting twice", err) + _, err = k.Validate(ctx, sectypes.HashKey(raw), "") + rejected(t, "deleted key", err) +} + +func (s *suite) loginAPIKey(t *testing.T) { + a := s.Provider.Auth + uid := s.register(t, "apikey").User.UserID + good, generic, jwtKey, off, old := s.name("ak-good"), s.name("ak-generic"), s.name("ak-jwt"), s.name("ak-off"), s.name("ak-old") + s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, good, nil) + s.createKey(t, uid, sectypes.KeyTypeGenericAPI, generic, nil) + s.createKey(t, uid, sectypes.KeyTypeJWTSecret, jwtKey, nil) + inactive := s.createKey(t, uid, sectypes.KeyTypeGenericAPI, off, nil) + if _, err := s.Provider.Keys.Delete(ctx, uid, inactive.ID); err != nil { + t.Fatal(err) + } + s.createKey(t, uid, sectypes.KeyTypeGenericAPI, old, ptr(time.Now().Add(-time.Hour))) + + for _, raw := range []string{good, generic} { + resp, err := a.LoginAPIKey(ctx, raw, map[string]any{"ip_address": "10.0.0.2"}) + if err != nil || resp.User.UserName != s.name("apikey") || resp.Token == "" { + t.Fatalf("api key login: %+v %v", resp, err) + } + if _, err := a.Session(ctx, resp.Token, ""); err != nil { + t.Fatalf("session from api key login: %v", err) + } + } + for _, raw := range []string{"", s.name("ak-missing"), jwtKey, off, old} { + _, err := a.LoginAPIKey(ctx, raw, nil) + if !errors.Is(err, lookup.ErrInvalidAPIKey) { + t.Fatalf("key %q: want ErrInvalidAPIKey, got %v", raw, err) + } + } +} + +func (s *suite) oauthClientAndCodes(t *testing.T) { + c := s.Provider.OAuthClient + cid := s.name("client") + reg, err := c.RegisterClient(ctx, §ypes.OAuthServerClient{ClientID: cid, RedirectURIs: []string{"https://app.example.test/cb"}, ClientName: "App"}) + if err != nil || reg.ClientID != cid { + t.Fatalf("register client: %+v %v", reg, err) + } + got, err := c.GetClient(ctx, cid) + if err != nil || got.ClientName != "App" || len(got.RedirectURIs) != 1 || got.RedirectURIs[0] != "https://app.example.test/cb" { + t.Fatalf("get client: %+v %v", got, err) + } + _, err = c.GetClient(ctx, s.name("no-client")) + rejected(t, "unknown client", err) + + code := §ypes.OAuthCode{Code: s.name("code1"), ClientID: cid, RedirectURI: "https://app.example.test/cb", + CodeChallenge: "challenge", SessionToken: s.name("sess"), Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)} + if err := c.SaveCode(ctx, code); err != nil { + t.Fatalf("save code: %v", err) + } + ex, err := c.ExchangeCode(ctx, code.Code) + if err != nil || ex.Code != code.Code || ex.ClientID != cid || ex.SessionToken != code.SessionToken || len(ex.Scopes) != 1 { + t.Fatalf("exchange: %+v %v", ex, err) + } + _, err = c.ExchangeCode(ctx, code.Code) + rejected(t, "code reuse", err) + + expired := *code + expired.Code, expired.ExpiresAt = s.name("code2"), time.Now().Add(-time.Minute) + if err := c.SaveCode(ctx, &expired); err != nil { + t.Fatal(err) + } + _, err = c.ExchangeCode(ctx, expired.Code) + rejected(t, "expired code", err) +} + +func (s *suite) oauthIntrospectRevoke(t *testing.T) { + c := s.Provider.OAuthClient + reg := s.register(t, "intro") + info, err := c.Introspect(ctx, reg.Token) + if err != nil || !info.Active || info.Username != s.name("intro") { + t.Fatalf("introspect: %+v %v", info, err) + } + if err := c.Revoke(ctx, reg.Token); err != nil { + t.Fatalf("revoke: %v", err) + } + if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active { + t.Fatalf("revoked token still active: %+v %v", info, err) + } + if err := c.Revoke(ctx, s.name("unknown-token")); err != nil { + t.Fatalf("revoking an unknown token must succeed (RFC 7009): %v", err) + } + if info, err := c.Introspect(ctx, s.name("unknown-token")); err != nil || info.Active { + t.Fatalf("unknown token: %+v %v", info, err) + } +} + +func (s *suite) oauthUsers(t *testing.T) { + o := s.Provider.OAuthUser + id, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github") + if err != nil || id == 0 { + t.Fatalf("get or create: %d %v", id, err) + } + again, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github") + if err != nil || again != id { + t.Fatalf("second login must return the same user: %d %v", again, err) + } + + exp := time.Now().Add(time.Hour) + sess := lookup.OAuthSession{SessionToken: s.name("os1"), UserID: id, AccessToken: "a1", RefreshToken: s.name("or1"), TokenType: "Bearer", ExpiresAt: exp, Provider: "github"} + if err := o.CreateSession(ctx, sess); err != nil { + t.Fatalf("create session: %v", err) + } + ref, err := o.GetByRefreshToken(ctx, sess.RefreshToken) + if err != nil || ref.UserID != id || ref.AccessToken != "a1" { + t.Fatalf("by refresh token: %+v %v", ref, err) + } + _, err = o.GetByRefreshToken(ctx, s.name("or-missing")) + rejected(t, "unknown refresh token", err) + if err := o.UpdateRefreshToken(ctx, id, sess.RefreshToken, s.name("os2"), "a2", s.name("or2"), exp); err != nil { + t.Fatalf("update refresh token: %v", err) + } + if _, err := o.GetByRefreshToken(ctx, s.name("or2")); err != nil { + t.Fatalf("rotated refresh token not found: %v", err) + } + u, err := o.GetUser(ctx, id) + if err != nil || u.UserName != s.name("gh") { + t.Fatalf("get user: %+v %v", u, err) + } + _, err = o.GetUser(ctx, id+1_000_000) + rejected(t, "unknown user", err) +} + +func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) } + +func (s *suite) passkey(t *testing.T) { + p := s.Provider.Passkey + reg := s.register(t, "pk") + uid := reg.User.UserID + c1, c2 := b64(s.name("cred1")), b64(s.name("cred2")) + + rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: c1, PublicKey: b64("pubkey"), AttestationType: "none", + Transports: []string{"usb", "nfc"}, Name: "Key 1"} + if id, err := p.Store(ctx, rec); err != nil || id == 0 { + t.Fatalf("store: %d %v", id, err) + } + _, err := p.Store(ctx, rec) + rejected(t, "duplicate credential", err) + rec.CredentialID, rec.Name = c2, "Key 2" + if _, err := p.Store(ctx, rec); err != nil { + t.Fatal(err) + } + + owner, count, err := p.Get(ctx, c1) + if err != nil || owner != uid || count != 0 { + t.Fatalf("get: %d %d %v", owner, count, err) + } + _, _, err = p.Get(ctx, b64(s.name("missing"))) + rejected(t, "unknown credential", err) + + if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || clone { + t.Fatalf("advance counter: clone=%v %v", clone, err) + } + if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || !clone { + t.Fatalf("replayed counter must raise a clone warning: clone=%v %v", clone, err) + } + + list, err := p.List(ctx, uid) + if err != nil || len(list) != 2 { + t.Fatalf("list: %d %v", len(list), err) + } + if err := p.Rename(ctx, uid, c1, "Renamed"); err != nil { + t.Fatalf("rename: %v", err) + } + rejected(t, "renaming another user's credential", p.Rename(ctx, uid+1_000_000, c1, "x")) + + gotID, refs, err := p.ByUsername(ctx, s.name("pk")) + if err != nil || gotID != uid || len(refs) != 2 { + t.Fatalf("by username: %d %+v %v", gotID, refs, err) + } + _, _, err = p.ByUsername(ctx, s.name("ghost")) + rejected(t, "unknown username", err) + + resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "10.0.0.3"}) + if err != nil || resp.Token == "" || resp.User.UserName != s.name("pk") { + t.Fatalf("passkey login: %+v %v", resp, err) + } + if _, err := s.Provider.Auth.Session(ctx, resp.Token, ""); err != nil { + t.Fatalf("session from passkey login: %v", err) + } + + rejected(t, "deleting another user's credential", p.Delete(ctx, uid+1_000_000, c1)) + if err := p.Delete(ctx, uid, c1); err != nil { + t.Fatalf("delete: %v", err) + } + rejected(t, "deleting twice", p.Delete(ctx, uid, c1)) +} + +func (s *suite) totp(t *testing.T) { + st := s.Provider.TOTP + uid := s.register(t, "totp").User.UserID + + if on, err := st.Status(ctx, uid); err != nil || on { + t.Fatalf("initial status: %v %v", on, err) + } + _, err := st.Secret(ctx, uid) + rejected(t, "secret without 2FA", err) + + if err := st.Enable(ctx, uid, "SECRET", []string{s.name("h1"), s.name("h2")}); err != nil { + t.Fatalf("enable: %v", err) + } + if on, _ := st.Status(ctx, uid); !on { + t.Fatal("not enabled") + } + if sec, err := st.Secret(ctx, uid); err != nil || sec != "SECRET" { + t.Fatalf("secret: %q %v", sec, err) + } + + if ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1")); err != nil || !ok { + t.Fatalf("backup code: %v %v", ok, err) + } + ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1")) + notOK(t, "backup code reuse", ok, err) + ok, err = st.ValidateBackupCode(ctx, uid, s.name("nope")) + notOK(t, "unknown backup code", ok, err) + + if err := st.RegenerateBackupCodes(ctx, uid, []string{s.name("n1")}); err != nil { + t.Fatalf("regenerate: %v", err) + } + ok, err = st.ValidateBackupCode(ctx, uid, s.name("h2")) + notOK(t, "old backup code after regenerate", ok, err) + if ok, err := st.ValidateBackupCode(ctx, uid, s.name("n1")); err != nil || !ok { + t.Fatalf("new backup code: %v %v", ok, err) + } + + if err := st.Disable(ctx, uid); err != nil { + t.Fatalf("disable: %v", err) + } + if on, _ := st.Status(ctx, uid); on { + t.Fatal("still enabled after disable") + } +} + +// seed inserts one row with dialect placeholders. Values are bound, booleans converted. +func (s *suite) seed(t *testing.T, table string, cols []string, vals ...any) { + t.Helper() + ph := make([]string, len(vals)) + args := make([]any, len(vals)) + for i, v := range vals { + ph[i] = s.Dialect.Placeholder(i + 1) + if b, ok := v.(bool); ok { + v = s.Dialect.Bool(b) + } + args[i] = v + } + q := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", table, strings.Join(cols, ", "), strings.Join(ph, ", ")) + if _, err := s.DB.ExecContext(ctx, q, args...); err != nil { + t.Fatalf("seed %s: %v", table, err) + } +} + +func (s *suite) policy(t *testing.T) { + p := s.Provider.Policy + u1 := s.register(t, "pol1").User.UserID + u2 := s.register(t, "pol2").User.UserID + group := 7_000_000 + u1 + schema, users, orders, secret := s.name("pub"), "Users", "orders", "secret" + + s.seed(t, "sec_group_members", []string{"group_id", "user_id"}, group, u1) + colCols := []string{"user_id", "group_id", "schema_name", "table_name", "column_path", "access_type", "is_active"} + s.seed(t, "sec_column_rules", colCols, u1, nil, schema, users, "email", "mask", true) + s.seed(t, "sec_column_rules", colCols, nil, group, schema, strings.ToLower(users), "profile.ssn", "hide", true) + s.seed(t, "sec_column_rules", colCols, u2, nil, schema, strings.ToLower(users), "other", "hide", true) + s.seed(t, "sec_column_rules", colCols, u1, nil, schema, strings.ToLower(users), "inactive", "hide", false) + s.seed(t, "sec_column_rules", colCols, u1, nil, schema, orders, "x", "hide", true) + s.seed(t, "sec_column_rules", colCols, u1, nil, schema, "users_archive", "y", "hide", true) + + rules, err := p.ColumnSecurity(ctx, u1, schema, "users") + if err != nil || len(rules) != 2 { + t.Fatalf("column rules (user + group, exact table, active only): %d %v %+v", len(rules), err, rules) + } + paths := map[string]bool{} + for _, r := range rules { + paths[strings.Join(r.Path, ".")] = true + } + if !paths["email"] || !paths["profile.ssn"] { + t.Fatalf("paths: %v", paths) + } + if r, err := p.ColumnSecurity(ctx, u2, schema, "users"); err != nil || len(r) != 1 { + t.Fatalf("other user's rules: %d %v", len(r), err) + } + if r, err := p.ColumnSecurity(ctx, u1+u2+1_000_000, schema, "users"); err != nil || len(r) != 0 { + t.Fatalf("no rules must be empty, not an error: %d %v", len(r), err) + } + + rowCols := []string{"user_id", "group_id", "schema_name", "table_name", "template", "has_block", "is_active"} + s.seed(t, "sec_row_rules", rowCols, u1, nil, schema, orders, "owner_id = {UserID}", false, true) + s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "region = 1", false, true) + s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "ignored = 1", false, false) + s.seed(t, "sec_row_rules", rowCols, u2, nil, schema, secret, nil, true, true) + s.seed(t, "sec_row_rules", rowCols, nil, group, schema, secret, "x = 1", false, true) + + rs, err := p.RowSecurity(ctx, u1, schema, orders) + if err != nil || rs.HasBlock || !strings.Contains(rs.Template, "owner_id = {UserID}") || !strings.Contains(rs.Template, "region = 1") || strings.Contains(rs.Template, "ignored") { + t.Fatalf("row template: %+v %v", rs, err) + } + if rs, err := p.RowSecurity(ctx, u2, schema, secret); err != nil || !rs.HasBlock { + t.Fatalf("blocking rule must win: %+v %v", rs, err) + } + if rs, err := p.RowSecurity(ctx, u1+u2+1_000_000, schema, orders); err != nil || rs.HasBlock || rs.Template != "" { + t.Fatalf("no rules: %+v %v", rs, err) + } + if _, err := p.RowSecurity(ctx, "not-a-number", schema, orders); err == nil { + t.Fatal("non-numeric user reference accepted (must fail closed)") + } +} + +func ptr[T any](v T) *T { return &v } diff --git a/pkg/security/lookup/database.go b/pkg/security/lookup/database.go new file mode 100644 index 0000000..adee2ea --- /dev/null +++ b/pkg/security/lookup/database.go @@ -0,0 +1,44 @@ +package lookup + +import ( + "database/sql" + "fmt" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" +) + +// FromDatabase extracts the *sql.DB and the dialect name from an application's +// common.Database (bun, gorm or pgsql adapter), so callers do not have to dig the +// connection out or set Config.Dialect by hand. The returned name is the adapter's +// normalised DriverName ("postgres", "sqlite", "mssql", "mysql") and is empty when +// the adapter reports a driver the dialect registry does not know; set +// Config.Dialect explicitly in that case. +// +// Transaction adapters do not expose a *sql.DB and are rejected. +func FromDatabase(db common.Database) (*sql.DB, string, error) { + if db == nil { + return nil, "", fmt.Errorf("lookup: nil database") + } + p, ok := db.(common.SQLDBProvider) + if !ok { + return nil, "", fmt.Errorf("lookup: %T does not expose a *sql.DB (transaction adapter or unsupported adapter)", db) + } + sqlDB := p.SQLDB() + if sqlDB == nil { + return nil, "", fmt.Errorf("lookup: %T has no *sql.DB", db) + } + name := db.DriverName() + if _, err := dialect.Get(name); err != nil { + name = "" + } + return sqlDB, name, nil +} + +// ResolveDialect returns the dialect for db: the configured one, or detected from the driver. +func (c Config) ResolveDialect(db *sql.DB) (dialect.Dialect, error) { + if c.Dialect != "" { + return dialect.Get(c.Dialect) + } + return dialect.Detect(db) +} diff --git a/pkg/security/database_schema.sql b/pkg/security/lookup/database_schema.sql similarity index 89% rename from pkg/security/database_schema.sql rename to pkg/security/lookup/database_schema.sql index 6000d38..214804f 100644 --- a/pkg/security/database_schema.sql +++ b/pkg/security/lookup/database_schema.sql @@ -428,30 +428,78 @@ EXCEPTION END; $$ LANGUAGE plpgsql; +-- ============================================ +-- Column / row security tables +-- ============================================ +-- A rule applies either to one user (user_id) or to every member of a group +-- (group_id via sec_group_members); exactly one of the two is set. + +CREATE TABLE IF NOT EXISTS sec_group_members ( + group_id INTEGER NOT NULL, + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + PRIMARY KEY (group_id, user_id) +); + +CREATE TABLE IF NOT EXISTS sec_column_rules ( + id SERIAL PRIMARY KEY, + user_id INTEGER REFERENCES users(id) ON DELETE CASCADE, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + column_path TEXT NOT NULL, -- dot path under the table: col or col.sub.field + access_type TEXT NOT NULL, -- e.g. mask, hide, read + mask_start INTEGER DEFAULT 0, + mask_end INTEGER DEFAULT 0, + mask_invert BOOLEAN DEFAULT false, + mask_char TEXT DEFAULT '*', + extra_filters TEXT, -- JSON object + is_active BOOLEAN NOT NULL DEFAULT true, + CHECK ((user_id IS NULL) <> (group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(lower(schema_name), lower(table_name)); + +CREATE TABLE IF NOT EXISTS sec_row_rules ( + id SERIAL PRIMARY KEY, + user_id INTEGER REFERENCES users(id) ON DELETE CASCADE, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + template TEXT, -- SQL fragment, e.g. 'user_id = {UserID}' + has_block BOOLEAN NOT NULL DEFAULT false, + is_active BOOLEAN NOT NULL DEFAULT true, + CHECK ((user_id IS NULL) <> (group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(lower(schema_name), lower(table_name)); + -- 8. resolvespec_column_security - Loads column security rules for user -- Input: user_id (int), schema (text), table_name (text) -- Output: p_success (bool), p_error (text), p_rules (array of security rules as jsonb) +-- Rules are the active sec_column_rules for the exact schema + table (case-insensitive) +-- that belong to the user or to a group the user is a member of. +-- 'control' is returned as schema.table.column_path. CREATE OR REPLACE FUNCTION resolvespec_column_security(p_user_id integer, p_schema text, p_table_name text) RETURNS TABLE(p_success boolean, p_error text, p_rules jsonb) AS $$ DECLARE v_rules jsonb; BEGIN - -- Query column security rules from core.secaccess SELECT jsonb_agg( jsonb_build_object( - 'control', control, - 'accesstype', accesstype, - 'jsonvalue', jsonvalue + 'control', r.schema_name || '.' || r.table_name || '.' || r.column_path, + 'accesstype', r.access_type, + 'jsonvalue', COALESCE(r.extra_filters, '') ) ) INTO v_rules - FROM core.secaccess - WHERE rid_hub IN ( - SELECT rid_hub_parent - FROM core.hub_link - WHERE rid_hub_child = p_user_id AND parent_hubtype = 'secgroup' - ) - AND control ILIKE (p_schema || '.' || p_table_name || '%'); + FROM sec_column_rules r + WHERE r.is_active = true + AND lower(r.schema_name) = lower(p_schema) + AND lower(r.table_name) = lower(p_table_name) + AND ( + r.user_id = p_user_id + OR r.group_id IN (SELECT m.group_id FROM sec_group_members m WHERE m.user_id = p_user_id) + ); IF v_rules IS NULL THEN v_rules := '[]'::jsonb; @@ -464,20 +512,36 @@ EXCEPTION END; $$ LANGUAGE plpgsql; --- 9. resolvespec_row_security - Loads row security template for user (replaces core.api_sec_rowtemplate) +-- 9. resolvespec_row_security - Loads the row security template for user -- Input: schema (text), table_name (text), user_id (int) -- Output: p_template (text), p_block (bool) +-- Applicable rules = active sec_row_rules of the user and of the user's groups for the exact +-- schema + table. Any has_block wins (template empty); otherwise templates are AND-combined, +-- each wrapped in parentheses. CREATE OR REPLACE FUNCTION resolvespec_row_security(p_schema text, p_table_name text, p_user_id integer) RETURNS TABLE(p_template text, p_block boolean) AS $$ +DECLARE + v_block boolean; + v_template text; BEGIN - -- Call the existing core function if it exists, or implement your own logic - -- This is a placeholder that you should customize based on your core.api_sec_rowtemplate logic - RETURN QUERY SELECT ''::text, false; + SELECT COALESCE(bool_or(r.has_block), false), + COALESCE(string_agg('(' || r.template || ')', ' AND ' ORDER BY r.id) + FILTER (WHERE r.template IS NOT NULL AND r.template <> ''), '') + INTO v_block, v_template + FROM sec_row_rules r + WHERE r.is_active = true + AND lower(r.schema_name) = lower(p_schema) + AND lower(r.table_name) = lower(p_table_name) + AND ( + r.user_id = p_user_id + OR r.group_id IN (SELECT m.group_id FROM sec_group_members m WHERE m.user_id = p_user_id) + ); - -- Example implementation: - -- RETURN QUERY SELECT template, has_block - -- FROM core.row_security_config - -- WHERE schema_name = p_schema AND table_name = p_table_name AND user_id = p_user_id; + IF v_block THEN + v_template := ''; + END IF; + + RETURN QUERY SELECT v_template, v_block; END; $$ LANGUAGE plpgsql; @@ -650,7 +714,7 @@ BEGIN v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2'); -- Convert roles array to comma-separated string - SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(p_user_data->'roles')), ',') + SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_user_data->'roles') = 'array' THEN p_user_data->'roles' ELSE '[]'::jsonb END)), ',') INTO v_roles; -- Try to find existing user by email @@ -701,7 +765,7 @@ BEGIN v_access_token := p_session_data->>'access_token'; v_refresh_token := p_session_data->>'refresh_token'; v_token_type := COALESCE(p_session_data->>'token_type', 'Bearer'); - v_expires_at := (p_session_data->>'expires_at')::timestamp; + v_expires_at := (p_session_data->>'expires_at')::timestamptz::timestamp; v_auth_provider := COALESCE(p_session_data->>'auth_provider', 'oauth2'); -- Insert or update session @@ -857,7 +921,7 @@ BEGIN v_new_session_token := p_update_data->>'new_session_token'; v_new_access_token := p_update_data->>'new_access_token'; v_new_refresh_token := p_update_data->>'new_refresh_token'; - v_expires_at := (p_update_data->>'expires_at')::timestamp; + v_expires_at := (p_update_data->>'expires_at')::timestamptz::timestamp; -- Update session in user_sessions table UPDATE user_sessions @@ -1214,7 +1278,7 @@ BEGIN -- Convert transports array IF p_credential->'transports' IS NOT NULL THEN - SELECT ARRAY(SELECT jsonb_array_elements_text(p_credential->'transports')) + SELECT ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_credential->'transports') = 'array' THEN p_credential->'transports' ELSE '[]'::jsonb END)) INTO v_transports; END IF; @@ -1304,12 +1368,11 @@ BEGIN 'name', name, 'created_at', created_at, 'last_used_at', last_used_at - ) + ) ORDER BY created_at DESC ), '[]'::jsonb) INTO v_credentials FROM user_passkey_credentials - WHERE user_id = p_user_id - ORDER BY created_at DESC; + WHERE user_id = p_user_id; RETURN QUERY SELECT true, NULL::text, v_credentials; EXCEPTION @@ -1451,6 +1514,64 @@ EXCEPTION END; $$ LANGUAGE plpgsql; +-- 8. resolvespec_passkey_login - Creates a session for a user whose passkey assertion was verified +-- Input: p_request (jsonb) {user_id: int, ip_address: string, user_agent: string} +-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb) +CREATE OR REPLACE FUNCTION resolvespec_passkey_login(p_request jsonb) +RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$ +DECLARE + v_user_id INTEGER; + v_username TEXT; + v_email TEXT; + v_user_level INTEGER; + v_roles TEXT; + v_program_user_id INTEGER; + v_program_user_table TEXT; + v_session_token TEXT; +BEGIN + v_user_id := (p_request->>'user_id')::integer; + + SELECT username, email, user_level, roles, program_user_id, program_user_table + INTO v_username, v_email, v_user_level, v_roles, v_program_user_id, v_program_user_table + FROM users + WHERE id = v_user_id AND is_active = true; + + IF NOT FOUND THEN + RETURN QUERY SELECT false, 'User not found'::text, NULL::jsonb; + RETURN; + END IF; + + v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text; + + INSERT INTO user_sessions (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at) + VALUES (v_session_token, v_user_id, now() + interval '24 hours', + p_request->>'ip_address', p_request->>'user_agent', now()); + + UPDATE users SET last_login_at = now() WHERE id = v_user_id; + + RETURN QUERY SELECT + true, + NULL::text, + jsonb_build_object( + 'token', v_session_token, + 'user', jsonb_build_object( + 'user_id', v_user_id, + 'user_name', v_username, + 'email', v_email, + 'user_level', v_user_level, + 'roles', string_to_array(COALESCE(v_roles, ''), ','), + 'session_id', v_session_token, + 'program_user_id', COALESCE(v_program_user_id, 0), + 'program_user_table', COALESCE(v_program_user_table, '') + ), + 'expires_in', 86400 + ); +EXCEPTION + WHEN OTHERS THEN + RETURN QUERY SELECT false, SQLERRM::text, NULL::jsonb; +END; +$$ LANGUAGE plpgsql; + -- ============================================ -- Example: Test Passkey stored procedures -- ============================================ @@ -1671,24 +1792,24 @@ CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at); -- OAuth2 Server Stored Procedures -- ============================================ -CREATE OR REPLACE FUNCTION resolvespec_oauth_register_client(p_data jsonb) +CREATE OR REPLACE FUNCTION resolvespec_oauth_register_client(p_request jsonb) RETURNS TABLE(p_success bool, p_error text, p_data jsonb) LANGUAGE plpgsql AS $$ DECLARE v_client_id text; v_row jsonb; BEGIN - v_client_id := p_data->>'client_id'; + v_client_id := p_request->>'client_id'; INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method) VALUES ( v_client_id, - ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')), - p_data->>'client_name', - COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'grant_types')), ARRAY['authorization_code']), - COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email']), - NULLIF(p_data->>'client_secret_hash', ''), - COALESCE(NULLIF(p_data->>'token_endpoint_auth_method', ''), 'none') + ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)), + p_request->>'client_name', + CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE ARRAY['authorization_code'] END, + CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE ARRAY['openid','profile','email'] END, + NULLIF(p_request->>'client_secret_hash', ''), + COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none') ) RETURNING to_jsonb(oauth_clients.*) INTO v_row; @@ -1717,22 +1838,22 @@ BEGIN END; $$; -CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_data jsonb) +CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_request jsonb) RETURNS TABLE(p_success bool, p_error text) LANGUAGE plpgsql AS $$ BEGIN INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at) VALUES ( - p_data->>'code', - p_data->>'client_id', - p_data->>'redirect_uri', - p_data->>'client_state', - p_data->>'code_challenge', - COALESCE(p_data->>'code_challenge_method', 'S256'), - p_data->>'session_token', - p_data->>'refresh_token', - ARRAY(SELECT jsonb_array_elements_text(p_data->'scopes')), - (p_data->>'expires_at')::timestamp + p_request->>'code', + p_request->>'client_id', + p_request->>'redirect_uri', + p_request->>'client_state', + p_request->>'code_challenge', + COALESCE(p_request->>'code_challenge_method', 'S256'), + p_request->>'session_token', + p_request->>'refresh_token', + ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)), + (p_request->>'expires_at')::timestamptz::timestamp ); RETURN QUERY SELECT true, null::text; diff --git a/pkg/security/lookup/database_test.go b/pkg/security/lookup/database_test.go new file mode 100644 index 0000000..13fef49 --- /dev/null +++ b/pkg/security/lookup/database_test.go @@ -0,0 +1,84 @@ +package lookup_test + +import ( + "database/sql" + "testing" + + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/sqlitedialect" + "github.com/uptrace/bun/driver/sqliteshim" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" +) + +func TestFromDatabase(t *testing.T) { + sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:?cache=shared") + if err != nil { + t.Fatal(err) + } + defer sqldb.Close() + + gdb, err := gorm.Open(sqlite.Open("file::memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + + cases := map[string]struct { + db common.Database + same *sql.DB // expected handle; nil = just must be non-nil + }{ + "pgsql": {database.NewPgSQLAdapter(sqldb, "sqlite"), sqldb}, + "bun": {database.NewBunAdapter(bun.NewDB(sqldb, sqlitedialect.New())), sqldb}, + "gorm": {database.NewGormAdapter(gdb), nil}, + } + for name, c := range cases { + got, dialectName, err := lookup.FromDatabase(c.db) + if err != nil { + t.Errorf("%s: %v", name, err) + continue + } + if got == nil || (c.same != nil && got != c.same) { + t.Errorf("%s: unexpected *sql.DB %v", name, got) + } + if dialectName != "sqlite" { + t.Errorf("%s: dialect = %q, want sqlite", name, dialectName) + } + if err := got.Ping(); err != nil { + t.Errorf("%s: handle not usable: %v", name, err) + } + } +} + +func TestFromDatabaseRejects(t *testing.T) { + if _, _, err := lookup.FromDatabase(nil); err == nil { + t.Error("nil database should fail") + } + // A database that does not expose a *sql.DB (the embedded interface is nil; only the type matters). + if _, _, err := lookup.FromDatabase(struct{ common.Database }{}); err == nil { + t.Error("adapter without SQLDB should fail") + } +} + +func TestResolveDialect(t *testing.T) { + sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:") + if err != nil { + t.Fatal(err) + } + defer sqldb.Close() + + d, err := lookup.Config{}.ResolveDialect(sqldb) + if err != nil || d.Name() != "sqlite" { + t.Errorf("detected = %v, %v", d, err) + } + d, err = lookup.Config{Dialect: "mysql"}.ResolveDialect(sqldb) + if err != nil || d.Name() != "mysql" { + t.Errorf("explicit dialect should win: %v, %v", d, err) + } + if _, err := (lookup.Config{Dialect: "oracle"}).Resolve(); err == nil { + t.Error("unknown dialect should fail Resolve") + } +} diff --git a/pkg/security/lookup/ddl/ddl.go b/pkg/security/lookup/ddl/ddl.go new file mode 100644 index 0000000..770476f --- /dev/null +++ b/pkg/security/lookup/ddl/ddl.go @@ -0,0 +1,50 @@ +// Package ddl holds the reference table schemas for the lookup direct backend, one per +// dialect. They use the lookup.DefaultSchema table and column names; copy and adapt them +// when you override names through lookup.Config.Schema. +// +// The Postgres file creates tables only. The stored-procedure schema +// (lookup/database_schema.sql) is a separate script with native bytea / text[] columns and +// must not be combined with it. +package ddl + +import ( + "embed" + "fmt" + "strings" +) + +//go:embed postgres.sql sqlite.sql mysql.sql mssql.sql +var files embed.FS + +// SQL returns the schema script for a dialect name ("postgres", "sqlite", "mysql", "mssql"). +func SQL(dialect string) (string, error) { + b, err := files.ReadFile(dialect + ".sql") + if err != nil { + return "", fmt.Errorf("ddl: no reference schema for dialect %q", dialect) + } + return string(b), nil +} + +// Statements returns the schema as separate statements, for drivers that reject +// multi-statement execution. Comment-only lines are dropped. +func Statements(dialect string) ([]string, error) { + s, err := SQL(dialect) + if err != nil { + return nil, err + } + var out []string + var cur strings.Builder + for _, line := range strings.Split(s, "\n") { + t := strings.TrimSpace(line) + if t == "" || strings.HasPrefix(t, "--") { + continue + } + cur.WriteString(line) + cur.WriteString("\n") + if strings.HasSuffix(t, ";") { + out = append(out, strings.TrimSpace(cur.String())) + cur.Reset() + } + } + return out, nil +} diff --git a/pkg/security/lookup/ddl/ddl_test.go b/pkg/security/lookup/ddl/ddl_test.go new file mode 100644 index 0000000..9c65d4c --- /dev/null +++ b/pkg/security/lookup/ddl/ddl_test.go @@ -0,0 +1,58 @@ +package ddl_test + +import ( + "database/sql" + "strings" + "testing" + + _ "github.com/glebarez/go-sqlite" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" +) + +func TestStatementsAllDialects(t *testing.T) { + for _, d := range []string{"postgres", "sqlite", "mysql", "mssql"} { + st, err := ddl.Statements(d) + if err != nil { + t.Fatal(err) + } + if len(st) < 12 { + t.Errorf("%s: %d statements", d, len(st)) + } + all := strings.Join(st, "\n") + for _, tbl := range []string{"users", "user_sessions", "user_keys", "oauth_codes", "sec_group_members", "sec_column_rules", "sec_row_rules"} { + if !strings.Contains(all, tbl+" (") { + t.Errorf("%s: missing table %s", d, tbl) + } + } + } + if _, err := ddl.SQL("oracle"); err == nil { + t.Error("unknown dialect must error") + } +} + +// Every table and column the default schema names must exist in the sqlite reference DDL. +func TestSQLiteMatchesDefaultSchema(t *testing.T) { + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + s, _ := ddl.SQL("sqlite") + if _, err := db.Exec(s); err != nil { + t.Fatal(err) + } + if _, err := db.Exec(s); err != nil { + t.Fatalf("script must be re-runnable: %v", err) + } + sc := lookup.DefaultSchema() + for _, tbl := range sc { + for _, c := range tbl.Columns { + if _, err := db.Exec("SELECT " + c + " FROM " + tbl.Name + " WHERE 1=0"); err != nil { + t.Errorf("%s.%s: %v", tbl.Name, c, err) + } + } + } +} diff --git a/pkg/security/lookup/ddl/mssql.sql b/pkg/security/lookup/ddl/mssql.sql new file mode 100644 index 0000000..7c3ec25 --- /dev/null +++ b/pkg/security/lookup/ddl/mssql.sql @@ -0,0 +1,259 @@ +-- Reference schema for the lookup direct backend: Microsoft SQL Server 2016+. Direct backend only. Run each statement separately (see ddl.Statements). +-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema. +-- Generated by the project, edit freely for your deployment (types, collations, extra columns). + +-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login +-- roles: comma-separated roles + +IF OBJECT_ID(N'users', N'U') IS NULL +CREATE TABLE users ( + id INT IDENTITY(1,1) PRIMARY KEY, + username NVARCHAR(255) NOT NULL UNIQUE, + email NVARCHAR(255) NOT NULL UNIQUE, + password NVARCHAR(255), + user_level INT DEFAULT 0, + roles NVARCHAR(500), + is_active BIT DEFAULT 1, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + updated_at DATETIME2 DEFAULT SYSUTCDATETIME(), + last_login_at DATETIME2, + program_user_id INT DEFAULT 0, + program_user_table NVARCHAR(255) DEFAULT '', + remote_id NVARCHAR(255), + auth_provider NVARCHAR(50), + totp_secret NVARCHAR(255), + totp_enabled BIT DEFAULT 0, + totp_enabled_at DATETIME2 +); + + +IF OBJECT_ID(N'user_sessions', N'U') IS NULL +CREATE TABLE user_sessions ( + id INT IDENTITY(1,1) PRIMARY KEY, + session_token NVARCHAR(450) NOT NULL UNIQUE, + user_id INT NOT NULL, + expires_at DATETIME2 NOT NULL, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + last_activity_at DATETIME2 DEFAULT SYSUTCDATETIME(), + ip_address NVARCHAR(45), + user_agent NVARCHAR(MAX), + access_token NVARCHAR(MAX), + refresh_token NVARCHAR(MAX), + token_type NVARCHAR(50) DEFAULT 'Bearer', + auth_provider NVARCHAR(50), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_user_id' AND object_id = OBJECT_ID(N'user_sessions')) +CREATE INDEX idx_user_sessions_user_id ON user_sessions(user_id); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_expires_at' AND object_id = OBJECT_ID(N'user_sessions')) +CREATE INDEX idx_user_sessions_expires_at ON user_sessions(expires_at); + + +IF OBJECT_ID(N'token_blacklist', N'U') IS NULL +CREATE TABLE token_blacklist ( + id INT IDENTITY(1,1) PRIMARY KEY, + token NVARCHAR(500) NOT NULL, + user_id INT, + expires_at DATETIME2 NOT NULL, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- code_hash: SHA-256 hex of the backup code + +IF OBJECT_ID(N'user_totp_backup_codes', N'U') IS NULL +CREATE TABLE user_totp_backup_codes ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT NOT NULL, + code_hash NVARCHAR(64) NOT NULL, + used BIT DEFAULT 0, + used_at DATETIME2, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_user_id' AND object_id = OBJECT_ID(N'user_totp_backup_codes')) +CREATE INDEX idx_totp_user_id ON user_totp_backup_codes(user_id); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_code_hash' AND object_id = OBJECT_ID(N'user_totp_backup_codes')) +CREATE INDEX idx_totp_code_hash ON user_totp_backup_codes(code_hash); + + +-- credential_id: base64 text +-- public_key: base64 text +-- aaguid: base64 text +-- transports: JSON-encoded array + +IF OBJECT_ID(N'user_passkey_credentials', N'U') IS NULL +CREATE TABLE user_passkey_credentials ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT NOT NULL, + credential_id VARCHAR(900) NOT NULL UNIQUE, + public_key NVARCHAR(MAX) NOT NULL, + attestation_type NVARCHAR(50) DEFAULT 'none', + aaguid NVARCHAR(MAX), + sign_count INT DEFAULT 0, + clone_warning BIT DEFAULT 0, + transports NVARCHAR(MAX), + backup_eligible BIT DEFAULT 0, + backup_state BIT DEFAULT 0, + name NVARCHAR(255), + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + last_used_at DATETIME2 DEFAULT SYSUTCDATETIME(), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_passkey_user_id' AND object_id = OBJECT_ID(N'user_passkey_credentials')) +CREATE INDEX idx_passkey_user_id ON user_passkey_credentials(user_id); + + +-- token_hash: SHA-256 hex of the raw token + +IF OBJECT_ID(N'user_password_resets', N'U') IS NULL +CREATE TABLE user_password_resets ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT NOT NULL, + token_hash NVARCHAR(64) NOT NULL UNIQUE, + expires_at DATETIME2 NOT NULL, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + used BIT DEFAULT 0, + used_at DATETIME2, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_user_id' AND object_id = OBJECT_ID(N'user_password_resets')) +CREATE INDEX idx_pw_reset_user_id ON user_password_resets(user_id); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_expires_at' AND object_id = OBJECT_ID(N'user_password_resets')) +CREATE INDEX idx_pw_reset_expires_at ON user_password_resets(expires_at); + + +-- redirect_uris: JSON-encoded array +-- grant_types: JSON-encoded array +-- allowed_scopes: JSON-encoded array +-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients + +IF OBJECT_ID(N'oauth_clients', N'U') IS NULL +CREATE TABLE oauth_clients ( + id INT IDENTITY(1,1) PRIMARY KEY, + client_id NVARCHAR(255) NOT NULL UNIQUE, + redirect_uris NVARCHAR(MAX) NOT NULL, + client_name NVARCHAR(255), + grant_types NVARCHAR(MAX), + allowed_scopes NVARCHAR(MAX), + client_secret_hash NVARCHAR(MAX), + token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none', + is_active BIT DEFAULT 1, + created_at DATETIME2 DEFAULT SYSUTCDATETIME() +); + + +-- scopes: JSON-encoded array + +IF OBJECT_ID(N'oauth_codes', N'U') IS NULL +CREATE TABLE oauth_codes ( + id INT IDENTITY(1,1) PRIMARY KEY, + code NVARCHAR(255) NOT NULL UNIQUE, + client_id NVARCHAR(255) NOT NULL, + redirect_uri NVARCHAR(MAX) NOT NULL, + client_state NVARCHAR(MAX), + code_challenge NVARCHAR(255) NOT NULL, + code_challenge_method NVARCHAR(10) DEFAULT 'S256', + session_token NVARCHAR(MAX) NOT NULL, + refresh_token NVARCHAR(MAX), + scopes NVARCHAR(MAX), + expires_at DATETIME2 NOT NULL, + created_at DATETIME2 DEFAULT SYSUTCDATETIME() +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_codes_expires' AND object_id = OBJECT_ID(N'oauth_codes')) +CREATE INDEX idx_oauth_codes_expires ON oauth_codes(expires_at); + + +-- key_hash: SHA-256 hex +-- scopes: JSON-encoded array +-- meta: JSON-encoded object + +IF OBJECT_ID(N'user_keys', N'U') IS NULL +CREATE TABLE user_keys ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT NOT NULL, + key_type NVARCHAR(50) NOT NULL, + key_hash NVARCHAR(64) NOT NULL UNIQUE, + name NVARCHAR(255) NOT NULL DEFAULT '', + scopes NVARCHAR(MAX), + meta NVARCHAR(MAX), + expires_at DATETIME2, + created_at DATETIME2 DEFAULT SYSUTCDATETIME(), + last_used_at DATETIME2, + is_active BIT DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_user_id' AND object_id = OBJECT_ID(N'user_keys')) +CREATE INDEX idx_user_keys_user_id ON user_keys(user_id); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_key_type' AND object_id = OBJECT_ID(N'user_keys')) +CREATE INDEX idx_user_keys_key_type ON user_keys(key_type); + + +-- Optional: omit to use per-user rules only. + +IF OBJECT_ID(N'sec_group_members', N'U') IS NULL +CREATE TABLE sec_group_members ( + group_id INT NOT NULL, + user_id INT NOT NULL, + PRIMARY KEY (group_id, user_id), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- column_path: dot path under the table: col or col.sub.field +-- access_type: mask, hide, read, ... +-- extra_filters: JSON object + +IF OBJECT_ID(N'sec_column_rules', N'U') IS NULL +CREATE TABLE sec_column_rules ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT, + group_id INT, + schema_name NVARCHAR(255) NOT NULL, + table_name NVARCHAR(255) NOT NULL, + column_path NVARCHAR(255) NOT NULL, + access_type NVARCHAR(50) NOT NULL, + mask_start INT DEFAULT 0, + mask_end INT DEFAULT 0, + mask_invert BIT DEFAULT 0, + mask_char NVARCHAR(10) DEFAULT '*', + extra_filters NVARCHAR(MAX), + is_active BIT NOT NULL DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_column_rules_table' AND object_id = OBJECT_ID(N'sec_column_rules')) +CREATE INDEX idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name); + + +-- template: SQL fragment, e.g. user_id = {UserID} + +IF OBJECT_ID(N'sec_row_rules', N'U') IS NULL +CREATE TABLE sec_row_rules ( + id INT IDENTITY(1,1) PRIMARY KEY, + user_id INT, + group_id INT, + schema_name NVARCHAR(255) NOT NULL, + table_name NVARCHAR(255) NOT NULL, + template NVARCHAR(MAX), + has_block BIT NOT NULL DEFAULT 0, + is_active BIT NOT NULL DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_row_rules_table' AND object_id = OBJECT_ID(N'sec_row_rules')) +CREATE INDEX idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name); + diff --git a/pkg/security/lookup/ddl/mysql.sql b/pkg/security/lookup/ddl/mysql.sql new file mode 100644 index 0000000..fd7959a --- /dev/null +++ b/pkg/security/lookup/ddl/mysql.sql @@ -0,0 +1,223 @@ +-- Reference schema for the lookup direct backend: MySQL 8.0.16+ / MariaDB 10.2+. Direct backend only. Run each statement separately (see ddl.Statements) unless multiStatements=true. +-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema. +-- Generated by the project, edit freely for your deployment (types, collations, extra columns). + +-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login +-- roles: comma-separated roles + +CREATE TABLE IF NOT EXISTS users ( + id INT AUTO_INCREMENT PRIMARY KEY, + username VARCHAR(255) NOT NULL UNIQUE, + email VARCHAR(255) NOT NULL UNIQUE, + password VARCHAR(255), + user_level INT DEFAULT 0, + roles VARCHAR(500), + is_active TINYINT(1) DEFAULT 1, + created_at DATETIME NULL, + updated_at DATETIME NULL, + last_login_at DATETIME, + program_user_id INT DEFAULT 0, + program_user_table VARCHAR(255) DEFAULT '', + remote_id VARCHAR(255), + auth_provider VARCHAR(50), + totp_secret VARCHAR(255), + totp_enabled TINYINT(1) DEFAULT 0, + totp_enabled_at DATETIME +); + + +CREATE TABLE IF NOT EXISTS user_sessions ( + id INT AUTO_INCREMENT PRIMARY KEY, + session_token VARCHAR(500) NOT NULL UNIQUE, + user_id INT NOT NULL, + expires_at DATETIME NOT NULL, + created_at DATETIME NULL, + last_activity_at DATETIME NULL, + ip_address VARCHAR(45), + user_agent TEXT, + access_token TEXT, + refresh_token TEXT, + token_type VARCHAR(50) DEFAULT 'Bearer', + auth_provider VARCHAR(50), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + INDEX idx_user_sessions_user_id (user_id), + INDEX idx_user_sessions_expires_at (expires_at) +); + + +CREATE TABLE IF NOT EXISTS token_blacklist ( + id INT AUTO_INCREMENT PRIMARY KEY, + token VARCHAR(500) NOT NULL, + user_id INT, + expires_at DATETIME NOT NULL, + created_at DATETIME NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- code_hash: SHA-256 hex of the backup code + +CREATE TABLE IF NOT EXISTS user_totp_backup_codes ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT NOT NULL, + code_hash VARCHAR(64) NOT NULL, + used TINYINT(1) DEFAULT 0, + used_at DATETIME, + created_at DATETIME NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + INDEX idx_totp_user_id (user_id), + INDEX idx_totp_code_hash (code_hash) +); + + +-- credential_id: base64 text +-- public_key: base64 text +-- aaguid: base64 text +-- transports: JSON-encoded array + +CREATE TABLE IF NOT EXISTS user_passkey_credentials ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT NOT NULL, + credential_id VARCHAR(1400) CHARACTER SET ascii NOT NULL UNIQUE, + public_key TEXT NOT NULL, + attestation_type VARCHAR(50) DEFAULT 'none', + aaguid TEXT, + sign_count INT DEFAULT 0, + clone_warning TINYINT(1) DEFAULT 0, + transports TEXT, + backup_eligible TINYINT(1) DEFAULT 0, + backup_state TINYINT(1) DEFAULT 0, + name VARCHAR(255), + created_at DATETIME NULL, + last_used_at DATETIME NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + INDEX idx_passkey_user_id (user_id) +); + + +-- token_hash: SHA-256 hex of the raw token + +CREATE TABLE IF NOT EXISTS user_password_resets ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT NOT NULL, + token_hash VARCHAR(64) NOT NULL UNIQUE, + expires_at DATETIME NOT NULL, + created_at DATETIME NULL, + used TINYINT(1) DEFAULT 0, + used_at DATETIME, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + INDEX idx_pw_reset_user_id (user_id), + INDEX idx_pw_reset_expires_at (expires_at) +); + + +-- redirect_uris: JSON-encoded array +-- grant_types: JSON-encoded array +-- allowed_scopes: JSON-encoded array +-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients + +CREATE TABLE IF NOT EXISTS oauth_clients ( + id INT AUTO_INCREMENT PRIMARY KEY, + client_id VARCHAR(255) NOT NULL UNIQUE, + redirect_uris TEXT NOT NULL, + client_name VARCHAR(255), + grant_types TEXT, + allowed_scopes TEXT, + client_secret_hash TEXT, + token_endpoint_auth_method VARCHAR(30) DEFAULT 'none', + is_active TINYINT(1) DEFAULT 1, + created_at DATETIME NULL +); + + +-- scopes: JSON-encoded array + +CREATE TABLE IF NOT EXISTS oauth_codes ( + id INT AUTO_INCREMENT PRIMARY KEY, + code VARCHAR(255) NOT NULL UNIQUE, + client_id VARCHAR(255) NOT NULL, + redirect_uri TEXT NOT NULL, + client_state TEXT, + code_challenge VARCHAR(255) NOT NULL, + code_challenge_method VARCHAR(10) DEFAULT 'S256', + session_token TEXT NOT NULL, + refresh_token TEXT, + scopes TEXT, + expires_at DATETIME NOT NULL, + created_at DATETIME NULL, + INDEX idx_oauth_codes_expires (expires_at) +); + + +-- key_hash: SHA-256 hex +-- scopes: JSON-encoded array +-- meta: JSON-encoded object + +CREATE TABLE IF NOT EXISTS user_keys ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT NOT NULL, + key_type VARCHAR(50) NOT NULL, + key_hash VARCHAR(64) NOT NULL UNIQUE, + name VARCHAR(255) NOT NULL DEFAULT '', + scopes TEXT, + meta TEXT, + expires_at DATETIME, + created_at DATETIME NULL, + last_used_at DATETIME, + is_active TINYINT(1) DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + INDEX idx_user_keys_user_id (user_id), + INDEX idx_user_keys_key_type (key_type) +); + + +-- Optional: omit to use per-user rules only. + +CREATE TABLE IF NOT EXISTS sec_group_members ( + group_id INT NOT NULL, + user_id INT NOT NULL, + PRIMARY KEY (group_id, user_id), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- column_path: dot path under the table: col or col.sub.field +-- access_type: mask, hide, read, ... +-- extra_filters: JSON object + +CREATE TABLE IF NOT EXISTS sec_column_rules ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT, + group_id INT, + schema_name VARCHAR(255) NOT NULL, + table_name VARCHAR(255) NOT NULL, + column_path VARCHAR(255) NOT NULL, + access_type VARCHAR(50) NOT NULL, + mask_start INT DEFAULT 0, + mask_end INT DEFAULT 0, + mask_invert TINYINT(1) DEFAULT 0, + mask_char VARCHAR(10) DEFAULT '*', + extra_filters TEXT, + is_active TINYINT(1) NOT NULL DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)), + INDEX idx_sec_column_rules_table (schema_name, table_name) +); + + +-- template: SQL fragment, e.g. user_id = {UserID} + +CREATE TABLE IF NOT EXISTS sec_row_rules ( + id INT AUTO_INCREMENT PRIMARY KEY, + user_id INT, + group_id INT, + schema_name VARCHAR(255) NOT NULL, + table_name VARCHAR(255) NOT NULL, + template TEXT, + has_block TINYINT(1) NOT NULL DEFAULT 0, + is_active TINYINT(1) NOT NULL DEFAULT 1, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)), + INDEX idx_sec_row_rules_table (schema_name, table_name) +); + diff --git a/pkg/security/lookup/ddl/postgres.sql b/pkg/security/lookup/ddl/postgres.sql new file mode 100644 index 0000000..813872b --- /dev/null +++ b/pkg/security/lookup/ddl/postgres.sql @@ -0,0 +1,239 @@ +-- Reference schema for the lookup direct backend: PostgreSQL (tables only, no stored procedures). Use with lookup.Config{Mode: lookup.ModeDirect}. +-- Do not combine with the procedure schema (database_schema.sql): that schema stores credential ids as bytea and +-- list columns as text[], this one stores base64 / JSON text which is what the direct backend reads and writes. +-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema. +-- Generated by the project, edit freely for your deployment (types, collations, extra columns). + +-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login +-- roles: comma-separated roles + +CREATE TABLE IF NOT EXISTS users ( + id SERIAL PRIMARY KEY, + username VARCHAR(255) NOT NULL UNIQUE, + email VARCHAR(255) NOT NULL UNIQUE, + password VARCHAR(255), + user_level INTEGER DEFAULT 0, + roles VARCHAR(500), + is_active BOOLEAN DEFAULT true, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + last_login_at TIMESTAMP, + program_user_id INTEGER DEFAULT 0, + program_user_table VARCHAR(255) DEFAULT '', + remote_id VARCHAR(255), + auth_provider VARCHAR(50), + totp_secret VARCHAR(255), + totp_enabled BOOLEAN DEFAULT false, + totp_enabled_at TIMESTAMP +); + + +CREATE TABLE IF NOT EXISTS user_sessions ( + id SERIAL PRIMARY KEY, + session_token VARCHAR(500) NOT NULL UNIQUE, + user_id INTEGER NOT NULL, + expires_at TIMESTAMP NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + ip_address VARCHAR(45), + user_agent TEXT, + access_token TEXT, + refresh_token TEXT, + token_type VARCHAR(50) DEFAULT 'Bearer', + auth_provider VARCHAR(50), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id); + +CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at); + +CREATE INDEX IF NOT EXISTS idx_user_sessions_refresh_token ON user_sessions(refresh_token); + + +CREATE TABLE IF NOT EXISTS token_blacklist ( + id SERIAL PRIMARY KEY, + token VARCHAR(500) NOT NULL, + user_id INTEGER, + expires_at TIMESTAMP NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- code_hash: SHA-256 hex of the backup code + +CREATE TABLE IF NOT EXISTS user_totp_backup_codes ( + id SERIAL PRIMARY KEY, + user_id INTEGER NOT NULL, + code_hash VARCHAR(64) NOT NULL, + used BOOLEAN DEFAULT false, + used_at TIMESTAMP, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +CREATE INDEX IF NOT EXISTS idx_totp_user_id ON user_totp_backup_codes(user_id); + +CREATE INDEX IF NOT EXISTS idx_totp_code_hash ON user_totp_backup_codes(code_hash); + + +-- credential_id: base64 text +-- public_key: base64 text +-- aaguid: base64 text +-- transports: JSON-encoded array + +CREATE TABLE IF NOT EXISTS user_passkey_credentials ( + id SERIAL PRIMARY KEY, + user_id INTEGER NOT NULL, + credential_id TEXT NOT NULL UNIQUE, + public_key TEXT NOT NULL, + attestation_type VARCHAR(50) DEFAULT 'none', + aaguid TEXT, + sign_count INTEGER DEFAULT 0, + clone_warning BOOLEAN DEFAULT false, + transports TEXT, + backup_eligible BOOLEAN DEFAULT false, + backup_state BOOLEAN DEFAULT false, + name VARCHAR(255), + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + last_used_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id); + + +-- token_hash: SHA-256 hex of the raw token + +CREATE TABLE IF NOT EXISTS user_password_resets ( + id SERIAL PRIMARY KEY, + user_id INTEGER NOT NULL, + token_hash VARCHAR(64) NOT NULL UNIQUE, + expires_at TIMESTAMP NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + used BOOLEAN DEFAULT false, + used_at TIMESTAMP, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +CREATE INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id); + +CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at); + + +-- redirect_uris: JSON-encoded array +-- grant_types: JSON-encoded array +-- allowed_scopes: JSON-encoded array +-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients + +CREATE TABLE IF NOT EXISTS oauth_clients ( + id SERIAL PRIMARY KEY, + client_id VARCHAR(255) NOT NULL UNIQUE, + redirect_uris TEXT NOT NULL, + client_name VARCHAR(255), + grant_types TEXT, + allowed_scopes TEXT, + client_secret_hash TEXT, + token_endpoint_auth_method VARCHAR(30) DEFAULT 'none', + is_active BOOLEAN DEFAULT true, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + + +-- scopes: JSON-encoded array + +CREATE TABLE IF NOT EXISTS oauth_codes ( + id SERIAL PRIMARY KEY, + code VARCHAR(255) NOT NULL UNIQUE, + client_id VARCHAR(255) NOT NULL, + redirect_uri TEXT NOT NULL, + client_state TEXT, + code_challenge VARCHAR(255) NOT NULL, + code_challenge_method VARCHAR(10) DEFAULT 'S256', + session_token TEXT NOT NULL, + refresh_token TEXT, + scopes TEXT, + expires_at TIMESTAMP NOT NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at); + + +-- key_hash: SHA-256 hex +-- scopes: JSON-encoded array +-- meta: JSON-encoded object + +CREATE TABLE IF NOT EXISTS user_keys ( + id SERIAL PRIMARY KEY, + user_id INTEGER NOT NULL, + key_type VARCHAR(50) NOT NULL, + key_hash VARCHAR(64) NOT NULL UNIQUE, + name VARCHAR(255) NOT NULL DEFAULT '', + scopes TEXT, + meta TEXT, + expires_at TIMESTAMP, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + last_used_at TIMESTAMP, + is_active BOOLEAN DEFAULT true, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + +CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id); + +CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type); + + +-- Optional: omit to use per-user rules only. + +CREATE TABLE IF NOT EXISTS sec_group_members ( + group_id INTEGER NOT NULL, + user_id INTEGER NOT NULL, + PRIMARY KEY (group_id, user_id), + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE +); + + +-- column_path: dot path under the table: col or col.sub.field +-- access_type: mask, hide, read, ... +-- extra_filters: JSON object + +CREATE TABLE IF NOT EXISTS sec_column_rules ( + id SERIAL PRIMARY KEY, + user_id INTEGER, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + column_path TEXT NOT NULL, + access_type VARCHAR(50) NOT NULL, + mask_start INTEGER DEFAULT 0, + mask_end INTEGER DEFAULT 0, + mask_invert BOOLEAN DEFAULT false, + mask_char VARCHAR(10) DEFAULT '*', + extra_filters TEXT, + is_active BOOLEAN NOT NULL DEFAULT true, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name); + + +-- template: SQL fragment, e.g. user_id = {UserID} + +CREATE TABLE IF NOT EXISTS sec_row_rules ( + id SERIAL PRIMARY KEY, + user_id INTEGER, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + template TEXT, + has_block BOOLEAN NOT NULL DEFAULT false, + is_active BOOLEAN NOT NULL DEFAULT true, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name); + diff --git a/pkg/security/database_schema_sqlite.sql b/pkg/security/lookup/ddl/sqlite.sql similarity index 51% rename from pkg/security/database_schema_sqlite.sql rename to pkg/security/lookup/ddl/sqlite.sql index 1904693..c3802bf 100644 --- a/pkg/security/database_schema_sqlite.sql +++ b/pkg/security/lookup/ddl/sqlite.sql @@ -1,13 +1,15 @@ --- Portable schema for Direct-mode (non-stored-procedure) operation. --- Plain CREATE TABLE statements only, no functions/triggers, using types --- understood by SQLite (and portable to MySQL). Used by Direct-mode tests --- and as a reference for deployments without Postgres. +-- Reference schema for the lookup direct backend: SQLite. Direct backend only. +-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema. +-- Generated by the project, edit freely for your deployment (types, collations, extra columns). + +-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login +-- roles: comma-separated roles CREATE TABLE IF NOT EXISTS users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username VARCHAR(255) NOT NULL UNIQUE, email VARCHAR(255) NOT NULL UNIQUE, - password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in) + password VARCHAR(255), user_level INTEGER DEFAULT 0, roles VARCHAR(500), is_active BOOLEAN DEFAULT 1, @@ -23,6 +25,7 @@ CREATE TABLE IF NOT EXISTS users ( totp_enabled_at TIMESTAMP ); + CREATE TABLE IF NOT EXISTS user_sessions ( id INTEGER PRIMARY KEY AUTOINCREMENT, session_token VARCHAR(500) NOT NULL UNIQUE, @@ -38,10 +41,12 @@ CREATE TABLE IF NOT EXISTS user_sessions ( auth_provider VARCHAR(50) ); -CREATE INDEX IF NOT EXISTS idx_session_token ON user_sessions(session_token); -CREATE INDEX IF NOT EXISTS idx_user_id ON user_sessions(user_id); -CREATE INDEX IF NOT EXISTS idx_expires_at ON user_sessions(expires_at); -CREATE INDEX IF NOT EXISTS idx_refresh_token ON user_sessions(refresh_token); +CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id); + +CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at); + +CREATE INDEX IF NOT EXISTS idx_user_sessions_refresh_token ON user_sessions(refresh_token); + CREATE TABLE IF NOT EXISTS token_blacklist ( id INTEGER PRIMARY KEY AUTOINCREMENT, @@ -51,6 +56,9 @@ CREATE TABLE IF NOT EXISTS token_blacklist ( created_at TIMESTAMP ); + +-- code_hash: SHA-256 hex of the backup code + CREATE TABLE IF NOT EXISTS user_totp_backup_codes ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, @@ -61,18 +69,25 @@ CREATE TABLE IF NOT EXISTS user_totp_backup_codes ( ); CREATE INDEX IF NOT EXISTS idx_totp_user_id ON user_totp_backup_codes(user_id); + CREATE INDEX IF NOT EXISTS idx_totp_code_hash ON user_totp_backup_codes(code_hash); + +-- credential_id: base64 text +-- public_key: base64 text +-- aaguid: base64 text +-- transports: JSON-encoded array + CREATE TABLE IF NOT EXISTS user_passkey_credentials ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, - credential_id TEXT NOT NULL UNIQUE, -- base64 text (Direct mode), not native bytea - public_key TEXT NOT NULL, -- base64 text + credential_id TEXT NOT NULL UNIQUE, + public_key TEXT NOT NULL, attestation_type VARCHAR(50) DEFAULT 'none', - aaguid TEXT, -- base64 text + aaguid TEXT, sign_count INTEGER DEFAULT 0, clone_warning BOOLEAN DEFAULT 0, - transports TEXT, -- JSON-encoded []string + transports TEXT, backup_eligible BOOLEAN DEFAULT 0, backup_state BOOLEAN DEFAULT 0, name VARCHAR(255), @@ -81,7 +96,9 @@ CREATE TABLE IF NOT EXISTS user_passkey_credentials ( ); CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id); -CREATE INDEX IF NOT EXISTS idx_passkey_credential_id ON user_passkey_credentials(credential_id); + + +-- token_hash: SHA-256 hex of the raw token CREATE TABLE IF NOT EXISTS user_password_resets ( id INTEGER PRIMARY KEY AUTOINCREMENT, @@ -93,19 +110,32 @@ CREATE TABLE IF NOT EXISTS user_password_resets ( used_at TIMESTAMP ); +CREATE INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id); + +CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at); + + +-- redirect_uris: JSON-encoded array +-- grant_types: JSON-encoded array +-- allowed_scopes: JSON-encoded array +-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients + CREATE TABLE IF NOT EXISTS oauth_clients ( id INTEGER PRIMARY KEY AUTOINCREMENT, client_id VARCHAR(255) NOT NULL UNIQUE, - redirect_uris TEXT NOT NULL, -- JSON-encoded []string + redirect_uris TEXT NOT NULL, client_name VARCHAR(255), - grant_types TEXT, -- JSON-encoded []string - allowed_scopes TEXT, -- JSON-encoded []string - client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients + grant_types TEXT, + allowed_scopes TEXT, + client_secret_hash TEXT, token_endpoint_auth_method VARCHAR(30) DEFAULT 'none', is_active BOOLEAN DEFAULT 1, created_at TIMESTAMP ); + +-- scopes: JSON-encoded array + CREATE TABLE IF NOT EXISTS oauth_codes ( id INTEGER PRIMARY KEY AUTOINCREMENT, code VARCHAR(255) NOT NULL UNIQUE, @@ -116,28 +146,83 @@ CREATE TABLE IF NOT EXISTS oauth_codes ( code_challenge_method VARCHAR(10) DEFAULT 'S256', session_token TEXT NOT NULL, refresh_token TEXT, - scopes TEXT, -- JSON-encoded []string + scopes TEXT, expires_at TIMESTAMP NOT NULL, created_at TIMESTAMP ); -CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code); CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at); + +-- key_hash: SHA-256 hex +-- scopes: JSON-encoded array +-- meta: JSON-encoded object + CREATE TABLE IF NOT EXISTS user_keys ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, key_type VARCHAR(50) NOT NULL, key_hash VARCHAR(64) NOT NULL UNIQUE, name VARCHAR(255) NOT NULL DEFAULT '', - scopes TEXT, -- JSON-encoded []string - meta TEXT, -- JSON-encoded map + scopes TEXT, + meta TEXT, expires_at TIMESTAMP, created_at TIMESTAMP, last_used_at TIMESTAMP, is_active BOOLEAN DEFAULT 1 ); -CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id); -CREATE INDEX IF NOT EXISTS idx_user_keys_key_hash ON user_keys(key_hash); +CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id); + CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type); + + +-- Optional: omit to use per-user rules only. + +CREATE TABLE IF NOT EXISTS sec_group_members ( + group_id INTEGER NOT NULL, + user_id INTEGER NOT NULL, + PRIMARY KEY (group_id, user_id) +); + + +-- column_path: dot path under the table: col or col.sub.field +-- access_type: mask, hide, read, ... +-- extra_filters: JSON object + +CREATE TABLE IF NOT EXISTS sec_column_rules ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + column_path TEXT NOT NULL, + access_type VARCHAR(50) NOT NULL, + mask_start INTEGER DEFAULT 0, + mask_end INTEGER DEFAULT 0, + mask_invert BOOLEAN DEFAULT 0, + mask_char VARCHAR(10) DEFAULT '*', + extra_filters TEXT, + is_active BOOLEAN NOT NULL DEFAULT 1, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name); + + +-- template: SQL fragment, e.g. user_id = {UserID} + +CREATE TABLE IF NOT EXISTS sec_row_rules ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER, + group_id INTEGER, + schema_name TEXT NOT NULL, + table_name TEXT NOT NULL, + template TEXT, + has_block BOOLEAN NOT NULL DEFAULT 0, + is_active BOOLEAN NOT NULL DEFAULT 1, + CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)) +); + +CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name); + diff --git a/pkg/security/lookup/deps_test.go b/pkg/security/lookup/deps_test.go new file mode 100644 index 0000000..4fad100 --- /dev/null +++ b/pkg/security/lookup/deps_test.go @@ -0,0 +1,34 @@ +package lookup_test + +import ( + "os/exec" + "strings" + "testing" +) + +// lookup and its backends sit above sectypes and below security: they must not import the +// core package (security imports lookup) or any sibling sub package. +func TestNoUpwardImports(t *testing.T) { + for _, pkg := range []string{".", "./procedure", "./direct", "./conformance"} { + checkNoUpwardImports(t, pkg) + } +} + +func checkNoUpwardImports(t *testing.T, pkg string) { + t.Helper() + out, err := exec.Command("go", "list", "-deps", "-f", "{{.ImportPath}}", pkg).Output() + if err != nil { + t.Skipf("go list unavailable: %v", err) + } + for _, p := range strings.Fields(string(out)) { + if strings.Contains(p, "uptrace/bun") || strings.Contains(p, "gorm.io") { + t.Errorf("%s must not depend on an ORM, found %s", pkg, p) + } + if !strings.Contains(p, "/pkg/security") || strings.HasSuffix(p, "/lookup") || + strings.HasSuffix(p, "/sectypes") || strings.HasSuffix(p, "/lookup/dialect") || + strings.HasSuffix(p, "/lookup/procedure") || strings.HasSuffix(p, "/lookup/direct") || strings.HasSuffix(p, "/lookup/conformance") { + continue + } + t.Errorf("%s must not import %s", pkg, p) + } +} diff --git a/pkg/security/lookup/dialect/builtin.go b/pkg/security/lookup/dialect/builtin.go new file mode 100644 index 0000000..9a61cc0 --- /dev/null +++ b/pkg/security/lookup/dialect/builtin.go @@ -0,0 +1,93 @@ +package dialect + +import ( + "strconv" + "strings" + "time" +) + +// --- postgres --------------------------------------------------------------- + +type postgres struct{} + +func (postgres) Name() string { return "postgres" } +func (postgres) Matches(driver string) bool { + return strings.Contains(driver, "pgx") || strings.Contains(driver, "lib/pq") || + strings.Contains(driver, "postgres") +} +func (postgres) Placeholder(n int) string { return "$" + strconv.Itoa(n) } +func (postgres) Quote(ident string) string { return quoteWith(ident, `"`, `"`) } +func (postgres) Bool(v bool) any { return v } +func (postgres) ScanBool(src any) (bool, error) { return scanBool(src) } +func (postgres) ScanTime(src any) (time.Time, error) { return scanTime(src) } +func (postgres) EncodeJSON(v any) (any, error) { return encodeJSON(v) } +func (postgres) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) } +func (d postgres) InsertReturningID(table string, cols []string, idCol string) Insert { + return Insert{SQL: insertSQL(d, table, cols, "", "RETURNING "+d.Quote(idCol), "DEFAULT VALUES"), Strategy: ReturningQuery} +} + +// --- sqlite ----------------------------------------------------------------- + +type sqlite struct{} + +func (sqlite) Name() string { return "sqlite" } +func (sqlite) Matches(driver string) bool { + return strings.Contains(driver, "sqlite") +} +func (sqlite) Placeholder(int) string { return "?" } +func (sqlite) Quote(ident string) string { return quoteWith(ident, `"`, `"`) } +func (sqlite) Bool(v bool) any { return boolInt(v) } +func (sqlite) ScanBool(src any) (bool, error) { return scanBool(src) } +func (sqlite) ScanTime(src any) (time.Time, error) { return scanTime(src) } +func (sqlite) EncodeJSON(v any) (any, error) { return encodeJSON(v) } +func (sqlite) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) } +func (d sqlite) InsertReturningID(table string, cols []string, _ string) Insert { + return Insert{SQL: insertSQL(d, table, cols, "", "", "DEFAULT VALUES"), Strategy: LastInsertID} +} + +// --- mysql / mariadb ---------------------------------------------------------- + +type mysql struct{} + +func (mysql) Name() string { return "mysql" } +func (mysql) Matches(driver string) bool { + return strings.Contains(driver, "mysql") || strings.Contains(driver, "mariadb") +} +func (mysql) Placeholder(int) string { return "?" } +func (mysql) Quote(ident string) string { return quoteWith(ident, "`", "`") } +func (mysql) Bool(v bool) any { return boolInt(v) } +func (mysql) ScanBool(src any) (bool, error) { return scanBool(src) } +func (mysql) ScanTime(src any) (time.Time, error) { return scanTime(src) } +func (mysql) EncodeJSON(v any) (any, error) { return encodeJSON(v) } +func (mysql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) } +func (d mysql) InsertReturningID(table string, cols []string, _ string) Insert { + // MySQL has no DEFAULT VALUES; an empty column list is spelled "() VALUES ()". + s := insertSQL(d, table, cols, "", "", "() VALUES ()") + return Insert{SQL: s, Strategy: LastInsertID} +} + +// --- mssql (SQL Server) -------------------------------------------------------- + +type mssql struct{} + +func (mssql) Name() string { return "mssql" } +func (mssql) Matches(driver string) bool { + return strings.Contains(driver, "mssql") || strings.Contains(driver, "sqlserver") +} +func (mssql) Placeholder(n int) string { return "@p" + strconv.Itoa(n) } +func (mssql) Quote(ident string) string { return quoteWith(ident, "[", "]") } +func (mssql) Bool(v bool) any { return v } +func (mssql) ScanBool(src any) (bool, error) { return scanBool(src) } +func (mssql) ScanTime(src any) (time.Time, error) { return scanTime(src) } +func (mssql) EncodeJSON(v any) (any, error) { return encodeJSON(v) } +func (mssql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) } +func (d mssql) InsertReturningID(table string, cols []string, idCol string) Insert { + return Insert{SQL: insertSQL(d, table, cols, "OUTPUT INSERTED."+d.Quote(idCol), "", "DEFAULT VALUES"), Strategy: ReturningQuery} +} + +func boolInt(v bool) int64 { + if v { + return 1 + } + return 0 +} diff --git a/pkg/security/lookup/dialect/dialect.go b/pkg/security/lookup/dialect/dialect.go new file mode 100644 index 0000000..0e9ed92 --- /dev/null +++ b/pkg/security/lookup/dialect/dialect.go @@ -0,0 +1,312 @@ +// Package dialect holds the per-database adaptors used by the lookup direct backend. +// An adaptor supplies only what differs between databases (placeholders, quoting, +// booleans, time and JSON handling, insert-returning-id); the backend builds queries from it. +// It imports only the standard library. +package dialect + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "reflect" + "sort" + "strings" + "sync" + "time" +) + +// Dialect is the adaptor for one database type. +type Dialect interface { + // Name is the registry name, e.g. "postgres". + Name() string + // Matches reports whether a driver type identifier (lowercased ".") + // belongs to this database. Used by Detect. + Matches(driver string) bool + // Placeholder returns the bind placeholder for the n-th (1-based) argument. + Placeholder(n int) string + // Quote quotes an identifier. A dotted name is quoted per part ("schema.table"). + // Embedded quote characters are escaped, never interpolated raw. + Quote(ident string) string + // Bool converts a Go bool to the value bound as a boolean column argument. + Bool(v bool) any + // ScanBool reads a boolean column value that the driver returned as bool, integer, string or bytes. + ScanBool(src any) (bool, error) + // ScanTime reads a time column value that the driver returned as time.Time, string or bytes. + // NULL (nil) yields the zero time. + ScanTime(src any) (time.Time, error) + // EncodeJSON converts a value to the argument bound to a JSON/TEXT column. A nil + // value, map or slice yields nil (SQL NULL). + EncodeJSON(v any) (any, error) + // DecodeJSON reads a JSON/TEXT column value into dst. NULL and empty values leave dst untouched. + DecodeJSON(src any, dst any) error + // InsertReturningID builds an INSERT of cols into table and describes how to read the new id. + // Arguments are bound positionally in cols order. + InsertReturningID(table string, cols []string, idCol string) Insert +} + +// InsertStrategy tells how the generated id is read after an Insert. +type InsertStrategy int + +const ( + // ReturningQuery means the statement returns the id as a single row (QueryRow + Scan). + ReturningQuery InsertStrategy = iota + // LastInsertID means the id is read from sql.Result.LastInsertId after Exec. + LastInsertID +) + +// Insert is a generated INSERT statement and how to read the id it creates. +type Insert struct { + SQL string + Strategy InsertStrategy +} + +// Querier is implemented by *sql.DB, *sql.Tx and *sql.Conn. +type Querier interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row +} + +// Run executes the insert and returns the generated id. +func (i Insert) Run(ctx context.Context, q Querier, args ...any) (int64, error) { + switch i.Strategy { + case ReturningQuery: + var id int64 + if err := q.QueryRowContext(ctx, i.SQL, args...).Scan(&id); err != nil { + return 0, err + } + return id, nil + case LastInsertID: + res, err := q.ExecContext(ctx, i.SQL, args...) + if err != nil { + return 0, err + } + return res.LastInsertId() + } + return 0, fmt.Errorf("dialect: unknown insert strategy %d", i.Strategy) +} + +// Factory creates a Dialect. +type Factory func() Dialect + +var ( + regMu sync.RWMutex + registry = map[string]Factory{} +) + +// Register adds a dialect under name. Registering a name twice replaces it, so applications +// can override a built-in. Adding a database = implementing Dialect and calling Register. +func Register(name string, f Factory) { + regMu.Lock() + defer regMu.Unlock() + registry[strings.ToLower(name)] = f +} + +// Get returns the dialect registered under name. +func Get(name string) (Dialect, error) { + regMu.RLock() + f, ok := registry[strings.ToLower(name)] + regMu.RUnlock() + if !ok { + return nil, fmt.Errorf("dialect: unknown dialect %q (registered: %s)", name, strings.Join(Names(), ", ")) + } + return f(), nil +} + +// Names lists the registered dialect names, sorted. +func Names() []string { + regMu.RLock() + defer regMu.RUnlock() + names := make([]string, 0, len(registry)) + for n := range registry { + names = append(names, n) + } + sort.Strings(names) + return names +} + +// Detect picks the dialect for db from its driver type. +func Detect(db *sql.DB) (Dialect, error) { + if db == nil { + return nil, fmt.Errorf("dialect: nil database") + } + return DetectDriver(driverID(db.Driver())) +} + +// DetectDriver picks the dialect for a driver type identifier (see Dialect.Matches). +func DetectDriver(driver string) (Dialect, error) { + driver = strings.ToLower(driver) + for _, n := range Names() { + d, _ := Get(n) + if d != nil && d.Matches(driver) { + return d, nil + } + } + return nil, fmt.Errorf("dialect: cannot detect a dialect for driver %q; set the dialect explicitly", driver) +} + +// driverID builds "." for a driver value, lowercased. +func driverID(drv any) string { + t := reflect.TypeOf(drv) + for t != nil && t.Kind() == reflect.Pointer { + t = t.Elem() + } + if t == nil { + return "" + } + return strings.ToLower(t.PkgPath() + "." + t.Name()) +} + +func init() { + Register("postgres", func() Dialect { return postgres{} }) + Register("sqlite", func() Dialect { return sqlite{} }) + Register("mysql", func() Dialect { return mysql{} }) + Register("mssql", func() Dialect { return mssql{} }) +} + +// --- shared helpers ------------------------------------------------------- + +// quoteWith quotes each dotted part of ident with open/close, doubling embedded close characters. +func quoteWith(ident, open, closeq string) string { + parts := strings.Split(ident, ".") + for i, p := range parts { + parts[i] = open + strings.ReplaceAll(p, closeq, closeq+closeq) + closeq + } + return strings.Join(parts, ".") +} + +func scanBool(src any) (bool, error) { + switch v := src.(type) { + case nil: + return false, nil + case bool: + return v, nil + case int64: + return v != 0, nil + case int: + return v != 0, nil + case int32: + return v != 0, nil + case float64: + return v != 0, nil + case []byte: + return scanBool(string(v)) + case string: + switch strings.ToLower(strings.TrimSpace(v)) { + case "1", "t", "true", "y", "yes", "on": + return true, nil + case "", "0", "f", "false", "n", "no", "off": + return false, nil + } + } + return false, fmt.Errorf("dialect: cannot read %T (%v) as bool", src, src) +} + +var timeLayouts = []string{ + time.RFC3339Nano, + "2006-01-02 15:04:05.999999999 -0700", + "2006-01-02 15:04:05.999999999-07:00", + "2006-01-02 15:04:05.999999999Z07:00", + "2006-01-02T15:04:05.999999999", + "2006-01-02 15:04:05.999999999", + "2006-01-02", +} + +func scanTime(src any) (time.Time, error) { + switch v := src.(type) { + case nil: + return time.Time{}, nil + case time.Time: + return v, nil + case []byte: + return scanTime(string(v)) + case string: + s := strings.TrimSpace(v) + if s == "" { + return time.Time{}, nil + } + // Some drivers append the Go monotonic/zone suffix ("+0000 UTC"); drop it. + if i := strings.Index(s, " m="); i >= 0 { + s = s[:i] + } + s = strings.TrimSuffix(s, " UTC") + s = strings.TrimSuffix(s, " +0000 +0000") + for _, l := range timeLayouts { + if t, err := time.Parse(l, s); err == nil { + return t, nil + } + } + } + return time.Time{}, fmt.Errorf("dialect: cannot read %T (%v) as time", src, src) +} + +func encodeJSON(v any) (any, error) { + if v == nil { + return nil, nil + } + rv := reflect.ValueOf(v) + switch rv.Kind() { + case reflect.Map, reflect.Slice, reflect.Pointer, reflect.Interface: + if rv.IsNil() { + return nil, nil + } + } + b, err := json.Marshal(v) + if err != nil { + return nil, fmt.Errorf("dialect: encode json: %w", err) + } + return string(b), nil +} + +func decodeJSON(src any, dst any) error { + var raw []byte + switch v := src.(type) { + case nil: + return nil + case string: + raw = []byte(v) + case []byte: + raw = v + default: + return fmt.Errorf("dialect: cannot read %T as json", src) + } + if len(strings.TrimSpace(string(raw))) == 0 { + return nil + } + if err := json.Unmarshal(raw, dst); err != nil { + return fmt.Errorf("dialect: decode json: %w", err) + } + return nil +} + +// insertSQL assembles "INSERT INTO t (cols) VALUES (...) " for a dialect. +func insertSQL(d Dialect, table string, cols []string, mid, tail, defaults string) string { + var b strings.Builder + b.WriteString("INSERT INTO ") + b.WriteString(d.Quote(table)) + if len(cols) == 0 { + if mid != "" { + b.WriteString(" " + mid) + } + b.WriteString(" " + defaults) + if tail != "" { + b.WriteString(" " + tail) + } + return b.String() + } + qc := make([]string, len(cols)) + ph := make([]string, len(cols)) + for i, c := range cols { + qc[i] = d.Quote(c) + ph[i] = d.Placeholder(i + 1) + } + b.WriteString(" (" + strings.Join(qc, ", ") + ")") + if mid != "" { + b.WriteString(" " + mid) + } + b.WriteString(" VALUES (" + strings.Join(ph, ", ") + ")") + if tail != "" { + b.WriteString(" " + tail) + } + return b.String() +} diff --git a/pkg/security/lookup/dialect/dialect_test.go b/pkg/security/lookup/dialect/dialect_test.go new file mode 100644 index 0000000..eb69ce9 --- /dev/null +++ b/pkg/security/lookup/dialect/dialect_test.go @@ -0,0 +1,270 @@ +package dialect_test + +import ( + "context" + "database/sql" + "strings" + "testing" + "time" + + _ "github.com/glebarez/go-sqlite" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" +) + +func get(t *testing.T, name string) dialect.Dialect { + t.Helper() + d, err := dialect.Get(name) + if err != nil { + t.Fatal(err) + } + return d +} + +func TestRegistryHasBuiltins(t *testing.T) { + got := strings.Join(dialect.Names(), ",") + if got != "mssql,mysql,postgres,sqlite" { + t.Errorf("names = %s", got) + } + if _, err := dialect.Get("oracle"); err == nil { + t.Error("unknown dialect should error") + } +} + +func TestRegisterCustomDialect(t *testing.T) { + // A new database is added by registering a dialect; no core change needed. + dialect.Register("custom", func() dialect.Dialect { return customDialect{get(t, "sqlite")} }) + d, err := dialect.Get("CUSTOM") + if err != nil || d.Name() != "custom" { + t.Fatalf("got %v, %v", d, err) + } + if got, _ := dialect.DetectDriver("example.com/driver.customdriver"); got == nil || got.Name() != "custom" { + t.Errorf("custom detection failed: %v", got) + } +} + +type customDialect struct{ dialect.Dialect } + +func (customDialect) Name() string { return "custom" } +func (customDialect) Matches(driver string) bool { return strings.Contains(driver, "customdriver") } + +func TestPlaceholderAndQuote(t *testing.T) { + cases := []struct { + name, ph1, ph3, quote, quoteDotted, quoteEscape string + }{ + {"postgres", "$1", "$3", `"users"`, `"auth"."users"`, `"a""b"`}, + {"sqlite", "?", "?", `"users"`, `"auth"."users"`, `"a""b"`}, + {"mysql", "?", "?", "`users`", "`auth`.`users`", "`a``b`"}, + {"mssql", "@p1", "@p3", "[users]", "[auth].[users]", "[a]]b]"}, + } + for _, c := range cases { + d := get(t, c.name) + if d.Placeholder(1) != c.ph1 || d.Placeholder(3) != c.ph3 { + t.Errorf("%s placeholders: %s %s", c.name, d.Placeholder(1), d.Placeholder(3)) + } + if d.Quote("users") != c.quote { + t.Errorf("%s quote: %s", c.name, d.Quote("users")) + } + if d.Quote("auth.users") != c.quoteDotted { + t.Errorf("%s dotted: %s", c.name, d.Quote("auth.users")) + } + raw := map[string]string{"postgres": `a"b`, "sqlite": `a"b`, "mysql": "a`b", "mssql": "a]b"}[c.name] + q := c.quoteEscape + if got := d.Quote(raw); got != q { + t.Errorf("%s escape: got %s want %s", c.name, got, q) + } + } +} + +func TestInsertReturningID(t *testing.T) { + cols := []string{"user_id", "name"} + cases := []struct { + name string + wantSQL string + strategy dialect.InsertStrategy + empty string + }{ + {"postgres", `INSERT INTO "t" ("user_id", "name") VALUES ($1, $2) RETURNING "id"`, dialect.ReturningQuery, `INSERT INTO "t" DEFAULT VALUES RETURNING "id"`}, + {"sqlite", `INSERT INTO "t" ("user_id", "name") VALUES (?, ?)`, dialect.LastInsertID, `INSERT INTO "t" DEFAULT VALUES`}, + {"mysql", "INSERT INTO `t` (`user_id`, `name`) VALUES (?, ?)", dialect.LastInsertID, "INSERT INTO `t` () VALUES ()"}, + {"mssql", `INSERT INTO [t] ([user_id], [name]) OUTPUT INSERTED.[id] VALUES (@p1, @p2)`, dialect.ReturningQuery, `INSERT INTO [t] OUTPUT INSERTED.[id] DEFAULT VALUES`}, + } + for _, c := range cases { + ins := get(t, c.name).InsertReturningID("t", cols, "id") + if ins.SQL != c.wantSQL || ins.Strategy != c.strategy { + t.Errorf("%s: got %q (%d)", c.name, ins.SQL, ins.Strategy) + } + if e := get(t, c.name).InsertReturningID("t", nil, "id"); e.SQL != c.empty { + t.Errorf("%s empty: got %q", c.name, e.SQL) + } + } +} + +func TestBoolRoundTrip(t *testing.T) { + for _, name := range dialect.Names() { + if name == "custom" { + continue + } + d := get(t, name) + for _, v := range []bool{true, false} { + got, err := d.ScanBool(d.Bool(v)) + if err != nil || got != v { + t.Errorf("%s: Bool(%v) round trip = %v, %v", name, v, got, err) + } + } + for _, c := range []struct { + in any + want bool + }{{int64(1), true}, {int64(0), false}, {"1", true}, {"false", false}, {[]byte("t"), true}, {nil, false}} { + if got, err := d.ScanBool(c.in); err != nil || got != c.want { + t.Errorf("%s: ScanBool(%v) = %v, %v", name, c.in, got, err) + } + } + if _, err := d.ScanBool("maybe"); err == nil { + t.Errorf("%s: ScanBool(maybe) should fail", name) + } + } +} + +func TestScanTime(t *testing.T) { + d := get(t, "sqlite") + want := time.Date(2026, 3, 4, 5, 6, 7, 0, time.UTC) + for _, in := range []any{ + want, + "2026-03-04 05:06:07+00:00", + "2026-03-04T05:06:07Z", + "2026-03-04 05:06:07", + []byte("2026-03-04 05:06:07"), + "2026-03-04 05:06:07 +0000 UTC", + } { + got, err := d.ScanTime(in) + if err != nil || !got.Equal(want) { + t.Errorf("ScanTime(%v) = %v, %v", in, got, err) + } + } + if got, err := d.ScanTime(nil); err != nil || !got.IsZero() { + t.Errorf("ScanTime(nil) = %v, %v", got, err) + } + if _, err := d.ScanTime("not a time"); err == nil { + t.Error("ScanTime(garbage) should fail") + } +} + +func TestJSON(t *testing.T) { + for _, name := range []string{"postgres", "sqlite", "mysql", "mssql"} { + d := get(t, name) + enc, err := d.EncodeJSON([]string{"a", "b"}) + if err != nil || enc != `["a","b"]` { + t.Errorf("%s encode = %v, %v", name, enc, err) + } + var nilSlice []string + var nilMap map[string]any + for _, v := range []any{nil, nilSlice, nilMap} { + if enc, err := d.EncodeJSON(v); err != nil || enc != nil { + t.Errorf("%s: EncodeJSON(%T nil) = %v, %v; want SQL NULL", name, v, enc, err) + } + } + var out []string + if err := d.DecodeJSON([]byte(`["x"]`), &out); err != nil || len(out) != 1 || out[0] != "x" { + t.Errorf("%s decode bytes = %v, %v", name, out, err) + } + out = []string{"keep"} + if err := d.DecodeJSON(nil, &out); err != nil || out[0] != "keep" { + t.Errorf("%s decode NULL should leave dst: %v, %v", name, out, err) + } + if err := d.DecodeJSON("", &out); err != nil || out[0] != "keep" { + t.Errorf("%s decode empty should leave dst: %v, %v", name, out, err) + } + if err := d.DecodeJSON("{bad", &out); err == nil { + t.Errorf("%s decode of invalid json should fail", name) + } + } +} + +func TestDetectDriver(t *testing.T) { + cases := map[string]string{ + "github.com/jackc/pgx/v5/stdlib.driver": "postgres", + "github.com/lib/pq.driver": "postgres", + "github.com/mattn/go-sqlite3.sqlitedriver": "sqlite", + "modernc.org/sqlite.driver": "sqlite", + "github.com/go-sql-driver/mysql.mysqldriver": "mysql", + "github.com/microsoft/go-mssqldb.driver": "mssql", + "github.com/denisenkom/go-mssqldb.driver": "mssql", + } + for drv, want := range cases { + d, err := dialect.DetectDriver(drv) + if err != nil || d.Name() != want { + t.Errorf("%s: got %v, %v; want %s", drv, d, err, want) + } + } + if _, err := dialect.DetectDriver("example.com/unknown.driver"); err == nil { + t.Error("unknown driver should fail with an explicit-dialect hint") + } + if _, err := dialect.Detect(nil); err == nil { + t.Error("nil db should fail") + } +} + +// TestSQLiteRoundTrip runs the sqlite dialect against a real in-memory database. +func TestSQLiteRoundTrip(t *testing.T) { + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + + d, err := dialect.Detect(db) + if err != nil || d.Name() != "sqlite" { + t.Fatalf("Detect = %v, %v", d, err) + } + ctx := context.Background() + if _, err := db.ExecContext(ctx, `CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT, active BOOLEAN, scopes TEXT, at TIMESTAMP)`); err != nil { + t.Fatal(err) + } + + now := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + scopes, _ := d.EncodeJSON([]string{"read", "write"}) + ins := d.InsertReturningID("t", []string{"name", "active", "scopes", "at"}, "id") + id, err := ins.Run(ctx, db, "k1", d.Bool(true), scopes, now) + if err != nil || id != 1 { + t.Fatalf("insert = %d, %v", id, err) + } + id, err = ins.Run(ctx, db, "k2", d.Bool(false), nil, now) + if err != nil || id != 2 { + t.Fatalf("second insert = %d, %v", id, err) + } + + q := "SELECT active, scopes, at FROM " + d.Quote("t") + " WHERE " + d.Quote("id") + " = " + d.Placeholder(1) + var active, scopesRaw, at any + if err := db.QueryRowContext(ctx, q, 1).Scan(&active, &scopesRaw, &at); err != nil { + t.Fatal(err) + } + if b, err := d.ScanBool(active); err != nil || !b { + t.Errorf("active = %v, %v", b, err) + } + var got []string + if err := d.DecodeJSON(scopesRaw, &got); err != nil || len(got) != 2 { + t.Errorf("scopes = %v, %v", got, err) + } + if ts, err := d.ScanTime(at); err != nil || !ts.Equal(now) { + t.Errorf("at = %v (%T), %v", ts, at, err) + } + + if err := db.QueryRowContext(ctx, q, 2).Scan(&active, &scopesRaw, &at); err != nil { + t.Fatal(err) + } + if b, _ := d.ScanBool(active); b { + t.Error("second row should be inactive") + } + if scopesRaw != nil { + t.Errorf("nil scopes should be NULL, got %v", scopesRaw) + } + + // Insert inside a transaction goes through the same Querier. + tx, _ := db.BeginTx(ctx, nil) + if id, err := d.InsertReturningID("t", nil, "id").Run(ctx, tx); err != nil || id != 3 { + t.Errorf("DEFAULT VALUES insert in tx = %d, %v", id, err) + } + _ = tx.Rollback() +} diff --git a/pkg/security/lookup/direct/auth.go b/pkg/security/lookup/direct/auth.go new file mode 100644 index 0000000..db36958 --- /dev/null +++ b/pkg/security/lookup/direct/auth.go @@ -0,0 +1,608 @@ +package direct + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "strings" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +const sessionLifetime = 24 * time.Hour + +// AuthOptions tunes Auth. +type AuthOptions struct { + // UpgradePasswordHash rewrites legacy cleartext passwords as bcrypt on a successful login. + UpgradePasswordHash bool +} + +// Auth implements lookup.AuthStore on the tables. Passwords are verified with bcrypt; +// legacy cleartext rows are accepted at login and only rewritten when UpgradePasswordHash +// is set. Registration never honours client-supplied user_level or roles. Multi-step writes +// (login, register, refresh, reset) run in one transaction. +type Auth struct { + *Base + opts AuthOptions +} + +var _ lookup.AuthStore = (*Auth)(nil) + +// NewAuth creates the direct AuthStore. +func NewAuth(b *Base, opts AuthOptions) *Auth { return &Auth{Base: b, opts: opts} } + +// GenerateSessionToken returns "sess_<64 hex>_". +func GenerateSessionToken() (string, error) { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return fmt.Sprintf("sess_%s_%d", hex.EncodeToString(buf), time.Now().Unix()), nil +} + +// ParseRoles splits the comma-separated roles column. +func ParseRoles(s string) []string { + if s == "" { + return []string{} + } + return strings.Split(s, ",") +} + +func claimStrings(claims map[string]any) (ip, ua string) { + if claims == nil { + return "", "" + } + if v, ok := claims["ip_address"].(string); ok { + ip = v + } + if v, ok := claims["user_agent"].(string); ok { + ua = v + } + return ip, ua +} + +func sha256Hex(s string) string { + h := sha256.Sum256([]byte(s)) + return hex.EncodeToString(h[:]) +} + +// userRow is the users columns every session-bearing response needs. +type userRow struct { + id int + username sql.NullString + email sql.NullString + roles sql.NullString + programUserTable sql.NullString + userLevel sql.NullInt64 + programUserID sql.NullInt64 +} + +func (u *userRow) context(sessionID string) *sectypes.UserContext { + return §ypes.UserContext{ + UserID: u.id, + UserName: u.username.String, + Email: u.email.String, + UserLevel: int(u.userLevel.Int64), + SessionID: sessionID, + Roles: ParseRoles(u.roles.String), + ProgramUserID: int(u.programUserID.Int64), + ProgramUserTable: u.programUserTable.String, + } +} + +// insertSession writes a session row and stamps the user's last login. +func (a *Auth) insertSession(ctx context.Context, q Querier, token string, userID int64, expiresAt time.Time, ip, ua string, now time.Time) error { + err := a.Insert(lookup.EntityUserSessions).Set( + Set(lookup.SessionsToken, token), + Set(lookup.SessionsUserID, userID), + Set(lookup.SessionsExpiresAt, expiresAt), + Set(lookup.SessionsIPAddress, ip), + Set(lookup.SessionsUserAgent, ua), + Set(lookup.SessionsLastActivityAt, now), + Set(lookup.SessionsCreatedAt, now), + ).Exec(ctx, q) + if err != nil { + return err + } + return a.touchLastLogin(ctx, q, userID, now) +} + +func (a *Auth) touchLastLogin(ctx context.Context, q Querier, userID int64, now time.Time) error { + _, err := a.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + return err +} + +// Login implements lookup.AuthStore. +func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) { + var userID int + var email, roles, programUserTable, storedPassword sql.NullString + var userLevel, programUserID sql.NullInt64 + + err := a.do(func(q Querier) error { + return a.From(lookup.EntityUsers). + Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, + lookup.UsersProgramUserID, lookup.UsersProgramUserTable, lookup.UsersPassword). + Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + BurnPasswordCheck(req.Password) + return nil, fmt.Errorf("invalid credentials") + } + return nil, fmt.Errorf("login query failed: %w", err) + } + + ok, needsRehash := VerifyPassword(storedPassword.String, req.Password) + if !ok { + if storedPassword.String == "" { + BurnPasswordCheck(req.Password) + } + return nil, fmt.Errorf("invalid credentials") + } + if needsRehash && a.opts.UpgradePasswordHash { + a.upgradePasswordHash(ctx, userID, req.Password) + } + + token, err := GenerateSessionToken() + if err != nil { + return nil, fmt.Errorf("failed to generate session token: %w", err) + } + now := a.Now() + ip, ua := claimStrings(req.Claims) + err = a.tx(ctx, func(q Querier) error { + return a.insertSession(ctx, q, token, int64(userID), now.Add(sessionLifetime), ip, ua, now) + }) + if err != nil { + return nil, fmt.Errorf("login query failed: %w", err) + } + + return §ypes.LoginResponse{ + Token: token, + User: §ypes.UserContext{ + UserID: userID, + UserName: req.Username, + Email: email.String, + UserLevel: int(userLevel.Int64), + Roles: ParseRoles(roles.String), + SessionID: token, + ProgramUserID: int(programUserID.Int64), + ProgramUserTable: programUserTable.String, + }, + ExpiresIn: int64(sessionLifetime.Seconds()), + }, nil +} + +// upgradePasswordHash replaces a legacy cleartext password with a bcrypt hash. Failure is +// logged and ignored: the login itself already succeeded. +func (a *Auth) upgradePasswordHash(ctx context.Context, userID int, password string) { + h, err := HashPassword(password) + if err != nil { + return + } + err = a.do(func(q Querier) error { + _, err := a.Update(lookup.EntityUsers). + Set(Set(lookup.UsersPassword, h), Set(lookup.UsersUpdatedAt, a.Now())). + Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + return err + }) + if err != nil { + logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err) + } +} + +// Register implements lookup.AuthStore. +func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) { + if req.Username == "" { + return nil, fmt.Errorf("username is required") + } + if req.Email == "" { + return nil, fmt.Errorf("email is required") + } + if req.Password == "" { + return nil, fmt.Errorf("password is required") + } + hash, err := HashPassword(req.Password) + if err != nil { + return nil, err + } + token, err := GenerateSessionToken() + if err != nil { + return nil, fmt.Errorf("failed to generate session token: %w", err) + } + + // Privileges are never taken from the request: self-registration always creates an + // unprivileged user. + const userLevel = 0 + now := a.Now() + ip, ua := claimStrings(req.Claims) + + var userID int64 + err = a.tx(ctx, func(q Querier) error { + exists, err := a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersUsername, req.Username)).Exists(ctx, q) + if err != nil { + return err + } + if exists { + return lookup.ErrUsernameExists + } + exists, err = a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersEmail, req.Email)).Exists(ctx, q) + if err != nil { + return err + } + if exists { + return lookup.ErrEmailExists + } + userID, err = a.Insert(lookup.EntityUsers).Set( + Set(lookup.UsersUsername, req.Username), + Set(lookup.UsersEmail, req.Email), + Set(lookup.UsersPassword, hash), + Set(lookup.UsersUserLevel, userLevel), + Set(lookup.UsersRoles, ""), + Set(lookup.UsersIsActive, true), + Set(lookup.UsersCreatedAt, now), + Set(lookup.UsersUpdatedAt, now), + Set(lookup.UsersProgramUserID, 0), + Set(lookup.UsersProgramUserTable, ""), + ).ExecID(ctx, q, lookup.UsersID) + if err != nil { + return err + } + return a.insertSession(ctx, q, token, userID, now.Add(sessionLifetime), ip, ua, now) + }) + if err != nil { + if errors.Is(err, lookup.ErrUsernameExists) || errors.Is(err, lookup.ErrEmailExists) { + return nil, err + } + return nil, fmt.Errorf("register query failed: %w", err) + } + + return §ypes.LoginResponse{ + Token: token, + User: §ypes.UserContext{ + UserID: int(userID), + UserName: req.Username, + Email: req.Email, + UserLevel: userLevel, + Roles: ParseRoles(""), + SessionID: token, + }, + ExpiresIn: int64(sessionLifetime.Seconds()), + }, nil +} + +// Logout implements lookup.AuthStore. +func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error { + token := strings.TrimPrefix(strings.TrimPrefix(req.Token, "Bearer "), "bearer ") + var rows int64 + err := a.do(func(q Querier) error { + var err error + rows, err = a.Delete(lookup.EntityUserSessions). + Where(Eq(lookup.SessionsToken, token), Eq(lookup.SessionsUserID, req.UserID)).Exec(ctx, q) + return err + }) + if err != nil { + return fmt.Errorf("logout query failed: %w", err) + } + if rows == 0 { + return fmt.Errorf("session not found") + } + return nil +} + +// sessionUser selects the user behind a live session token. +func (a *Auth) sessionUser(ctx context.Context, q Querier, token string, extra ...lookup.Column) (*userRow, []any, error) { + var u userRow + dest := []any{&u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable} + cols := []lookup.Column{lookup.SessionsUserID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, + lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable} + extras := make([]any, len(extra)) + for i, c := range extra { + cols = append(cols, c) + extras[i] = new(sql.NullString) + dest = append(dest, extras[i]) + } + err := a.From(lookup.EntityUserSessions).Cols(cols...). + Join(lookup.EntityUsers, EqCol(lookup.SessionsUserID, lookup.UsersID)). + Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, a.Now()), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, dest...) + return &u, extras, err +} + +// Session implements lookup.AuthStore. reference is only meaningful to the procedure backend. +func (a *Auth) Session(ctx context.Context, token, _ string) (*sectypes.UserContext, error) { + var u *userRow + err := a.do(func(q Querier) error { + var err error + u, _, err = a.sessionUser(ctx, q, token) + return err + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("invalid or expired session") + } + return nil, fmt.Errorf("session query failed: %w", err) + } + return u.context(token), nil +} + +// TouchSession implements lookup.AuthStore. +func (a *Auth) TouchSession(ctx context.Context, token string, _ *sectypes.UserContext) error { + return a.do(func(q Querier) error { + now := a.Now() + _, err := a.Update(lookup.EntityUserSessions).Set(Set(lookup.SessionsLastActivityAt, now)). + Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, now)).Exec(ctx, q) + return err + }) +} + +// Refresh implements lookup.AuthStore: the old session is replaced by a new one. +func (a *Auth) Refresh(ctx context.Context, oldToken string) (*sectypes.LoginResponse, error) { + var u *userRow + var extras []any + err := a.do(func(q Querier) error { + var err error + u, extras, err = a.sessionUser(ctx, q, oldToken, lookup.SessionsIPAddress, lookup.SessionsUserAgent) + return err + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("invalid or expired refresh token") + } + return nil, fmt.Errorf("refresh token query failed: %w", err) + } + ip := extras[0].(*sql.NullString).String + ua := extras[1].(*sql.NullString).String + + newToken, err := GenerateSessionToken() + if err != nil { + return nil, fmt.Errorf("failed to generate session token: %w", err) + } + now := a.Now() + err = a.tx(ctx, func(q Querier) error { + err := a.Insert(lookup.EntityUserSessions).Set( + Set(lookup.SessionsToken, newToken), + Set(lookup.SessionsUserID, u.id), + Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)), + Set(lookup.SessionsIPAddress, ip), + Set(lookup.SessionsUserAgent, ua), + Set(lookup.SessionsLastActivityAt, now), + Set(lookup.SessionsCreatedAt, now), + ).Exec(ctx, q) + if err != nil { + return err + } + _, err = a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, oldToken)).Exec(ctx, q) + return err + }) + if err != nil { + return nil, fmt.Errorf("refresh token generation failed: %w", err) + } + return §ypes.LoginResponse{ + Token: newToken, + User: u.context(newToken), + ExpiresIn: int64(sessionLifetime.Seconds()), + }, nil +} + +// apiKeyTypes are the key types accepted by LoginAPIKey. +var apiKeyTypes = []any{string(sectypes.KeyTypeHeaderAPI), string(sectypes.KeyTypeGenericAPI)} + +// LoginAPIKey implements lookup.AuthStore. Unknown, expired, inactive and wrong-type keys +// (and inactive users) 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 + } + now := a.Now() + var keyID int64 + var u userRow + err := a.do(func(q Querier) error { + return a.From(lookup.EntityUserKeys). + Cols(lookup.KeysID, lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, + lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable). + Join(lookup.EntityUsers, EqCol(lookup.KeysUserID, lookup.UsersID)). + Where( + Eq(lookup.KeysKeyHash, sectypes.HashKey(rawKey)), + In(lookup.KeysKeyType, apiKeyTypes...), + Eq(lookup.KeysIsActive, true), + Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, now)), + Eq(lookup.UsersIsActive, true), + ).QueryRow(ctx, q, &keyID, &u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, lookup.ErrInvalidAPIKey + } + return nil, fmt.Errorf("api key login query failed: %w", err) + } + + token, err := GenerateSessionToken() + if err != nil { + return nil, fmt.Errorf("failed to generate session token: %w", err) + } + ip, ua := claimStrings(claims) + err = a.tx(ctx, func(q Querier) error { + if err := a.insertSession(ctx, q, token, int64(u.id), now.Add(sessionLifetime), ip, ua, now); err != nil { + return err + } + _, err := a.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, keyID)).Exec(ctx, q) + return err + }) + if err != nil { + return nil, fmt.Errorf("api key login query failed: %w", err) + } + return §ypes.LoginResponse{ + Token: token, + User: u.context(token), + ExpiresIn: int64(sessionLifetime.Seconds()), + }, nil +} + +// JWTLogin implements lookup.AuthStore (mirrors resolvespec_jwt_login). 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 userID int + var email, roles, storedPassword sql.NullString + var userLevel sql.NullInt64 + err := a.do(func(q Querier) error { + return a.From(lookup.EntityUsers). + Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, lookup.UsersPassword). + Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &storedPassword) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + BurnPasswordCheck(req.Password) + return nil, fmt.Errorf("invalid credentials") + } + return nil, fmt.Errorf("login query failed: %w", err) + } + ok, needsRehash := VerifyPassword(storedPassword.String, req.Password) + if !ok { + if storedPassword.String == "" { + BurnPasswordCheck(req.Password) + } + return nil, fmt.Errorf("invalid credentials") + } + if needsRehash && a.opts.UpgradePasswordHash { + a.upgradePasswordHash(ctx, userID, req.Password) + } + expiresAt := a.Now().Add(sessionLifetime) + return §ypes.LoginResponse{ + Token: fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()), + User: §ypes.UserContext{ + UserID: userID, + UserName: req.Username, + Email: email.String, + UserLevel: int(userLevel.Int64), + Roles: ParseRoles(roles.String), + }, + ExpiresIn: int64(sessionLifetime.Seconds()), + }, nil +} + +// JWTLogout implements lookup.AuthStore: the token goes on the blacklist. +func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error { + now := a.Now() + err := a.do(func(q Querier) error { + return a.Insert(lookup.EntityTokenBlacklist).Set( + Set(lookup.BlacklistToken, req.Token), + Set(lookup.BlacklistUserID, req.UserID), + Set(lookup.BlacklistExpiresAt, now.Add(sessionLifetime)), + Set(lookup.BlacklistCreatedAt, now), + ).Exec(ctx, q) + }) + if err != nil { + return fmt.Errorf("logout query failed: %w", err) + } + return nil +} + +// ResetRequest implements lookup.AuthStore. An unknown user yields a generic empty success +// so accounts cannot be enumerated. +func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) { + if req.Email == "" && req.Username == "" { + return nil, fmt.Errorf("email or username is required") + } + var userID int + err := a.do(func(q Querier) error { + lookupCol, val := lookup.UsersUsername, req.Username + if req.Email != "" { + lookupCol, val = lookup.UsersEmail, req.Email + } + return a.From(lookup.EntityUsers).Cols(lookup.UsersID). + Where(Eq(lookupCol, val), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return §ypes.PasswordResetResponse{Token: "", ExpiresIn: 0}, nil + } + return nil, fmt.Errorf("password reset request query failed: %w", err) + } + + raw := make([]byte, 32) + if _, err := rand.Read(raw); err != nil { + return nil, fmt.Errorf("failed to generate reset token: %w", err) + } + rawToken := hex.EncodeToString(raw) + now := a.Now() + err = a.tx(ctx, func(q Querier) error { + if _, err := a.Delete(lookup.EntityUserPasswordResets). + Where(Eq(lookup.ResetsUserID, userID), Eq(lookup.ResetsUsed, false)).Exec(ctx, q); err != nil { + return err + } + return a.Insert(lookup.EntityUserPasswordResets).Set( + Set(lookup.ResetsUserID, userID), + Set(lookup.ResetsTokenHash, sha256Hex(rawToken)), + Set(lookup.ResetsExpiresAt, now.Add(time.Hour)), + Set(lookup.ResetsCreatedAt, now), + Set(lookup.ResetsUsed, false), + ).Exec(ctx, q) + }) + if err != nil { + return nil, fmt.Errorf("password reset request query failed: %w", err) + } + return §ypes.PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil +} + +// ResetComplete implements lookup.AuthStore: sets the new password, ends every session of +// the user and consumes the reset token, atomically. +func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error { + if req.Token == "" { + return fmt.Errorf("token is required") + } + if req.NewPassword == "" { + return fmt.Errorf("new_password is required") + } + newHash, err := HashPassword(req.NewPassword) + if err != nil { + return err + } + tokenHash := sha256Hex(req.Token) + + now := a.Now() + var resetID, userID int + var expiresAt time.Time + err = a.do(func(q Querier) error { + return a.From(lookup.EntityUserPasswordResets). + Cols(lookup.ResetsID, lookup.ResetsUserID, lookup.ResetsExpiresAt). + Where(Eq(lookup.ResetsTokenHash, tokenHash), Eq(lookup.ResetsUsed, false)). + QueryRow(ctx, q, &resetID, &userID, a.timeDest(&expiresAt)) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("invalid or expired token") + } + return fmt.Errorf("password reset complete query failed: %w", err) + } + if !expiresAt.After(now) { + return fmt.Errorf("invalid or expired token") + } + + err = a.tx(ctx, func(q Querier) error { + if _, err := a.Update(lookup.EntityUsers). + Set(Set(lookup.UsersPassword, newHash), Set(lookup.UsersUpdatedAt, now)). + Where(Eq(lookup.UsersID, userID)).Exec(ctx, q); err != nil { + return err + } + if _, err := a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsUserID, userID)).Exec(ctx, q); err != nil { + return err + } + _, err := a.Update(lookup.EntityUserPasswordResets). + Set(Set(lookup.ResetsUsed, true), Set(lookup.ResetsUsedAt, now)). + Where(Eq(lookup.ResetsID, resetID)).Exec(ctx, q) + return err + }) + if err != nil { + return fmt.Errorf("password reset complete query failed: %w", err) + } + return nil +} diff --git a/pkg/security/lookup/direct/auth_test.go b/pkg/security/lookup/direct/auth_test.go new file mode 100644 index 0000000..c5bd38a --- /dev/null +++ b/pkg/security/lookup/direct/auth_test.go @@ -0,0 +1,268 @@ +package direct + +import ( + "context" + "database/sql" + "errors" + "strings" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +func newAuth(t *testing.T, opts AuthOptions) (*Auth, *sql.DB) { + db := newTestDB(t) + return NewAuth(newTestBase(t, db, nil), opts), db +} + +func registerUser(t *testing.T, a *Auth, name string) *sectypes.LoginResponse { + t.Helper() + resp, err := a.Register(context.Background(), sectypes.RegisterRequest{Username: name, Email: name + "@x.io", Password: "pw-" + name}) + if err != nil { + t.Fatal(err) + } + return resp +} + +func TestRegisterLoginSessionFlow(t *testing.T) { + ctx := context.Background() + a, db := newAuth(t, AuthOptions{}) + + reg, err := a.Register(ctx, sectypes.RegisterRequest{ + Username: "ann", Email: "ann@x.io", Password: "secret", + UserLevel: 99, Roles: []string{"admin"}, // must be ignored + }) + if err != nil { + t.Fatal(err) + } + if reg.User.UserLevel != 0 || len(reg.User.Roles) != 0 { + t.Fatalf("register honoured privileges: %+v", reg.User) + } + var stored string + if err := db.QueryRow(`SELECT password FROM users WHERE username='ann'`).Scan(&stored); err != nil || !strings.HasPrefix(stored, "$2") { + t.Fatalf("password not bcrypt: %q %v", stored, err) + } + + if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "ann", Email: "other@x.io", Password: "x"}); !errors.Is(err, lookup.ErrUsernameExists) { + t.Fatalf("got %v", err) + } + if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "bob", Email: "ann@x.io", Password: "x"}); !errors.Is(err, lookup.ErrEmailExists) { + t.Fatalf("got %v", err) + } + var n int + _ = db.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&n) + if n != 1 { + t.Fatalf("failed register left a row: %d", n) + } + + login, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "secret", Claims: map[string]any{"ip_address": "1.2.3.4", "user_agent": "ua"}}) + if err != nil { + t.Fatal(err) + } + if !strings.HasPrefix(login.Token, "sess_") || login.ExpiresIn != 86400 || login.User.Email != "ann@x.io" { + t.Fatalf("login: %+v", login) + } + var ip string + _ = db.QueryRow(`SELECT ip_address FROM user_sessions WHERE session_token=?`, login.Token).Scan(&ip) + if ip != "1.2.3.4" { + t.Fatalf("ip %q", ip) + } + + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "wrong"}); err == nil || err.Error() != "invalid credentials" { + t.Fatalf("got %v", err) + } + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "nobody", Password: "x"}); err == nil || err.Error() != "invalid credentials" { + t.Fatalf("got %v", err) + } + + u, err := a.Session(ctx, login.Token, "authenticate") + if err != nil || u.UserName != "ann" || u.SessionID != login.Token { + t.Fatalf("session: %+v %v", u, err) + } + if err := a.TouchSession(ctx, login.Token, u); err != nil { + t.Fatal(err) + } + if _, err := a.Session(ctx, "nope", "authenticate"); err == nil || err.Error() != "invalid or expired session" { + t.Fatalf("got %v", err) + } + + ref, err := a.Refresh(ctx, login.Token) + if err != nil || ref.Token == login.Token { + t.Fatalf("refresh: %+v %v", ref, err) + } + if _, err := a.Session(ctx, login.Token, ""); err == nil { + t.Fatal("old session still valid after refresh") + } + if _, err := a.Refresh(ctx, login.Token); err == nil || err.Error() != "invalid or expired refresh token" { + t.Fatalf("got %v", err) + } + + if err := a.Logout(ctx, sectypes.LogoutRequest{Token: "Bearer " + ref.Token, UserID: ref.User.UserID}); err != nil { + t.Fatal(err) + } + if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err == nil || err.Error() != "session not found" { + t.Fatalf("got %v", err) + } +} + +func TestExpiredSessionRejected(t *testing.T) { + ctx := context.Background() + a, _ := newAuth(t, AuthOptions{}) + resp := registerUser(t, a, "eve") + a.Now = func() time.Time { return time.Now().Add(48 * time.Hour) } + if _, err := a.Session(ctx, resp.Token, ""); err == nil { + t.Fatal("expired session accepted") + } +} + +func TestLegacyPasswordUpgradeIsOptIn(t *testing.T) { + ctx := context.Background() + for _, upgrade := range []bool{false, true} { + a, db := newAuth(t, AuthOptions{UpgradePasswordHash: upgrade}) + _, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active) VALUES ('old','o@x.io','clear',1,'a,b',1)`) + if err != nil { + t.Fatal(err) + } + resp, err := a.Login(ctx, sectypes.LoginRequest{Username: "old", Password: "clear"}) + if err != nil || len(resp.User.Roles) != 2 { + t.Fatalf("login: %+v %v", resp, err) + } + var stored string + _ = db.QueryRow(`SELECT password FROM users WHERE username='old'`).Scan(&stored) + if got := strings.HasPrefix(stored, "$2"); got != upgrade { + t.Fatalf("upgrade=%v stored=%q", upgrade, stored) + } + } +} + +func TestInactiveUserCannotLogin(t *testing.T) { + ctx := context.Background() + a, db := newAuth(t, AuthOptions{}) + resp := registerUser(t, a, "ian") + _, _ = db.Exec(`UPDATE users SET is_active = 0`) + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ian", Password: "pw-ian"}); err == nil { + t.Fatal("inactive login accepted") + } + if _, err := a.Session(ctx, resp.Token, ""); err == nil { + t.Fatal("inactive session accepted") + } +} + +func TestLoginAPIKey(t *testing.T) { + ctx := context.Background() + a, db := newAuth(t, AuthOptions{}) + reg := registerUser(t, a, "kim") + insert := func(raw, typ string, active int, expires any) { + t.Helper() + _, err := db.Exec(`INSERT INTO user_keys (user_id, key_type, key_hash, name, is_active, expires_at) VALUES (?,?,?,?,?,?)`, + reg.User.UserID, typ, sectypes.HashKey(raw), "k", active, expires) + if err != nil { + t.Fatal(err) + } + } + insert("good", "header_api", 1, nil) + insert("generic", "api", 1, nil) + insert("jwt", "jwt_secret", 1, nil) + insert("off", "api", 0, nil) + insert("old", "api", 1, time.Now().Add(-time.Hour)) + + for _, k := range []string{"good", "generic"} { + resp, err := a.LoginAPIKey(ctx, k, map[string]any{"ip_address": "9.9.9.9"}) + if err != nil || resp.User.UserName != "kim" || !strings.HasPrefix(resp.Token, "sess_") { + t.Fatalf("%s: %+v %v", k, resp, err) + } + if _, err := a.Session(ctx, resp.Token, ""); err != nil { + t.Fatal(err) + } + } + var used sql.NullString + _ = db.QueryRow(`SELECT last_used_at FROM user_keys WHERE key_hash = ?`, sectypes.HashKey("good")).Scan(&used) + if !used.Valid { + t.Fatal("last_used_at not stamped") + } + for _, k := range []string{"", "missing", "jwt", "off", "old"} { + if _, err := a.LoginAPIKey(ctx, k, nil); !errors.Is(err, lookup.ErrInvalidAPIKey) { + t.Fatalf("%q: got %v", k, err) + } + } + _, _ = db.Exec(`UPDATE users SET is_active = 0`) + if _, err := a.LoginAPIKey(ctx, "good", nil); !errors.Is(err, lookup.ErrInvalidAPIKey) { + t.Fatalf("inactive user: got %v", err) + } +} + +func TestPasswordReset(t *testing.T) { + ctx := context.Background() + a, _ := newAuth(t, AuthOptions{}) + reg := registerUser(t, a, "rae") + + empty, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "none@x.io"}) + if err != nil || empty.Token != "" { + t.Fatalf("enumeration leak: %+v %v", empty, err) + } + if _, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{}); err == nil { + t.Fatal("expected error") + } + + r1, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "rae@x.io"}) + if err != nil || r1.Token == "" { + t.Fatal(err) + } + r2, _ := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Username: "rae"}) + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r1.Token, NewPassword: "n"}); err == nil { + t.Fatal("superseded token accepted") + } + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "newpw"}); err != nil { + t.Fatal(err) + } + if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "again"}); err == nil { + t.Fatal("token reused") + } + if _, err := a.Session(ctx, reg.Token, ""); err == nil { + t.Fatal("sessions survived reset") + } + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "rae", Password: "newpw"}); err != nil { + t.Fatal(err) + } +} + +func TestJWTLoginLogout(t *testing.T) { + ctx := context.Background() + a, db := newAuth(t, AuthOptions{}) + reg := registerUser(t, a, "jay") + resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: "jay", Password: "pw-jay"}) + if err != nil || !strings.HasPrefix(resp.Token, "token_") { + t.Fatalf("%+v %v", resp, err) + } + if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: "tok", UserID: reg.User.UserID}); err != nil { + t.Fatal(err) + } + var n int + _ = db.QueryRow(`SELECT COUNT(*) FROM token_blacklist WHERE token='tok'`).Scan(&n) + if n != 1 { + t.Fatal("token not blacklisted") + } +} + +func TestCustomSchemaNames(t *testing.T) { + db := newTestDB(t, ` +CREATE TABLE app_users (uid INTEGER PRIMARY KEY AUTOINCREMENT, login TEXT, email TEXT, password TEXT, + user_level INTEGER, roles TEXT, is_active INTEGER, created_at DATETIME, updated_at DATETIME, + last_login_at DATETIME, program_user_id INTEGER, program_user_table TEXT, remote_id TEXT, auth_provider TEXT, + totp_secret TEXT, totp_enabled INTEGER, totp_enabled_at DATETIME);`) + schema := lookup.Schema{lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"id": "uid", "username": "login"}}} + a := NewAuth(newTestBase(t, db, schema), AuthOptions{}) + ctx := context.Background() + if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "zed", Email: "z@x.io", Password: "p"}); err != nil { + t.Fatal(err) + } + if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "zed", Password: "p"}); err != nil { + t.Fatal(err) + } + var login string + if err := db.QueryRow(`SELECT login FROM app_users`).Scan(&login); err != nil || login != "zed" { + t.Fatalf("%q %v", login, err) + } +} diff --git a/pkg/security/lookup/direct/base.go b/pkg/security/lookup/direct/base.go new file mode 100644 index 0000000..0f97395 --- /dev/null +++ b/pkg/security/lookup/direct/base.go @@ -0,0 +1,462 @@ +// Package direct is the table-backed implementation of the lookup stores. SQL is built +// from the configured Schema (table and column names) and Dialect (placeholders, quoting, +// booleans, insert-returning-id); no statement is written per database and no ORM is used. +package direct + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" +) + +// Runner runs a database operation, reconnecting once when the *sql.DB has been closed. +// procedure.Runner (and procedure.DB) satisfy it. +type Runner interface { + Run(run func(*sql.DB) error) error +} + +// Querier is implemented by *sql.DB and *sql.Tx. +type Querier interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) + QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row +} + +// Base is the state shared by every direct store: the runner, dialect, schema and clock. +type Base struct { + run Runner + d dialect.Dialect + schema lookup.Schema + // Now is the clock; tests replace it. + Now func() time.Time +} + +// NewBase creates the shared state. The schema is merged with the defaults and validated. +func NewBase(run Runner, d dialect.Dialect, schema lookup.Schema) (*Base, error) { + if run == nil { + return nil, fmt.Errorf("direct: nil runner") + } + if d == nil { + return nil, fmt.Errorf("direct: nil dialect") + } + merged := lookup.DefaultSchema().Merge(schema) + if err := merged.Validate(); err != nil { + return nil, err + } + return &Base{run: run, d: d, schema: merged, Now: time.Now}, nil +} + +// Dialect returns the dialect in use. +func (b *Base) Dialect() dialect.Dialect { return b.d } + +// do runs fn against the database without a transaction. +func (b *Base) do(fn func(q Querier) error) error { + return b.run.Run(func(db *sql.DB) error { return fn(db) }) +} + +// tx runs fn in one transaction; an error rolls back. +func (b *Base) tx(ctx context.Context, fn func(q Querier) error) error { + return b.run.Run(func(db *sql.DB) error { + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return err + } + if err := fn(tx); err != nil { + _ = tx.Rollback() + return err + } + return tx.Commit() + }) +} + +// tableRef returns the (possibly schema-qualified) physical table name of an entity. +func (b *Base) tableRef(e lookup.Entity) string { + t := b.schema[e] + name := t.Name + if name == "" { + name = string(e) + } + if t.Schema != "" { + return t.Schema + "." + name + } + return name +} + +// colName returns the physical column name of a logical column. +func (b *Base) colName(c lookup.Column) string { + if t, ok := b.schema[c.Entity]; ok { + if n := t.Columns[c.Name]; n != "" { + return n + } + } + return c.Name +} + +// arg converts a Go value to a bind argument (booleans go through the dialect). +func (b *Base) arg(v any) any { + if bv, ok := v.(bool); ok { + return b.d.Bool(bv) + } + return v +} + +// timeDest scans a time column through the dialect, so drivers returning strings work. +type timeDest struct { + d dialect.Dialect + v *time.Time +} + +func (t timeDest) Scan(src any) error { + v, err := t.d.ScanTime(src) + if err != nil { + return err + } + *t.v = v + return nil +} + +type boolDest struct { + d dialect.Dialect + v *bool +} + +func (t boolDest) Scan(src any) error { + v, err := t.d.ScanBool(src) + if err != nil { + return err + } + *t.v = v + return nil +} + +func (b *Base) timeDest(v *time.Time) sql.Scanner { return timeDest{d: b.d, v: v} } +func (b *Base) boolDest(v *bool) sql.Scanner { return boolDest{d: b.d, v: v} } + +// --- query builder -------------------------------------------------------- + +// builder accumulates bind arguments and renders column references. +type builder struct { + b *Base + args []any + aliases map[lookup.Entity]string + nalias int +} + +func (bl *builder) ph(v any) string { + bl.args = append(bl.args, bl.b.arg(v)) + return bl.b.d.Placeholder(len(bl.args)) +} + +// col renders a column; with aliases set (select queries) it is qualified by its table alias. +func (bl *builder) col(c lookup.Column) string { + name := bl.b.d.Quote(bl.b.colName(c)) + if bl.aliases != nil { + if a, ok := bl.aliases[c.Entity]; ok { + return a + "." + name + } + } + return name +} + +// Cond renders one boolean condition. +type Cond func(*builder) string + +// Eq is `col = value`. +func Eq(c lookup.Column, v any) Cond { + return func(bl *builder) string { return bl.col(c) + " = " + bl.ph(v) } +} + +// EqFold is a case-insensitive `LOWER(col) = value` match (the value is lowered in Go). +func EqFold(c lookup.Column, v string) Cond { + return func(bl *builder) string { return "LOWER(" + bl.col(c) + ") = " + bl.ph(strings.ToLower(v)) } +} + +// Ne is `col <> value`. +func Ne(c lookup.Column, v any) Cond { + return func(bl *builder) string { return bl.col(c) + " <> " + bl.ph(v) } +} + +// Gt is `col > value`. +func Gt(c lookup.Column, v any) Cond { + return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) } +} + +// IsNull is `col IS NULL`. +func IsNull(c lookup.Column) Cond { + return func(bl *builder) string { return bl.col(c) + " IS NULL" } +} + +// EqCol is `a = b` between two columns (join conditions). +func EqCol(a, c lookup.Column) Cond { + return func(bl *builder) string { return bl.col(a) + " = " + bl.col(c) } +} + +// In is `col IN (v...)`; an empty list renders a condition that is never true. +func In(c lookup.Column, vs ...any) Cond { + return func(bl *builder) string { + if len(vs) == 0 { + return "1 = 0" + } + ph := make([]string, len(vs)) + for i, v := range vs { + ph[i] = bl.ph(v) + } + return bl.col(c) + " IN (" + strings.Join(ph, ", ") + ")" + } +} + +// InSelect is `col IN (subselect)`; the subselect's arguments share the outer numbering. +func InSelect(c lookup.Column, sub *Select) Cond { + return func(bl *builder) string { return bl.col(c) + " IN (" + sub.render(bl) + ")" } +} + +// Or joins conditions with OR inside parentheses. +func Or(cs ...Cond) Cond { return joinConds("OR", cs) } + +// And joins conditions with AND inside parentheses. +func And(cs ...Cond) Cond { return joinConds("AND", cs) } + +func joinConds(op string, cs []Cond) Cond { + return func(bl *builder) string { + parts := make([]string, len(cs)) + for i, c := range cs { + parts[i] = c(bl) + } + return "(" + strings.Join(parts, " "+op+" ") + ")" + } +} + +func (bl *builder) where(cs []Cond) string { + if len(cs) == 0 { + return "" + } + parts := make([]string, len(cs)) + for i, c := range cs { + parts[i] = c(bl) + } + return " WHERE " + strings.Join(parts, " AND ") +} + +// Stmt is a rendered statement. +type Stmt struct { + SQL string + Args []any +} + +// Select builds a SELECT. +type Select struct { + b *Base + from lookup.Entity + joins []join + cols []lookup.Column + conds []Cond + order []lookup.Column +} + +type join struct { + e lookup.Entity + on Cond +} + +// From starts a SELECT on e. +func (b *Base) From(e lookup.Entity) *Select { return &Select{b: b, from: e} } + +// Cols sets the selected columns. +func (s *Select) Cols(cs ...lookup.Column) *Select { s.cols = cs; return s } + +// Join adds `JOIN e ON on`. +func (s *Select) Join(e lookup.Entity, on Cond) *Select { + s.joins = append(s.joins, join{e: e, on: on}) + return s +} + +// Where adds AND-ed conditions. +func (s *Select) Where(cs ...Cond) *Select { s.conds = append(s.conds, cs...); return s } + +// OrderBy adds ascending order columns. +func (s *Select) OrderBy(cs ...lookup.Column) *Select { s.order = append(s.order, cs...); return s } + +// Build renders the statement. +func (s *Select) Build() Stmt { + bl := &builder{b: s.b} + sqlText := s.render(bl) + return Stmt{SQL: sqlText, Args: bl.args} +} + +// render writes the select into bl, giving every table a fresh alias so a subselect cannot +// clash with the statement around it. +func (s *Select) render(bl *builder) string { + saved := bl.aliases + defer func() { bl.aliases = saved }() + bl.aliases = map[lookup.Entity]string{} + alias := func() string { a := fmt.Sprintf("t%d", bl.nalias); bl.nalias++; return a } + bl.aliases[s.from] = alias() + for _, j := range s.joins { + bl.aliases[j.e] = alias() + } + sel := make([]string, len(s.cols)) + for i, c := range s.cols { + sel[i] = bl.col(c) + } + var sb strings.Builder + sb.WriteString("SELECT " + strings.Join(sel, ", ")) + sb.WriteString(" FROM " + s.b.d.Quote(s.b.tableRef(s.from)) + " " + bl.aliases[s.from]) + for _, j := range s.joins { + sb.WriteString(" JOIN " + s.b.d.Quote(s.b.tableRef(j.e)) + " " + bl.aliases[j.e] + " ON " + j.on(bl)) + } + sb.WriteString(bl.where(s.conds)) + if len(s.order) > 0 { + o := make([]string, len(s.order)) + for i, c := range s.order { + o[i] = bl.col(c) + } + sb.WriteString(" ORDER BY " + strings.Join(o, ", ")) + } + return sb.String() +} + +// QueryRow runs the select and scans the first row into dest. +func (s *Select) QueryRow(ctx context.Context, q Querier, dest ...any) error { + st := s.Build() + return q.QueryRowContext(ctx, st.SQL, st.Args...).Scan(dest...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters +} + +// Query runs the select. +func (s *Select) Query(ctx context.Context, q Querier) (*sql.Rows, error) { + st := s.Build() + return q.QueryContext(ctx, st.SQL, st.Args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters +} + +// Exists reports whether the select returns at least one row. +func (s *Select) Exists(ctx context.Context, q Querier) (bool, error) { + s.cols = []lookup.Column{s.firstCol()} + rows, err := s.Query(ctx, q) + if err != nil { + return false, err + } + defer func() { _ = rows.Close() }() + ok := rows.Next() + return ok, rows.Err() +} + +func (s *Select) firstCol() lookup.Column { + if len(s.cols) > 0 { + return s.cols[0] + } + return lookup.Column{Entity: s.from, Name: lookup.FirstColumn(s.from)} +} + +// Assignment is one `col = value` of an UPDATE or INSERT. +type Assignment struct { + Col lookup.Column + Val any +} + +// Set builds an Assignment. +func Set(c lookup.Column, v any) Assignment { return Assignment{Col: c, Val: v} } + +// Update builds an UPDATE. +type Update struct { + b *Base + e lookup.Entity + sets []Assignment + conds []Cond +} + +// Update starts an UPDATE of e. +func (b *Base) Update(e lookup.Entity) *Update { return &Update{b: b, e: e} } + +// Set adds assignments. +func (u *Update) Set(as ...Assignment) *Update { u.sets = append(u.sets, as...); return u } + +// Where adds AND-ed conditions. +func (u *Update) Where(cs ...Cond) *Update { u.conds = append(u.conds, cs...); return u } + +// Exec runs the update and returns the affected row count. +func (u *Update) Exec(ctx context.Context, q Querier) (int64, error) { + bl := &builder{b: u.b} + set := make([]string, len(u.sets)) + for i, a := range u.sets { + set[i] = bl.col(a.Col) + " = " + bl.ph(a.Val) + } + sqlText := "UPDATE " + u.b.d.Quote(u.b.tableRef(u.e)) + " SET " + strings.Join(set, ", ") + bl.where(u.conds) + res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// Delete builds a DELETE. +type Delete struct { + b *Base + e lookup.Entity + conds []Cond +} + +// Delete starts a DELETE on e. +func (b *Base) Delete(e lookup.Entity) *Delete { return &Delete{b: b, e: e} } + +// Where adds AND-ed conditions. +func (d *Delete) Where(cs ...Cond) *Delete { d.conds = append(d.conds, cs...); return d } + +// Exec runs the delete and returns the affected row count. +func (d *Delete) Exec(ctx context.Context, q Querier) (int64, error) { + bl := &builder{b: d.b} + sqlText := "DELETE FROM " + d.b.d.Quote(d.b.tableRef(d.e)) + bl.where(d.conds) + res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// Insert builds an INSERT. +type Insert struct { + b *Base + e lookup.Entity + sets []Assignment +} + +// Insert starts an INSERT into e. +func (b *Base) Insert(e lookup.Entity) *Insert { return &Insert{b: b, e: e} } + +// Set adds assignments. +func (i *Insert) Set(as ...Assignment) *Insert { i.sets = append(i.sets, as...); return i } + +func (i *Insert) colsAndArgs() ([]string, []any) { + cols := make([]string, len(i.sets)) + args := make([]any, len(i.sets)) + for n, a := range i.sets { + cols[n] = i.b.colName(a.Col) + args[n] = i.b.arg(a.Val) + } + return cols, args +} + +// Exec runs the insert. +func (i *Insert) Exec(ctx context.Context, q Querier) error { + cols, args := i.colsAndArgs() + ph := make([]string, len(cols)) + qc := make([]string, len(cols)) + for n, c := range cols { + qc[n] = i.b.d.Quote(c) + ph[n] = i.b.d.Placeholder(n + 1) + } + sqlText := "INSERT INTO " + i.b.d.Quote(i.b.tableRef(i.e)) + " (" + strings.Join(qc, ", ") + ") VALUES (" + strings.Join(ph, ", ") + ")" + _, err := q.ExecContext(ctx, sqlText, args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters + return err +} + +// ExecID runs the insert and returns the generated value of idCol, using the dialect's +// insert-returning-id strategy. +func (i *Insert) ExecID(ctx context.Context, q Querier, idCol lookup.Column) (int64, error) { + cols, args := i.colsAndArgs() + ins := i.b.d.InsertReturningID(i.b.tableRef(i.e), cols, i.b.colName(idCol)) + return ins.Run(ctx, q, args...) +} diff --git a/pkg/security/lookup/direct/builder_test.go b/pkg/security/lookup/direct/builder_test.go new file mode 100644 index 0000000..dd2347f --- /dev/null +++ b/pkg/security/lookup/direct/builder_test.go @@ -0,0 +1,75 @@ +package direct + +import ( + "database/sql" + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" +) + +type noRun struct{} + +func (noRun) Run(func(*sql.DB) error) error { return nil } + +func TestBuilderRendersPerDialect(t *testing.T) { + schema := lookup.Schema{lookup.EntityUserSessions: {Schema: "auth", Name: "sessions"}} + cases := map[string]string{ + "postgres": `SELECT t0."session_token" FROM "auth"."sessions" t0`, + "mysql": "SELECT t0.`session_token` FROM `auth`.`sessions` t0", + "mssql": `SELECT t0.[session_token] FROM [auth].[sessions] t0`, + } + tails := map[string]string{ + "postgres": ` WHERE t0."user_id" = $1 AND t0."session_token" IN ($2, $3)`, + "mysql": " WHERE t0.`user_id` = ? AND t0.`session_token` IN (?, ?)", + "mssql": ` WHERE t0.[user_id] = @p1 AND t0.[session_token] IN (@p2, @p3)`, + } + for name, head := range cases { + d, err := dialect.Get(name) + if err != nil { + t.Fatal(err) + } + b, err := NewBase(noRun{}, d, schema) + if err != nil { + t.Fatal(err) + } + st := b.From(lookup.EntityUserSessions).Cols(lookup.SessionsToken). + Where(Eq(lookup.SessionsUserID, 7), In(lookup.SessionsToken, "a", "b")).Build() + if st.SQL != head+tails[name] || len(st.Args) != 3 { + t.Errorf("%s:\n got %s\n want %s", name, st.SQL, head+tails[name]) + } + } +} + +func TestBuilderBoolsGoThroughDialect(t *testing.T) { + d, _ := dialect.Get("sqlite") + b, _ := NewBase(noRun{}, d, nil) + st := b.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersIsActive, true)).Build() + if st.Args[0] != d.Bool(true) { + t.Fatalf("bool not converted: %#v", st.Args[0]) + } +} + +func TestBuilderSubselectSharesArguments(t *testing.T) { + d, _ := dialect.Get("postgres") + b, _ := NewBase(noRun{}, d, nil) + sub := b.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, 5)) + st := b.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesID). + Where(Eq(lookup.RowRulesTableName, "t"), InSelect(lookup.RowRulesGroupID, sub), Eq(lookup.RowRulesSchemaName, "s")).Build() + want := `SELECT t0."id" FROM "sec_row_rules" t0 WHERE t0."table_name" = $1 AND t0."group_id" IN (SELECT t1."group_id" FROM "sec_group_members" t1 WHERE t1."user_id" = $2) AND t0."schema_name" = $3` + if st.SQL != want || len(st.Args) != 3 { + t.Fatalf("got %s\nwant %s", st.SQL, want) + } +} + +func TestSchemaRejectsUnsafeIdentifiers(t *testing.T) { + d, _ := dialect.Get("postgres") + bad := lookup.Schema{lookup.EntityUsers: {Name: `users"; DROP TABLE x; --`}} + if _, err := NewBase(noRun{}, d, bad); err == nil { + t.Fatal("unsafe table name accepted") + } + bad = lookup.Schema{lookup.EntityUsers: {Columns: map[string]string{"username": "a b"}}} + if _, err := NewBase(noRun{}, d, bad); err == nil { + t.Fatal("unsafe column name accepted") + } +} diff --git a/pkg/security/lookup/direct/helpers_test.go b/pkg/security/lookup/direct/helpers_test.go new file mode 100644 index 0000000..2ea0206 --- /dev/null +++ b/pkg/security/lookup/direct/helpers_test.go @@ -0,0 +1,46 @@ +package direct + +import ( + "database/sql" + "testing" + + _ "github.com/glebarez/go-sqlite" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure" +) + +func newTestDB(t *testing.T, extraDDL ...string) *sql.DB { + t.Helper() + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatal(err) + } + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + ref, err := ddl.SQL("sqlite") + if err != nil { + t.Fatal(err) + } + for _, s := range append([]string{ref}, extraDDL...) { + if _, err := db.Exec(s); err != nil { + t.Fatalf("ddl: %v", err) + } + } + return db +} + +func newTestBase(t *testing.T, db *sql.DB, schema lookup.Schema) *Base { + t.Helper() + d, err := dialect.Detect(db) + if err != nil { + t.Fatal(err) + } + b, err := NewBase(procedure.NewDB(db, nil, nil), d, schema) + if err != nil { + t.Fatal(err) + } + return b +} diff --git a/pkg/security/lookup/direct/keys.go b/pkg/security/lookup/direct/keys.go new file mode 100644 index 0000000..0cb98fa --- /dev/null +++ b/pkg/security/lookup/direct/keys.go @@ -0,0 +1,197 @@ +package direct + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Keys implements lookup.KeyStore on the user keys table. scopes and meta are stored as +// JSON through the dialect (native JSON column or TEXT). +type Keys struct{ *Base } + +var _ lookup.KeyStore = (*Keys)(nil) + +// NewKeys creates the direct KeyStore. +func NewKeys(b *Base) *Keys { return &Keys{Base: b} } + +// Create implements lookup.KeyStore. +func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) { + scopes, err := k.d.EncodeJSON(req.Scopes) + if err != nil { + return nil, fmt.Errorf("failed to marshal scopes: %w", err) + } + meta, err := k.d.EncodeJSON(req.Meta) + if err != nil { + return nil, fmt.Errorf("failed to marshal meta: %w", err) + } + now := k.Now() + var id int64 + err = k.do(func(q Querier) error { + var err error + id, err = k.Insert(lookup.EntityUserKeys).Set( + Set(lookup.KeysUserID, req.UserID), + Set(lookup.KeysKeyType, string(req.KeyType)), + Set(lookup.KeysKeyHash, keyHash), + Set(lookup.KeysName, req.Name), + Set(lookup.KeysScopes, scopes), + Set(lookup.KeysMeta, meta), + Set(lookup.KeysExpiresAt, req.ExpiresAt), + Set(lookup.KeysCreatedAt, now), + Set(lookup.KeysIsActive, true), + ).ExecID(ctx, q, lookup.KeysID) + return err + }) + if err != nil { + return nil, fmt.Errorf("create key query failed: %w", err) + } + return §ypes.UserKey{ + ID: id, + UserID: req.UserID, + KeyType: req.KeyType, + KeyHash: keyHash, + Name: req.Name, + Scopes: req.Scopes, + Meta: req.Meta, + ExpiresAt: req.ExpiresAt, + CreatedAt: now, + IsActive: true, + }, nil +} + +// keyScan holds the destinations for one key row. +type keyScan struct { + k sectypes.UserKey + keyType string + scopes, meta any + expiresAt, created, lastU time.Time + active bool +} + +func (k *Keys) keyCols(withLastUsed bool) []lookup.Column { + cols := []lookup.Column{lookup.KeysID, lookup.KeysUserID, lookup.KeysKeyType, lookup.KeysName, lookup.KeysScopes, + lookup.KeysMeta, lookup.KeysExpiresAt, lookup.KeysCreatedAt, lookup.KeysIsActive} + if withLastUsed { + cols = append(cols, lookup.KeysLastUsedAt) + } + return cols +} + +func (k *Keys) dest(s *keyScan, withLastUsed bool) []any { + d := []any{&s.k.ID, &s.k.UserID, &s.keyType, &s.k.Name, &s.scopes, &s.meta, + k.timeDest(&s.expiresAt), k.timeDest(&s.created), k.boolDest(&s.active)} + if withLastUsed { + d = append(d, k.timeDest(&s.lastU)) + } + return d +} + +func (k *Keys) finish(s *keyScan) sectypes.UserKey { + out := s.k + out.KeyType = sectypes.KeyType(s.keyType) + out.CreatedAt = s.created + out.IsActive = s.active + _ = k.d.DecodeJSON(s.scopes, &out.Scopes) + _ = k.d.DecodeJSON(s.meta, &out.Meta) + if !s.expiresAt.IsZero() { + t := s.expiresAt + out.ExpiresAt = &t + } + if !s.lastU.IsZero() { + t := s.lastU + out.LastUsedAt = &t + } + return out +} + +// List implements lookup.KeyStore: active, non-expired keys; an empty keyType means all types. +func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) { + keys := []sectypes.UserKey{} + conds := []Cond{ + Eq(lookup.KeysUserID, userID), + Eq(lookup.KeysIsActive, true), + Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())), + } + if keyType != "" { + conds = append(conds, Eq(lookup.KeysKeyType, string(keyType))) + } + err := k.do(func(q Querier) error { + keys = keys[:0] + rows, err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(true)...).Where(conds...).OrderBy(lookup.KeysID).Query(ctx, q) + if err != nil { + return err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var s keyScan + if err := rows.Scan(k.dest(&s, true)...); err != nil { + return err + } + keys = append(keys, k.finish(&s)) + } + return rows.Err() + }) + if err != nil { + return nil, fmt.Errorf("get user keys query failed: %w", err) + } + return keys, nil +} + +// Delete implements lookup.KeyStore: soft-deletes the key after checking ownership and +// returns its hash. +func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) { + var keyHash string + err := k.tx(ctx, func(q Querier) error { + match := []Cond{Eq(lookup.KeysID, keyID), Eq(lookup.KeysUserID, userID), Eq(lookup.KeysIsActive, true)} + if err := k.From(lookup.EntityUserKeys).Cols(lookup.KeysKeyHash).Where(match...).QueryRow(ctx, q, &keyHash); err != nil { + return err + } + _, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysIsActive, false)).Where(match...).Exec(ctx, q) + return err + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", errors.New("key not found or already deleted") + } + return "", fmt.Errorf("delete key query failed: %w", err) + } + return keyHash, nil +} + +// Validate implements lookup.KeyStore: finds an active, non-expired key by hash (optionally of +// one type) and stamps last_used_at. +func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) { + conds := []Cond{ + Eq(lookup.KeysKeyHash, keyHash), + Eq(lookup.KeysIsActive, true), + Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())), + } + if keyType != "" { + conds = append(conds, Eq(lookup.KeysKeyType, string(keyType))) + } + var s keyScan + now := k.Now() + err := k.tx(ctx, func(q Querier) error { + s = keyScan{} + if err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(false)...).Where(conds...).QueryRow(ctx, q, k.dest(&s, false)...); err != nil { + return err + } + _, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, s.k.ID)).Exec(ctx, q) + return err + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, errors.New("invalid or expired key") + } + return nil, fmt.Errorf("validate key query failed: %w", err) + } + out := k.finish(&s) + out.KeyHash = keyHash + out.LastUsedAt = &now + return &out, nil +} diff --git a/pkg/security/lookup/direct/keys_test.go b/pkg/security/lookup/direct/keys_test.go new file mode 100644 index 0000000..0ea1942 --- /dev/null +++ b/pkg/security/lookup/direct/keys_test.go @@ -0,0 +1,71 @@ +package direct + +import ( + "context" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +func TestKeysLifecycle(t *testing.T) { + ctx := context.Background() + db := newTestDB(t) + k := NewKeys(newTestBase(t, db, nil)) + _, _ = db.Exec(`INSERT INTO users (username,email,password,is_active) VALUES ('u','u@x.io','x',1)`) + + exp := time.Now().Add(time.Hour) + created, err := k.Create(ctx, sectypes.CreateKeyRequest{ + UserID: 1, KeyType: sectypes.KeyTypeHeaderAPI, Name: "ci", + Scopes: []string{"read", "write"}, Meta: map[string]any{"env": "prod"}, ExpiresAt: &exp, + }, sectypes.HashKey("raw1")) + if err != nil || created.ID == 0 { + t.Fatalf("%+v %v", created, err) + } + // nil scopes/meta must not store the JSON text "null" + if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeJWTSecret, Name: "j"}, sectypes.HashKey("raw2")); err != nil { + t.Fatal(err) + } + if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeGenericAPI, Name: "old", ExpiresAt: ptr(time.Now().Add(-time.Hour))}, sectypes.HashKey("raw3")); err != nil { + t.Fatal(err) + } + + all, err := k.List(ctx, 1, "") + if err != nil || len(all) != 2 { + t.Fatalf("list all: %d %v", len(all), err) + } + one, _ := k.List(ctx, 1, sectypes.KeyTypeHeaderAPI) + if len(one) != 1 || one[0].Name != "ci" || len(one[0].Scopes) != 2 || one[0].Meta["env"] != "prod" || one[0].ExpiresAt == nil { + t.Fatalf("list typed: %+v", one) + } + if other, _ := k.List(ctx, 2, ""); len(other) != 0 { + t.Fatal("other user's keys listed") + } + + got, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeHeaderAPI) + if err != nil || got.UserID != 1 || got.KeyHash != sectypes.HashKey("raw1") || got.LastUsedAt == nil { + t.Fatalf("%+v %v", got, err) + } + if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeGenericAPI); err == nil || err.Error() != "invalid or expired key" { + t.Fatalf("wrong type: %v", err) + } + if _, err := k.Validate(ctx, sectypes.HashKey("raw3"), ""); err == nil { + t.Fatal("expired key validated") + } + + if _, err := k.Delete(ctx, 2, created.ID); err == nil || err.Error() != "key not found or already deleted" { + t.Fatalf("foreign delete: %v", err) + } + hash, err := k.Delete(ctx, 1, created.ID) + if err != nil || hash != sectypes.HashKey("raw1") { + t.Fatalf("%q %v", hash, err) + } + if _, err := k.Delete(ctx, 1, created.ID); err == nil { + t.Fatal("double delete succeeded") + } + if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), ""); err == nil { + t.Fatal("deleted key validated") + } +} + +func ptr[T any](v T) *T { return &v } diff --git a/pkg/security/lookup/direct/oauth.go b/pkg/security/lookup/direct/oauth.go new file mode 100644 index 0000000..132f239 --- /dev/null +++ b/pkg/security/lookup/direct/oauth.go @@ -0,0 +1,380 @@ +package direct + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strings" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// nullIfEmpty keeps optional TEXT columns (e.g. client_secret_hash of a public client) NULL +// rather than "". +func nullIfEmpty(s string) any { + if s == "" { + return nil + } + return s +} + +// OAuthClients implements lookup.OAuthClientStore. Array columns (redirect_uris, grant_types, +// allowed_scopes, scopes) are JSON through the dialect. +type OAuthClients struct{ *Base } + +var _ lookup.OAuthClientStore = (*OAuthClients)(nil) + +// NewOAuthClients creates the direct OAuthClientStore. +func NewOAuthClients(b *Base) *OAuthClients { return &OAuthClients{Base: b} } + +// RegisterClient implements lookup.OAuthClientStore. +func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) { + grantTypes := client.GrantTypes + if len(grantTypes) == 0 { + grantTypes = []string{"authorization_code"} + } + allowedScopes := client.AllowedScopes + if len(allowedScopes) == 0 { + allowedScopes = []string{"openid", "profile", "email"} + } + authMethod := client.TokenEndpointAuthMethod + if authMethod == "" { + authMethod = "none" + } + redirects, err := o.d.EncodeJSON(client.RedirectURIs) + if err != nil { + return nil, fmt.Errorf("failed to marshal redirect_uris: %w", err) + } + if redirects == nil { // the column is NOT NULL + redirects = "[]" + } + grants, err := o.d.EncodeJSON(grantTypes) + if err != nil { + return nil, fmt.Errorf("failed to marshal grant_types: %w", err) + } + scopes, err := o.d.EncodeJSON(allowedScopes) + if err != nil { + return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err) + } + + err = o.do(func(q Querier) error { + return o.Insert(lookup.EntityOAuthClients).Set( + Set(lookup.OAuthClientsClientID, client.ClientID), + Set(lookup.OAuthClientsRedirectURIs, redirects), + Set(lookup.OAuthClientsClientName, client.ClientName), + Set(lookup.OAuthClientsGrantTypes, grants), + Set(lookup.OAuthClientsAllowedScopes, scopes), + Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)), + Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod), + Set(lookup.OAuthClientsIsActive, true), + Set(lookup.OAuthClientsCreatedAt, o.Now()), + ).Exec(ctx, q) + }) + if err != nil { + return nil, fmt.Errorf("failed to register client: %w", err) + } + return §ypes.OAuthServerClient{ + ClientID: client.ClientID, + RedirectURIs: client.RedirectURIs, + ClientName: client.ClientName, + GrantTypes: grantTypes, + AllowedScopes: allowedScopes, + ClientSecretHash: client.ClientSecretHash, + TokenEndpointAuthMethod: authMethod, + }, nil +} + +// GetClient implements lookup.OAuthClientStore. +func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) { + var redirects, grants, scopes any + var name, secret, method sql.NullString + err := o.do(func(q Querier) error { + return o.From(lookup.EntityOAuthClients). + Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes, + lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod). + Where(Eq(lookup.OAuthClientsClientID, clientID), Eq(lookup.OAuthClientsIsActive, true)). + QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("client not found") + } + return nil, fmt.Errorf("failed to get client: %w", err) + } + res := §ypes.OAuthServerClient{ + ClientID: clientID, + ClientName: name.String, + ClientSecretHash: secret.String, + TokenEndpointAuthMethod: method.String, + } + _ = o.d.DecodeJSON(redirects, &res.RedirectURIs) + _ = o.d.DecodeJSON(grants, &res.GrantTypes) + _ = o.d.DecodeJSON(scopes, &res.AllowedScopes) + return res, nil +} + +// SaveCode implements lookup.OAuthClientStore. +func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error { + scopes, err := o.d.EncodeJSON(code.Scopes) + if err != nil { + return fmt.Errorf("failed to marshal scopes: %w", err) + } + method := code.CodeChallengeMethod + if method == "" { + method = "S256" + } + return o.do(func(q Querier) error { + return o.Insert(lookup.EntityOAuthCodes).Set( + Set(lookup.OAuthCodesCode, code.Code), + Set(lookup.OAuthCodesClientID, code.ClientID), + Set(lookup.OAuthCodesRedirectURI, code.RedirectURI), + Set(lookup.OAuthCodesClientState, code.ClientState), + Set(lookup.OAuthCodesCodeChallenge, code.CodeChallenge), + Set(lookup.OAuthCodesCodeChallengeMethod, method), + Set(lookup.OAuthCodesSessionToken, code.SessionToken), + Set(lookup.OAuthCodesRefreshToken, code.RefreshToken), + Set(lookup.OAuthCodesScopes, scopes), + Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt), + Set(lookup.OAuthCodesCreatedAt, o.Now()), + ).Exec(ctx, q) + }) +} + +// ExchangeCode implements lookup.OAuthClientStore: the code is consumed in a transaction and +// only the caller whose delete removes the row gets it, so a code cannot be redeemed twice. +func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) { + var res sectypes.OAuthCode + var state, refresh sql.NullString + var scopes any + err := o.tx(ctx, func(q Querier) error { + err := o.From(lookup.EntityOAuthCodes). + Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState, + lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken, + lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes). + Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())). + QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod, + &res.SessionToken, &refresh, &scopes) + if err != nil { + return err + } + n, err := o.Delete(lookup.EntityOAuthCodes).Where(Eq(lookup.OAuthCodesCode, code)).Exec(ctx, q) + if err != nil { + return err + } + if n == 0 { + return sql.ErrNoRows + } + return nil + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("invalid or expired code") + } + return nil, fmt.Errorf("failed to exchange code: %w", err) + } + res.Code = code + res.ClientState = state.String + res.RefreshToken = refresh.String + _ = o.d.DecodeJSON(scopes, &res.Scopes) + return &res, nil +} + +// Introspect implements lookup.OAuthClientStore (RFC 7662). An unknown or expired token is +// {active:false}, not an error. +func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) { + var info sectypes.OAuthTokenInfo + var userID int + var username, email, roles sql.NullString + var level sql.NullInt64 + var exp, iat time.Time + err := o.do(func(q Querier) error { + return o.From(lookup.EntityUserSessions). + Cols(lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, + lookup.SessionsExpiresAt, lookup.SessionsCreatedAt). + Join(lookup.EntityUsers, EqCol(lookup.UsersID, lookup.SessionsUserID)). + Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, o.Now()), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, &userID, &username, &email, &level, &roles, o.timeDest(&exp), o.timeDest(&iat)) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return §ypes.OAuthTokenInfo{Active: false}, nil + } + return nil, fmt.Errorf("failed to introspect token: %w", err) + } + info.Active = true + info.Sub = fmt.Sprintf("%d", userID) + info.Username = username.String + info.Email = email.String + info.UserLevel = int(level.Int64) + info.Roles = ParseRoles(roles.String) + if !exp.IsZero() { + info.Exp = exp.Unix() + } + if !iat.IsZero() { + info.Iat = iat.Unix() + } + return &info, nil +} + +// Revoke implements lookup.OAuthClientStore (RFC 7009): the session is deleted; an unknown +// token is not an error. +func (o *OAuthClients) Revoke(ctx context.Context, token string) error { + return o.do(func(q Querier) error { + _, err := o.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, token)).Exec(ctx, q) + return err + }) +} + +// OAuthUsers implements lookup.OAuthUserStore. +type OAuthUsers struct{ *Base } + +var _ lookup.OAuthUserStore = (*OAuthUsers)(nil) + +// NewOAuthUsers creates the direct OAuthUserStore. +func NewOAuthUsers(b *Base) *OAuthUsers { return &OAuthUsers{Base: b} } + +// GetOrCreateUser implements lookup.OAuthUserStore: select by email, then update or insert, +// in one transaction (no upsert). +func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) { + roles := strings.Join(user.Roles, ",") + var userID int + err := o.tx(ctx, func(q Querier) error { + now := o.Now() + var remoteID, authProvider sql.NullString + err := o.From(lookup.EntityUsers).Cols(lookup.UsersID, lookup.UsersRemoteID, lookup.UsersAuthProvider). + Where(Eq(lookup.UsersEmail, user.Email)).QueryRow(ctx, q, &userID, &remoteID, &authProvider) + if err == nil { + // remote_id and auth_provider are only filled when still unset. + sets := []Assignment{Set(lookup.UsersLastLoginAt, now), Set(lookup.UsersUpdatedAt, now)} + if !remoteID.Valid { + sets = append(sets, Set(lookup.UsersRemoteID, user.RemoteID)) + } + if !authProvider.Valid { + sets = append(sets, Set(lookup.UsersAuthProvider, provider)) + } + _, err := o.Update(lookup.EntityUsers).Set(sets...).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + return err + } + if !errors.Is(err, sql.ErrNoRows) { + return err + } + id, err := o.Insert(lookup.EntityUsers).Set( + Set(lookup.UsersUsername, user.UserName), + Set(lookup.UsersEmail, user.Email), + Set(lookup.UsersPassword, nil), + Set(lookup.UsersUserLevel, user.UserLevel), + Set(lookup.UsersRoles, roles), + Set(lookup.UsersIsActive, true), + Set(lookup.UsersCreatedAt, now), + Set(lookup.UsersUpdatedAt, now), + Set(lookup.UsersLastLoginAt, now), + Set(lookup.UsersRemoteID, user.RemoteID), + Set(lookup.UsersAuthProvider, provider), + ).ExecID(ctx, q, lookup.UsersID) + userID = int(id) + return err + }) + if err != nil { + return 0, fmt.Errorf("failed to get or create user: %w", err) + } + return userID, nil +} + +// CreateSession implements lookup.OAuthUserStore: insert, or update when the token exists. +func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error { + return o.tx(ctx, func(q Querier) error { + now := o.Now() + exists, err := o.From(lookup.EntityUserSessions).Cols(lookup.SessionsID).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exists(ctx, q) + if err != nil { + return err + } + if exists { + _, err := o.Update(lookup.EntityUserSessions).Set( + Set(lookup.SessionsAccessToken, s.AccessToken), + Set(lookup.SessionsRefreshToken, s.RefreshToken), + Set(lookup.SessionsTokenType, s.TokenType), + Set(lookup.SessionsExpiresAt, s.ExpiresAt), + Set(lookup.SessionsLastActivityAt, now), + ).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exec(ctx, q) + return err + } + return o.Insert(lookup.EntityUserSessions).Set( + Set(lookup.SessionsToken, s.SessionToken), + Set(lookup.SessionsUserID, s.UserID), + Set(lookup.SessionsExpiresAt, s.ExpiresAt), + Set(lookup.SessionsCreatedAt, now), + Set(lookup.SessionsLastActivityAt, now), + Set(lookup.SessionsAccessToken, s.AccessToken), + Set(lookup.SessionsRefreshToken, s.RefreshToken), + Set(lookup.SessionsTokenType, s.TokenType), + Set(lookup.SessionsAuthProvider, s.Provider), + ).Exec(ctx, q) + }) +} + +// GetByRefreshToken implements lookup.OAuthUserStore. +func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) { + var s lookup.OAuthRefreshSession + var access, tokenType sql.NullString + err := o.do(func(q Querier) error { + return o.From(lookup.EntityUserSessions). + Cols(lookup.SessionsUserID, lookup.SessionsAccessToken, lookup.SessionsTokenType, lookup.SessionsExpiresAt). + Where(Eq(lookup.SessionsRefreshToken, refreshToken), Gt(lookup.SessionsExpiresAt, o.Now())). + QueryRow(ctx, q, &s.UserID, &access, &tokenType, o.timeDest(&s.Expiry)) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("refresh token not found or expired") + } + return nil, fmt.Errorf("failed to get session by refresh token: %w", err) + } + s.AccessToken = access.String + s.TokenType = tokenType.String + return &s, nil +} + +// UpdateRefreshToken implements lookup.OAuthUserStore. +func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error { + var rows int64 + err := o.do(func(q Querier) error { + var err error + rows, err = o.Update(lookup.EntityUserSessions).Set( + Set(lookup.SessionsToken, newSessionToken), + Set(lookup.SessionsAccessToken, newAccessToken), + Set(lookup.SessionsRefreshToken, newRefreshToken), + Set(lookup.SessionsExpiresAt, expiresAt), + Set(lookup.SessionsLastActivityAt, o.Now()), + ).Where(Eq(lookup.SessionsUserID, userID), Eq(lookup.SessionsRefreshToken, oldRefreshToken)).Exec(ctx, q) + return err + }) + if err != nil { + return fmt.Errorf("failed to update session: %w", err) + } + if rows == 0 { + return fmt.Errorf("session not found") + } + return nil +} + +// GetUser implements lookup.OAuthUserStore. +func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) { + var u userRow + err := o.do(func(q Querier) error { + return o.From(lookup.EntityUsers). + Cols(lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, + lookup.UsersProgramUserID, lookup.UsersProgramUserTable). + Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("user not found") + } + return nil, fmt.Errorf("failed to get user data: %w", err) + } + u.id = userID + return u.context(""), nil +} diff --git a/pkg/security/lookup/direct/oauth_test.go b/pkg/security/lookup/direct/oauth_test.go new file mode 100644 index 0000000..6dbe531 --- /dev/null +++ b/pkg/security/lookup/direct/oauth_test.go @@ -0,0 +1,130 @@ +package direct + +import ( + "context" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +func TestOAuthClientAndCodeFlow(t *testing.T) { + ctx := context.Background() + db := newTestDB(t) + b := newTestBase(t, db, nil) + c := NewOAuthClients(b) + + reg, err := c.RegisterClient(ctx, §ypes.OAuthServerClient{ClientID: "cid", RedirectURIs: []string{"https://a/cb"}, ClientName: "App"}) + if err != nil || reg.TokenEndpointAuthMethod != "none" || len(reg.GrantTypes) != 1 || len(reg.AllowedScopes) != 3 { + t.Fatalf("%+v %v", reg, err) + } + got, err := c.GetClient(ctx, "cid") + if err != nil || got.ClientName != "App" || got.RedirectURIs[0] != "https://a/cb" || got.ClientSecretHash != "" { + t.Fatalf("%+v %v", got, err) + } + if _, err := c.GetClient(ctx, "nope"); err == nil || err.Error() != "client not found" { + t.Fatalf("got %v", err) + } + _, _ = db.Exec(`UPDATE oauth_clients SET is_active = 0`) + if _, err := c.GetClient(ctx, "cid"); err == nil { + t.Fatal("inactive client returned") + } + + code := §ypes.OAuthCode{Code: "c1", ClientID: "cid", RedirectURI: "https://a/cb", CodeChallenge: "ch", + SessionToken: "st", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)} + if err := c.SaveCode(ctx, code); err != nil { + t.Fatal(err) + } + ex, err := c.ExchangeCode(ctx, "c1") + if err != nil || ex.Code != "c1" || ex.CodeChallengeMethod != "S256" || ex.SessionToken != "st" || len(ex.Scopes) != 1 { + t.Fatalf("%+v %v", ex, err) + } + if _, err := c.ExchangeCode(ctx, "c1"); err == nil || err.Error() != "invalid or expired code" { + t.Fatalf("code reused: %v", err) + } + code.Code, code.ExpiresAt = "c2", time.Now().Add(-time.Minute) + _ = c.SaveCode(ctx, code) + if _, err := c.ExchangeCode(ctx, "c2"); err == nil { + t.Fatal("expired code exchanged") + } +} + +func TestOAuthIntrospectRevoke(t *testing.T) { + ctx := context.Background() + a, db := newAuth(t, AuthOptions{}) + reg := registerUser(t, a, "oli") + _, _ = db.Exec(`UPDATE users SET roles='r1,r2', user_level=3`) + c := NewOAuthClients(a.Base) + + info, err := c.Introspect(ctx, reg.Token) + if err != nil || !info.Active || info.Username != "oli" || info.UserLevel != 3 || len(info.Roles) != 2 || info.Exp == 0 || info.Iat == 0 || info.Sub != "1" { + t.Fatalf("%+v %v", info, err) + } + if err := c.Revoke(ctx, reg.Token); err != nil { + t.Fatal(err) + } + if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active { + t.Fatalf("%+v %v", info, err) + } + if err := c.Revoke(ctx, "unknown"); err != nil { + t.Fatal(err) + } +} + +func TestOAuthUsers(t *testing.T) { + ctx := context.Background() + db := newTestDB(t) + o := NewOAuthUsers(newTestBase(t, db, nil)) + + id, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-1", Roles: []string{"a"}}, "github") + if err != nil || id == 0 { + t.Fatalf("%d %v", id, err) + } + // second login: same user; existing remote_id/auth_provider are kept + id2, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-2"}, "google") + if err != nil || id2 != id { + t.Fatalf("%d %v", id2, err) + } + var remote, prov string + _ = db.QueryRow(`SELECT remote_id, auth_provider FROM users WHERE id=?`, id).Scan(&remote, &prov) + if remote != "r-1" || prov != "github" { + t.Fatalf("overwrote: %q %q", remote, prov) + } + + exp := time.Now().Add(time.Hour) + s := lookup.OAuthSession{SessionToken: "s1", UserID: id, AccessToken: "a1", RefreshToken: "r1", TokenType: "Bearer", ExpiresAt: exp, Provider: "github"} + if err := o.CreateSession(ctx, s); err != nil { + t.Fatal(err) + } + s.AccessToken = "a1b" // same token: updated, not duplicated + if err := o.CreateSession(ctx, s); err != nil { + t.Fatal(err) + } + var n int + _ = db.QueryRow(`SELECT COUNT(*) FROM user_sessions`).Scan(&n) + if n != 1 { + t.Fatalf("sessions: %d", n) + } + + ref, err := o.GetByRefreshToken(ctx, "r1") + if err != nil || ref.UserID != id || ref.AccessToken != "a1b" || ref.TokenType != "Bearer" || ref.Expiry.IsZero() { + t.Fatalf("%+v %v", ref, err) + } + if _, err := o.GetByRefreshToken(ctx, "zzz"); err == nil { + t.Fatal("unknown refresh token accepted") + } + if err := o.UpdateRefreshToken(ctx, id, "r1", "s2", "a2", "r2", exp); err != nil { + t.Fatal(err) + } + if err := o.UpdateRefreshToken(ctx, id, "r1", "s3", "a3", "r3", exp); err == nil || err.Error() != "session not found" { + t.Fatalf("got %v", err) + } + u, err := o.GetUser(ctx, id) + if err != nil || u.UserName != "gh" || u.UserID != id { + t.Fatalf("%+v %v", u, err) + } + if _, err := o.GetUser(ctx, 999); err == nil || err.Error() != "user not found" { + t.Fatalf("got %v", err) + } +} diff --git a/pkg/security/lookup/direct/passkey.go b/pkg/security/lookup/direct/passkey.go new file mode 100644 index 0000000..a8dbc01 --- /dev/null +++ b/pkg/security/lookup/direct/passkey.go @@ -0,0 +1,294 @@ +package direct + +import ( + "context" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Passkey implements lookup.PasskeyStore. credential_id, public_key and aaguid are base64 +// TEXT (not native bytea) and transports is JSON, so one schema works on every dialect. +type Passkey struct{ *Base } + +var _ lookup.PasskeyStore = (*Passkey)(nil) + +// NewPasskey creates the direct PasskeyStore. +func NewPasskey(b *Base) *Passkey { return &Passkey{Base: b} } + +// Store implements lookup.PasskeyStore. +func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) { + transports, err := p.d.EncodeJSON(rec.Transports) + if err != nil { + return 0, fmt.Errorf("failed to marshal transports: %w", err) + } + var id int64 + err = p.tx(ctx, func(q Querier) error { + exists, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyID). + Where(Eq(lookup.PasskeyCredentialID, rec.CredentialID)).Exists(ctx, q) + if err != nil { + return err + } + if exists { + return fmt.Errorf("credential already exists") + } + userExists, err := p.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersID, rec.UserID)).Exists(ctx, q) + if err != nil { + return err + } + if !userExists { + return fmt.Errorf("user not found") + } + now := p.Now() + id, err = p.Insert(lookup.EntityUserPasskeyCredentials).Set( + Set(lookup.PasskeyUserID, rec.UserID), + Set(lookup.PasskeyCredentialID, rec.CredentialID), + Set(lookup.PasskeyPublicKey, rec.PublicKey), + Set(lookup.PasskeyAttestationType, rec.AttestationType), + Set(lookup.PasskeyAAGUID, ""), + Set(lookup.PasskeySignCount, int64(rec.SignCount)), + Set(lookup.PasskeyTransports, transports), + Set(lookup.PasskeyBackupEligible, rec.BackupEligible), + Set(lookup.PasskeyBackupState, rec.BackupState), + Set(lookup.PasskeyName, rec.Name), + Set(lookup.PasskeyCreatedAt, now), + Set(lookup.PasskeyLastUsedAt, now), + ).ExecID(ctx, q, lookup.PasskeyID) + return err + }) + if err != nil { + return 0, err + } + return id, nil +} + +// Get implements lookup.PasskeyStore. +func (p *Passkey) Get(ctx context.Context, credentialID string) (int, uint32, error) { + var userID int + var count int64 + err := p.do(func(q Querier) error { + return p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyUserID, lookup.PasskeySignCount). + Where(Eq(lookup.PasskeyCredentialID, credentialID)).QueryRow(ctx, q, &userID, &count) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, 0, fmt.Errorf("credential not found") + } + return 0, 0, fmt.Errorf("failed to get credential: %w", err) + } + return userID, uint32(count), nil //nolint:gosec // sign counters are stored from uint32 values +} + +// UpdateCounter implements lookup.PasskeyStore. A counter that did not advance flags the +// credential as possibly cloned and leaves the stored counter unchanged. +func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) { + var clone bool + err := p.tx(ctx, func(q Querier) error { + match := Eq(lookup.PasskeyCredentialID, credentialID) + var old int64 + if err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeySignCount).Where(match).QueryRow(ctx, q, &old); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("credential not found") + } + return err + } + if int64(newCounter) <= old { + clone = true + _, err := p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyCloneWarning, true)).Where(match).Exec(ctx, q) + return err + } + _, err := p.Update(lookup.EntityUserPasskeyCredentials). + Set(Set(lookup.PasskeySignCount, int64(newCounter)), Set(lookup.PasskeyLastUsedAt, p.Now())). + Where(match).Exec(ctx, q) + return err + }) + return clone, err +} + +// List implements lookup.PasskeyStore, newest first. +func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) { + var out []sectypes.PasskeyCredential + err := p.do(func(q Querier) error { + out = make([]sectypes.PasskeyCredential, 0) + rows, err := p.From(lookup.EntityUserPasskeyCredentials). + Cols(lookup.PasskeyID, lookup.PasskeyUserID, lookup.PasskeyCredentialID, lookup.PasskeyPublicKey, + lookup.PasskeyAttestationType, lookup.PasskeyAAGUID, lookup.PasskeySignCount, lookup.PasskeyCloneWarning, + lookup.PasskeyTransports, lookup.PasskeyBackupEligible, lookup.PasskeyBackupState, lookup.PasskeyName, + lookup.PasskeyCreatedAt, lookup.PasskeyLastUsedAt). + Where(Eq(lookup.PasskeyUserID, userID)).OrderBy(lookup.PasskeyCreatedAt).Query(ctx, q) + if err != nil { + return err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var id, uid int + var credB64, pubB64 string + var attestation, aaguidB64, name sql.NullString + var count sql.NullInt64 + var clone, eligible, state bool + var transports any + var created, last time.Time + if err := rows.Scan(&id, &uid, &credB64, &pubB64, &attestation, &aaguidB64, &count, p.boolDest(&clone), + &transports, p.boolDest(&eligible), p.boolDest(&state), &name, p.timeDest(&created), p.timeDest(&last)); err != nil { + return err + } + credID, err := base64.StdEncoding.DecodeString(credB64) + if err != nil { + continue + } + pub, err := base64.StdEncoding.DecodeString(pubB64) + if err != nil { + continue + } + aaguid, _ := base64.StdEncoding.DecodeString(aaguidB64.String) + c := sectypes.PasskeyCredential{ + ID: fmt.Sprintf("%d", id), + UserID: uid, + CredentialID: credID, + PublicKey: pub, + AttestationType: attestation.String, + AAGUID: aaguid, + SignCount: uint32(count.Int64), //nolint:gosec // stored from uint32 values + CloneWarning: clone, + BackupEligible: eligible, + BackupState: state, + Name: name.String, + CreatedAt: created, + LastUsedAt: last, + } + _ = p.d.DecodeJSON(transports, &c.Transports) + out = append(out, c) + } + return rows.Err() + }) + if err != nil { + return nil, fmt.Errorf("failed to get credentials: %w", err) + } + // newest first + for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 { + out[i], out[j] = out[j], out[i] + } + return out, nil +} + +// Delete implements lookup.PasskeyStore. +func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error { + var rows int64 + err := p.do(func(q Querier) error { + var err error + rows, err = p.Base.Delete(lookup.EntityUserPasskeyCredentials). + Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q) + return err + }) + if err != nil { + return err + } + if rows == 0 { + return fmt.Errorf("credential not found") + } + return nil +} + +// Rename implements lookup.PasskeyStore. +func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error { + var rows int64 + err := p.do(func(q Querier) error { + var err error + rows, err = p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyName, name)). + Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q) + return err + }) + if err != nil { + return err + } + if rows == 0 { + return fmt.Errorf("credential not found") + } + return nil +} + +// ByUsername implements lookup.PasskeyStore. +func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) { + var userID int + var creds []lookup.PasskeyCredentialRef + err := p.do(func(q Querier) error { + creds = make([]lookup.PasskeyCredentialRef, 0) + if err := p.From(lookup.EntityUsers).Cols(lookup.UsersID). + Where(Eq(lookup.UsersUsername, username), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID); err != nil { + return err + } + rows, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyCredentialID, lookup.PasskeyTransports). + Where(Eq(lookup.PasskeyUserID, userID)).Query(ctx, q) + if err != nil { + return err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var ref lookup.PasskeyCredentialRef + var transports any + if err := rows.Scan(&ref.CredentialID, &transports); err != nil { + return err + } + _ = p.d.DecodeJSON(transports, &ref.Transports) + creds = append(creds, ref) + } + return rows.Err() + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, nil, fmt.Errorf("user not found") + } + return 0, nil, fmt.Errorf("failed to get credentials: %w", err) + } + return userID, 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) { + var u userRow + err := p.do(func(q Querier) error { + return p.From(lookup.EntityUsers). + Cols(lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, + lookup.UsersProgramUserID, lookup.UsersProgramUserTable). + Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersIsActive, true)). + QueryRow(ctx, q, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("user not found") + } + return nil, fmt.Errorf("passkey login query failed: %w", err) + } + u.id = userID + token, err := GenerateSessionToken() + if err != nil { + return nil, fmt.Errorf("failed to generate session token: %w", err) + } + now := p.Now() + ip, ua := claimStrings(claims) + err = p.tx(ctx, func(q Querier) error { + if err := p.Insert(lookup.EntityUserSessions).Set( + Set(lookup.SessionsToken, token), + Set(lookup.SessionsUserID, userID), + Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)), + Set(lookup.SessionsIPAddress, ip), + Set(lookup.SessionsUserAgent, ua), + Set(lookup.SessionsLastActivityAt, now), + Set(lookup.SessionsCreatedAt, now), + ).Exec(ctx, q); err != nil { + return err + } + _, err := p.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + return err + }) + if err != nil { + return nil, fmt.Errorf("passkey login query failed: %w", err) + } + return §ypes.LoginResponse{Token: token, User: u.context(token), ExpiresIn: int64(sessionLifetime.Seconds())}, nil +} diff --git a/pkg/security/lookup/direct/passkey_totp_test.go b/pkg/security/lookup/direct/passkey_totp_test.go new file mode 100644 index 0000000..6a1b027 --- /dev/null +++ b/pkg/security/lookup/direct/passkey_totp_test.go @@ -0,0 +1,162 @@ +package direct + +import ( + "context" + "encoding/base64" + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" +) + +func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) } + +func TestPasskeyLifecycle(t *testing.T) { + ctx := context.Background() + a, _ := newAuth(t, AuthOptions{}) + reg := registerUser(t, a, "pat") + uid := reg.User.UserID + p := NewPasskey(a.Base) + + rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: b64("cred1"), PublicKey: b64("pk"), AttestationType: "none", + Transports: []string{"usb", "nfc"}, Name: "Key 1"} + id, err := p.Store(ctx, rec) + if err != nil || id == 0 { + t.Fatalf("%d %v", id, err) + } + if _, err := p.Store(ctx, rec); err == nil || err.Error() != "credential already exists" { + t.Fatalf("dup: %v", err) + } + rec.CredentialID, rec.UserID = b64("cred2"), 999 + if _, err := p.Store(ctx, rec); err == nil || err.Error() != "user not found" { + t.Fatalf("no user: %v", err) + } + rec.UserID, rec.Name = uid, "Key 2" + if _, err := p.Store(ctx, rec); err != nil { + t.Fatal(err) + } + + owner, count, err := p.Get(ctx, b64("cred1")) + if err != nil || owner != uid || count != 0 { + t.Fatalf("%d %d %v", owner, count, err) + } + if _, _, err := p.Get(ctx, b64("zzz")); err == nil || err.Error() != "credential not found" { + t.Fatalf("got %v", err) + } + + if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || clone { + t.Fatalf("%v %v", clone, err) + } + if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || !clone { + t.Fatalf("replayed counter must flag clone: %v %v", clone, err) + } + if _, count, _ = p.Get(ctx, b64("cred1")); count != 5 { + t.Fatalf("counter changed on clone: %d", count) + } + if _, err := p.UpdateCounter(ctx, b64("missing"), 1); err == nil { + t.Fatal("expected not found") + } + + list, err := p.List(ctx, uid) + if err != nil || len(list) != 2 { + t.Fatalf("%d %v", len(list), err) + } + for _, c := range list { + if string(c.CredentialID) == "cred1" { + if !c.CloneWarning || c.SignCount != 5 || len(c.Transports) != 2 || c.Name != "Key 1" { + t.Fatalf("%+v", c) + } + } + } + + if err := p.Rename(ctx, uid, b64("cred1"), "Renamed"); err != nil { + t.Fatal(err) + } + if err := p.Rename(ctx, uid+1, b64("cred1"), "x"); err == nil { + t.Fatal("renamed another user's credential") + } + + gotID, refs, err := p.ByUsername(ctx, "pat") + if err != nil || gotID != uid || len(refs) != 2 { + t.Fatalf("%d %+v %v", gotID, refs, err) + } + if _, _, err := p.ByUsername(ctx, "ghost"); err == nil || err.Error() != "user not found" { + t.Fatalf("got %v", err) + } + + resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "1.1.1.1"}) + if err != nil || resp.User.UserName != "pat" || resp.ExpiresIn != 86400 { + t.Fatalf("%+v %v", resp, err) + } + if _, err := a.Session(ctx, resp.Token, ""); err != nil { + t.Fatal(err) + } + + if err := p.Delete(ctx, uid+1, b64("cred1")); err == nil { + t.Fatal("deleted another user's credential") + } + if err := p.Delete(ctx, uid, b64("cred1")); err != nil { + t.Fatal(err) + } + if err := p.Delete(ctx, uid, b64("cred1")); err == nil || err.Error() != "credential not found" { + t.Fatalf("got %v", err) + } +} + +func TestTOTPLifecycle(t *testing.T) { + ctx := context.Background() + a, _ := newAuth(t, AuthOptions{}) + uid := registerUser(t, a, "tom").User.UserID + s := NewTOTP(a.Base) + + if on, err := s.Status(ctx, uid); err != nil || on { + t.Fatalf("%v %v", on, err) + } + if _, err := s.Secret(ctx, uid); err == nil || err.Error() != "TOTP not enabled for user" { + t.Fatalf("got %v", err) + } + if err := s.RegenerateBackupCodes(ctx, uid, []string{"h"}); err == nil { + t.Fatal("regenerate without 2FA") + } + if err := s.Enable(ctx, 999, "S", nil); err == nil || err.Error() != "user not found" { + t.Fatalf("got %v", err) + } + + if err := s.Enable(ctx, uid, "SECRET", []string{"h1", "h2"}); err != nil { + t.Fatal(err) + } + if on, _ := s.Status(ctx, uid); !on { + t.Fatal("not enabled") + } + if sec, err := s.Secret(ctx, uid); err != nil || sec != "SECRET" { + t.Fatalf("%q %v", sec, err) + } + + if ok, err := s.ValidateBackupCode(ctx, uid, "h1"); err != nil || !ok { + t.Fatalf("%v %v", ok, err) + } + if _, err := s.ValidateBackupCode(ctx, uid, "h1"); err == nil || err.Error() != "backup code already used" { + t.Fatalf("reuse: %v", err) + } + if ok, err := s.ValidateBackupCode(ctx, uid, "nope"); err != nil || ok { + t.Fatalf("%v %v", ok, err) + } + if err := s.RegenerateBackupCodes(ctx, uid, []string{"n1"}); err != nil { + t.Fatal(err) + } + if ok, _ := s.ValidateBackupCode(ctx, uid, "h2"); ok { + t.Fatal("old code survived regenerate") + } + if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); !ok { + t.Fatal("new code rejected") + } + + if err := s.Disable(ctx, uid); err != nil { + t.Fatal(err) + } + if on, _ := s.Status(ctx, uid); on { + t.Fatal("still enabled") + } + if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); ok { + t.Fatal("codes survived disable") + } +} diff --git a/pkg/security/password.go b/pkg/security/lookup/direct/password.go similarity index 66% rename from pkg/security/password.go rename to pkg/security/lookup/direct/password.go index 6fabd4b..a44cf95 100644 --- a/pkg/security/password.go +++ b/pkg/security/lookup/direct/password.go @@ -1,4 +1,4 @@ -package security +package direct import ( "crypto/subtle" @@ -15,7 +15,8 @@ const maxPasswordBytes = 72 var errPasswordTooLong = errors.New("password must be at most 72 bytes") -func hashPassword(password string) (string, error) { +// HashPassword returns the bcrypt hash of password. +func HashPassword(password string) (string, error) { if len(password) > maxPasswordBytes { return "", errPasswordTooLong } @@ -30,12 +31,11 @@ func isBcryptHash(s string) bool { return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$") } -// verifyPassword checks supplied against the stored value. A stored bcrypt hash -// is compared with bcrypt. A legacy cleartext value (written before hashing was -// implemented) is compared in constant time and, on a match, needsRehash is true -// so the caller can upgrade the row to a bcrypt hash. An empty stored value -// (e.g. an OAuth2-only user) never matches. -func verifyPassword(stored, supplied string) (ok, needsRehash bool) { +// VerifyPassword checks supplied against the stored value. A stored bcrypt hash is compared +// with bcrypt. A legacy cleartext value (written before hashing was implemented) is compared +// in constant time and, on a match, needsRehash is true so the caller can upgrade the row. +// An empty stored value (e.g. an OAuth2-only user) never matches. +func VerifyPassword(stored, supplied string) (ok, needsRehash bool) { if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes { return false, false } @@ -53,9 +53,9 @@ var ( dummyHash string ) -// burnPasswordCheck spends roughly one bcrypt comparison so an unknown username -// costs about the same as a wrong password. -func burnPasswordCheck(supplied string) { +// BurnPasswordCheck spends roughly one bcrypt comparison so an unknown username costs about +// the same as a wrong password. +func BurnPasswordCheck(supplied string) { dummyHashOnce.Do(func() { h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost) dummyHash = string(h) diff --git a/pkg/security/lookup/direct/password_test.go b/pkg/security/lookup/direct/password_test.go new file mode 100644 index 0000000..10c37c2 --- /dev/null +++ b/pkg/security/lookup/direct/password_test.go @@ -0,0 +1,19 @@ +package direct + +import "testing" + +func TestVerifyPasswordEdgeCases(t *testing.T) { + h, _ := HashPassword("pw") + if ok, _ := VerifyPassword(h, "pw"); !ok { + t.Error("bcrypt match failed") + } + if ok, _ := VerifyPassword("", "pw"); ok { + t.Error("empty stored must not match") + } + if ok, _ := VerifyPassword("pw", ""); ok { + t.Error("empty supplied must not match") + } + if _, err := HashPassword(string(make([]byte, 73))); err == nil { + t.Error("73-byte password must be rejected") + } +} diff --git a/pkg/security/lookup/direct/policy.go b/pkg/security/lookup/direct/policy.go new file mode 100644 index 0000000..8ff7971 --- /dev/null +++ b/pkg/security/lookup/direct/policy.go @@ -0,0 +1,200 @@ +package direct + +import ( + "context" + "database/sql" + "fmt" + "strconv" + "strings" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// PolicyOptions tunes Policy. +type PolicyOptions struct { + // NoGroups skips the group membership table: only rules addressed to the user directly + // apply. Use it when the sec_group_members table is not deployed. + NoGroups bool +} + +// Policy implements lookup.PolicyStore on the rule tables. +// +// Applicable rules are the active rules whose user_id is the caller plus the rules of every +// group the caller belongs to; schema and table match case-insensitively and exactly (never a +// prefix). Column security returns the union of the matching rules. Row security: any +// applicable has_block rule wins, otherwise the templates are combined with AND, each in +// parentheses. No rule is an empty result; failures are errors so callers fail closed. +type Policy struct { + *Base + opts PolicyOptions +} + +var _ lookup.PolicyStore = (*Policy)(nil) + +// NewPolicy creates the direct PolicyStore. +func NewPolicy(b *Base, opts PolicyOptions) *Policy { return &Policy{Base: b, opts: opts} } + +// applicable restricts a rule query to the rules that apply to userID. +func (p *Policy) applicable(userCol, groupCol lookup.Column, userID int64) Cond { + if p.opts.NoGroups { + return Eq(userCol, userID) + } + members := p.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, userID)) + return Or(Eq(userCol, userID), InSelect(groupCol, members)) +} + +// ColumnSecurity implements lookup.PolicyStore. +func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) { + var rules []sectypes.ColumnSecurity + err := p.do(func(q Querier) error { + rules = nil + rows, err := p.From(lookup.EntitySecColumnRules). + Cols(lookup.ColRulesID, lookup.ColRulesColumnPath, lookup.ColRulesAccessType, lookup.ColRulesMaskStart, + lookup.ColRulesMaskEnd, lookup.ColRulesMaskInvert, lookup.ColRulesMaskChar, lookup.ColRulesExtraFilters). + Where( + Eq(lookup.ColRulesIsActive, true), + EqFold(lookup.ColRulesSchemaName, schema), + EqFold(lookup.ColRulesTableName, table), + p.applicable(lookup.ColRulesUserID, lookup.ColRulesGroupID, int64(userID)), + ).OrderBy(lookup.ColRulesID).Query(ctx, q) + if err != nil { + return err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var id int + var path, access string + var start, end sql.NullInt64 + var invert sql.NullBool + var maskChar sql.NullString + var extra any + var inv any + if err := rows.Scan(&id, &path, &access, &start, &end, &inv, &maskChar, &extra); err != nil { + return err + } + if inv != nil { + b, err := p.d.ScanBool(inv) + if err != nil { + return err + } + invert = sql.NullBool{Bool: b, Valid: true} + } + rule := sectypes.ColumnSecurity{ + ID: id, + Schema: schema, + Tablename: table, + Path: strings.Split(path, "."), + Accesstype: access, + UserID: userID, + MaskStart: int(start.Int64), + MaskEnd: int(end.Int64), + MaskInvert: invert.Bool, + MaskChar: "*", + Control: schema + "." + table + "." + path, + } + if maskChar.Valid && maskChar.String != "" { + rule.MaskChar = maskChar.String + } + if err := p.d.DecodeJSON(extra, &rule.ExtraFilters); err != nil { + return err + } + rules = append(rules, rule) + } + return rows.Err() + }) + if err != nil { + return nil, fmt.Errorf("failed to load column security: %w", err) + } + return rules, nil +} + +// numericUser reduces a user reference to the integer the rule tables key on. Structured +// values are rejected, and so are non-numeric strings: a reference that cannot be matched +// must fail closed rather than silently load no rules. +func numericUser(ref any) (int64, error) { + switch v := ref.(type) { + case *sectypes.UserContext: + if v == nil { + return 0, fmt.Errorf("row security: nil user context") + } + return int64(v.UserID), nil + case sectypes.UserContext: + return int64(v.UserID), nil + case int: + return int64(v), nil + case int8: + return int64(v), nil + case int16: + return int64(v), nil + case int32: + return int64(v), nil + case int64: + return v, nil + case uint: + return int64(v), nil //nolint:gosec // user ids fit int64 + case uint8: + return int64(v), nil + case uint16: + return int64(v), nil + case uint32: + return int64(v), nil + case uint64: + return int64(v), nil //nolint:gosec // user ids fit int64 + case string: + n, err := strconv.ParseInt(strings.TrimSpace(v), 10, 64) + if err != nil { + return 0, fmt.Errorf("row security: user reference %q is not a numeric id", v) + } + return n, nil + case nil: + return 0, fmt.Errorf("row security: no user reference") + } + return 0, fmt.Errorf("row security: unsupported user reference type %T", ref) +} + +// RowSecurity implements lookup.PolicyStore. +func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) { + uid, err := numericUser(userRef) + if err != nil { + return sectypes.RowSecurity{}, err + } + var templates []string + block := false + err = p.do(func(q Querier) error { + templates, block = nil, false + rows, err := p.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesTemplate, lookup.RowRulesHasBlock). + Where( + Eq(lookup.RowRulesIsActive, true), + EqFold(lookup.RowRulesSchemaName, schema), + EqFold(lookup.RowRulesTableName, table), + p.applicable(lookup.RowRulesUserID, lookup.RowRulesGroupID, uid), + ).OrderBy(lookup.RowRulesID).Query(ctx, q) + if err != nil { + return err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var tpl sql.NullString + var hb bool + if err := rows.Scan(&tpl, p.boolDest(&hb)); err != nil { + return err + } + if hb { + block = true + } + if t := strings.TrimSpace(tpl.String); t != "" { + templates = append(templates, "("+t+")") + } + } + return rows.Err() + }) + if err != nil { + return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err) + } + rs := sectypes.RowSecurity{Schema: schema, Tablename: table, UserID: userRef, HasBlock: block} + if !block { + rs.Template = strings.Join(templates, " AND ") + } + return rs, nil +} diff --git a/pkg/security/lookup/direct/policy_test.go b/pkg/security/lookup/direct/policy_test.go new file mode 100644 index 0000000..5c9e89d --- /dev/null +++ b/pkg/security/lookup/direct/policy_test.go @@ -0,0 +1,93 @@ +package direct + +import ( + "context" + "testing" +) + +func TestPolicyColumnAndRowSecurity(t *testing.T) { + ctx := context.Background() + db := newTestDB(t) + p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{}) + exec := func(q string, args ...any) { + t.Helper() + if _, err := db.Exec(q, args...); err != nil { + t.Fatal(err) + } + } + exec(`INSERT INTO sec_group_members (group_id, user_id) VALUES (10, 1), (10, 3)`) + + // column rules: user 1 direct, group 10 (user 1 and 3), another user, inactive, other table, prefix table + exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, mask_start, mask_end, mask_invert, mask_char, extra_filters, is_active) + VALUES (1, NULL, 'Public', 'Users', 'email', 'mask', 2, 1, 1, '#', '{"k":"v"}', 1), + (NULL, 10, 'public', 'users', 'profile.ssn', 'hide', NULL, NULL, NULL, NULL, NULL, 1), + (2, NULL, 'public', 'users', 'other', 'hide', 0, 0, 0, '*', NULL, 1), + (1, NULL, 'public', 'users', 'off', 'hide', 0, 0, 0, '*', NULL, 0), + (1, NULL, 'public', 'orders', 'x', 'hide', 0, 0, 0, '*', NULL, 1), + (1, NULL, 'public', 'users_archive', 'y', 'hide', 0, 0, 0, '*', NULL, 1)`) + + rules, err := p.ColumnSecurity(ctx, 1, "public", "users") + if err != nil || len(rules) != 2 { + t.Fatalf("%d %v %+v", len(rules), err, rules) + } + m := rules[0] + if m.Accesstype != "mask" || m.MaskStart != 2 || m.MaskEnd != 1 || !m.MaskInvert || m.MaskChar != "#" || + m.ExtraFilters["k"] != "v" || len(m.Path) != 1 || m.Path[0] != "email" || m.UserID != 1 { + t.Fatalf("%+v", m) + } + h := rules[1] + if len(h.Path) != 2 || h.Path[1] != "ssn" || h.MaskChar != "*" || h.Accesstype != "hide" { + t.Fatalf("%+v", h) + } + if r, err := p.ColumnSecurity(ctx, 3, "public", "users"); err != nil || len(r) != 1 { + t.Fatalf("group member: %d %v", len(r), err) + } + if r, err := p.ColumnSecurity(ctx, 99, "public", "users"); err != nil || len(r) != 0 { + t.Fatalf("no rules must be empty: %d %v", len(r), err) + } + + // row rules + exec(`INSERT INTO sec_row_rules (user_id, group_id, schema_name, table_name, template, has_block, is_active) VALUES + (1, NULL, 'public', 'orders', 'owner_id = {UserID}', 0, 1), + (NULL, 10, 'public', 'orders', 'region = 1', 0, 1), + (NULL, 10, 'public', 'orders', 'ignored', 0, 0), + (3, NULL, 'public', 'secret', NULL, 1, 1), + (NULL, 10, 'public', 'secret', 'x = 1', 0, 1)`) + rs, err := p.RowSecurity(ctx, 1, "public", "orders") + if err != nil || rs.Template != "(owner_id = {UserID}) AND (region = 1)" || rs.HasBlock { + t.Fatalf("%+v %v", rs, err) + } + if rs, err := p.RowSecurity(ctx, "3", "PUBLIC", "Secret"); err != nil || !rs.HasBlock || rs.Template != "" { + t.Fatalf("block must win: %+v %v", rs, err) + } + if rs, err := p.RowSecurity(ctx, 99, "public", "orders"); err != nil || rs.Template != "" || rs.HasBlock { + t.Fatalf("%+v %v", rs, err) + } + for _, bad := range []any{nil, "abc", []int{1}, 1.5} { + if _, err := p.RowSecurity(ctx, bad, "public", "orders"); err == nil { + t.Fatalf("user ref %#v accepted", bad) + } + } +} + +func TestPolicyNoGroups(t *testing.T) { + ctx := context.Background() + db := newTestDB(t) + p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{NoGroups: true}) + _, _ = db.Exec(`DROP TABLE sec_group_members`) + _, _ = db.Exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, is_active) VALUES + (1, NULL, 's', 't', 'a', 'hide', 1), (NULL, 5, 's', 't', 'b', 'hide', 1)`) + rules, err := p.ColumnSecurity(ctx, 1, "s", "t") + if err != nil || len(rules) != 1 || rules[0].Path[0] != "a" { + t.Fatalf("%+v %v", rules, err) + } +} + +func TestPolicyFailsClosedOnMissingTable(t *testing.T) { + db := newTestDB(t) + p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{}) + _, _ = db.Exec(`DROP TABLE sec_row_rules`) + if _, err := p.RowSecurity(context.Background(), 1, "s", "t"); err == nil { + t.Fatal("expected error for missing table") + } +} diff --git a/pkg/security/lookup/direct/totp.go b/pkg/security/lookup/direct/totp.go new file mode 100644 index 0000000..4ad849a --- /dev/null +++ b/pkg/security/lookup/direct/totp.go @@ -0,0 +1,159 @@ +package direct + +import ( + "context" + "database/sql" + "errors" + "fmt" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" +) + +// TOTP implements lookup.TOTPStore: the secret and enabled flag live on the users table, +// backup code hashes in their own table. +type TOTP struct{ *Base } + +var _ lookup.TOTPStore = (*TOTP)(nil) + +// NewTOTP creates the direct TOTPStore. +func NewTOTP(b *Base) *TOTP { return &TOTP{Base: b} } + +func (t *TOTP) replaceBackupCodes(ctx context.Context, q Querier, userID int, hashed []string) error { + if _, err := t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q); err != nil { + return err + } + now := t.Now() + for _, h := range hashed { + if err := t.Insert(lookup.EntityUserTOTPBackupCodes).Set( + Set(lookup.BackupCodesUserID, userID), + Set(lookup.BackupCodesCodeHash, h), + Set(lookup.BackupCodesUsed, false), + Set(lookup.BackupCodesCreatedAt, now), + ).Exec(ctx, q); err != nil { + return err + } + } + return nil +} + +// Enable implements lookup.TOTPStore. +func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error { + return t.tx(ctx, func(q Querier) error { + n, err := t.Update(lookup.EntityUsers).Set( + Set(lookup.UsersTOTPSecret, secret), + Set(lookup.UsersTOTPEnabled, true), + Set(lookup.UsersTOTPEnabledAt, t.Now()), + ).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + if err != nil { + return err + } + if n == 0 { + return fmt.Errorf("user not found") + } + return t.replaceBackupCodes(ctx, q, userID, hashedCodes) + }) +} + +// Disable implements lookup.TOTPStore. +func (t *TOTP) Disable(ctx context.Context, userID int) error { + return t.tx(ctx, func(q Querier) error { + n, err := t.Update(lookup.EntityUsers).Set(Set(lookup.UsersTOTPSecret, nil), Set(lookup.UsersTOTPEnabled, false)). + Where(Eq(lookup.UsersID, userID)).Exec(ctx, q) + if err != nil { + return err + } + if n == 0 { + return fmt.Errorf("user not found") + } + _, err = t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q) + return err + }) +} + +// Status implements lookup.TOTPStore. +func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) { + var enabled bool + err := t.do(func(q Querier) error { + return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)). + QueryRow(ctx, q, t.boolDest(&enabled)) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return false, fmt.Errorf("user not found") + } + return false, fmt.Errorf("get 2FA status query failed: %w", err) + } + return enabled, nil +} + +// Secret implements lookup.TOTPStore. +func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) { + var secret sql.NullString + var enabled bool + err := t.do(func(q Querier) error { + return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPSecret, lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)). + QueryRow(ctx, q, &secret, t.boolDest(&enabled)) + }) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", fmt.Errorf("user not found") + } + return "", fmt.Errorf("get 2FA secret query failed: %w", err) + } + if !enabled { + return "", fmt.Errorf("TOTP not enabled for user") + } + return secret.String, nil +} + +// RegenerateBackupCodes implements lookup.TOTPStore. +func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error { + return t.tx(ctx, func(q Querier) error { + ok, err := t.From(lookup.EntityUsers).Cols(lookup.UsersID). + Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersTOTPEnabled, true)).Exists(ctx, q) + if err != nil { + return err + } + if !ok { + return fmt.Errorf("user not found or TOTP not enabled") + } + return t.replaceBackupCodes(ctx, q, userID, hashedCodes) + }) +} + +// ValidateBackupCode implements lookup.TOTPStore. An unknown code is (false, nil); a used +// code is an error. The code is consumed with a conditional update so it cannot be spent twice. +func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) { + var valid bool + err := t.tx(ctx, func(q Querier) error { + var id int64 + var used bool + err := t.From(lookup.EntityUserTOTPBackupCodes).Cols(lookup.BackupCodesID, lookup.BackupCodesUsed). + Where(Eq(lookup.BackupCodesUserID, userID), Eq(lookup.BackupCodesCodeHash, codeHash)). + QueryRow(ctx, q, &id, t.boolDest(&used)) + if errors.Is(err, sql.ErrNoRows) { + return nil + } + if err != nil { + return err + } + if used { + return fmt.Errorf("backup code already used") + } + n, err := t.Update(lookup.EntityUserTOTPBackupCodes). + Set(Set(lookup.BackupCodesUsed, true), Set(lookup.BackupCodesUsedAt, t.Now())). + Where(Eq(lookup.BackupCodesID, id), Eq(lookup.BackupCodesUsed, false)).Exec(ctx, q) + if err != nil { + return err + } + if n == 0 { + return fmt.Errorf("backup code already used") + } + valid = true + return nil + }) + if err != nil { + return false, err + } + return valid, nil +} diff --git a/pkg/security/keystore_schema.sql b/pkg/security/lookup/keystore_schema.sql similarity index 67% rename from pkg/security/keystore_schema.sql rename to pkg/security/lookup/keystore_schema.sql index 5e527ef..83bbd86 100644 --- a/pkg/security/keystore_schema.sql +++ b/pkg/security/lookup/keystore_schema.sql @@ -83,7 +83,7 @@ BEGIN p_request->>'scopes', p_request->'meta', CASE WHEN p_request->>'expires_at' IS NOT NULL - THEN (p_request->>'expires_at')::TIMESTAMP + THEN (p_request->>'expires_at')::timestamptz::timestamp ELSE NULL END ) @@ -185,3 +185,81 @@ EXCEPTION WHEN OTHERS THEN RETURN QUERY SELECT false, SQLERRM, NULL::JSONB; END; $$; + +-- resolvespec_login_api_key - Exchanges a raw API key for a session +-- Input: p_request jsonb {api_key: string, claims: object} +-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb) +-- Requires pgcrypto (digest) and the user_keys table (keystore_schema.sql). +-- Only header_api / api keys are accepted. Unknown, expired and inactive keys +-- all return the same generic error. +CREATE OR REPLACE FUNCTION resolvespec_login_api_key(p_request jsonb) +RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$ +DECLARE + v_raw_key TEXT; + v_key_id BIGINT; + v_user_id INTEGER; + v_username TEXT; + v_email TEXT; + v_user_level INTEGER; + v_roles TEXT; + v_program_user_id INTEGER; + v_program_user_table TEXT; + v_session_token TEXT; + v_expires_at TIMESTAMP; + v_ip_address TEXT; + v_user_agent TEXT; +BEGIN + v_raw_key := p_request->>'api_key'; + v_ip_address := p_request->'claims'->>'ip_address'; + v_user_agent := p_request->'claims'->>'user_agent'; + + IF v_raw_key IS NULL OR v_raw_key = '' THEN + RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb; + RETURN; + END IF; + + SELECT k.id, u.id, u.username, u.email, u.user_level, u.roles, u.program_user_id, u.program_user_table + INTO v_key_id, v_user_id, v_username, v_email, v_user_level, v_roles, v_program_user_id, v_program_user_table + FROM user_keys k + JOIN users u ON u.id = k.user_id + WHERE k.key_hash = encode(digest(v_raw_key, 'sha256'), 'hex') + AND k.key_type IN ('header_api', 'api') + AND k.is_active = true + AND (k.expires_at IS NULL OR k.expires_at > now()) + AND u.is_active = true; + + IF NOT FOUND THEN + RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb; + RETURN; + END IF; + + v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text; + v_expires_at := now() + interval '24 hours'; + + INSERT INTO user_sessions (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at) + VALUES (v_session_token, v_user_id, v_expires_at, v_ip_address, v_user_agent, now()); + + UPDATE user_keys SET last_used_at = now() WHERE id = v_key_id; + UPDATE users SET last_login_at = now() WHERE id = v_user_id; + + RETURN QUERY SELECT + true, + NULL::text, + jsonb_build_object( + 'token', v_session_token, + 'user', jsonb_build_object( + 'user_id', v_user_id, + 'user_name', v_username, + 'email', v_email, + 'user_level', v_user_level, + 'roles', string_to_array(COALESCE(v_roles, ''), ','), + 'session_id', v_session_token, + 'program_user_id', COALESCE(v_program_user_id, 0), + 'program_user_table', COALESCE(v_program_user_table, '') + ), + 'expires_in', 86400 + ); +EXCEPTION WHEN OTHERS THEN + RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb; +END; +$$ LANGUAGE plpgsql; diff --git a/pkg/security/lookup/lookup.go b/pkg/security/lookup/lookup.go new file mode 100644 index 0000000..7daf022 --- /dev/null +++ b/pkg/security/lookup/lookup.go @@ -0,0 +1,208 @@ +// Package lookup owns every database read and write the security package needs. +// pkg/security itself contains no SQL: it calls the store interfaces defined here. +// +// Each store has a procedure implementation (stored procedures, the Postgres default) +// and a direct implementation (tables through a dialect-driven query builder). Which +// one runs is decided per operation by Config.EffectiveMode. +package lookup + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// AuthStore covers sessions, login, registration and password reset. +type AuthStore interface { + Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) + Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) + Logout(ctx context.Context, req sectypes.LogoutRequest) error + // Session resolves a session token to its user. reference says where the token came + // from ("authenticate", "cookie", "refresh"); the procedure backend passes it through. + Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) + // TouchSession records last activity for a session token. user is the context the + // session resolved to; the procedure backend passes it to the update procedure. + TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error + Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) + // LoginAPIKey logs in with a raw header/generic API key. Unknown, expired, inactive and + // wrong-type keys all return the same error. + LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) + JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) + JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error + ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) + ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error +} + +// ErrInvalidAPIKey is the single error LoginAPIKey returns for unknown, expired, inactive +// and wrong-type keys, so callers cannot tell them apart. +var ErrInvalidAPIKey = errors.New("invalid api key") + +// KeyStore persists per-user auth keys. Hashing and raw-key generation happen in Go, +// so the store only sees key hashes. +type KeyStore interface { + Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) + List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) + // Delete soft-deletes a key after verifying ownership and returns its hash so callers can + // invalidate caches. The hash is empty when the backend cannot report it. + Delete(ctx context.Context, userID int, keyID int64) (keyHash string, err error) + Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) +} + +// OAuthClientStore persists the OAuth2 authorization server state (RFC 7591 clients, +// authorization codes, token introspection and revocation). +type OAuthClientStore interface { + RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) + GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) + SaveCode(ctx context.Context, code *sectypes.OAuthCode) error + // ExchangeCode atomically consumes an authorization code. + ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) + Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) + Revoke(ctx context.Context, token string) error +} + +// OAuthSession is the session row written after an OAuth2 client login. +type OAuthSession struct { + SessionToken string + UserID int + AccessToken string + RefreshToken string + TokenType string + ExpiresAt time.Time + Provider string +} + +// OAuthRefreshSession is the stored token state needed to refresh an OAuth2 login. +type OAuthRefreshSession struct { + UserID int `json:"user_id"` + AccessToken string `json:"access_token"` + TokenType string `json:"token_type"` + Expiry time.Time `json:"expiry"` +} + +// OAuthUserStore persists users and sessions created through OAuth2 client login. +type OAuthUserStore interface { + GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) + CreateSession(ctx context.Context, session OAuthSession) error + GetByRefreshToken(ctx context.Context, refreshToken string) (*OAuthRefreshSession, error) + UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error + GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) +} + +// PasskeyCredentialRecord is a credential as persisted: byte fields are base64 text. +type PasskeyCredentialRecord struct { + UserID int + CredentialID string // base64 + PublicKey string // base64 + AttestationType string + SignCount uint32 + Transports []string + BackupEligible bool + BackupState bool + Name string +} + +// PasskeyCredentialRef is a credential id and transports, as returned for a username lookup. +type PasskeyCredentialRef struct { + CredentialID string `json:"credential_id"` + Transports []string `json:"transports"` +} + +// PasskeyStore persists WebAuthn credentials. +type PasskeyStore interface { + Store(ctx context.Context, rec PasskeyCredentialRecord) (int64, error) + // Get returns the owner and signature counter of a credential. + Get(ctx context.Context, credentialID string) (userID int, signCount uint32, err error) + UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (cloneWarning bool, err error) + List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) + Delete(ctx context.Context, userID int, credentialID string) error + Rename(ctx context.Context, userID int, credentialID, name string) error + ByUsername(ctx context.Context, username string) (userID int, creds []PasskeyCredentialRef, err error) + Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) +} + +// TOTPStore persists two-factor state. Backup codes arrive already hashed. +type TOTPStore interface { + Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error + Disable(ctx context.Context, userID int) error + Status(ctx context.Context, userID int) (bool, error) + Secret(ctx context.Context, userID int) (string, error) + RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error + ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) +} + +// PolicyStore loads column and row security rules. No rules is an empty result, never +// an error; failures are errors so callers fail closed. +type PolicyStore interface { + ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) + RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) +} + +// Provider bundles every store. Security constructors take a Provider. +type Provider struct { + Auth AuthStore + Keys KeyStore + OAuthClient OAuthClientStore + OAuthUser OAuthUserStore + Passkey PasskeyStore + TOTP TOTPStore + Policy PolicyStore +} + +// Config selects dialect, mode and naming. The zero value is valid: dialect detected from +// the driver, default mode per dialect, default procedure/table/column names. +type Config struct { + // Dialect names a registered dialect ("postgres", "sqlite", "mysql", "mssql", or one added + // with dialect.Register). Empty = detect from the driver. + Dialect string + // Mode is the default mode for every operation. See ModeDefault. + Mode Mode + // Overrides sets the mode per operation, e.g. direct for OpSession, procedure for OpLogin. + Overrides map[Op]Mode + // Procs overrides procedure names; empty fields keep the default. + Procs ProcNames + // Schema overrides table and column names; missing entries keep the default. + Schema Schema +} + +// Resolved is a Config merged with defaults and validated. +type Resolved struct { + Config + Procs ProcNames + Schema Schema +} + +// Resolve merges c with the defaults and validates the result. +func (c Config) Resolve() (*Resolved, error) { + if !c.Mode.valid() { + return nil, fmt.Errorf("lookup: invalid mode %q", c.Mode) + } + for op, m := range c.Overrides { + if !m.valid() { + return nil, fmt.Errorf("lookup: invalid mode %q for %s", m, op) + } + } + if c.Dialect != "" { + if _, err := dialect.Get(c.Dialect); err != nil { + return nil, fmt.Errorf("lookup: %w", err) + } + } + procs := DefaultProcNames().Merge(c.Procs) + if err := procs.Validate(); err != nil { + return nil, err + } + schema := DefaultSchema().Merge(c.Schema) + if err := schema.Validate(); err != nil { + return nil, err + } + return &Resolved{Config: c, Procs: procs, Schema: schema}, nil +} + +// Registration conflicts reported by AuthStore.Register in direct mode. +var ( + ErrUsernameExists = errors.New("username already exists") + ErrEmailExists = errors.New("email already exists") +) diff --git a/pkg/security/lookup/lookup_test.go b/pkg/security/lookup/lookup_test.go new file mode 100644 index 0000000..f0f30be --- /dev/null +++ b/pkg/security/lookup/lookup_test.go @@ -0,0 +1,165 @@ +package lookup + +import ( + "regexp" + "strings" + "testing" + + ddlpkg "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl" +) + +// TestDefaultSchemaMatchesSQLiteDDL keeps the default schema in step with the reference DDL: +// every table and column in the DDL must be a known logical column and vice versa. +func TestDefaultSchemaMatchesSQLiteDDL(t *testing.T) { + ddl, err := ddlpkg.SQL("sqlite") + if err != nil { + t.Fatal(err) + } + tableRe := regexp.MustCompile(`(?s)CREATE TABLE IF NOT EXISTS (\w+) \((.*?)\n\);`) + colRe := regexp.MustCompile(`^\s*(\w+)\s+[A-Z]+`) + def := DefaultSchema() + seen := map[Entity]bool{} + for _, m := range tableRe.FindAllStringSubmatch(string(ddl), -1) { + e := Entity(m[1]) + tbl, ok := def[e] + if !ok { + t.Errorf("DDL table %s has no entity", e) + continue + } + seen[e] = true + cols := map[string]bool{} + for _, line := range strings.Split(m[2], "\n") { + if cm := colRe.FindStringSubmatch(line); cm != nil && cm[1] != "PRIMARY" && cm[1] != "CHECK" && cm[1] != "FOREIGN" { + cols[cm[1]] = true + if _, ok := tbl.Columns[cm[1]]; !ok { + t.Errorf("DDL column %s.%s is not a logical column", e, cm[1]) + } + } + } + for c := range tbl.Columns { + if !cols[c] { + t.Errorf("logical column %s.%s is not in the DDL", e, c) + } + } + } + for e := range def { + if !seen[e] { + t.Errorf("entity %s is not in the DDL", e) + } + } +} + +func TestDefaultSchemaValid(t *testing.T) { + if err := DefaultSchema().Validate(); err != nil { + t.Fatal(err) + } + if err := DefaultProcNames().Validate(); err != nil { + t.Fatal(err) + } +} + +func TestSchemaMergeAndLookup(t *testing.T) { + cfg := Config{Schema: Schema{ + EntityUsers: {Name: "app_users", Schema: "auth", Columns: map[string]string{"username": "login_name"}}, + }} + r, err := cfg.Resolve() + if err != nil { + t.Fatal(err) + } + if got := r.Schema.TableName(EntityUsers); got != "app_users" { + t.Errorf("table = %q", got) + } + if got := r.Schema.SchemaName(EntityUsers); got != "auth" { + t.Errorf("schema = %q", got) + } + if got := r.Schema.Col(UsersUsername); got != "login_name" { + t.Errorf("username col = %q", got) + } + if got := r.Schema.Col(UsersEmail); got != "email" { + t.Errorf("email col = %q, want default", got) + } + // Merge must not mutate the defaults. + if DefaultSchema().Col(UsersUsername) != "username" { + t.Error("default schema was mutated") + } +} + +func TestSchemaZeroValueUsesDefaults(t *testing.T) { + var s Schema + if s.TableName(EntityUserKeys) != "user_keys" || s.Col(KeysKeyHash) != "key_hash" { + t.Error("zero schema should fall back to defaults") + } +} + +func TestResolveRejectsBadConfig(t *testing.T) { + bad := map[string]Config{ + "table injection": {Schema: Schema{EntityUsers: {Name: "users; DROP TABLE users"}}}, + "column injection": {Schema: Schema{EntityUsers: {Columns: map[string]string{"id": "id) --"}}}}, + "schema injection": {Schema: Schema{EntityUsers: {Schema: "a.b"}}}, + "unknown entity": {Schema: Schema{"nope": {Name: "x"}}}, + "unknown column": {Schema: Schema{EntityUsers: {Columns: map[string]string{"nope": "x"}}}}, + "proc injection": {Procs: ProcNames{Login: "f(); --"}}, + "bad mode": {Mode: "sometimes"}, + "bad override": {Overrides: map[Op]Mode{OpLogin: "x"}}, + "bad dialect": {Dialect: "oracle"}, + } + for name, cfg := range bad { + if _, err := cfg.Resolve(); err == nil { + t.Errorf("%s: expected error", name) + } + } + // A single schema qualifier is allowed on table and procedure names. + ok := Config{ + Schema: Schema{EntityUsers: {Name: "auth.users"}}, + Procs: ProcNames{Login: "auth.resolvespec_login"}, + } + if _, err := ok.Resolve(); err != nil { + t.Errorf("qualified names should be valid: %v", err) + } +} + +func TestProcNamesMerge(t *testing.T) { + m := DefaultProcNames().Merge(ProcNames{Login: "custom_login"}) + if m.Login != "custom_login" { + t.Errorf("Login = %q", m.Login) + } + if m.Register != "resolvespec_register" { + t.Errorf("Register = %q, want default", m.Register) + } + if DefaultProcNames().LoginAPIKey != "resolvespec_login_api_key" { + t.Error("LoginAPIKey default missing") + } + if DefaultProcNames().KeystoreValidateKey != "resolvespec_keystore_validate_key" { + t.Error("keystore defaults missing") + } +} + +func TestEffectiveMode(t *testing.T) { + cases := []struct { + name string + cfg Config + dialect string + want Mode + wantErr bool + }{ + {"pg default", Config{}, DialectPostgres, ModeProcedure, false}, + {"sqlite default", Config{}, DialectSQLite, ModeDirect, false}, + {"mysql default", Config{}, DialectMySQL, ModeDirect, false}, + {"pg direct", Config{Mode: ModeDirect}, DialectPostgres, ModeDirect, false}, + {"pg auto probes", Config{Mode: ModeAuto}, DialectPostgres, ModeAuto, false}, + {"sqlite auto is direct", Config{Mode: ModeAuto}, DialectSQLite, ModeDirect, false}, + {"sqlite procedure rejected", Config{Mode: ModeProcedure}, DialectSQLite, "", true}, + {"override wins", Config{Overrides: map[Op]Mode{OpSession: ModeDirect}}, DialectPostgres, ModeDirect, false}, + {"override only for its op", Config{Overrides: map[Op]Mode{OpSession: ModeDirect}}, DialectPostgres, ModeProcedure, false}, + } + for _, c := range cases { + op := OpLogin + if strings.HasPrefix(c.name, "override wins") { + op = OpSession + } + got, err := c.cfg.EffectiveMode(op, c.dialect) + if (err != nil) != c.wantErr || got != c.want { + t.Errorf("%s: got (%q, %v), want (%q, err=%v)", c.name, got, err, c.want, c.wantErr) + } + } +} diff --git a/pkg/security/lookup/mode.go b/pkg/security/lookup/mode.go new file mode 100644 index 0000000..6214073 --- /dev/null +++ b/pkg/security/lookup/mode.go @@ -0,0 +1,169 @@ +package lookup + +import "fmt" + +// Dialect names understood by Config.Dialect. The dialect package (step 2) owns the +// implementations; the names are defined here so Config can be validated without it. +const ( + DialectPostgres = "postgres" + DialectSQLite = "sqlite" + DialectMySQL = "mysql" + DialectMSSQL = "mssql" +) + +// Mode selects how a store talks to the database. +type Mode string + +const ( + // ModeDefault (the zero value) resolves to ModeProcedure on Postgres and ModeDirect elsewhere. + ModeDefault Mode = "" + // ModeProcedure always calls the configured stored procedure; a missing procedure is an error. + ModeProcedure Mode = "procedure" + // ModeDirect always works on the tables through the dialect builder. + ModeDirect Mode = "direct" + // ModeAuto probes the procedure once per operation on Postgres (cached) and uses it when + // present, otherwise direct. Other dialects resolve to ModeDirect. + ModeAuto Mode = "auto" +) + +func (m Mode) valid() bool { + switch m { + case ModeDefault, ModeProcedure, ModeDirect, ModeAuto: + return true + } + return false +} + +// Op names one store operation so its mode can be overridden individually. +type Op string + +const ( + OpLogin Op = "login" + OpRegister Op = "register" + OpLogout Op = "logout" + OpSession Op = "session" + OpTouchSession Op = "touch_session" + OpRefresh Op = "refresh" + OpLoginAPIKey Op = "login_api_key" + OpJWTLogin Op = "jwt_login" + OpJWTLogout Op = "jwt_logout" + OpResetRequest Op = "reset_request" + OpResetComplete Op = "reset_complete" + + OpKeyCreate Op = "key_create" + OpKeyList Op = "key_list" + OpKeyDelete Op = "key_delete" + OpKeyValidate Op = "key_validate" + + OpOAuthRegisterClient Op = "oauth_register_client" + OpOAuthGetClient Op = "oauth_get_client" + OpOAuthSaveCode Op = "oauth_save_code" + OpOAuthExchangeCode Op = "oauth_exchange_code" + OpOAuthIntrospect Op = "oauth_introspect" + OpOAuthRevoke Op = "oauth_revoke" + + OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user" + OpOAuthCreateSession Op = "oauth_create_session" + OpOAuthGetRefreshToken Op = "oauth_get_refresh_token" + OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" + OpOAuthGetUser Op = "oauth_get_user" + + OpPasskeyStore Op = "passkey_store" + OpPasskeyGet Op = "passkey_get" + OpPasskeyUpdateCounter Op = "passkey_update_counter" + OpPasskeyList Op = "passkey_list" + OpPasskeyDelete Op = "passkey_delete" + OpPasskeyRename Op = "passkey_rename" + OpPasskeyByUsername Op = "passkey_by_username" + OpPasskeyLogin Op = "passkey_login" + + OpTOTPEnable Op = "totp_enable" + OpTOTPDisable Op = "totp_disable" + OpTOTPStatus Op = "totp_status" + OpTOTPSecret Op = "totp_secret" + OpTOTPRegenerateBackup Op = "totp_regenerate_backup" + OpTOTPValidateBackupCode Op = "totp_validate_backup_code" + + OpColumnSecurity Op = "column_security" + OpRowSecurity Op = "row_security" +) + +// EffectiveMode resolves the mode for one operation: a per-operation override wins over +// Config.Mode, and ModeDefault is replaced by the dialect default. The result is +// ModeProcedure, ModeDirect or ModeAuto; ModeAuto only survives on Postgres, where the +// caller must probe the procedure. dialect is the resolved dialect name. +func (c Config) EffectiveMode(op Op, dialect string) (Mode, error) { + m := c.Mode + if o, ok := c.Overrides[op]; ok && o != ModeDefault { + m = o + } + if !m.valid() { + return "", fmt.Errorf("lookup: invalid mode %q for %s", m, op) + } + pg := dialect == DialectPostgres + switch m { + case ModeDefault: + if pg { + return ModeProcedure, nil + } + return ModeDirect, nil + case ModeAuto: + if pg { + return ModeAuto, nil + } + return ModeDirect, nil + case ModeProcedure: + if !pg { + return "", fmt.Errorf("lookup: procedure mode for %s requires the postgres dialect, got %q", op, dialect) + } + } + return m, nil +} + +// AllOps lists every operation, so callers can resolve or validate modes up front. +func AllOps() []Op { + return []Op{ + OpLogin, + OpRegister, + OpLogout, + OpSession, + OpTouchSession, + OpRefresh, + OpLoginAPIKey, + OpJWTLogin, + OpJWTLogout, + OpResetRequest, + OpResetComplete, + OpKeyCreate, + OpKeyList, + OpKeyDelete, + OpKeyValidate, + OpOAuthRegisterClient, + OpOAuthGetClient, + OpOAuthSaveCode, + OpOAuthExchangeCode, + OpOAuthIntrospect, + OpOAuthRevoke, + OpOAuthGetOrCreateUser, + OpOAuthCreateSession, + OpOAuthGetRefreshToken, + OpOAuthUpdateRefreshToken, + OpOAuthGetUser, + OpPasskeyStore, + OpPasskeyGet, + OpPasskeyUpdateCounter, + OpPasskeyList, + OpPasskeyDelete, + OpPasskeyRename, + OpPasskeyByUsername, + OpPasskeyLogin, + OpTOTPEnable, + OpTOTPDisable, + OpTOTPStatus, + OpTOTPSecret, + OpTOTPRegenerateBackup, + OpTOTPValidateBackupCode, + OpColumnSecurity, + OpRowSecurity, + } +} diff --git a/pkg/security/lookup/procedure/auth.go b/pkg/security/lookup/procedure/auth.go new file mode 100644 index 0000000..80d6947 --- /dev/null +++ b/pkg/security/lookup/procedure/auth.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/flextime.go b/pkg/security/lookup/procedure/flextime.go new file mode 100644 index 0000000..8dcfe2b --- /dev/null +++ b/pkg/security/lookup/procedure/flextime.go @@ -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 + } + } + } +} diff --git a/pkg/security/lookup/procedure/keys.go b/pkg/security/lookup/procedure/keys.go new file mode 100644 index 0000000..4bc8359 --- /dev/null +++ b/pkg/security/lookup/procedure/keys.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/oauth.go b/pkg/security/lookup/procedure/oauth.go new file mode 100644 index 0000000..7ea810f --- /dev/null +++ b/pkg/security/lookup/procedure/oauth.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/passkey.go b/pkg/security/lookup/procedure/passkey.go new file mode 100644 index 0000000..089d32c --- /dev/null +++ b/pkg/security/lookup/procedure/passkey.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/policy.go b/pkg/security/lookup/procedure/policy.go new file mode 100644 index 0000000..a883453 --- /dev/null +++ b/pkg/security/lookup/procedure/policy.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/procedure.go b/pkg/security/lookup/procedure/procedure.go new file mode 100644 index 0000000..03e557c --- /dev/null +++ b/pkg/security/lookup/procedure/procedure.go @@ -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) +} diff --git a/pkg/security/lookup/procedure/procedure_test.go b/pkg/security/lookup/procedure/procedure_test.go new file mode 100644 index 0000000..d0b4b7b --- /dev/null +++ b/pkg/security/lookup/procedure/procedure_test.go @@ -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) + } +} diff --git a/pkg/security/lookup/procedure/totp.go b/pkg/security/lookup/procedure/totp.go new file mode 100644 index 0000000..a45e7f3 --- /dev/null +++ b/pkg/security/lookup/procedure/totp.go @@ -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 +} diff --git a/pkg/security/lookup/procs.go b/pkg/security/lookup/procs.go new file mode 100644 index 0000000..4db0e64 --- /dev/null +++ b/pkg/security/lookup/procs.go @@ -0,0 +1,143 @@ +package lookup + +import ( + "fmt" + "reflect" +) + +// ProcNames holds the stored procedure (function) names used by the procedure backend. +// It replaces security.SQLNames and security.KeyStoreSQLNames. Zero fields mean "default" +// when merged with DefaultProcNames. +type ProcNames struct { + + // Auth procedures (DatabaseAuthenticator) + Login string // default: "resolvespec_login" + Register string // default: "resolvespec_register" + Logout string // default: "resolvespec_logout" + Session string // default: "resolvespec_session" + SessionUpdate string // default: "resolvespec_session_update" + RefreshToken string // default: "resolvespec_refresh_token" + LoginAPIKey string // default: "resolvespec_login_api_key" + + // JWT procedures (JWTAuthenticator) + JWTLogin string // default: "resolvespec_jwt_login" + JWTLogout string // default: "resolvespec_jwt_logout" + + // Security policy procedures + ColumnSecurity string // default: "resolvespec_column_security" + RowSecurity string // default: "resolvespec_row_security" + + // TOTP procedures (DatabaseTwoFactorProvider) + TOTPEnable string // default: "resolvespec_totp_enable" + TOTPDisable string // default: "resolvespec_totp_disable" + TOTPGetStatus string // default: "resolvespec_totp_get_status" + TOTPGetSecret string // default: "resolvespec_totp_get_secret" + TOTPRegenerateBackup string // default: "resolvespec_totp_regenerate_backup_codes" + TOTPValidateBackupCode string // default: "resolvespec_totp_validate_backup_code" + + // Passkey procedures (DatabasePasskeyProvider) + PasskeyStoreCredential string // default: "resolvespec_passkey_store_credential" + PasskeyGetCredsByUsername string // default: "resolvespec_passkey_get_credentials_by_username" + PasskeyGetCredential string // default: "resolvespec_passkey_get_credential" + PasskeyUpdateCounter string // default: "resolvespec_passkey_update_counter" + PasskeyGetUserCredentials string // default: "resolvespec_passkey_get_user_credentials" + PasskeyDeleteCredential string // default: "resolvespec_passkey_delete_credential" + PasskeyUpdateName string // default: "resolvespec_passkey_update_name" + PasskeyLogin string // default: "resolvespec_passkey_login" + + // Password reset procedures (DatabaseAuthenticator) + PasswordResetRequest string // default: "resolvespec_password_reset_request" + PasswordResetComplete string // default: "resolvespec_password_reset" + + // OAuth2 procedures (DatabaseAuthenticator OAuth2 methods) + OAuthGetOrCreateUser string // default: "resolvespec_oauth_getorcreateuser" + OAuthCreateSession string // default: "resolvespec_oauth_createsession" + OAuthGetRefreshToken string // default: "resolvespec_oauth_getrefreshtoken" + OAuthUpdateRefreshToken string // default: "resolvespec_oauth_updaterefreshtoken" + OAuthGetUser string // default: "resolvespec_oauth_getuser" + + // OAuth2 server procedures (OAuthServer persistence) + OAuthRegisterClient string // default: "resolvespec_oauth_register_client" + OAuthGetClient string // default: "resolvespec_oauth_get_client" + OAuthSaveCode string // default: "resolvespec_oauth_save_code" + OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code" + OAuthIntrospect string // default: "resolvespec_oauth_introspect" + OAuthRevoke string // default: "resolvespec_oauth_revoke" + + // Keystore procedures (KeyStore) + KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys" + KeystoreCreateKey string // default: "resolvespec_keystore_create_key" + KeystoreDeleteKey string // default: "resolvespec_keystore_delete_key" + KeystoreValidateKey string // default: "resolvespec_keystore_validate_key" +} + +// DefaultProcNames returns the default resolvespec_* procedure names. +func DefaultProcNames() ProcNames { + return ProcNames{ //nolint:gosec // G101: false positive: identifiers, not credentials + Login: "resolvespec_login", + Register: "resolvespec_register", + Logout: "resolvespec_logout", + Session: "resolvespec_session", + SessionUpdate: "resolvespec_session_update", + RefreshToken: "resolvespec_refresh_token", + LoginAPIKey: "resolvespec_login_api_key", + JWTLogin: "resolvespec_jwt_login", + JWTLogout: "resolvespec_jwt_logout", + ColumnSecurity: "resolvespec_column_security", + RowSecurity: "resolvespec_row_security", + TOTPEnable: "resolvespec_totp_enable", + TOTPDisable: "resolvespec_totp_disable", + TOTPGetStatus: "resolvespec_totp_get_status", + TOTPGetSecret: "resolvespec_totp_get_secret", + TOTPRegenerateBackup: "resolvespec_totp_regenerate_backup_codes", + TOTPValidateBackupCode: "resolvespec_totp_validate_backup_code", + PasskeyStoreCredential: "resolvespec_passkey_store_credential", + PasskeyGetCredsByUsername: "resolvespec_passkey_get_credentials_by_username", + PasskeyGetCredential: "resolvespec_passkey_get_credential", + PasskeyUpdateCounter: "resolvespec_passkey_update_counter", + PasskeyGetUserCredentials: "resolvespec_passkey_get_user_credentials", + PasskeyDeleteCredential: "resolvespec_passkey_delete_credential", + PasskeyUpdateName: "resolvespec_passkey_update_name", + PasskeyLogin: "resolvespec_passkey_login", + PasswordResetRequest: "resolvespec_password_reset_request", + PasswordResetComplete: "resolvespec_password_reset", + OAuthGetOrCreateUser: "resolvespec_oauth_getorcreateuser", + OAuthCreateSession: "resolvespec_oauth_createsession", + OAuthGetRefreshToken: "resolvespec_oauth_getrefreshtoken", + OAuthUpdateRefreshToken: "resolvespec_oauth_updaterefreshtoken", + OAuthGetUser: "resolvespec_oauth_getuser", + OAuthRegisterClient: "resolvespec_oauth_register_client", + OAuthGetClient: "resolvespec_oauth_get_client", + OAuthSaveCode: "resolvespec_oauth_save_code", + OAuthExchangeCode: "resolvespec_oauth_exchange_code", + OAuthIntrospect: "resolvespec_oauth_introspect", + OAuthRevoke: "resolvespec_oauth_revoke", + KeystoreGetUserKeys: "resolvespec_keystore_get_user_keys", + KeystoreCreateKey: "resolvespec_keystore_create_key", + KeystoreDeleteKey: "resolvespec_keystore_delete_key", + KeystoreValidateKey: "resolvespec_keystore_validate_key", + } +} + +// Merge returns a copy of p with every non-empty field of override applied. +func (p ProcNames) Merge(override ProcNames) ProcNames { + merged := p + mv, ov := reflect.ValueOf(&merged).Elem(), reflect.ValueOf(override) + for i := 0; i < ov.NumField(); i++ { + if v := ov.Field(i).String(); v != "" { + mv.Field(i).SetString(v) + } + } + return merged +} + +// Validate checks that every name is a safe (optionally schema-qualified) identifier. +func (p ProcNames) Validate() error { + v, t := reflect.ValueOf(p), reflect.TypeOf(p) + for i := 0; i < v.NumField(); i++ { + if name := v.Field(i).String(); !validQualifiedIdent(name) { + return fmt.Errorf("lookup: invalid procedure name %q for %s", name, t.Field(i).Name) + } + } + return nil +} diff --git a/pkg/security/lookup/schema.go b/pkg/security/lookup/schema.go new file mode 100644 index 0000000..b7482c1 --- /dev/null +++ b/pkg/security/lookup/schema.go @@ -0,0 +1,346 @@ +package lookup + +import ( + "fmt" + "regexp" + "sort" +) + +var ( + identRe = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`) + qualifiedIdentRe = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*(\.[a-zA-Z_][a-zA-Z0-9_]*)?$`) +) + +func validIdent(s string) bool { return identRe.MatchString(s) } +func validQualifiedIdent(s string) bool { return qualifiedIdentRe.MatchString(s) } + +// Entity identifies one table the direct backend reads or writes. +type Entity string + +const ( + EntityUsers Entity = "users" + EntityUserSessions Entity = "user_sessions" + EntityTokenBlacklist Entity = "token_blacklist" + EntityUserTOTPBackupCodes Entity = "user_totp_backup_codes" + EntityUserPasskeyCredentials Entity = "user_passkey_credentials" + EntityUserPasswordResets Entity = "user_password_resets" + EntityOAuthClients Entity = "oauth_clients" + EntityOAuthCodes Entity = "oauth_codes" + EntityUserKeys Entity = "user_keys" + EntitySecGroupMembers Entity = "sec_group_members" + EntitySecColumnRules Entity = "sec_column_rules" + EntitySecRowRules Entity = "sec_row_rules" +) + +// Column is a typed key naming one logical column of an entity. The physical column +// name is looked up in the Schema, so every column of every entity is configurable. +type Column struct { + Entity Entity + Name string +} + +func (c Column) String() string { return string(c.Entity) + "." + c.Name } + +func col(e Entity, name string) Column { return Column{Entity: e, Name: name} } + +// Logical columns. The names are the default physical column names. +var ( + UsersID = col(EntityUsers, "id") + UsersUsername = col(EntityUsers, "username") + UsersEmail = col(EntityUsers, "email") + UsersPassword = col(EntityUsers, "password") + UsersUserLevel = col(EntityUsers, "user_level") + UsersRoles = col(EntityUsers, "roles") + UsersIsActive = col(EntityUsers, "is_active") + UsersCreatedAt = col(EntityUsers, "created_at") + UsersUpdatedAt = col(EntityUsers, "updated_at") + UsersLastLoginAt = col(EntityUsers, "last_login_at") + UsersProgramUserID = col(EntityUsers, "program_user_id") + UsersProgramUserTable = col(EntityUsers, "program_user_table") + UsersRemoteID = col(EntityUsers, "remote_id") + UsersAuthProvider = col(EntityUsers, "auth_provider") + UsersTOTPSecret = col(EntityUsers, "totp_secret") + UsersTOTPEnabled = col(EntityUsers, "totp_enabled") + UsersTOTPEnabledAt = col(EntityUsers, "totp_enabled_at") + + SessionsID = col(EntityUserSessions, "id") + SessionsToken = col(EntityUserSessions, "session_token") + SessionsUserID = col(EntityUserSessions, "user_id") + SessionsExpiresAt = col(EntityUserSessions, "expires_at") + SessionsCreatedAt = col(EntityUserSessions, "created_at") + SessionsLastActivityAt = col(EntityUserSessions, "last_activity_at") + SessionsIPAddress = col(EntityUserSessions, "ip_address") + SessionsUserAgent = col(EntityUserSessions, "user_agent") + SessionsAccessToken = col(EntityUserSessions, "access_token") + SessionsRefreshToken = col(EntityUserSessions, "refresh_token") + SessionsTokenType = col(EntityUserSessions, "token_type") + SessionsAuthProvider = col(EntityUserSessions, "auth_provider") + + BlacklistID = col(EntityTokenBlacklist, "id") + BlacklistToken = col(EntityTokenBlacklist, "token") + BlacklistUserID = col(EntityTokenBlacklist, "user_id") + BlacklistExpiresAt = col(EntityTokenBlacklist, "expires_at") + BlacklistCreatedAt = col(EntityTokenBlacklist, "created_at") + + BackupCodesID = col(EntityUserTOTPBackupCodes, "id") + BackupCodesUserID = col(EntityUserTOTPBackupCodes, "user_id") + BackupCodesCodeHash = col(EntityUserTOTPBackupCodes, "code_hash") + BackupCodesUsed = col(EntityUserTOTPBackupCodes, "used") + BackupCodesUsedAt = col(EntityUserTOTPBackupCodes, "used_at") + BackupCodesCreatedAt = col(EntityUserTOTPBackupCodes, "created_at") + + PasskeyID = col(EntityUserPasskeyCredentials, "id") + PasskeyUserID = col(EntityUserPasskeyCredentials, "user_id") + PasskeyCredentialID = col(EntityUserPasskeyCredentials, "credential_id") + PasskeyPublicKey = col(EntityUserPasskeyCredentials, "public_key") + PasskeyAttestationType = col(EntityUserPasskeyCredentials, "attestation_type") + PasskeyAAGUID = col(EntityUserPasskeyCredentials, "aaguid") + PasskeySignCount = col(EntityUserPasskeyCredentials, "sign_count") + PasskeyCloneWarning = col(EntityUserPasskeyCredentials, "clone_warning") + PasskeyTransports = col(EntityUserPasskeyCredentials, "transports") + PasskeyBackupEligible = col(EntityUserPasskeyCredentials, "backup_eligible") + PasskeyBackupState = col(EntityUserPasskeyCredentials, "backup_state") + PasskeyName = col(EntityUserPasskeyCredentials, "name") + PasskeyCreatedAt = col(EntityUserPasskeyCredentials, "created_at") + PasskeyLastUsedAt = col(EntityUserPasskeyCredentials, "last_used_at") + + ResetsID = col(EntityUserPasswordResets, "id") + ResetsUserID = col(EntityUserPasswordResets, "user_id") + ResetsTokenHash = col(EntityUserPasswordResets, "token_hash") + ResetsExpiresAt = col(EntityUserPasswordResets, "expires_at") + ResetsCreatedAt = col(EntityUserPasswordResets, "created_at") + ResetsUsed = col(EntityUserPasswordResets, "used") + ResetsUsedAt = col(EntityUserPasswordResets, "used_at") + + OAuthClientsID = col(EntityOAuthClients, "id") + OAuthClientsClientID = col(EntityOAuthClients, "client_id") + OAuthClientsRedirectURIs = col(EntityOAuthClients, "redirect_uris") + OAuthClientsClientName = col(EntityOAuthClients, "client_name") + OAuthClientsGrantTypes = col(EntityOAuthClients, "grant_types") + OAuthClientsAllowedScopes = col(EntityOAuthClients, "allowed_scopes") + OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash") + OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method") + OAuthClientsIsActive = col(EntityOAuthClients, "is_active") + OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at") + + OAuthCodesID = col(EntityOAuthCodes, "id") + OAuthCodesCode = col(EntityOAuthCodes, "code") + OAuthCodesClientID = col(EntityOAuthCodes, "client_id") + OAuthCodesRedirectURI = col(EntityOAuthCodes, "redirect_uri") + OAuthCodesClientState = col(EntityOAuthCodes, "client_state") + OAuthCodesCodeChallenge = col(EntityOAuthCodes, "code_challenge") + OAuthCodesCodeChallengeMethod = col(EntityOAuthCodes, "code_challenge_method") + OAuthCodesSessionToken = col(EntityOAuthCodes, "session_token") + OAuthCodesRefreshToken = col(EntityOAuthCodes, "refresh_token") + OAuthCodesScopes = col(EntityOAuthCodes, "scopes") + OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at") + OAuthCodesCreatedAt = col(EntityOAuthCodes, "created_at") + + KeysID = col(EntityUserKeys, "id") + KeysUserID = col(EntityUserKeys, "user_id") + KeysKeyType = col(EntityUserKeys, "key_type") + KeysKeyHash = col(EntityUserKeys, "key_hash") + KeysName = col(EntityUserKeys, "name") + KeysScopes = col(EntityUserKeys, "scopes") + KeysMeta = col(EntityUserKeys, "meta") + KeysExpiresAt = col(EntityUserKeys, "expires_at") + KeysCreatedAt = col(EntityUserKeys, "created_at") + KeysLastUsedAt = col(EntityUserKeys, "last_used_at") + KeysIsActive = col(EntityUserKeys, "is_active") + + GroupMembersGroupID = col(EntitySecGroupMembers, "group_id") + GroupMembersUserID = col(EntitySecGroupMembers, "user_id") + + ColRulesID = col(EntitySecColumnRules, "id") + ColRulesUserID = col(EntitySecColumnRules, "user_id") + ColRulesGroupID = col(EntitySecColumnRules, "group_id") + ColRulesSchemaName = col(EntitySecColumnRules, "schema_name") + ColRulesTableName = col(EntitySecColumnRules, "table_name") + ColRulesColumnPath = col(EntitySecColumnRules, "column_path") + ColRulesAccessType = col(EntitySecColumnRules, "access_type") + ColRulesMaskStart = col(EntitySecColumnRules, "mask_start") + ColRulesMaskEnd = col(EntitySecColumnRules, "mask_end") + ColRulesMaskInvert = col(EntitySecColumnRules, "mask_invert") + ColRulesMaskChar = col(EntitySecColumnRules, "mask_char") + ColRulesExtraFilters = col(EntitySecColumnRules, "extra_filters") + ColRulesIsActive = col(EntitySecColumnRules, "is_active") + + RowRulesID = col(EntitySecRowRules, "id") + RowRulesUserID = col(EntitySecRowRules, "user_id") + RowRulesGroupID = col(EntitySecRowRules, "group_id") + RowRulesSchemaName = col(EntitySecRowRules, "schema_name") + RowRulesTableName = col(EntitySecRowRules, "table_name") + RowRulesTemplate = col(EntitySecRowRules, "template") + RowRulesHasBlock = col(EntitySecRowRules, "has_block") + RowRulesIsActive = col(EntitySecRowRules, "is_active") +) + +// allColumns lists every logical column; it defines the default schema. +var allColumns = []Column{ + UsersID, UsersUsername, UsersEmail, UsersPassword, UsersUserLevel, UsersRoles, UsersIsActive, + UsersCreatedAt, UsersUpdatedAt, UsersLastLoginAt, UsersProgramUserID, UsersProgramUserTable, + UsersRemoteID, UsersAuthProvider, UsersTOTPSecret, UsersTOTPEnabled, UsersTOTPEnabledAt, + SessionsID, SessionsToken, SessionsUserID, SessionsExpiresAt, SessionsCreatedAt, SessionsLastActivityAt, + SessionsIPAddress, SessionsUserAgent, SessionsAccessToken, SessionsRefreshToken, SessionsTokenType, SessionsAuthProvider, + BlacklistID, BlacklistToken, BlacklistUserID, BlacklistExpiresAt, BlacklistCreatedAt, + BackupCodesID, BackupCodesUserID, BackupCodesCodeHash, BackupCodesUsed, BackupCodesUsedAt, BackupCodesCreatedAt, + PasskeyID, PasskeyUserID, PasskeyCredentialID, PasskeyPublicKey, PasskeyAttestationType, PasskeyAAGUID, + PasskeySignCount, PasskeyCloneWarning, PasskeyTransports, PasskeyBackupEligible, PasskeyBackupState, + PasskeyName, PasskeyCreatedAt, PasskeyLastUsedAt, + ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt, + OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes, + OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod, + OAuthClientsIsActive, OAuthClientsCreatedAt, + OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState, + OAuthCodesCodeChallenge, OAuthCodesCodeChallengeMethod, OAuthCodesSessionToken, OAuthCodesRefreshToken, + OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt, + KeysID, KeysUserID, KeysKeyType, KeysKeyHash, KeysName, KeysScopes, KeysMeta, KeysExpiresAt, + KeysCreatedAt, KeysLastUsedAt, KeysIsActive, + GroupMembersGroupID, GroupMembersUserID, + ColRulesID, ColRulesUserID, ColRulesGroupID, ColRulesSchemaName, ColRulesTableName, ColRulesColumnPath, + ColRulesAccessType, ColRulesMaskStart, ColRulesMaskEnd, ColRulesMaskInvert, ColRulesMaskChar, + ColRulesExtraFilters, ColRulesIsActive, + RowRulesID, RowRulesUserID, RowRulesGroupID, RowRulesSchemaName, RowRulesTableName, RowRulesTemplate, + RowRulesHasBlock, RowRulesIsActive, +} + +// Table maps one entity to a physical table and its columns. +type Table struct { + // Schema optionally qualifies the table (schema.table). Empty = unqualified. + Schema string + // Name is the physical table name. Empty = default (the entity name). + Name string + // Columns maps logical column name -> physical column name. Missing = default. + Columns map[string]string +} + +// Schema maps every entity to its physical table and columns. The zero value is +// valid and means "all defaults"; use DefaultSchema for the explicit baseline. +type Schema map[Entity]Table + +// DefaultSchema returns the baseline schema: every entity and column under its default name. +func DefaultSchema() Schema { + s := Schema{} + for _, c := range allColumns { + t, ok := s[c.Entity] + if !ok { + t = Table{Name: string(c.Entity), Columns: map[string]string{}} + } + t.Columns[c.Name] = c.Name + s[c.Entity] = t + } + return s +} + +// Merge returns a copy of s with every non-empty field of override applied. +// Unknown entities or columns in override are kept so Validate can report them. +func (s Schema) Merge(override Schema) Schema { + merged := Schema{} + for e, t := range s { + merged[e] = cloneTable(t) + } + for e, ot := range override { + t, ok := merged[e] + if !ok { + merged[e] = cloneTable(ot) + continue + } + if ot.Schema != "" { + t.Schema = ot.Schema + } + if ot.Name != "" { + t.Name = ot.Name + } + for k, v := range ot.Columns { + if v != "" { + t.Columns[k] = v + } + } + merged[e] = t + } + return merged +} + +func cloneTable(t Table) Table { + c := t + c.Columns = make(map[string]string, len(t.Columns)) + for k, v := range t.Columns { + c.Columns[k] = v + } + return c +} + +// Validate checks that the schema only names known entities and columns and that +// every identifier is safe. It is meant to run on the merged (default + override) schema. +func (s Schema) Validate() error { + known := map[Column]bool{} + for _, c := range allColumns { + known[c] = true + } + entities := make([]string, 0, len(s)) + for e := range s { + entities = append(entities, string(e)) + } + sort.Strings(entities) + for _, en := range entities { + e := Entity(en) + t := s[e] + if firstKnownColumn(e) == "" { + return fmt.Errorf("lookup: unknown entity %q", e) + } + if !validQualifiedIdent(t.Name) { + return fmt.Errorf("lookup: invalid table name %q for %s", t.Name, e) + } + if t.Schema != "" && !validIdent(t.Schema) { + return fmt.Errorf("lookup: invalid schema name %q for %s", t.Schema, e) + } + names := make([]string, 0, len(t.Columns)) + for n := range t.Columns { + names = append(names, n) + } + sort.Strings(names) + for _, n := range names { + if !known[Column{Entity: e, Name: n}] { + return fmt.Errorf("lookup: unknown column %q for %s", n, e) + } + if !validIdent(t.Columns[n]) { + return fmt.Errorf("lookup: invalid column name %q for %s.%s", t.Columns[n], e, n) + } + } + } + return nil +} + +func firstKnownColumn(e Entity) string { + for _, c := range allColumns { + if c.Entity == e { + return c.Name + } + } + return "" +} + +// TableName returns the physical table name of an entity (unqualified, unquoted). +func (s Schema) TableName(e Entity) string { + if t, ok := s[e]; ok && t.Name != "" { + return t.Name + } + return string(e) +} + +// SchemaName returns the optional schema qualifier of an entity. +func (s Schema) SchemaName(e Entity) string { return s[e].Schema } + +// Col returns the physical column name for a logical column (unquoted). +func (s Schema) Col(c Column) string { + if t, ok := s[c.Entity]; ok { + if n := t.Columns[c.Name]; n != "" { + return n + } + } + return c.Name +} + +// FirstColumn returns the name of the first logical column of an entity (used by the +// direct backend for existence checks). +func FirstColumn(e Entity) string { return firstKnownColumn(e) } diff --git a/pkg/security/lookup/txsettings.go b/pkg/security/lookup/txsettings.go new file mode 100644 index 0000000..91a99b5 --- /dev/null +++ b/pkg/security/lookup/txsettings.go @@ -0,0 +1,47 @@ +package lookup + +import ( + "context" + "encoding/hex" + "fmt" + "regexp" + "sort" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// settingNameRE matches a custom GUC name: two or more dot-separated identifiers. +var settingNameRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)+$`) + +// ApplyTxSettings sets each entry as a transaction-local setting on tx, in name +// order. Postgres only; any other driver with a non-empty map is an error so a +// missing RLS stamp fails closed. +func ApplyTxSettings(ctx context.Context, tx common.Database, settings map[string]string) error { + if len(settings) == 0 { + return nil + } + if tx == nil { + return fmt.Errorf("tx settings: no transaction") + } + if drv := tx.DriverName(); drv != "postgres" && drv != "pgsql" { + return fmt.Errorf("tx settings: unsupported driver %q", drv) + } + names := make([]string, 0, len(settings)) + for name := range settings { + if !settingNameRE.MatchString(name) { + return fmt.Errorf("tx settings: invalid setting name %q", name) + } + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + // The value is hex-encoded so it needs no quoting and cannot be read as a + // bind placeholder by any adapter. + query := fmt.Sprintf("SELECT set_config('%s', convert_from(decode('%s', 'hex'), 'UTF8'), true)", + name, hex.EncodeToString([]byte(settings[name]))) + if _, err := tx.Exec(ctx, query); err != nil { + return fmt.Errorf("tx settings: set %s: %w", name, err) + } + } + return nil +} diff --git a/pkg/security/lookup/txsettings_test.go b/pkg/security/lookup/txsettings_test.go new file mode 100644 index 0000000..9ca8a42 --- /dev/null +++ b/pkg/security/lookup/txsettings_test.go @@ -0,0 +1,93 @@ +package lookup_test + +import ( + "context" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" +) + +func txSettingsDB(t *testing.T) (common.Database, 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 database.NewPgSQLAdapter(db), mock +} + +func TestApplyTxSettingsStampsInNameOrderOnTx(t *testing.T) { + pool, mock := txSettingsDB(t) + ctx := context.Background() + + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.tenant', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec(`set_config\('app\.user_id', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + + err := pool.RunInTransaction(context.Background(), func(tx common.Database) error { + return lookup.ApplyTxSettings(ctx, tx, map[string]string{"app.user_id": "7", "app.tenant": "o'x?"}) + }) + if err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestApplyTxSettingsValueIsNeverInlined(t *testing.T) { + pool, mock := txSettingsDB(t) + ctx := context.Background() + var seen string + mock.ExpectBegin() + mock.ExpectExec(`set_config`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + _ = pool.RunInTransaction(context.Background(), func(tx common.Database) error { + // Capture via a wrapper so the raw statement can be inspected. + return lookup.ApplyTxSettings(ctx, &queryRecorder{Database: tx, got: &seen}, map[string]string{"app.v": "'; DROP TABLE x; --?"}) + }) + if strings.Contains(seen, "DROP") || strings.Contains(seen, "?") { + t.Fatalf("value leaked into SQL text: %s", seen) + } +} + +type queryRecorder struct { + common.Database + got *string +} + +func (q *queryRecorder) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) { + *q.got = query + return q.Database.Exec(ctx, query, args...) +} + +func TestApplyTxSettingsRejectsBadNameAndDriver(t *testing.T) { + pool, _ := txSettingsDB(t) + ctx := context.Background() + + for _, name := range []string{"user_id", "app.x'); DROP", "app..x", "app.x y"} { + if err := lookup.ApplyTxSettings(ctx, pool, map[string]string{name: "1"}); err == nil { + t.Fatalf("name %q must be rejected", name) + } + } + if err := lookup.ApplyTxSettings(ctx, pool, nil); err != nil { + t.Fatalf("empty settings must be a no-op: %v", err) + } + if err := lookup.ApplyTxSettings(ctx, &driverStub{Database: pool, name: "sqlite"}, map[string]string{"app.x": "1"}); err == nil { + t.Fatal("non-postgres driver must fail closed") + } +} + +type driverStub struct { + common.Database + name string +} + +func (d *driverStub) DriverName() string { return d.name } diff --git a/pkg/security/lookup_bridge.go b/pkg/security/lookup_bridge.go new file mode 100644 index 0000000..61011d0 --- /dev/null +++ b/pkg/security/lookup_bridge.go @@ -0,0 +1,52 @@ +package security + +import ( + "database/sql" + "sync" + + "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends" +) + +// lookupSource builds the lookup.Provider a security component uses, on first use, so the +// With* builders can still change the configuration after construction. An explicit Provider +// bypasses the build. +type lookupSource struct { + db *sql.DB + cfg lookup.Config + opts backends.Options + + provider *lookup.Provider + + once sync.Once + built *lookup.Provider +} + +func newLookupSource(db *sql.DB) *lookupSource { return &lookupSource{db: db} } + +// get returns the provider. A bad configuration is logged once and yields a provider that +// returns the error from every call, so the component fails closed. +func (s *lookupSource) get() *lookup.Provider { + if s.provider != nil { + return s.provider + } + s.once.Do(func() { s.built = resolveLookup(s.db, s.cfg, s.opts) }) + return s.built +} + +// resolveLookup builds a provider for db. When the dialect is not configured and cannot be +// detected from the driver it falls back to postgres, the stored-procedure default. +func resolveLookup(db *sql.DB, cfg lookup.Config, opts backends.Options) *lookup.Provider { + if db != nil && cfg.Dialect == "" { + if _, err := cfg.ResolveDialect(db); err != nil { + cfg.Dialect = "postgres" + } + } + p, err := backends.New(db, cfg, opts) + if err != nil { + logger.Error("security: lookup configuration invalid: %v", err) + return backends.Failed(err) + } + return p +} diff --git a/pkg/security/nosql_test.go b/pkg/security/nosql_test.go new file mode 100644 index 0000000..2cd373c --- /dev/null +++ b/pkg/security/nosql_test.go @@ -0,0 +1,42 @@ +package security + +import ( + "os" + "path/filepath" + "regexp" + "strings" + "testing" +) + +// The core package must contain no SQL: every database access goes through pkg/security/lookup. +func TestCoreContainsNoSQL(t *testing.T) { + forbidden := []*regexp.Regexp{ + regexp.MustCompile(`(?i)"[^"\n]*\b(select\s.+\sfrom|insert\s+into|delete\s+from|update\s+\w+\s+set|create\s+table)\b`), + regexp.MustCompile("(?i)`[^`]*\\b(select\\s.+\\sfrom|insert\\s+into|delete\\s+from|update\\s+\\w+\\s+set)\\b"), + regexp.MustCompile(`\b(db|sqlDB|conn)\.(Query|QueryRow|Exec|Prepare|Begin)(Context|Tx)?\(`), + } + files, err := filepath.Glob("*.go") + if err != nil { + t.Fatal(err) + } + for _, f := range files { + // Example files are documentation, not library code. + if strings.HasSuffix(f, "_test.go") || strings.HasPrefix(f, "examples") || strings.HasSuffix(f, "_examples.go") { + continue + } + b, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + for i, line := range strings.Split(string(b), "\n") { + if c := strings.TrimSpace(line); strings.HasPrefix(c, "//") { + continue + } + for _, re := range forbidden { + if re.MatchString(line) { + t.Errorf("%s:%d looks like SQL / direct database access: %s", f, i+1, strings.TrimSpace(line)) + } + } + } + } +} diff --git a/pkg/security/oauth2_examples.go b/pkg/security/oauth2_examples.go index a4673d7..f167e8c 100644 --- a/pkg/security/oauth2_examples.go +++ b/pkg/security/oauth2_examples.go @@ -441,7 +441,7 @@ func ExampleOAuth2Complete() { } func setupOAuth2Tables(db *sql.DB) { - // Create tables from database_schema.sql + // Create tables from lookup/database_schema.sql // This is a helper function - in production, use migrations ctx := context.Background() diff --git a/pkg/security/oauth2_methods.go b/pkg/security/oauth2_methods.go index bc4fe9c..40da412 100644 --- a/pkg/security/oauth2_methods.go +++ b/pkg/security/oauth2_methods.go @@ -13,6 +13,7 @@ import ( "time" "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" "golang.org/x/oauth2" ) @@ -234,92 +235,20 @@ func (a *DatabaseAuthenticator) getOAuth2Provider(providerName string) (*OAuth2P // oauth2GetOrCreateUser finds or creates a user based on OAuth2 info using stored procedure func (a *DatabaseAuthenticator) oauth2GetOrCreateUser(ctx context.Context, userCtx *UserContext, providerName string) (int, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetOrCreateUser) { - return a.oauth2GetOrCreateUserDirect(ctx, userCtx, providerName) - } - - userData := map[string]interface{}{ - "username": userCtx.UserName, - "email": userCtx.Email, - "remote_id": userCtx.RemoteID, - "user_level": userCtx.UserLevel, - "roles": userCtx.Roles, - "auth_provider": providerName, - } - - userJSON, err := json.Marshal(userData) - if err != nil { - return 0, fmt.Errorf("failed to marshal user data: %w", err) - } - - var success bool - var errMsg *string - var userID *int - - err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_user_id - FROM %s($1::jsonb) - `, a.sqlNames.OAuthGetOrCreateUser), userJSON).Scan(&success, &errMsg, &userID) - - if err != nil { - return 0, fmt.Errorf("failed to get or create user: %w", err) - } - - if !success { - if errMsg != nil { - return 0, fmt.Errorf("%s", *errMsg) - } - return 0, fmt.Errorf("failed to get or create user") - } - - if userID == nil { - return 0, fmt.Errorf("user ID not returned") - } - - return *userID, nil + return a.src.get().OAuthUser.GetOrCreateUser(ctx, userCtx, providerName) } // oauth2CreateSession creates a new OAuth2 session using stored procedure func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthCreateSession) { - return a.oauth2CreateSessionDirect(ctx, sessionToken, userID, token, expiresAt, providerName) - } - - sessionData := map[string]interface{}{ - "session_token": sessionToken, - "user_id": userID, - "access_token": token.AccessToken, - "refresh_token": token.RefreshToken, - "token_type": token.TokenType, - "expires_at": expiresAt, - "auth_provider": providerName, - } - - sessionJSON, err := json.Marshal(sessionData) - if err != nil { - return fmt.Errorf("failed to marshal session data: %w", err) - } - - var success bool - var errMsg *string - - err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error - FROM %s($1::jsonb) - `, a.sqlNames.OAuthCreateSession), sessionJSON).Scan(&success, &errMsg) - - if err != nil { - return fmt.Errorf("failed to create session: %w", err) - } - - if !success { - if errMsg != nil { - return fmt.Errorf("%s", *errMsg) - } - return fmt.Errorf("failed to create session") - } - - return nil + return a.src.get().OAuthUser.CreateSession(ctx, lookup.OAuthSession{ + SessionToken: sessionToken, + UserID: userID, + AccessToken: token.AccessToken, + RefreshToken: token.RefreshToken, + TokenType: token.TokenType, + ExpiresAt: expiresAt, + Provider: providerName, + }) } // validateState validates state using in-memory storage @@ -420,7 +349,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT } // Get session by refresh token from database - session, err := a.oauthGetByRefreshToken(ctx, refreshToken) + session, err := a.src.get().OAuthUser.GetByRefreshToken(ctx, refreshToken) if err != nil { return nil, err } @@ -447,12 +376,12 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT } // Update session in database with new tokens - if err := a.oauthUpdateRefreshTokenRecord(ctx, session.UserID, refreshToken, newSessionToken, newToken.AccessToken, newToken.RefreshToken, newToken.Expiry); err != nil { + if err := a.src.get().OAuthUser.UpdateRefreshToken(ctx, session.UserID, refreshToken, newSessionToken, newToken.AccessToken, newToken.RefreshToken, newToken.Expiry); err != nil { return nil, err } // Get user data - userCtx, err := a.oauthGetUserByID(ctx, session.UserID) + userCtx, err := a.src.get().OAuthUser.GetUser(ctx, session.UserID) if err != nil { return nil, err } diff --git a/pkg/security/oauth2_methods_direct.go b/pkg/security/oauth2_methods_direct.go deleted file mode 100644 index abfc448..0000000 --- a/pkg/security/oauth2_methods_direct.go +++ /dev/null @@ -1,242 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "encoding/json" - "errors" - "fmt" - "strings" - "time" - - "golang.org/x/oauth2" -) - -// oauthRefreshSession is the session data needed to refresh an OAuth2 token, -// shared by both the stored-procedure and Direct-mode code paths. -type oauthRefreshSession struct { - UserID int `json:"user_id"` - AccessToken string `json:"access_token"` - TokenType string `json:"token_type"` - Expiry time.Time `json:"expiry"` -} - -// oauth2GetOrCreateUserDirect mirrors resolvespec_oauth_getorcreateuser. -func (a *DatabaseAuthenticator) oauth2GetOrCreateUserDirect(ctx context.Context, userCtx *UserContext, providerName string) (int, error) { - rolesStr := strings.Join(userCtx.Roles, ",") - var userID int - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE email = ?`, a.tableNames.Users)) - err := db.QueryRowContext(ctx, query, userCtx.Email).Scan(&userID) - if err == nil { - now := time.Now() - updQuery := rewritePlaceholders(db, fmt.Sprintf( - `UPDATE %s SET last_login_at = ?, updated_at = ?, remote_id = COALESCE(remote_id, ?), auth_provider = COALESCE(auth_provider, ?) WHERE id = ?`, - a.tableNames.Users)) - _, err := db.ExecContext(ctx, updQuery, now, now, userCtx.RemoteID, providerName, userID) - return err - } - if !errors.Is(err, sql.ErrNoRows) { - return err - } - - now := time.Now() - insQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, last_login_at, remote_id, auth_provider) VALUES (?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.Users)) - res, err := db.ExecContext(ctx, insQuery, userCtx.UserName, userCtx.Email, userCtx.UserLevel, rolesStr, true, now, now, now, userCtx.RemoteID, providerName) - if err != nil { - return err - } - id, err := res.LastInsertId() - if err != nil { - return err - } - userID = int(id) - return nil - }) - if err != nil { - return 0, fmt.Errorf("failed to get or create user: %w", err) - } - return userID, nil -} - -// oauth2CreateSessionDirect mirrors resolvespec_oauth_createsession (insert-or-update by session_token). -func (a *DatabaseAuthenticator) oauth2CreateSessionDirect(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error { - return a.runDBOpWithReconnect(func(db *sql.DB) error { - var exists int - checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE session_token = ?`, a.tableNames.UserSessions)) - err := db.QueryRowContext(ctx, checkQuery, sessionToken).Scan(&exists) - now := time.Now() - if err == nil { - updQuery := rewritePlaceholders(db, fmt.Sprintf( - `UPDATE %s SET access_token = ?, refresh_token = ?, token_type = ?, expires_at = ?, last_activity_at = ? WHERE session_token = ?`, - a.tableNames.UserSessions)) - _, err := db.ExecContext(ctx, updQuery, token.AccessToken, token.RefreshToken, token.TokenType, expiresAt, now, sessionToken) - return err - } - if !errors.Is(err, sql.ErrNoRows) { - return err - } - - insQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (session_token, user_id, expires_at, created_at, last_activity_at, access_token, refresh_token, token_type, auth_provider) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.UserSessions)) - _, err = db.ExecContext(ctx, insQuery, sessionToken, userID, expiresAt, now, now, token.AccessToken, token.RefreshToken, token.TokenType, providerName) - return err - }) -} - -// oauthGetByRefreshToken retrieves the session for a refresh token, dispatching between -// the resolvespec_oauth_getrefreshtoken stored procedure and Direct-mode SQL. -func (a *DatabaseAuthenticator) oauthGetByRefreshToken(ctx context.Context, refreshToken string) (*oauthRefreshSession, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetRefreshToken) { - var session oauthRefreshSession - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT user_id, access_token, token_type, expires_at FROM %s WHERE refresh_token = ? AND expires_at > ?`, - a.tableNames.UserSessions)) - return db.QueryRowContext(ctx, query, refreshToken, time.Now()).Scan(&session.UserID, &session.AccessToken, &session.TokenType, &session.Expiry) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("refresh token not found or expired") - } - return nil, fmt.Errorf("failed to get session by refresh token: %w", err) - } - return &session, nil - } - - var success bool - var errMsg *string - var sessionData []byte - - err := a.getDB().QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1) - `, a.sqlNames.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &sessionData) - if err != nil { - return nil, fmt.Errorf("failed to get session by refresh token: %w", err) - } - if !success { - if errMsg != nil { - return nil, fmt.Errorf("%s", *errMsg) - } - return nil, fmt.Errorf("invalid or expired refresh token") - } - - var session oauthRefreshSession - if err := json.Unmarshal(sessionData, &session); err != nil { - return nil, fmt.Errorf("failed to parse session data: %w", err) - } - return &session, nil -} - -// oauthUpdateRefreshTokenRecord updates a session with new tokens, dispatching between -// the resolvespec_oauth_updaterefreshtoken stored procedure and Direct-mode SQL. -func (a *DatabaseAuthenticator) oauthUpdateRefreshTokenRecord(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthUpdateRefreshToken) { - var rows int64 - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `UPDATE %s SET session_token = ?, access_token = ?, refresh_token = ?, expires_at = ?, last_activity_at = ? WHERE user_id = ? AND refresh_token = ?`, - a.tableNames.UserSessions)) - res, err := db.ExecContext(ctx, query, newSessionToken, newAccessToken, newRefreshToken, expiresAt, time.Now(), userID, oldRefreshToken) - if err != nil { - return err - } - rows, err = res.RowsAffected() - return err - }) - if err != nil { - return fmt.Errorf("failed to update session: %w", err) - } - if rows == 0 { - return fmt.Errorf("session not found") - } - return nil - } - - updateData := map[string]interface{}{ - "user_id": userID, - "old_refresh_token": oldRefreshToken, - "new_session_token": newSessionToken, - "new_access_token": newAccessToken, - "new_refresh_token": newRefreshToken, - "expires_at": expiresAt, - } - updateJSON, err := json.Marshal(updateData) - if err != nil { - return fmt.Errorf("failed to marshal update data: %w", err) - } - - var updateSuccess bool - var updateErrMsg *string - err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error - FROM %s($1::jsonb) - `, a.sqlNames.OAuthUpdateRefreshToken), updateJSON).Scan(&updateSuccess, &updateErrMsg) - if err != nil { - return fmt.Errorf("failed to update session: %w", err) - } - if !updateSuccess { - if updateErrMsg != nil { - return fmt.Errorf("%s", *updateErrMsg) - } - return fmt.Errorf("failed to update session") - } - return nil -} - -// oauthGetUserByID retrieves user data by ID, dispatching between the -// resolvespec_oauth_getuser stored procedure and Direct-mode SQL. -func (a *DatabaseAuthenticator) oauthGetUserByID(ctx context.Context, userID int) (*UserContext, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetUser) { - var username, email, roles, programUserTable sql.NullString - var userLevel, programUserID sql.NullInt64 - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT username, email, user_level, roles, program_user_id, program_user_table FROM %s WHERE id = ? AND is_active = ?`, - a.tableNames.Users)) - return db.QueryRowContext(ctx, query, userID, true).Scan(&username, &email, &userLevel, &roles, &programUserID, &programUserTable) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("user not found") - } - return nil, fmt.Errorf("failed to get user data: %w", err) - } - return &UserContext{ - UserID: userID, - UserName: username.String, - Email: email.String, - UserLevel: int(userLevel.Int64), - Roles: parseRoles(roles.String), - ProgramUserID: int(programUserID.Int64), - ProgramUserTable: programUserTable.String, - }, nil - } - - var userSuccess bool - var userErrMsg *string - var userData []byte - err := a.getDB().QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1) - `, a.sqlNames.OAuthGetUser), userID).Scan(&userSuccess, &userErrMsg, &userData) - if err != nil { - return nil, fmt.Errorf("failed to get user data: %w", err) - } - if !userSuccess { - if userErrMsg != nil { - return nil, fmt.Errorf("%s", *userErrMsg) - } - return nil, fmt.Errorf("failed to get user data") - } - var userCtx UserContext - if err := json.Unmarshal(userData, &userCtx); err != nil { - return nil, fmt.Errorf("failed to parse user context: %w", err) - } - return &userCtx, nil -} diff --git a/pkg/security/oauth_server_db.go b/pkg/security/oauth_server_db.go index 3462733..330dc4f 100644 --- a/pkg/security/oauth_server_db.go +++ b/pkg/security/oauth_server_db.go @@ -2,229 +2,34 @@ package security import ( "context" - "encoding/json" - "fmt" - "time" ) -// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client. -type OAuthServerClient struct { - ClientID string `json:"client_id"` - RedirectURIs []string `json:"redirect_uris"` - ClientName string `json:"client_name,omitempty"` - GrantTypes []string `json:"grant_types"` - AllowedScopes []string `json:"allowed_scopes,omitempty"` - ClientSecretHash string `json:"client_secret_hash,omitempty"` - TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"` -} - -// OAuthCode is a short-lived authorization code. -type OAuthCode struct { - Code string `json:"code"` - ClientID string `json:"client_id"` - RedirectURI string `json:"redirect_uri"` - ClientState string `json:"client_state,omitempty"` - CodeChallenge string `json:"code_challenge"` - CodeChallengeMethod string `json:"code_challenge_method"` - SessionToken string `json:"session_token"` - RefreshToken string `json:"refresh_token,omitempty"` - Scopes []string `json:"scopes,omitempty"` - ExpiresAt time.Time `json:"expires_at"` -} - -// OAuthTokenInfo is the RFC 7662 token introspection response. -type OAuthTokenInfo struct { - Active bool `json:"active"` - Sub string `json:"sub,omitempty"` - Username string `json:"username,omitempty"` - Email string `json:"email,omitempty"` - UserLevel int `json:"user_level,omitempty"` - Roles []string `json:"roles,omitempty"` - Exp int64 `json:"exp,omitempty"` - Iat int64 `json:"iat,omitempty"` -} - // OAuthRegisterClient persists an OAuth2 client registration. func (a *DatabaseAuthenticator) OAuthRegisterClient(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthRegisterClient) { - return a.oauthRegisterClientDirect(ctx, client) - } - - input, err := json.Marshal(client) - if err != nil { - return nil, fmt.Errorf("failed to marshal client: %w", err) - } - - var success bool - var errMsg *string - var data []byte - - err = a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1::jsonb) - `, a.sqlNames.OAuthRegisterClient), input).Scan(&success, &errMsg, &data) - if err != nil { - return nil, fmt.Errorf("failed to register client: %w", err) - } - if !success { - if errMsg != nil { - return nil, fmt.Errorf("%s", *errMsg) - } - return nil, fmt.Errorf("failed to register client") - } - - var result OAuthServerClient - if err := json.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("failed to parse registered client: %w", err) - } - return &result, nil + return a.src.get().OAuthClient.RegisterClient(ctx, client) } // OAuthGetClient retrieves a registered client by ID. func (a *DatabaseAuthenticator) OAuthGetClient(ctx context.Context, clientID string) (*OAuthServerClient, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetClient) { - return a.oauthGetClientDirect(ctx, clientID) - } - - var success bool - var errMsg *string - var data []byte - - err := a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1) - `, a.sqlNames.OAuthGetClient), clientID).Scan(&success, &errMsg, &data) - if err != nil { - return nil, fmt.Errorf("failed to get client: %w", err) - } - if !success { - if errMsg != nil { - return nil, fmt.Errorf("%s", *errMsg) - } - return nil, fmt.Errorf("client not found") - } - - var result OAuthServerClient - if err := json.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("failed to parse client: %w", err) - } - return &result, nil + return a.src.get().OAuthClient.GetClient(ctx, clientID) } // OAuthSaveCode persists an authorization code. func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCode) error { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthSaveCode) { - return a.oauthSaveCodeDirect(ctx, code) - } - - 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 *string - - err = a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error - FROM %s($1::jsonb) - `, a.sqlNames.OAuthSaveCode), input).Scan(&success, &errMsg) - if err != nil { - return fmt.Errorf("failed to save code: %w", err) - } - if !success { - if errMsg != nil { - return fmt.Errorf("%s", *errMsg) - } - return fmt.Errorf("failed to save code") - } - return nil + return a.src.get().OAuthClient.SaveCode(ctx, code) } // OAuthExchangeCode retrieves and deletes an authorization code (single use). func (a *DatabaseAuthenticator) OAuthExchangeCode(ctx context.Context, code string) (*OAuthCode, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthExchangeCode) { - return a.oauthExchangeCodeDirect(ctx, code) - } - - var success bool - var errMsg *string - var data []byte - - err := a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1) - `, a.sqlNames.OAuthExchangeCode), code).Scan(&success, &errMsg, &data) - if err != nil { - return nil, fmt.Errorf("failed to exchange code: %w", err) - } - if !success { - if errMsg != nil { - return nil, fmt.Errorf("%s", *errMsg) - } - return nil, fmt.Errorf("invalid or expired code") - } - - var result OAuthCode - if err := json.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("failed to parse code data: %w", err) - } - result.Code = code - return &result, nil + return a.src.get().OAuthClient.ExchangeCode(ctx, code) } // OAuthIntrospectToken validates a token and returns its metadata (RFC 7662). func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token string) (*OAuthTokenInfo, error) { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthIntrospect) { - return a.oauthIntrospectTokenDirect(ctx, token) - } - - var success bool - var errMsg *string - var data []byte - - err := a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error, p_data::text - FROM %s($1) - `, a.sqlNames.OAuthIntrospect), token).Scan(&success, &errMsg, &data) - if err != nil { - return nil, fmt.Errorf("failed to introspect token: %w", err) - } - if !success { - if errMsg != nil { - return nil, fmt.Errorf("%s", *errMsg) - } - return nil, fmt.Errorf("introspection failed") - } - - var result OAuthTokenInfo - if err := json.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("failed to parse token info: %w", err) - } - return &result, nil + return a.src.get().OAuthClient.Introspect(ctx, token) } // OAuthRevokeToken revokes a token by deleting the session (RFC 7009). func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error { - if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthRevoke) { - return a.oauthRevokeTokenDirect(ctx, token) - } - - var success bool - var errMsg *string - - err := a.db.QueryRowContext(ctx, fmt.Sprintf(` - SELECT p_success, p_error - FROM %s($1) - `, a.sqlNames.OAuthRevoke), token).Scan(&success, &errMsg) - if err != nil { - return fmt.Errorf("failed to revoke token: %w", err) - } - if !success { - if errMsg != nil { - return fmt.Errorf("%s", *errMsg) - } - return fmt.Errorf("failed to revoke token") - } - return nil + return a.src.get().OAuthClient.Revoke(ctx, token) } diff --git a/pkg/security/oauth_server_db_direct.go b/pkg/security/oauth_server_db_direct.go deleted file mode 100644 index b25a6bb..0000000 --- a/pkg/security/oauth_server_db_direct.go +++ /dev/null @@ -1,209 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "encoding/json" - "errors" - "fmt" - "time" -) - -// Direct-mode implementations mirroring the OAuth2 server stored procedures -// (resolvespec_oauth_register_client, etc.) in database_schema.sql, using -// plain SQL against TableNames.OAuthClients / TableNames.OAuthCodes. -// Array columns (redirect_uris, grant_types, allowed_scopes, scopes) are -// JSON-encoded TEXT instead of native Postgres arrays. - -// nullIfEmpty converts an empty string to a SQL NULL so optional TEXT columns -// (e.g. client_secret_hash for public clients) stay unset rather than "". -func nullIfEmpty(s string) any { - if s == "" { - return nil - } - return s -} - -func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) { - grantTypes := client.GrantTypes - if len(grantTypes) == 0 { - grantTypes = []string{"authorization_code"} - } - allowedScopes := client.AllowedScopes - if len(allowedScopes) == 0 { - allowedScopes = []string{"openid", "profile", "email"} - } - - redirectURIsJSON, err := json.Marshal(client.RedirectURIs) - if err != nil { - return nil, fmt.Errorf("failed to marshal redirect_uris: %w", err) - } - grantTypesJSON, err := json.Marshal(grantTypes) - if err != nil { - return nil, fmt.Errorf("failed to marshal grant_types: %w", err) - } - allowedScopesJSON, err := json.Marshal(allowedScopes) - if err != nil { - return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err) - } - - authMethod := client.TokenEndpointAuthMethod - if authMethod == "" { - authMethod = "none" - } - - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, is_active, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.OAuthClients)) - _, err := db.ExecContext(ctx, query, client.ClientID, string(redirectURIsJSON), client.ClientName, string(grantTypesJSON), string(allowedScopesJSON), nullIfEmpty(client.ClientSecretHash), authMethod, true, time.Now()) - return err - }) - if err != nil { - return nil, fmt.Errorf("failed to register client: %w", err) - } - - return &OAuthServerClient{ - ClientID: client.ClientID, - RedirectURIs: client.RedirectURIs, - ClientName: client.ClientName, - GrantTypes: grantTypes, - AllowedScopes: allowedScopes, - ClientSecretHash: client.ClientSecretHash, - TokenEndpointAuthMethod: authMethod, - }, nil -} - -func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, clientID string) (*OAuthServerClient, error) { - var redirectURIsJSON, grantTypesJSON, allowedScopesJSON sql.NullString - var clientName, clientSecretHash, authMethod sql.NullString - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method FROM %s WHERE client_id = ? AND is_active = ?`, - a.tableNames.OAuthClients)) - return db.QueryRowContext(ctx, query, clientID, true).Scan(&redirectURIsJSON, &clientName, &grantTypesJSON, &allowedScopesJSON, &clientSecretHash, &authMethod) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("client not found") - } - return nil, fmt.Errorf("failed to get client: %w", err) - } - - result := &OAuthServerClient{ - ClientID: clientID, - ClientName: clientName.String, - ClientSecretHash: clientSecretHash.String, - TokenEndpointAuthMethod: authMethod.String, - } - if redirectURIsJSON.Valid { - _ = json.Unmarshal([]byte(redirectURIsJSON.String), &result.RedirectURIs) - } - if grantTypesJSON.Valid { - _ = json.Unmarshal([]byte(grantTypesJSON.String), &result.GrantTypes) - } - if allowedScopesJSON.Valid { - _ = json.Unmarshal([]byte(allowedScopesJSON.String), &result.AllowedScopes) - } - return result, nil -} - -func (a *DatabaseAuthenticator) oauthSaveCodeDirect(ctx context.Context, code *OAuthCode) error { - scopesJSON, err := json.Marshal(code.Scopes) - if err != nil { - return fmt.Errorf("failed to marshal scopes: %w", err) - } - method := code.CodeChallengeMethod - if method == "" { - method = "S256" - } - - return a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, a.tableNames.OAuthCodes)) - _, err := db.ExecContext(ctx, query, code.Code, code.ClientID, code.RedirectURI, code.ClientState, code.CodeChallenge, - method, code.SessionToken, code.RefreshToken, string(scopesJSON), code.ExpiresAt, time.Now()) - return err - }) -} - -func (a *DatabaseAuthenticator) oauthExchangeCodeDirect(ctx context.Context, code string) (*OAuthCode, error) { - var result OAuthCode - var clientState, refreshToken sql.NullString - var scopesJSON sql.NullString - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes - FROM %s WHERE code = ? AND expires_at > ?`, a.tableNames.OAuthCodes)) - err := db.QueryRowContext(ctx, query, code, time.Now()).Scan( - &result.ClientID, &result.RedirectURI, &clientState, &result.CodeChallenge, &result.CodeChallengeMethod, &result.SessionToken, &refreshToken, &scopesJSON) - if err != nil { - return err - } - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE code = ?`, a.tableNames.OAuthCodes)) - _, err = db.ExecContext(ctx, delQuery, code) - return err - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("invalid or expired code") - } - return nil, fmt.Errorf("failed to exchange code: %w", err) - } - - result.Code = code - result.ClientState = clientState.String - result.RefreshToken = refreshToken.String - if scopesJSON.Valid { - _ = json.Unmarshal([]byte(scopesJSON.String), &result.Scopes) - } - return &result, nil -} - -func (a *DatabaseAuthenticator) oauthIntrospectTokenDirect(ctx context.Context, token string) (*OAuthTokenInfo, error) { - var info OAuthTokenInfo - var roles sql.NullString - var exp, iat sql.NullTime - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT u.id, u.username, u.email, u.user_level, u.roles, s.expires_at, s.created_at - FROM %s s JOIN %s u ON u.id = s.user_id - WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`, - a.tableNames.UserSessions, a.tableNames.Users)) - var userID int - err := db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &info.Username, &info.Email, &info.UserLevel, &roles, &exp, &iat) - if err != nil { - return err - } - info.Sub = fmt.Sprintf("%d", userID) - return nil - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return &OAuthTokenInfo{Active: false}, nil - } - return nil, fmt.Errorf("failed to introspect token: %w", err) - } - - info.Active = true - info.Roles = parseRoles(roles.String) - if exp.Valid { - info.Exp = exp.Time.Unix() - } - if iat.Valid { - info.Iat = iat.Time.Unix() - } - return &info, nil -} - -func (a *DatabaseAuthenticator) oauthRevokeTokenDirect(ctx context.Context, token string) error { - return a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ?`, a.tableNames.UserSessions)) - _, err := db.ExecContext(ctx, query, token) - return err - }) -} diff --git a/pkg/security/oauth_server_test.go b/pkg/security/oauth_server_test.go index 13bb838..6c6bb0a 100644 --- a/pkg/security/oauth_server_test.go +++ b/pkg/security/oauth_server_test.go @@ -19,7 +19,7 @@ import ( func newTestOAuthServer(t *testing.T) (*OAuthServer, *DatabaseAuthenticator) { t.Helper() db := newDirectTestDB(t) - auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) + auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig}) srv := NewOAuthServer(OAuthServerConfig{Issuer: "https://auth.example.com", PersistCodes: true}, auth) t.Cleanup(srv.Close) return srv, auth diff --git a/pkg/security/passkey.go b/pkg/security/passkey.go index e0c49a9..54de595 100644 --- a/pkg/security/passkey.go +++ b/pkg/security/passkey.go @@ -3,118 +3,8 @@ package security import ( "context" "encoding/json" - "time" ) -// PasskeyCredential represents a stored WebAuthn/FIDO2 credential -type PasskeyCredential struct { - ID string `json:"id"` - UserID int `json:"user_id"` - CredentialID []byte `json:"credential_id"` // Raw credential ID from authenticator - PublicKey []byte `json:"public_key"` // COSE public key - AttestationType string `json:"attestation_type"` // none, indirect, direct - AAGUID []byte `json:"aaguid"` // Authenticator AAGUID - SignCount uint32 `json:"sign_count"` // Signature counter - CloneWarning bool `json:"clone_warning"` // True if cloning detected - Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal - BackupEligible bool `json:"backup_eligible"` // Credential can be backed up - BackupState bool `json:"backup_state"` // Credential is currently backed up - Name string `json:"name,omitempty"` // User-friendly name - CreatedAt time.Time `json:"created_at"` - LastUsedAt time.Time `json:"last_used_at"` -} - -// PasskeyRegistrationOptions contains options for beginning passkey registration -type PasskeyRegistrationOptions struct { - Challenge []byte `json:"challenge"` - RelyingParty PasskeyRelyingParty `json:"rp"` - User PasskeyUser `json:"user"` - PubKeyCredParams []PasskeyCredentialParam `json:"pubKeyCredParams"` - Timeout int64 `json:"timeout,omitempty"` // Milliseconds - ExcludeCredentials []PasskeyCredentialDescriptor `json:"excludeCredentials,omitempty"` - AuthenticatorSelection *PasskeyAuthenticatorSelection `json:"authenticatorSelection,omitempty"` - Attestation string `json:"attestation,omitempty"` // none, indirect, direct, enterprise - Extensions map[string]any `json:"extensions,omitempty"` -} - -// PasskeyAuthenticationOptions contains options for beginning passkey authentication -type PasskeyAuthenticationOptions struct { - Challenge []byte `json:"challenge"` - Timeout int64 `json:"timeout,omitempty"` - RelyingPartyID string `json:"rpId,omitempty"` - AllowCredentials []PasskeyCredentialDescriptor `json:"allowCredentials,omitempty"` - UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged - Extensions map[string]any `json:"extensions,omitempty"` -} - -// PasskeyRelyingParty identifies the relying party -type PasskeyRelyingParty struct { - ID string `json:"id"` // Domain (e.g., "example.com") - Name string `json:"name"` // Display name -} - -// PasskeyUser identifies the user -type PasskeyUser struct { - ID []byte `json:"id"` // User handle (unique, persistent) - Name string `json:"name"` // Username - DisplayName string `json:"displayName"` // Display name -} - -// PasskeyCredentialParam specifies supported public key algorithm -type PasskeyCredentialParam struct { - Type string `json:"type"` // "public-key" - Alg int `json:"alg"` // COSE algorithm identifier (e.g., -7 for ES256, -257 for RS256) -} - -// PasskeyCredentialDescriptor describes a credential -type PasskeyCredentialDescriptor struct { - Type string `json:"type"` // "public-key" - ID []byte `json:"id"` // Credential ID - Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal -} - -// PasskeyAuthenticatorSelection specifies authenticator requirements -type PasskeyAuthenticatorSelection struct { - AuthenticatorAttachment string `json:"authenticatorAttachment,omitempty"` // platform, cross-platform - RequireResidentKey bool `json:"requireResidentKey,omitempty"` - ResidentKey string `json:"residentKey,omitempty"` // discouraged, preferred, required - UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged -} - -// PasskeyRegistrationResponse contains the client's registration response -type PasskeyRegistrationResponse struct { - ID string `json:"id"` // Base64URL encoded credential ID - RawID []byte `json:"rawId"` // Raw credential ID - Type string `json:"type"` // "public-key" - Response PasskeyAuthenticatorAttestationResponse `json:"response"` - ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"` - Transports []string `json:"transports,omitempty"` -} - -// PasskeyAuthenticatorAttestationResponse contains attestation data -type PasskeyAuthenticatorAttestationResponse struct { - ClientDataJSON []byte `json:"clientDataJSON"` - AttestationObject []byte `json:"attestationObject"` - Transports []string `json:"transports,omitempty"` -} - -// PasskeyAuthenticationResponse contains the client's authentication response -type PasskeyAuthenticationResponse struct { - ID string `json:"id"` // Base64URL encoded credential ID - RawID []byte `json:"rawId"` // Raw credential ID - Type string `json:"type"` // "public-key" - Response PasskeyAuthenticatorAssertionResponse `json:"response"` - ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"` -} - -// PasskeyAuthenticatorAssertionResponse contains assertion data -type PasskeyAuthenticatorAssertionResponse struct { - ClientDataJSON []byte `json:"clientDataJSON"` - AuthenticatorData []byte `json:"authenticatorData"` - Signature []byte `json:"signature"` - UserHandle []byte `json:"userHandle,omitempty"` -} - // PasskeyProvider handles passkey registration and authentication type PasskeyProvider interface { // BeginRegistration creates registration options for a new passkey diff --git a/pkg/security/passkey_provider.go b/pkg/security/passkey_provider.go index 003e968..5c05c35 100644 --- a/pkg/security/passkey_provider.go +++ b/pkg/security/passkey_provider.go @@ -5,26 +5,21 @@ import ( "crypto/rand" "database/sql" "encoding/base64" - "encoding/json" "fmt" - "sync" "time" + + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends" ) -// DatabasePasskeyProvider implements PasskeyProvider using database storage -// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults) +// DatabasePasskeyProvider implements PasskeyProvider on top of the lookup package +// (stored procedures on Postgres by default, direct SQL elsewhere). type DatabasePasskeyProvider struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - rpID string // Relying Party ID (domain) - rpName string // Relying Party display name - rpOrigin string // Expected origin for WebAuthn - timeout int64 // Timeout in milliseconds (default: 60000) - sqlNames *SQLNames - tableNames *TableNames - queryMode QueryMode - capability *dbCapability + src *lookupSource + rpID string // Relying Party ID (domain) + rpName string // Relying Party display name + rpOrigin string // Expected origin for WebAuthn + timeout int64 // Timeout in milliseconds (default: 60000) } // DatabasePasskeyProviderOptions configures the passkey provider @@ -37,12 +32,10 @@ type DatabasePasskeyProviderOptions struct { RPOrigin string // Timeout is the timeout for operations in milliseconds (default: 60000) Timeout int64 - // SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames(). - 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 + // Lookup selects dialect, query mode and procedure/table/column names. + Lookup lookup.Config + // LookupProvider, when set, is used instead of building one from Lookup and the db. + LookupProvider *lookup.Provider // 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) @@ -53,60 +46,20 @@ func NewDatabasePasskeyProvider(db *sql.DB, opts DatabasePasskeyProviderOptions) if opts.Timeout == 0 { opts.Timeout = 60000 // 60 seconds default } - - sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames) - tableNames := resolveTableNames(opts.TableNames) - + src := newLookupSource(db) + src.cfg = opts.Lookup + src.provider = opts.LookupProvider + src.opts = backends.Options{DBFactory: opts.DBFactory} return &DatabasePasskeyProvider{ - db: db, - dbFactory: opts.DBFactory, - rpID: opts.RPID, - rpName: opts.RPName, - rpOrigin: opts.RPOrigin, - timeout: opts.Timeout, - sqlNames: sqlNames, - tableNames: tableNames, - queryMode: opts.QueryMode, - capability: newDBCapability(), + src: src, + rpID: opts.RPID, + rpName: opts.RPName, + rpOrigin: opts.RPOrigin, + timeout: opts.Timeout, } } -func (p *DatabasePasskeyProvider) getDB() *sql.DB { - p.dbMu.RLock() - defer p.dbMu.RUnlock() - return p.db -} - -func (p *DatabasePasskeyProvider) 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 *DatabasePasskeyProvider) runDBOpWithReconnect(run func(*sql.DB) error) error { - db := p.getDB() - if db == nil { - return fmt.Errorf("database connection is nil") - } - err := run(db) - if isDBClosed(err) { - if reconnErr := p.reconnectDB(); reconnErr == nil { - err = run(p.getDB()) - } - } - return err -} +func (p *DatabasePasskeyProvider) store() lookup.PasskeyStore { return p.src.get().Passkey } // BeginRegistration creates registration options for a new passkey func (p *DatabasePasskeyProvider) BeginRegistration(ctx context.Context, userID int, username, displayName string) (*PasskeyRegistrationOptions, error) { @@ -176,69 +129,20 @@ func (p *DatabasePasskeyProvider) CompleteRegistration(ctx context.Context, user credIDB64 := base64.StdEncoding.EncodeToString(response.RawID) pubKeyB64 := base64.StdEncoding.EncodeToString(response.Response.AttestationObject) - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyStoreCredential) { - credentialID, err := p.storeCredentialDirect(ctx, storeCredentialParams{ - UserID: userID, - CredentialID: credIDB64, - PublicKey: pubKeyB64, - AttestationType: "none", - SignCount: 0, - Transports: response.Transports, - BackupEligible: false, - BackupState: false, - Name: "Passkey", - }) - if err != nil { - return nil, err - } - return &PasskeyCredential{ - ID: fmt.Sprintf("%d", credentialID), - UserID: userID, - CredentialID: response.RawID, - PublicKey: response.Response.AttestationObject, - AttestationType: "none", - Transports: response.Transports, - CreatedAt: time.Now(), - LastUsedAt: time.Now(), - }, nil - } - - credData := map[string]any{ - "user_id": userID, - "credential_id": credIDB64, - "public_key": pubKeyB64, - "attestation_type": "none", - "sign_count": 0, - "transports": response.Transports, - "backup_eligible": false, - "backup_state": false, - "name": "Passkey", - } - - credJSON, err := json.Marshal(credData) + credentialID, err := p.store().Store(ctx, lookup.PasskeyCredentialRecord{ + UserID: userID, + CredentialID: credIDB64, + PublicKey: pubKeyB64, + AttestationType: "none", + Transports: response.Transports, + Name: "Passkey", + }) if err != nil { - return nil, fmt.Errorf("failed to marshal credential data: %w", err) - } - - var success bool - var errorMsg sql.NullString - var credentialID sql.NullInt64 - - query := fmt.Sprintf(`SELECT p_success, p_error, p_credential_id FROM %s($1::jsonb)`, p.sqlNames.PasskeyStoreCredential) - err = p.getDB().QueryRowContext(ctx, query, string(credJSON)).Scan(&success, &errorMsg, &credentialID) - if err != nil { - return nil, fmt.Errorf("failed to store credential: %w", err) - } - - if !success { - if errorMsg.Valid { - return nil, fmt.Errorf("%s", errorMsg.String) - } - return nil, fmt.Errorf("failed to store credential") + return nil, err } return &PasskeyCredential{ - ID: fmt.Sprintf("%d", credentialID.Int64), + ID: fmt.Sprintf("%d", credentialID), UserID: userID, CredentialID: response.RawID, PublicKey: response.Response.AttestationObject, @@ -260,41 +164,15 @@ func (p *DatabasePasskeyProvider) BeginAuthentication(ctx context.Context, usern // If username is provided, get user's credentials var allowCredentials []PasskeyCredentialDescriptor if username != "" { - var creds []passkeyCredential - - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) { - _, directCreds, err := p.getCredsByUsernameDirect(ctx, username) - if err != nil { - return nil, err - } - creds = directCreds - } else { - var success bool - var errorMsg sql.NullString - var userID sql.NullInt64 - var credentialsJSON sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetCredsByUsername) - err := p.getDB().QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON) - if err != nil { - return nil, fmt.Errorf("failed to get credentials: %w", err) - } - - if !success { - if errorMsg.Valid { - return nil, fmt.Errorf("%s", errorMsg.String) - } - return nil, fmt.Errorf("failed to get credentials") - } - - if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil { - return nil, fmt.Errorf("failed to parse credentials: %w", err) - } + _, refs, err := p.store().ByUsername(ctx, username) + if err != nil { + return nil, err } + creds := refs allowCredentials = make([]PasskeyCredentialDescriptor, 0, len(creds)) for _, cred := range creds { - credID, err := base64.StdEncoding.DecodeString(cred.ID) + credID, err := base64.StdEncoding.DecodeString(cred.CredentialID) if err != nil { continue } @@ -327,214 +205,47 @@ func (p *DatabasePasskeyProvider) CompleteAuthentication(ctx context.Context, re credIDB64 := base64.StdEncoding.EncodeToString(response.RawID) - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredential) { - userID, signCount, err := p.getCredentialDirect(ctx, credIDB64) - if err != nil { - return 0, err - } - newCounter := signCount + 1 - cloneWarning, err := p.updateCounterDirect(ctx, credIDB64, newCounter) - if err != nil { - return 0, fmt.Errorf("failed to update counter: %w", err) - } - if cloneWarning { - return 0, fmt.Errorf("credential cloning detected") - } - return userID, nil - } - - // Get credential from database - var success bool - var errorMsg sql.NullString - var credentialJSON sql.NullString - - runQuery := func() error { - query := fmt.Sprintf(`SELECT p_success, p_error, p_credential::text FROM %s($1)`, p.sqlNames.PasskeyGetCredential) - return p.getDB().QueryRowContext(ctx, query, response.RawID).Scan(&success, &errorMsg, &credentialJSON) - } - err := runQuery() - if isDBClosed(err) { - if reconnErr := p.reconnectDB(); reconnErr == nil { - err = runQuery() - } - } - if err != nil { - return 0, fmt.Errorf("failed to get credential: %w", err) - } - - if !success { - if errorMsg.Valid { - return 0, fmt.Errorf("%s", errorMsg.String) - } - return 0, fmt.Errorf("credential not found") - } - - // Parse credential - var cred struct { - UserID int `json:"user_id"` - SignCount uint32 `json:"sign_count"` - } - if err := json.Unmarshal([]byte(credentialJSON.String), &cred); err != nil { - return 0, fmt.Errorf("failed to parse credential: %w", err) - } - // TODO: Verify signature here // For now, we'll just update the counter as a placeholder + store := p.store() + userID, signCount, err := store.Get(ctx, credIDB64) + if err != nil { + return 0, err + } // Update counter (in production, this should be done after successful verification) - newCounter := cred.SignCount + 1 - var updateSuccess bool - var updateError sql.NullString - var cloneWarning sql.NullBool - - updateQuery := fmt.Sprintf(`SELECT p_success, p_error, p_clone_warning FROM %s($1, $2)`, p.sqlNames.PasskeyUpdateCounter) - err = p.getDB().QueryRowContext(ctx, updateQuery, response.RawID, newCounter).Scan(&updateSuccess, &updateError, &cloneWarning) + cloneWarning, err := store.UpdateCounter(ctx, credIDB64, signCount+1) if err != nil { return 0, fmt.Errorf("failed to update counter: %w", err) } - - if cloneWarning.Valid && cloneWarning.Bool { + if cloneWarning { return 0, fmt.Errorf("credential cloning detected") } - return cred.UserID, nil + return userID, nil } // GetCredentials returns all passkey credentials for a user func (p *DatabasePasskeyProvider) GetCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error) { - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetUserCredentials) { - return p.getUserCredentialsDirect(ctx, userID) - } - - var success bool - var errorMsg sql.NullString - var credentialsJSON sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetUserCredentials) - err := p.getDB().QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &credentialsJSON) - if err != nil { - return nil, fmt.Errorf("failed to get credentials: %w", err) - } - - if !success { - if errorMsg.Valid { - return nil, fmt.Errorf("%s", errorMsg.String) - } - return nil, fmt.Errorf("failed to get credentials") - } - - // Parse 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([]byte(credentialsJSON.String), &rawCreds); err != nil { - return nil, fmt.Errorf("failed to parse credentials: %w", err) - } - - credentials := make([]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, 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 + return p.store().List(ctx, userID) } // DeleteCredential removes a passkey credential func (p *DatabasePasskeyProvider) DeleteCredential(ctx context.Context, userID int, credentialID string) error { - credID, err := base64.StdEncoding.DecodeString(credentialID) + _, err := base64.StdEncoding.DecodeString(credentialID) if err != nil { return fmt.Errorf("invalid credential ID: %w", err) } - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyDeleteCredential) { - return p.deleteCredentialDirect(ctx, userID, credentialID) - } - - var success bool - var errorMsg sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, p.sqlNames.PasskeyDeleteCredential) - err = p.getDB().QueryRowContext(ctx, query, userID, credID).Scan(&success, &errorMsg) - if err != nil { - return fmt.Errorf("failed to delete credential: %w", err) - } - - if !success { - if errorMsg.Valid { - return fmt.Errorf("%s", errorMsg.String) - } - return fmt.Errorf("failed to delete credential") - } - - return nil + return p.store().Delete(ctx, userID, credentialID) } // UpdateCredentialName updates the friendly name of a credential func (p *DatabasePasskeyProvider) UpdateCredentialName(ctx context.Context, userID int, credentialID string, name string) error { - credID, err := base64.StdEncoding.DecodeString(credentialID) + _, err := base64.StdEncoding.DecodeString(credentialID) if err != nil { return fmt.Errorf("invalid credential ID: %w", err) } - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyUpdateName) { - return p.updateNameDirect(ctx, userID, credentialID, name) - } - - var success bool - var errorMsg sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3)`, p.sqlNames.PasskeyUpdateName) - err = p.getDB().QueryRowContext(ctx, query, userID, credID, name).Scan(&success, &errorMsg) - if err != nil { - return fmt.Errorf("failed to update credential name: %w", err) - } - - if !success { - if errorMsg.Valid { - return fmt.Errorf("%s", errorMsg.String) - } - return fmt.Errorf("failed to update credential name") - } - - return nil + return p.store().Rename(ctx, userID, credentialID, name) } diff --git a/pkg/security/passkey_provider_direct.go b/pkg/security/passkey_provider_direct.go deleted file mode 100644 index 229121f..0000000 --- a/pkg/security/passkey_provider_direct.go +++ /dev/null @@ -1,256 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "encoding/base64" - "encoding/json" - "errors" - "fmt" - "time" -) - -// Direct-mode implementations mirroring the resolvespec_passkey_* stored -// procedures in database_schema.sql using plain SQL against TableNames. -// credential_id/public_key/aaguid are stored as base64 TEXT (not native -// bytea) and transports as JSON-encoded TEXT, so the same schema works on -// SQLite, MySQL, and Postgres. - -type storeCredentialParams struct { - UserID int - CredentialID string // base64 - PublicKey string // base64 - AttestationType string - SignCount int - Transports []string - BackupEligible bool - BackupState bool - Name string -} - -func (p *DatabasePasskeyProvider) storeCredentialDirect(ctx context.Context, params storeCredentialParams) (int64, error) { - transportsJSON, err := json.Marshal(params.Transports) - if err != nil { - return 0, fmt.Errorf("failed to marshal transports: %w", err) - } - - var credentialID int64 - err = p.runDBOpWithReconnect(func(db *sql.DB) error { - var exists int - checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - if err := db.QueryRowContext(ctx, checkQuery, params.CredentialID).Scan(&exists); err == nil { - return fmt.Errorf("credential already exists") - } else if !errors.Is(err, sql.ErrNoRows) { - return err - } - - var userExists int - userCheckQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE id = ?`, p.tableNames.Users)) - if err := db.QueryRowContext(ctx, userCheckQuery, params.UserID).Scan(&userExists); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return fmt.Errorf("user not found") - } - return err - } - - now := time.Now() - insertQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (user_id, credential_id, public_key, attestation_type, aaguid, sign_count, transports, backup_eligible, backup_state, name, created_at, last_used_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, p.tableNames.UserPasskeyCredentials)) - res, err := db.ExecContext(ctx, insertQuery, params.UserID, params.CredentialID, params.PublicKey, params.AttestationType, - "", params.SignCount, string(transportsJSON), params.BackupEligible, params.BackupState, params.Name, now, now) - if err != nil { - return err - } - credentialID, err = res.LastInsertId() - return err - }) - if err != nil { - return 0, err - } - return credentialID, nil -} - -func (p *DatabasePasskeyProvider) getCredentialDirect(ctx context.Context, credentialIDB64 string) (userID int, signCount uint32, err error) { - err = p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT user_id, sign_count FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - return db.QueryRowContext(ctx, query, credentialIDB64).Scan(&userID, &signCount) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return 0, 0, fmt.Errorf("credential not found") - } - return 0, 0, fmt.Errorf("failed to get credential: %w", err) - } - return userID, signCount, nil -} - -func (p *DatabasePasskeyProvider) updateCounterDirect(ctx context.Context, credentialIDB64 string, newCounter uint32) (cloneWarning bool, err error) { - err = p.runDBOpWithReconnect(func(db *sql.DB) error { - var oldCounter int - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT sign_count FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - if err := db.QueryRowContext(ctx, query, credentialIDB64).Scan(&oldCounter); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return fmt.Errorf("credential not found") - } - return err - } - - if int(newCounter) <= oldCounter { - cloneWarning = true - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET clone_warning = ? WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - _, err := db.ExecContext(ctx, updQuery, true, credentialIDB64) - return err - } - - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET sign_count = ?, last_used_at = ? WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - _, err := db.ExecContext(ctx, updQuery, newCounter, time.Now(), credentialIDB64) - return err - }) - return cloneWarning, err -} - -func (p *DatabasePasskeyProvider) getUserCredentialsDirect(ctx context.Context, userID int) ([]PasskeyCredential, error) { - var credentials []PasskeyCredential - err := p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, user_id, credential_id, public_key, attestation_type, aaguid, sign_count, clone_warning, transports, backup_eligible, backup_state, name, created_at, last_used_at - FROM %s WHERE user_id = ? ORDER BY created_at DESC`, p.tableNames.UserPasskeyCredentials)) - rows, err := db.QueryContext(ctx, query, userID) - if err != nil { - return err - } - defer rows.Close() - - credentials = make([]PasskeyCredential, 0) - for rows.Next() { - var id, uid int - var credIDB64, pubKeyB64, attestationType, aaguidB64, name string - var signCount uint32 - var cloneWarning, backupEligible, backupState bool - var transportsJSON sql.NullString - var createdAt, lastUsedAt time.Time - - if err := rows.Scan(&id, &uid, &credIDB64, &pubKeyB64, &attestationType, &aaguidB64, &signCount, - &cloneWarning, &transportsJSON, &backupEligible, &backupState, &name, &createdAt, &lastUsedAt); err != nil { - return err - } - - credID, err := base64.StdEncoding.DecodeString(credIDB64) - if err != nil { - continue - } - pubKey, err := base64.StdEncoding.DecodeString(pubKeyB64) - if err != nil { - continue - } - aaguid, _ := base64.StdEncoding.DecodeString(aaguidB64) - - var transports []string - if transportsJSON.Valid && transportsJSON.String != "" { - _ = json.Unmarshal([]byte(transportsJSON.String), &transports) - } - - credentials = append(credentials, PasskeyCredential{ - ID: fmt.Sprintf("%d", id), - UserID: uid, - CredentialID: credID, - PublicKey: pubKey, - AttestationType: attestationType, - AAGUID: aaguid, - SignCount: signCount, - CloneWarning: cloneWarning, - Transports: transports, - BackupEligible: backupEligible, - BackupState: backupState, - Name: name, - CreatedAt: createdAt, - LastUsedAt: lastUsedAt, - }) - } - return rows.Err() - }) - if err != nil { - return nil, fmt.Errorf("failed to get credentials: %w", err) - } - return credentials, nil -} - -func (p *DatabasePasskeyProvider) deleteCredentialDirect(ctx context.Context, userID int, credentialIDB64 string) error { - return p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ? AND credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - res, err := db.ExecContext(ctx, query, userID, credentialIDB64) - if err != nil { - return err - } - rows, err := res.RowsAffected() - if err != nil { - return err - } - if rows == 0 { - return fmt.Errorf("credential not found") - } - return nil - }) -} - -func (p *DatabasePasskeyProvider) updateNameDirect(ctx context.Context, userID int, credentialIDB64, name string) error { - return p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET name = ? WHERE user_id = ? AND credential_id = ?`, p.tableNames.UserPasskeyCredentials)) - res, err := db.ExecContext(ctx, query, name, userID, credentialIDB64) - if err != nil { - return err - } - rows, err := res.RowsAffected() - if err != nil { - return err - } - if rows == 0 { - return fmt.Errorf("credential not found") - } - return nil - }) -} - -type passkeyCredential struct { - ID string `json:"credential_id"` - Transports []string `json:"transports"` -} - -func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context, username string) (userID int, creds []passkeyCredential, err error) { - err = p.runDBOpWithReconnect(func(db *sql.DB) error { - userQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE username = ? AND is_active = ?`, p.tableNames.Users)) - if err := db.QueryRowContext(ctx, userQuery, username, true).Scan(&userID); err != nil { - return err - } - - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT credential_id, transports FROM %s WHERE user_id = ?`, p.tableNames.UserPasskeyCredentials)) - rows, err := db.QueryContext(ctx, query, userID) - if err != nil { - return err - } - defer rows.Close() - - creds = make([]passkeyCredential, 0) - for rows.Next() { - var credID string - var transportsJSON sql.NullString - if err := rows.Scan(&credID, &transportsJSON); err != nil { - return err - } - var transports []string - if transportsJSON.Valid && transportsJSON.String != "" { - _ = json.Unmarshal([]byte(transportsJSON.String), &transports) - } - creds = append(creds, passkeyCredential{ID: credID, Transports: transports}) - } - return rows.Err() - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return 0, nil, fmt.Errorf("user not found") - } - return 0, nil, fmt.Errorf("failed to get credentials: %w", err) - } - return userID, creds, nil -} diff --git a/pkg/security/provider.go b/pkg/security/provider.go index 547e70b..d6d488f 100644 --- a/pkg/security/provider.go +++ b/pkg/security/provider.go @@ -5,7 +5,6 @@ import ( "errors" "fmt" "reflect" - "regexp" "strings" "sync" "time" @@ -18,36 +17,6 @@ import ( "golang.org/x/sync/singleflight" ) -type ColumnSecurity struct { - Schema string `json:"schema"` - Tablename string `json:"tablename"` - Path []string `json:"path"` - ExtraFilters map[string]string `json:"extra_filters"` - UserID int `json:"user_id"` - Accesstype string `json:"accesstype"` - MaskStart int `json:"mask_start"` - MaskEnd int `json:"mask_end"` - MaskInvert bool `json:"mask_invert"` - MaskChar string `json:"mask_char"` - Control string `json:"control"` - ID int `json:"id"` -} - -type RowSecurity struct { - Schema string `json:"schema"` - Tablename string `json:"tablename"` - Template string `json:"template"` - HasBlock bool `json:"has_block"` - // UserID is the opaque user reference the security rules were loaded for. - // It may be an int, a string/UUID, or a *UserContext, depending on what the - // RowSecurityProvider/SecurityContext.GetUserRef implementation returns. - UserID any `json:"user_id"` -} - -// safeIdentRe matches an unquoted SQL identifier. Identifiers substituted into a -// row-security template must match it; anything else is rejected. -var safeIdentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) - // ErrNoRowSecurity is returned by GetRowSecurityTemplate when no row security // entry is loaded for the user and table. It means "no rules", as opposed to a // failure, which callers must treat as fatal. @@ -56,66 +25,6 @@ var ErrNoRowSecurity = errors.New("no row security data") // ErrNoColumnSecurity is the column-security equivalent of ErrNoRowSecurity. var ErrNoColumnSecurity = errors.New("no column security data") -// userIDScalar reduces the opaque user reference to a scalar that is safe to -// bind as a query argument. A *UserContext is reduced to its UserID; other -// structured values are rejected rather than stringified into SQL. -func userIDScalar(ref any) (any, error) { - switch v := ref.(type) { - case nil: - return nil, fmt.Errorf("row security: no user reference") - case *UserContext: - if v == nil { - return nil, fmt.Errorf("row security: nil user context") - } - return v.UserID, nil - case UserContext: - return v.UserID, nil - case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64: - return v, nil - case string: - return v, nil - default: - return nil, fmt.Errorf("row security: unsupported user reference type %T", ref) - } -} - -// GetTemplate expands the row-security template into a WHERE clause and its -// bind arguments. {PrimaryKeyName}, {TableName} and {SchemaName} are validated -// identifiers substituted in place; every {UserID} becomes a `?` placeholder -// with the user reference bound as an argument, so user data never reaches the -// SQL text. -func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) (clause string, args []any, err error) { - str := m.Template - - for placeholder, ident := range map[string]string{ - "{PrimaryKeyName}": pPrimaryKeyName, - "{TableName}": m.Tablename, - "{SchemaName}": m.Schema, - } { - if !strings.Contains(str, placeholder) { - continue - } - if !safeIdentRe.MatchString(ident) { - return "", nil, fmt.Errorf("row security: invalid identifier %q for %s", ident, placeholder) - } - str = strings.ReplaceAll(str, placeholder, ident) - } - - n := strings.Count(str, "{UserID}") - if n == 0 { - return str, nil, nil - } - uid, err := userIDScalar(m.UserID) - if err != nil { - return "", nil, err - } - args = make([]any, n) - for i := range args { - args[i] = uid - } - return strings.ReplaceAll(str, "{UserID}", "?"), args, nil -} - // SecurityList manages security state and caching // It wraps a SecurityProvider and provides caching and utility methods type SecurityList struct { diff --git a/pkg/security/providers.go b/pkg/security/providers.go index 1b5b3d1..e10f50b 100644 --- a/pkg/security/providers.go +++ b/pkg/security/providers.go @@ -3,7 +3,6 @@ package security import ( "context" "database/sql" - "encoding/json" "fmt" "net/http" "strconv" @@ -16,57 +15,13 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/dbtrace" "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends" ) // 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 @@ -109,24 +64,14 @@ func (t *activityThrottle) allow(token string, now time.Time) bool { // 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 +// Procedure names and modes are configured through lookup.Config (see lookup.DefaultProcNames) +// See lookup/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 + src *lookupSource + cache *cache.Cache + cacheTTL time.Duration // activityWG tracks in-flight asynchronous session activity updates activityWG sync.WaitGroup @@ -159,13 +104,11 @@ type DatabaseAuthenticatorOptions struct { 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 + // Lookup selects dialect, query mode and procedure/table/column names. + // The zero value uses stored procedures on Postgres and direct SQL elsewhere. + Lookup lookup.Config + // LookupProvider, when set, is used instead of building one from Lookup and the db. + LookupProvider *lookup.Provider // 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) @@ -204,173 +147,49 @@ func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorO cacheInstance = cache.GetDefaultCache() } - sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames) - tableNames := resolveTableNames(opts.TableNames) + src := newLookupSource(db) + src.cfg = opts.Lookup + src.provider = opts.LookupProvider + src.opts = backends.Options{DBFactory: opts.DBFactory, UpgradePasswordHash: opts.UpgradePasswordHash} return &DatabaseAuthenticator{ - db: db, - dbFactory: opts.DBFactory, + src: src, 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) auth() lookup.AuthStore { return a.src.get().Auth } 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) - } + return a.auth().Login(ctx, req) +} - 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 +// LoginWithAPIKey implements APIKeyLoginable. It validates a raw header/generic +// API key and creates a session for the key's user. Unknown, expired and +// inactive keys all return errInvalidAPIKey; the raw key is never logged. +// Procedure-only: the key and user lookup live in resolvespec_login_api_key so +// the underlying schema can differ per database; there is no direct-SQL path. +func (a *DatabaseAuthenticator) LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) { + return a.auth().LoginAPIKey(ctx, rawKey, claims) } // 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 + return a.auth().Register(ctx, req) } 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") + if err := a.auth().Logout(ctx, req); err != nil { + return err } // Clear cache for this token @@ -467,40 +286,8 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err err := a.cache.GetOrSet(r.Context(), cacheKey, &loaded, a.cacheTTL, func() (any, error) { // This function is called only if cache miss dbtrace.Raw(r.Context(), "auth.session") - 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 + return a.auth().Session(r.Context(), token, reference) }) if err != nil { return nil, err @@ -566,262 +353,61 @@ func (a *DatabaseAuthenticator) ClearUserCache(userID int) error { // updateSessionActivity updates the last activity timestamp for the session func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) { dbtrace.Raw(ctx, "auth.activity") - 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 - }) + _ = a.auth().TouchSession(ctx, sessionToken, userCtx) } // 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 + return a.auth().Refresh(ctx, refreshToken) } // JWTAuthenticator provides JWT token-based authentication // All database operations go through stored procedures -// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults) +// Procedure names and modes are configured through lookup.Config (see lookup.DefaultProcNames) // 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 + secretKey []byte + src *lookupSource } // 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 + a.src.opts.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(), - } +func NewJWTAuthenticator(secretKey string, db *sql.DB) *JWTAuthenticator { + return &JWTAuthenticator{secretKey: []byte(secretKey), src: newLookupSource(db)} } // 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 + a.src.opts.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) +// WithLookup configures dialect, query mode and names. Call before first use. +func (a *JWTAuthenticator) WithLookup(cfg lookup.Config) *JWTAuthenticator { + a.src.cfg = cfg return a } -// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto). -func (a *JWTAuthenticator) WithQueryMode(mode QueryMode) *JWTAuthenticator { - a.queryMode = mode +// WithLookupProvider uses an existing provider instead of building one. +func (a *JWTAuthenticator) WithLookupProvider(p *lookup.Provider) *JWTAuthenticator { + a.src.provider = p 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) auth() lookup.AuthStore { return a.src.get().Auth } 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 + return a.auth().JWTLogin(ctx, req) } 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 + return a.auth().JWTLogout(ctx, req) } func (a *JWTAuthenticator) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) { @@ -850,279 +436,85 @@ func (a *JWTAuthenticator) Authenticate(r *http.Request) (*UserContext, error) { // 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) +// DatabaseColumnSecurityProvider loads column security through the lookup package +// (stored procedure on Postgres by default, direct SQL elsewhere). type DatabaseColumnSecurityProvider struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - sqlNames *SQLNames - queryMode QueryMode - capability *dbCapability + src *lookupSource } -func NewDatabaseColumnSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseColumnSecurityProvider { - return &DatabaseColumnSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()} +func NewDatabaseColumnSecurityProvider(db *sql.DB) *DatabaseColumnSecurityProvider { + return &DatabaseColumnSecurityProvider{src: newLookupSource(db)} } -// 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 +// WithLookup configures dialect, query mode and names. Call before first use. +func (p *DatabaseColumnSecurityProvider) WithLookup(cfg lookup.Config) *DatabaseColumnSecurityProvider { + p.src.cfg = cfg + return p +} + +// WithLookupProvider uses an existing provider instead of building one. +func (p *DatabaseColumnSecurityProvider) WithLookupProvider(lp *lookup.Provider) *DatabaseColumnSecurityProvider { + p.src.provider = lp + return p +} + +// WithNoGroupTables skips group membership when loading rules in direct mode. +func (p *DatabaseColumnSecurityProvider) WithNoGroupTables() *DatabaseColumnSecurityProvider { + p.src.opts.NoGroupTables = true return p } func (p *DatabaseColumnSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseColumnSecurityProvider { - p.dbFactory = factory + p.src.opts.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 - } dbtrace.Raw(ctx, "security.column") - - 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 + return p.src.get().Policy.ColumnSecurity(ctx, userID, schema, table) } -// DatabaseRowSecurityProvider loads row security from database -// All database operations go through stored procedures -// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults) +// DatabaseRowSecurityProvider loads row security through the lookup package +// (stored procedure on Postgres by default, direct SQL elsewhere). type DatabaseRowSecurityProvider struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - sqlNames *SQLNames - queryMode QueryMode - capability *dbCapability + src *lookupSource } -func NewDatabaseRowSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseRowSecurityProvider { - return &DatabaseRowSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()} +func NewDatabaseRowSecurityProvider(db *sql.DB) *DatabaseRowSecurityProvider { + return &DatabaseRowSecurityProvider{src: newLookupSource(db)} } -// 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 +// WithLookup configures dialect, query mode and names. Call before first use. +func (p *DatabaseRowSecurityProvider) WithLookup(cfg lookup.Config) *DatabaseRowSecurityProvider { + p.src.cfg = cfg + return p +} + +// WithLookupProvider uses an existing provider instead of building one. +func (p *DatabaseRowSecurityProvider) WithLookupProvider(lp *lookup.Provider) *DatabaseRowSecurityProvider { + p.src.provider = lp + return p +} + +// WithNoGroupTables skips group membership when loading rules in direct mode. +func (p *DatabaseRowSecurityProvider) WithNoGroupTables() *DatabaseRowSecurityProvider { + p.src.opts.NoGroupTables = true return p } func (p *DatabaseRowSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseRowSecurityProvider { - p.dbFactory = factory + p.src.opts.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 - } dbtrace.Raw(ctx, "security.row") - - // 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 + return p.src.get().Policy.RowSecurity(ctx, userRef, schema, table) } // 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{} @@ -1169,73 +561,13 @@ func generateRandomString(length int) string { // 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 + return a.auth().ResetRequest(ctx, req) } // 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 + return a.auth().ResetComplete(ctx, req) } // Passkey authentication methods @@ -1294,55 +626,7 @@ func (a *DatabaseAuthenticator) LoginWithPasskey(ctx context.Context, req Passke 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 + return a.src.get().Passkey.Login(ctx, userID, req.Claims) } // GetPasskeyCredentials returns all passkey credentials for a user diff --git a/pkg/security/providers/header.go b/pkg/security/providers/header.go new file mode 100644 index 0000000..1e14e06 --- /dev/null +++ b/pkg/security/providers/header.go @@ -0,0 +1,76 @@ +package providers + +import ( + "context" + "fmt" + "net/http" + "strconv" + "strings" + + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// 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 sectypes.LoginRequest) (*sectypes.LoginResponse, error) { + return nil, fmt.Errorf("header authentication does not support login") +} + +func (a *HeaderAuthenticator) LoginWithCookie(ctx context.Context, req sectypes.LoginRequest, w http.ResponseWriter) (*sectypes.LoginResponse, error) { + return a.Login(ctx, req) +} + +func (a *HeaderAuthenticator) Logout(ctx context.Context, req sectypes.LogoutRequest) error { + return nil +} + +func (a *HeaderAuthenticator) LogoutWithCookie(ctx context.Context, req sectypes.LogoutRequest, w http.ResponseWriter) error { + return a.Logout(ctx, req) +} + +func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*sectypes.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 §ypes.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 +} + +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 +} diff --git a/pkg/security/keystore_authenticator.go b/pkg/security/providers/keystore_authenticator.go similarity index 76% rename from pkg/security/keystore_authenticator.go rename to pkg/security/providers/keystore_authenticator.go index 567261e..0437f5e 100644 --- a/pkg/security/keystore_authenticator.go +++ b/pkg/security/providers/keystore_authenticator.go @@ -1,8 +1,10 @@ -package security +package providers import ( "context" "fmt" + "github.com/bitechdev/ResolveSpec/pkg/security" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" "net/http" "strings" ) @@ -17,45 +19,45 @@ import ( // 2. Authorization: ApiKey // 3. X-API-Key header type KeyStoreAuthenticator struct { - keyStore KeyStore - keyType KeyType // empty = accept any type - authenticateCallback func(r *http.Request) (*UserContext, error) + keyStore security.KeyStore + keyType sectypes.KeyType // empty = accept any type + authenticateCallback func(r *http.Request) (*sectypes.UserContext, error) } // NewKeyStoreAuthenticator creates a KeyStoreAuthenticator. // Pass an empty keyType to accept keys of any type. -func NewKeyStoreAuthenticator(ks KeyStore, keyType KeyType) *KeyStoreAuthenticator { +func NewKeyStoreAuthenticator(ks security.KeyStore, keyType sectypes.KeyType) *KeyStoreAuthenticator { return &KeyStoreAuthenticator{keyStore: ks, keyType: keyType} } // Login is not supported for keystore authentication. -func (a *KeyStoreAuthenticator) Login(_ context.Context, _ LoginRequest) (*LoginResponse, error) { +func (a *KeyStoreAuthenticator) Login(_ context.Context, _ sectypes.LoginRequest) (*sectypes.LoginResponse, error) { return nil, fmt.Errorf("keystore authenticator does not support login") } // LoginWithCookie is not supported for keystore authentication. -func (a *KeyStoreAuthenticator) LoginWithCookie(_ context.Context, _ LoginRequest, _ http.ResponseWriter) (*LoginResponse, error) { +func (a *KeyStoreAuthenticator) LoginWithCookie(_ context.Context, _ sectypes.LoginRequest, _ http.ResponseWriter) (*sectypes.LoginResponse, error) { return nil, fmt.Errorf("keystore authenticator does not support login") } // Logout is not supported for keystore authentication. -func (a *KeyStoreAuthenticator) Logout(_ context.Context, _ LogoutRequest) error { +func (a *KeyStoreAuthenticator) Logout(_ context.Context, _ sectypes.LogoutRequest) error { return nil } // LogoutWithCookie is not supported for keystore authentication. -func (a *KeyStoreAuthenticator) LogoutWithCookie(_ context.Context, _ LogoutRequest, _ http.ResponseWriter) error { +func (a *KeyStoreAuthenticator) LogoutWithCookie(_ context.Context, _ sectypes.LogoutRequest, _ http.ResponseWriter) error { return nil } // SetAuthenticateCallback registers a fallback called when key authentication fails. -func (a *KeyStoreAuthenticator) SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error)) { +func (a *KeyStoreAuthenticator) SetAuthenticateCallback(fn func(r *http.Request) (*sectypes.UserContext, error)) { a.authenticateCallback = fn } // Authenticate extracts an API key from the request and validates it against the KeyStore. // Returns a UserContext built from the matching UserKey on success. -func (a *KeyStoreAuthenticator) Authenticate(r *http.Request) (*UserContext, error) { +func (a *KeyStoreAuthenticator) Authenticate(r *http.Request) (*sectypes.UserContext, error) { rawKey := extractAPIKey(r) if rawKey == "" { if a.authenticateCallback != nil { @@ -93,7 +95,7 @@ func extractAPIKey(r *http.Request) string { // userKeyToUserContext converts a UserKey into a UserContext. // Scopes are mapped to Roles. Key type and name are stored in Claims. -func userKeyToUserContext(k *UserKey) *UserContext { +func userKeyToUserContext(k *sectypes.UserKey) *sectypes.UserContext { claims := map[string]any{ "key_type": string(k.KeyType), "key_name": k.Name, @@ -109,7 +111,7 @@ func userKeyToUserContext(k *UserKey) *UserContext { roles = []string{} } - return &UserContext{ + return §ypes.UserContext{ UserID: k.UserID, SessionID: fmt.Sprintf("key:%d", k.ID), Roles: roles, diff --git a/pkg/security/keystore_config.go b/pkg/security/providers/keystore_config.go similarity index 84% rename from pkg/security/keystore_config.go rename to pkg/security/providers/keystore_config.go index 353cf14..efed6e1 100644 --- a/pkg/security/keystore_config.go +++ b/pkg/security/providers/keystore_config.go @@ -1,4 +1,4 @@ -package security +package providers import ( "context" @@ -7,6 +7,7 @@ import ( "encoding/base64" "encoding/hex" "fmt" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" "sync" "sync/atomic" "time" @@ -20,16 +21,16 @@ import ( // Keys created at runtime via CreateKey are held in memory only and lost on restart. type ConfigKeyStore struct { mu sync.RWMutex - keys []UserKey + keys []sectypes.UserKey next int64 // monotonic ID counter for runtime-created keys (atomic) } // NewConfigKeyStore creates a ConfigKeyStore seeded with the provided keys. // Pass nil or an empty slice to start with no pre-loaded keys. // Zero-value entries (CreatedAt is zero) are treated as active and assigned the current time. -func NewConfigKeyStore(keys []UserKey) *ConfigKeyStore { +func NewConfigKeyStore(keys []sectypes.UserKey) *ConfigKeyStore { var maxID int64 - copied := make([]UserKey, len(keys)) + copied := make([]sectypes.UserKey, len(keys)) copy(copied, keys) for i := range copied { if copied[i].CreatedAt.IsZero() { @@ -44,16 +45,16 @@ func NewConfigKeyStore(keys []UserKey) *ConfigKeyStore { } // CreateKey generates a new raw key, stores its SHA-256 hash, and returns the raw key once. -func (s *ConfigKeyStore) CreateKey(_ context.Context, req CreateKeyRequest) (*CreateKeyResponse, error) { +func (s *ConfigKeyStore) CreateKey(_ context.Context, req sectypes.CreateKeyRequest) (*sectypes.CreateKeyResponse, error) { rawBytes := make([]byte, 32) if _, err := rand.Read(rawBytes); err != nil { return nil, fmt.Errorf("failed to generate key material: %w", err) } rawKey := base64.RawURLEncoding.EncodeToString(rawBytes) - hash := hashSHA256Hex(rawKey) + hash := sectypes.HashKey(rawKey) id := atomic.AddInt64(&s.next, 1) - key := UserKey{ + key := sectypes.UserKey{ ID: id, UserID: req.UserID, KeyType: req.KeyType, @@ -70,17 +71,17 @@ func (s *ConfigKeyStore) CreateKey(_ context.Context, req CreateKeyRequest) (*Cr s.keys = append(s.keys, key) s.mu.Unlock() - return &CreateKeyResponse{Key: key, RawKey: rawKey}, nil + return §ypes.CreateKeyResponse{Key: key, RawKey: rawKey}, nil } // GetUserKeys returns all active, non-expired keys for the given user. // Pass an empty KeyType to return all types. -func (s *ConfigKeyStore) GetUserKeys(_ context.Context, userID int, keyType KeyType) ([]UserKey, error) { +func (s *ConfigKeyStore) GetUserKeys(_ context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) { now := time.Now() s.mu.RLock() defer s.mu.RUnlock() - var result []UserKey + var result []sectypes.UserKey for i := range s.keys { k := &s.keys[i] if k.UserID != userID || !k.IsActive { @@ -117,8 +118,8 @@ func (s *ConfigKeyStore) DeleteKey(_ context.Context, userID int, keyID int64) e // ValidateKey hashes the raw key and finds a matching, active, non-expired entry. // Uses constant-time comparison to prevent timing side-channels. // Pass an empty KeyType to accept any type. -func (s *ConfigKeyStore) ValidateKey(_ context.Context, rawKey string, keyType KeyType) (*UserKey, error) { - hash := hashSHA256Hex(rawKey) +func (s *ConfigKeyStore) ValidateKey(_ context.Context, rawKey string, keyType sectypes.KeyType) (*sectypes.UserKey, error) { + hash := sectypes.HashKey(rawKey) hashBytes, _ := hex.DecodeString(hash) now := time.Now() diff --git a/pkg/security/providers/policy_config.go b/pkg/security/providers/policy_config.go new file mode 100644 index 0000000..c1e9ab8 --- /dev/null +++ b/pkg/security/providers/policy_config.go @@ -0,0 +1,61 @@ +package providers + +import ( + "context" + "fmt" + + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// ConfigColumnSecurityProvider provides static column security configuration +type ConfigColumnSecurityProvider struct { + rules map[string][]sectypes.ColumnSecurity +} + +func NewConfigColumnSecurityProvider(rules map[string][]sectypes.ColumnSecurity) *ConfigColumnSecurityProvider { + return &ConfigColumnSecurityProvider{rules: rules} +} + +func (p *ConfigColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) { + key := fmt.Sprintf("%s.%s", schema, table) + rules, ok := p.rules[key] + if !ok { + return []sectypes.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) (sectypes.RowSecurity, error) { + key := fmt.Sprintf("%s.%s", schema, table) + + if p.blocked[key] { + return sectypes.RowSecurity{ + Schema: schema, + Tablename: table, + UserID: userRef, + HasBlock: true, + }, nil + } + + template := p.templates[key] + return sectypes.RowSecurity{ + Schema: schema, + Tablename: table, + UserID: userRef, + Template: template, + HasBlock: false, + }, nil +} diff --git a/pkg/security/providers/providers_test.go b/pkg/security/providers/providers_test.go new file mode 100644 index 0000000..bb8a15f --- /dev/null +++ b/pkg/security/providers/providers_test.go @@ -0,0 +1,180 @@ +package providers + +import ( + "context" + "net/http/httptest" + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Test HeaderAuthenticator +func TestHeaderAuthenticator(t *testing.T) { + auth := NewHeaderAuthenticator() + + t.Run("successful authentication", func(t *testing.T) { + req := httptest.NewRequest("GET", "/test", nil) + req.Header.Set("X-User-ID", "123") + req.Header.Set("X-User-Name", "testuser") + req.Header.Set("X-User-Level", "5") + req.Header.Set("X-Session-ID", "session123") + req.Header.Set("X-Remote-ID", "remote456") + req.Header.Set("X-User-Email", "test@example.com") + req.Header.Set("X-User-Roles", "admin,user") + + userCtx, err := auth.Authenticate(req) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if userCtx.UserID != 123 { + t.Errorf("expected UserID 123, got %d", userCtx.UserID) + } + if userCtx.UserName != "testuser" { + t.Errorf("expected UserName testuser, got %s", userCtx.UserName) + } + if userCtx.UserLevel != 5 { + t.Errorf("expected UserLevel 5, got %d", userCtx.UserLevel) + } + if userCtx.SessionID != "session123" { + t.Errorf("expected SessionID session123, got %s", userCtx.SessionID) + } + if userCtx.Email != "test@example.com" { + t.Errorf("expected Email test@example.com, got %s", userCtx.Email) + } + if len(userCtx.Roles) != 2 { + t.Errorf("expected 2 roles, got %d", len(userCtx.Roles)) + } + }) + + t.Run("missing user ID header", func(t *testing.T) { + req := httptest.NewRequest("GET", "/test", nil) + req.Header.Set("X-User-Name", "testuser") + + _, err := auth.Authenticate(req) + if err == nil { + t.Fatal("expected error when X-User-ID is missing") + } + }) + + t.Run("invalid user ID", func(t *testing.T) { + req := httptest.NewRequest("GET", "/test", nil) + req.Header.Set("X-User-ID", "invalid") + + _, err := auth.Authenticate(req) + if err == nil { + t.Fatal("expected error with invalid user ID") + } + }) + + t.Run("login not supported", func(t *testing.T) { + ctx := context.Background() + req := sectypes.LoginRequest{Username: "test", Password: "pass"} + + _, err := auth.Login(ctx, req) + if err == nil { + t.Fatal("expected error for unsupported login") + } + }) + + t.Run("logout always succeeds", func(t *testing.T) { + ctx := context.Background() + req := sectypes.LogoutRequest{Token: "token", UserID: 1} + + err := auth.Logout(ctx, req) + if err != nil { + t.Errorf("expected no error, got %v", err) + } + }) +} + +// Test ConfigColumnSecurityProvider +func TestConfigColumnSecurityProvider(t *testing.T) { + rules := map[string][]sectypes.ColumnSecurity{ + "public.users": { + { + Schema: "public", + Tablename: "users", + Path: []string{"email"}, + Accesstype: "mask", + }, + }, + } + + provider := NewConfigColumnSecurityProvider(rules) + ctx := context.Background() + + t.Run("get existing rules", func(t *testing.T) { + result, err := provider.GetColumnSecurity(ctx, 1, "public", "users") + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if len(result) != 1 { + t.Errorf("expected 1 rule, got %d", len(result)) + } + }) + + t.Run("get non-existent rules returns empty", func(t *testing.T) { + result, err := provider.GetColumnSecurity(ctx, 1, "public", "orders") + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if len(result) != 0 { + t.Errorf("expected 0 rules, got %d", len(result)) + } + }) +} + +// Test ConfigRowSecurityProvider +func TestConfigRowSecurityProvider(t *testing.T) { + templates := map[string]string{ + "public.orders": "user_id = {UserID}", + } + blocked := map[string]bool{ + "public.secrets": true, + } + + provider := NewConfigRowSecurityProvider(templates, blocked) + ctx := context.Background() + + t.Run("get template for allowed table", func(t *testing.T) { + result, err := provider.GetRowSecurity(ctx, 1, "public", "orders") + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if result.Template != "user_id = {UserID}" { + t.Errorf("expected template 'user_id = {UserID}', got %s", result.Template) + } + if result.HasBlock { + t.Error("expected HasBlock to be false") + } + }) + + t.Run("get blocked table", func(t *testing.T) { + result, err := provider.GetRowSecurity(ctx, 1, "public", "secrets") + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if !result.HasBlock { + t.Error("expected HasBlock to be true") + } + }) + + t.Run("get non-existent table returns empty template", func(t *testing.T) { + result, err := provider.GetRowSecurity(ctx, 1, "public", "unknown") + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + + if result.Template != "" { + t.Errorf("expected empty template, got %s", result.Template) + } + if result.HasBlock { + t.Error("expected HasBlock to be false") + } + }) +} diff --git a/pkg/security/providers_direct.go b/pkg/security/providers_direct.go deleted file mode 100644 index d86d5af..0000000 --- a/pkg/security/providers_direct.go +++ /dev/null @@ -1,535 +0,0 @@ -package security - -import ( - "context" - "crypto/rand" - "crypto/sha256" - "database/sql" - "encoding/hex" - "errors" - "fmt" - "strings" - "time" - - "github.com/bitechdev/ResolveSpec/pkg/logger" -) - -// Direct-mode implementations for DatabaseAuthenticator and JWTAuthenticator. -// These mirror the plpgsql bodies in database_schema.sql using plain -// parameterized SQL against the configured TableNames, so they work on -// SQLite, MySQL, or Postgres without the resolvespec_* functions installed. -// -// Passwords are verified with bcrypt (see password.go). Legacy cleartext rows -// are still accepted at login; they are only rewritten as bcrypt when the -// upgrade is explicitly enabled (UpgradePasswordHash). Registration -// never honours client-supplied user_level/roles. - -var ( - errUsernameExists = errors.New("username already exists") - errEmailExists = errors.New("email already exists") -) - -func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) { - var userID int - var email, roles, programUserTable, storedPassword sql.NullString - var userLevel, programUserID sql.NullInt64 - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT id, email, user_level, roles, program_user_id, program_user_table, password FROM %s WHERE username = ? AND is_active = ?`, - a.tableNames.Users)) - return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - burnPasswordCheck(req.Password) - return nil, fmt.Errorf("invalid credentials") - } - return nil, fmt.Errorf("login query failed: %w", err) - } - - ok, needsRehash := verifyPassword(storedPassword.String, req.Password) - if !ok { - if storedPassword.String == "" { - burnPasswordCheck(req.Password) - } - return nil, fmt.Errorf("invalid credentials") - } - if needsRehash && a.upgradePasswordHash { - a.upgradePasswordHashFor(ctx, userID, req.Password) - } - - sessionToken, err := generateSessionToken() - if err != nil { - return nil, fmt.Errorf("failed to generate session token: %w", err) - } - now := time.Now() - expiresAt := now.Add(24 * time.Hour) - ipAddress, userAgent := claimStrings(req.Claims) - - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - insertQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.UserSessions)) - if _, err := db.ExecContext(ctx, insertQuery, sessionToken, userID, expiresAt, ipAddress, userAgent, now, now); err != nil { - return err - } - updateQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_login_at = ? WHERE id = ?`, a.tableNames.Users)) - _, err := db.ExecContext(ctx, updateQuery, now, userID) - return err - }) - if err != nil { - return nil, fmt.Errorf("login query failed: %w", err) - } - - userCtx := &UserContext{ - UserID: userID, - UserName: req.Username, - Email: email.String, - UserLevel: int(userLevel.Int64), - Roles: parseRoles(roles.String), - SessionID: sessionToken, - ProgramUserID: int(programUserID.Int64), - ProgramUserTable: programUserTable.String, - } - - return &LoginResponse{ - Token: sessionToken, - User: userCtx, - ExpiresIn: 86400, - }, nil -} - -// upgradePasswordHashFor replaces a legacy cleartext password with a bcrypt hash. -// Only called when the upgrade has been explicitly enabled. Failure is logged -// and ignored: the login itself already succeeded. -func (a *DatabaseAuthenticator) upgradePasswordHashFor(ctx context.Context, userID int, password string) { - h, err := hashPassword(password) - if err != nil { - return - } - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - q := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users)) - _, err := db.ExecContext(ctx, q, h, time.Now(), userID) - return err - }) - if err != nil { - logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err) - } -} - -func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*LoginResponse, error) { - if req.Username == "" { - return nil, fmt.Errorf("username is required") - } - if req.Email == "" { - return nil, fmt.Errorf("email is required") - } - if req.Password == "" { - return nil, fmt.Errorf("password is required") - } - - passwordHash, err := hashPassword(req.Password) - if err != nil { - return nil, err - } - - // Privileges are never taken from the request: self-registration always - // creates an unprivileged user. - const userLevel = 0 - const rolesStr = "" - now := time.Now() - ipAddress, userAgent := claimStrings(req.Claims) - - var userID int64 - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - var count int - checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE username = ?`, a.tableNames.Users)) - if err := db.QueryRowContext(ctx, checkQuery, req.Username).Scan(&count); err != nil { - return err - } - if count > 0 { - return errUsernameExists - } - checkQuery2 := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE email = ?`, a.tableNames.Users)) - if err := db.QueryRowContext(ctx, checkQuery2, req.Email).Scan(&count); err != nil { - return err - } - if count > 0 { - return errEmailExists - } - - insertQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.Users)) - res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, passwordHash, userLevel, rolesStr, true, now, now, 0, "") - if err != nil { - return err - } - userID, err = res.LastInsertId() - return err - }) - if err != nil { - if errors.Is(err, errUsernameExists) { - return nil, errUsernameExists - } - if errors.Is(err, errEmailExists) { - return nil, errEmailExists - } - return nil, fmt.Errorf("register query failed: %w", err) - } - - sessionToken, err := generateSessionToken() - if err != nil { - return nil, fmt.Errorf("failed to generate session token: %w", err) - } - expiresAt := now.Add(24 * time.Hour) - - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - insertSession := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.UserSessions)) - if _, err := db.ExecContext(ctx, insertSession, sessionToken, userID, expiresAt, ipAddress, userAgent, now, now); err != nil { - return err - } - updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_login_at = ? WHERE id = ?`, a.tableNames.Users)) - _, err := db.ExecContext(ctx, updUser, now, userID) - return err - }) - if err != nil { - return nil, fmt.Errorf("register query failed: %w", err) - } - - userCtx := &UserContext{ - UserID: int(userID), - UserName: req.Username, - Email: req.Email, - UserLevel: userLevel, - Roles: parseRoles(rolesStr), - SessionID: sessionToken, - ProgramUserID: 0, - ProgramUserTable: "", - } - - return &LoginResponse{ - Token: sessionToken, - User: userCtx, - ExpiresIn: 86400, - }, nil -} - -func (a *DatabaseAuthenticator) logoutDirect(ctx context.Context, req LogoutRequest) error { - token := req.Token - token = strings.TrimPrefix(token, "Bearer ") - token = strings.TrimPrefix(token, "bearer ") - - var rows int64 - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ? AND user_id = ?`, a.tableNames.UserSessions)) - res, err := db.ExecContext(ctx, query, token, req.UserID) - if err != nil { - return err - } - rows, err = res.RowsAffected() - return err - }) - if err != nil { - return fmt.Errorf("logout query failed: %w", err) - } - if rows == 0 { - return fmt.Errorf("session not found") - } - - if req.Token != "" { - cacheKey := fmt.Sprintf("auth:session:%s", req.Token) - _ = a.cache.Delete(ctx, cacheKey) - } - return nil -} - -func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string) (*UserContext, error) { - var userID int - var username, email, roles, programUserTable sql.NullString - var userLevel, programUserID sql.NullInt64 - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT s.user_id, u.username, u.email, u.user_level, u.roles, u.program_user_id, u.program_user_table - FROM %s s JOIN %s u ON s.user_id = u.id - WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`, - a.tableNames.UserSessions, a.tableNames.Users)) - return db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &programUserID, &programUserTable) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("invalid or expired session") - } - return nil, fmt.Errorf("session query failed: %w", err) - } - - return &UserContext{ - UserID: userID, - UserName: username.String, - Email: email.String, - UserLevel: int(userLevel.Int64), - SessionID: token, - Roles: parseRoles(roles.String), - ProgramUserID: int(programUserID.Int64), - ProgramUserTable: programUserTable.String, - }, nil -} - -func (a *DatabaseAuthenticator) updateSessionActivityDirect(ctx context.Context, sessionToken string) error { - return a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_activity_at = ? WHERE session_token = ? AND expires_at > ?`, a.tableNames.UserSessions)) - _, err := db.ExecContext(ctx, query, time.Now(), sessionToken, time.Now()) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters - return err - }) -} - -func (a *DatabaseAuthenticator) refreshTokenDirect(ctx context.Context, oldToken string) (*LoginResponse, error) { - var userID int - var username, email, roles, ipAddress, userAgent, programUserTable sql.NullString - var userLevel, programUserID sql.NullInt64 - - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf( - `SELECT s.user_id, u.username, u.email, u.user_level, u.roles, s.ip_address, s.user_agent, u.program_user_id, u.program_user_table - FROM %s s JOIN %s u ON s.user_id = u.id - WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`, - a.tableNames.UserSessions, a.tableNames.Users)) - return db.QueryRowContext(ctx, query, oldToken, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &ipAddress, &userAgent, &programUserID, &programUserTable) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, fmt.Errorf("invalid or expired refresh token") - } - return nil, fmt.Errorf("refresh token query failed: %w", err) - } - - newToken, err := generateSessionToken() - if err != nil { - return nil, fmt.Errorf("failed to generate session token: %w", err) - } - now := time.Now() - expiresAt := now.Add(24 * time.Hour) - - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - insertQuery := rewritePlaceholders(db, fmt.Sprintf( - `INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, - a.tableNames.UserSessions)) - if _, err := db.ExecContext(ctx, insertQuery, newToken, userID, expiresAt, ipAddress.String, userAgent.String, now, now); err != nil { - return err - } - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ?`, a.tableNames.UserSessions)) - _, err := db.ExecContext(ctx, delQuery, oldToken) - return err - }) - if err != nil { - return nil, fmt.Errorf("refresh token generation failed: %w", err) - } - - userCtx := &UserContext{ - UserID: userID, - UserName: username.String, - Email: email.String, - UserLevel: int(userLevel.Int64), - SessionID: newToken, - Roles: parseRoles(roles.String), - ProgramUserID: int(programUserID.Int64), - ProgramUserTable: programUserTable.String, - } - - return &LoginResponse{ - Token: newToken, - User: userCtx, - ExpiresIn: int64(24 * time.Hour.Seconds()), - }, nil -} - -func (a *DatabaseAuthenticator) requestPasswordResetDirect(ctx context.Context, req PasswordResetRequest) (*PasswordResetResponse, error) { - if req.Email == "" && req.Username == "" { - return nil, fmt.Errorf("email or username is required") - } - - var userID int - err := a.runDBOpWithReconnect(func(db *sql.DB) error { - var query string - var arg string - if req.Email != "" { - query = rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE email = ? AND is_active = ?`, a.tableNames.Users)) - arg = req.Email - } else { - query = rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users)) - arg = req.Username - } - return db.QueryRowContext(ctx, query, arg, true).Scan(&userID) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - // Return generic success even when user not found to avoid user enumeration. - return &PasswordResetResponse{Token: "", ExpiresIn: 0}, nil - } - return nil, fmt.Errorf("password reset request query failed: %w", err) - } - - rawBytes := make([]byte, 32) - if _, err := rand.Read(rawBytes); err != nil { - return nil, fmt.Errorf("failed to generate reset token: %w", err) - } - rawToken := hex.EncodeToString(rawBytes) - hash := sha256.Sum256([]byte(rawToken)) - tokenHash := hex.EncodeToString(hash[:]) - now := time.Now() - expiresAt := now.Add(1 * time.Hour) - - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ? AND used = ?`, a.tableNames.UserPasswordResets)) - if _, err := db.ExecContext(ctx, delQuery, userID, false); err != nil { - return err - } - insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, token_hash, expires_at, created_at, used) VALUES (?, ?, ?, ?, ?)`, a.tableNames.UserPasswordResets)) - _, err := db.ExecContext(ctx, insQuery, userID, tokenHash, expiresAt, now, false) - return err - }) - if err != nil { - return nil, fmt.Errorf("password reset request query failed: %w", err) - } - - return &PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil -} - -func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context, req PasswordResetCompleteRequest) error { - if req.Token == "" { - return fmt.Errorf("token is required") - } - if req.NewPassword == "" { - return fmt.Errorf("new_password is required") - } - - newHash, err := hashPassword(req.NewPassword) - if err != nil { - return err - } - - hash := sha256.Sum256([]byte(req.Token)) - tokenHash := hex.EncodeToString(hash[:]) - - var resetID, userID int - var expiresAt time.Time - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, user_id, expires_at FROM %s WHERE token_hash = ? AND used = ?`, a.tableNames.UserPasswordResets)) - return db.QueryRowContext(ctx, query, tokenHash, false).Scan(&resetID, &userID, &expiresAt) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return fmt.Errorf("invalid or expired token") - } - return fmt.Errorf("password reset complete query failed: %w", err) - } - if !expiresAt.After(time.Now()) { - return fmt.Errorf("invalid or expired token") - } - - now := time.Now() - err = a.runDBOpWithReconnect(func(db *sql.DB) error { - updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users)) - if _, err := db.ExecContext(ctx, updUser, newHash, now, userID); err != nil { - return err - } - delSessions := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, a.tableNames.UserSessions)) - if _, err := db.ExecContext(ctx, delSessions, userID); err != nil { - return err - } - updReset := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET used = ?, used_at = ? WHERE id = ?`, a.tableNames.UserPasswordResets)) - _, err := db.ExecContext(ctx, updReset, true, now, resetID) - return err - }) - if err != nil { - return fmt.Errorf("password reset complete query failed: %w", err) - } - return nil -} - -// claimStrings extracts ip_address/user_agent from a request's Claims map, mirroring -// p_request->'claims'->>'ip_address' / 'user_agent' in the plpgsql procedures. -func claimStrings(claims map[string]any) (ipAddress, userAgent string) { - if claims == nil { - return "", "" - } - if v, ok := claims["ip_address"].(string); ok { - ipAddress = v - } - if v, ok := claims["user_agent"].(string); ok { - userAgent = v - } - return ipAddress, userAgent -} - -// jwtLoginDirect mirrors resolvespec_jwt_login. -func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) { - var userID int - var email, roles, storedPassword sql.NullString - var userLevel sql.NullInt64 - - runQuery := func() error { - query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles, password FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users)) - return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &storedPassword) - } - err := runQuery() - if isDBClosed(err) { - if reconnErr := a.reconnectDB(); reconnErr == nil { - err = runQuery() - } - } - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - burnPasswordCheck(req.Password) - return nil, fmt.Errorf("invalid credentials") - } - return nil, fmt.Errorf("login query failed: %w", err) - } - - ok, needsRehash := verifyPassword(storedPassword.String, req.Password) - if !ok { - if storedPassword.String == "" { - burnPasswordCheck(req.Password) - } - return nil, fmt.Errorf("invalid credentials") - } - if needsRehash && a.upgradePasswordHash { - if h, herr := hashPassword(req.Password); herr == nil { - q := rewritePlaceholders(a.getDB(), fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users)) - if _, uerr := a.getDB().ExecContext(ctx, q, h, time.Now(), userID); uerr != nil { - logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, uerr) - } - } - } - - expiresAt := time.Now().Add(24 * time.Hour) - tokenString := fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()) - - return &LoginResponse{ - Token: tokenString, - User: &UserContext{ - UserID: userID, - UserName: req.Username, - Email: email.String, - UserLevel: int(userLevel.Int64), - Roles: parseRoles(roles.String), - }, - ExpiresIn: int64(24 * time.Hour.Seconds()), - }, nil -} - -// jwtLogoutDirect mirrors resolvespec_jwt_logout (adds token to the blacklist table). -func (a *JWTAuthenticator) jwtLogoutDirect(ctx context.Context, req LogoutRequest) error { - db := a.getDB() - expiresAt := time.Now().Add(24 * time.Hour) - query := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (token, user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, a.tableNames.TokenBlacklist)) - _, err := db.ExecContext(ctx, query, req.Token, req.UserID, expiresAt, time.Now()) - if err != nil { - return fmt.Errorf("logout query failed: %w", err) - } - return nil -} diff --git a/pkg/security/providers_test.go b/pkg/security/providers_test.go index a5af9bb..f4480a6 100644 --- a/pkg/security/providers_test.go +++ b/pkg/security/providers_test.go @@ -13,86 +13,6 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/cache" ) -// Test HeaderAuthenticator -func TestHeaderAuthenticator(t *testing.T) { - auth := NewHeaderAuthenticator() - - t.Run("successful authentication", func(t *testing.T) { - req := httptest.NewRequest("GET", "/test", nil) - req.Header.Set("X-User-ID", "123") - req.Header.Set("X-User-Name", "testuser") - req.Header.Set("X-User-Level", "5") - req.Header.Set("X-Session-ID", "session123") - req.Header.Set("X-Remote-ID", "remote456") - req.Header.Set("X-User-Email", "test@example.com") - req.Header.Set("X-User-Roles", "admin,user") - - userCtx, err := auth.Authenticate(req) - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if userCtx.UserID != 123 { - t.Errorf("expected UserID 123, got %d", userCtx.UserID) - } - if userCtx.UserName != "testuser" { - t.Errorf("expected UserName testuser, got %s", userCtx.UserName) - } - if userCtx.UserLevel != 5 { - t.Errorf("expected UserLevel 5, got %d", userCtx.UserLevel) - } - if userCtx.SessionID != "session123" { - t.Errorf("expected SessionID session123, got %s", userCtx.SessionID) - } - if userCtx.Email != "test@example.com" { - t.Errorf("expected Email test@example.com, got %s", userCtx.Email) - } - if len(userCtx.Roles) != 2 { - t.Errorf("expected 2 roles, got %d", len(userCtx.Roles)) - } - }) - - t.Run("missing user ID header", func(t *testing.T) { - req := httptest.NewRequest("GET", "/test", nil) - req.Header.Set("X-User-Name", "testuser") - - _, err := auth.Authenticate(req) - if err == nil { - t.Fatal("expected error when X-User-ID is missing") - } - }) - - t.Run("invalid user ID", func(t *testing.T) { - req := httptest.NewRequest("GET", "/test", nil) - req.Header.Set("X-User-ID", "invalid") - - _, err := auth.Authenticate(req) - if err == nil { - t.Fatal("expected error with invalid user ID") - } - }) - - t.Run("login not supported", func(t *testing.T) { - ctx := context.Background() - req := LoginRequest{Username: "test", Password: "pass"} - - _, err := auth.Login(ctx, req) - if err == nil { - t.Fatal("expected error for unsupported login") - } - }) - - t.Run("logout always succeeds", func(t *testing.T) { - ctx := context.Background() - req := LogoutRequest{Token: "token", UserID: 1} - - err := auth.Logout(ctx, req) - if err != nil { - t.Errorf("expected no error, got %v", err) - } - }) -} - // Test parseRoles helper func TestParseRoles(t *testing.T) { tests := []struct { @@ -1238,97 +1158,6 @@ func TestDatabaseRowSecurityProvider(t *testing.T) { }) } -// Test ConfigColumnSecurityProvider -func TestConfigColumnSecurityProvider(t *testing.T) { - rules := map[string][]ColumnSecurity{ - "public.users": { - { - Schema: "public", - Tablename: "users", - Path: []string{"email"}, - Accesstype: "mask", - }, - }, - } - - provider := NewConfigColumnSecurityProvider(rules) - ctx := context.Background() - - t.Run("get existing rules", func(t *testing.T) { - result, err := provider.GetColumnSecurity(ctx, 1, "public", "users") - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if len(result) != 1 { - t.Errorf("expected 1 rule, got %d", len(result)) - } - }) - - t.Run("get non-existent rules returns empty", func(t *testing.T) { - result, err := provider.GetColumnSecurity(ctx, 1, "public", "orders") - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if len(result) != 0 { - t.Errorf("expected 0 rules, got %d", len(result)) - } - }) -} - -// Test ConfigRowSecurityProvider -func TestConfigRowSecurityProvider(t *testing.T) { - templates := map[string]string{ - "public.orders": "user_id = {UserID}", - } - blocked := map[string]bool{ - "public.secrets": true, - } - - provider := NewConfigRowSecurityProvider(templates, blocked) - ctx := context.Background() - - t.Run("get template for allowed table", func(t *testing.T) { - result, err := provider.GetRowSecurity(ctx, 1, "public", "orders") - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if result.Template != "user_id = {UserID}" { - t.Errorf("expected template 'user_id = {UserID}', got %s", result.Template) - } - if result.HasBlock { - t.Error("expected HasBlock to be false") - } - }) - - t.Run("get blocked table", func(t *testing.T) { - result, err := provider.GetRowSecurity(ctx, 1, "public", "secrets") - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if !result.HasBlock { - t.Error("expected HasBlock to be true") - } - }) - - t.Run("get non-existent table returns empty template", func(t *testing.T) { - result, err := provider.GetRowSecurity(ctx, 1, "public", "unknown") - if err != nil { - t.Fatalf("expected no error, got %v", err) - } - - if result.Template != "" { - t.Errorf("expected empty template, got %s", result.Template) - } - if result.HasBlock { - t.Error("expected HasBlock to be false") - } - }) -} - // authenticateSync authenticates and waits for the asynchronous session // activity update so sqlmock expectations are never touched concurrently. func authenticateSync(auth *DatabaseAuthenticator, req *http.Request) (*UserContext, error) { diff --git a/pkg/security/query_mode.go b/pkg/security/query_mode.go deleted file mode 100644 index 0a32bb6..0000000 --- a/pkg/security/query_mode.go +++ /dev/null @@ -1,172 +0,0 @@ -package security - -import ( - "context" - "crypto/rand" - "database/sql" - "encoding/hex" - "fmt" - "strconv" - "strings" - "sync" - "time" - - "github.com/bitechdev/ResolveSpec/pkg/dbtrace" -) - -// QueryMode selects how a provider talks to the database: via the configured -// resolvespec_* stored procedure (ModeProcedure), via portable Go/SQL logic -// (ModeDirect), or auto-detected per connection (ModeAuto, the default). -type QueryMode int - -const ( - // ModeAuto probes whether the configured stored procedure exists on a - // Postgres connection and uses it if so, otherwise falls back to Direct mode. - ModeAuto QueryMode = iota - // ModeProcedure always calls the configured resolvespec_* stored procedure. - ModeProcedure - // ModeDirect always uses the portable Go/SQL implementation, never the - // stored procedure. - ModeDirect -) - -// dbCapability probes and caches whether a given *sql.DB is Postgres and -// whether specific stored procedures exist on it. One instance is shared by -// a provider (DatabaseAuthenticator, DatabaseTwoFactorProvider, etc.) across -// all of its operations. -type dbCapability struct { - funcExists sync.Map // procName (string) -> exists (bool) -} - -// newDBCapability creates a new, empty capability cache. -func newDBCapability() *dbCapability { - return &dbCapability{} -} - -// reset clears all cached probe results. Call after reconnecting to a -// (possibly different) database. -func (c *dbCapability) reset() { - c.funcExists.Range(func(key, _ any) bool { - c.funcExists.Delete(key) - return true - }) -} - -// probeFunctionExists checks, via a Postgres-specific system catalog query, -// whether a function named procName exists. Any error (wrong dialect, -// placeholder syntax rejected, relation missing, etc.) is treated as "does -// not exist" rather than propagated - the probe must never be able to panic -// or block resolution of the query mode. -func probeFunctionExists(ctx context.Context, db *sql.DB, procName string) bool { - if db == nil { - return false - } - dbtrace.Raw(ctx, "probe.pg_proc") - var exists bool - defer func() { - // Guard against any unexpected panic from a misbehaving driver. - _ = recover() - }() - row := db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_proc WHERE proname = $1 LIMIT 1)`, procName) - if err := row.Scan(&exists); err != nil { - return false - } - return exists -} - -// ShouldUseProcedure decides whether the stored-procedure code path should be -// used for procName given mode. ModeProcedure/ModeDirect are unconditional. -// -// ModeAuto resolves as follows: -// - A recognized portable-only driver (SQLite, MySQL) always uses Direct -// mode; no query is issued. -// - A recognized Postgres driver (lib/pq, pgx) probes pg_proc for procName -// and uses the stored procedure only if it actually exists there. -// - Any other/unrecognized driver (including test doubles such as -// sqlmock) cannot be safely dialect-probed without risking an -// unexpected query against a strictly-ordered mock, so it defaults to -// the stored-procedure path, preserving pre-existing behavior for -// callers that configure Postgres-flavored access without an -// identifiable driver type. -func (c *dbCapability) ShouldUseProcedure(ctx context.Context, mode QueryMode, db *sql.DB, procName string) bool { - switch mode { - case ModeProcedure: - return true - case ModeDirect: - return false - default: - if cached, ok := c.funcExists.Load(procName); ok { - return cached.(bool) - } - - var exists bool - switch { - case driverIsPortableOnly(db): - exists = false - case driverIsPostgres(db): - exists = probeFunctionExists(ctx, db, procName) - default: - exists = true - } - - c.funcExists.Store(procName, exists) - return exists - } -} - -// driverIsPostgres reports whether db's underlying driver looks like a -// Postgres driver (lib/pq or pgx), based on the driver's Go type name. -func driverIsPostgres(db *sql.DB) bool { - if db == nil { - return false - } - t := strings.ToLower(fmt.Sprintf("%T", db.Driver())) - return strings.Contains(t, "pq.") || strings.Contains(t, "pgx") || strings.Contains(t, "postgres") -} - -// driverIsPortableOnly reports whether db's underlying driver is a dialect -// that never has the resolvespec_* Postgres functions available (SQLite, -// MySQL), so ModeAuto can skip probing entirely and go straight to Direct. -func driverIsPortableOnly(db *sql.DB) bool { - if db == nil { - return false - } - t := strings.ToLower(fmt.Sprintf("%T", db.Driver())) - return strings.Contains(t, "sqlite") || strings.Contains(t, "mysql") -} - -// rewritePlaceholders converts a query written with "?" placeholders -// (SQLite/MySQL style) to Postgres "$1", "$2", ... style when db's driver is -// Postgres. All Direct-mode SQL in this package is written with "?" and -// passed through this helper before execution so the same query source works -// against SQLite, MySQL, and (in the rare fallback case) Postgres without the -// resolvespec_* functions installed. -func rewritePlaceholders(db *sql.DB, query string) string { - if !driverIsPostgres(db) { - return query - } - var b strings.Builder - n := 0 - for _, r := range query { - if r == '?' { - n++ - b.WriteString("$") - b.WriteString(strconv.Itoa(n)) - } else { - b.WriteRune(r) - } - } - return b.String() -} - -// generateSessionToken produces a session token in the same shape the -// plpgsql stored procedures generate: "sess_" + hex(32 random bytes) + "_" + -// unix timestamp, so downstream code that parses/displays tokens is -// unaffected by which mode created them. -func generateSessionToken() (string, error) { - buf := make([]byte, 32) - if _, err := rand.Read(buf); err != nil { - return "", err - } - return fmt.Sprintf("sess_%s_%d", hex.EncodeToString(buf), time.Now().Unix()), nil -} diff --git a/pkg/security/query_mode_test.go b/pkg/security/query_mode_test.go deleted file mode 100644 index fdfe649..0000000 --- a/pkg/security/query_mode_test.go +++ /dev/null @@ -1,165 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "testing" - - "github.com/DATA-DOG/go-sqlmock" - _ "github.com/mattn/go-sqlite3" -) - -func openTestSQLite(t *testing.T) *sql.DB { - t.Helper() - db, err := sql.Open("sqlite3", ":memory:") - if err != nil { - t.Fatalf("failed to open sqlite db: %v", err) - } - t.Cleanup(func() { _ = db.Close() }) - return db -} - -func TestShouldUseProcedure_ModeProcedure_AlwaysTrue(t *testing.T) { - db := openTestSQLite(t) - c := newDBCapability() - if !c.ShouldUseProcedure(context.Background(), ModeProcedure, db, "resolvespec_login") { - t.Error("ModeProcedure should always resolve to true") - } -} - -func TestShouldUseProcedure_ModeDirect_AlwaysFalse(t *testing.T) { - db := openTestSQLite(t) - c := newDBCapability() - if c.ShouldUseProcedure(context.Background(), ModeDirect, db, "resolvespec_login") { - t.Error("ModeDirect should always resolve to false") - } -} - -func TestShouldUseProcedure_Auto_SQLite_ResolvesDirect(t *testing.T) { - db := openTestSQLite(t) - c := newDBCapability() - if c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") { - t.Error("ModeAuto against a SQLite connection should resolve to Direct (false)") - } -} - -func TestShouldUseProcedure_Auto_UnknownDriver_DefaultsToProcedure(t *testing.T) { - db, mock, err := sqlmock.New() - if err != nil { - t.Fatalf("failed to create mock db: %v", err) - } - defer db.Close() - - c := newDBCapability() - if !c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") { - t.Error("ModeAuto against an unrecognized driver (e.g. a test double) should default to Procedure (true) to preserve existing behavior") - } - if err := mock.ExpectationsWereMet(); err != nil { - t.Errorf("unfulfilled expectations: %v", err) - } -} - -func TestShouldUseProcedure_CachesResult(t *testing.T) { - db := openTestSQLite(t) - c := newDBCapability() - first := c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") - if _, ok := c.funcExists.Load("resolvespec_login"); !ok { - t.Error("expected result to be cached") - } - second := c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") - if first != second { - t.Errorf("cached result changed: first=%v second=%v", first, second) - } -} - -func TestDBCapability_Reset_ClearsCache(t *testing.T) { - c := newDBCapability() - c.funcExists.Store("resolvespec_login", true) - c.reset() - if _, ok := c.funcExists.Load("resolvespec_login"); ok { - t.Error("reset should clear all cached entries") - } -} - -func TestProbeFunctionExists_Postgres_FunctionExists(t *testing.T) { - db, mock, err := sqlmock.New() - if err != nil { - t.Fatalf("failed to create mock db: %v", err) - } - defer db.Close() - - rows := sqlmock.NewRows([]string{"exists"}).AddRow(true) - mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`). - WithArgs("resolvespec_login"). - WillReturnRows(rows) - - if !probeFunctionExists(context.Background(), db, "resolvespec_login") { - t.Error("expected probeFunctionExists to return true") - } - if err := mock.ExpectationsWereMet(); err != nil { - t.Errorf("unfulfilled expectations: %v", err) - } -} - -func TestProbeFunctionExists_Postgres_FunctionMissing(t *testing.T) { - db, mock, err := sqlmock.New() - if err != nil { - t.Fatalf("failed to create mock db: %v", err) - } - defer db.Close() - - rows := sqlmock.NewRows([]string{"exists"}).AddRow(false) - mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`). - WithArgs("resolvespec_login"). - WillReturnRows(rows) - - if probeFunctionExists(context.Background(), db, "resolvespec_login") { - t.Error("expected probeFunctionExists to return false") - } - if err := mock.ExpectationsWereMet(); err != nil { - t.Errorf("unfulfilled expectations: %v", err) - } -} - -func TestProbeFunctionExists_QueryError_ReturnsFalse(t *testing.T) { - db, mock, err := sqlmock.New() - if err != nil { - t.Fatalf("failed to create mock db: %v", err) - } - defer db.Close() - - mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`). - WithArgs("resolvespec_login"). - WillReturnError(sql.ErrConnDone) - - if probeFunctionExists(context.Background(), db, "resolvespec_login") { - t.Error("expected probeFunctionExists to return false on query error") - } -} - -func TestProbeFunctionExists_NilDB(t *testing.T) { - if probeFunctionExists(context.Background(), nil, "resolvespec_login") { - t.Error("expected probeFunctionExists(nil) to return false") - } -} - -func TestRewritePlaceholders_NonPostgres_NoOp(t *testing.T) { - db := openTestSQLite(t) - q := "SELECT * FROM users WHERE id = ? AND name = ?" - if got := rewritePlaceholders(db, q); got != q { - t.Errorf("rewritePlaceholders on non-Postgres = %q, want unchanged %q", got, q) - } -} - -func TestGenerateSessionToken_Format(t *testing.T) { - token, err := generateSessionToken() - if err != nil { - t.Fatalf("generateSessionToken() error = %v", err) - } - if len(token) < len("sess_")+64+2 { - t.Errorf("generateSessionToken() = %q, unexpected length", token) - } - if token[:5] != "sess_" { - t.Errorf("generateSessionToken() = %q, want prefix 'sess_'", token) - } -} diff --git a/pkg/security/sectypes/auth.go b/pkg/security/sectypes/auth.go new file mode 100644 index 0000000..47cfeb7 --- /dev/null +++ b/pkg/security/sectypes/auth.go @@ -0,0 +1,76 @@ +package sectypes + +// UserContext holds authenticated user information +type UserContext struct { + UserID int `json:"user_id"` + UserName string `json:"user_name"` + UserLevel int `json:"user_level"` + SessionID string `json:"session_id"` + SessionRID int64 `json:"session_rid"` + RemoteID string `json:"remote_id"` + Roles []string `json:"roles"` + Email string `json:"email"` + Claims map[string]any `json:"claims"` + Meta map[string]any `json:"meta"` // Additional metadata that can hold any JSON-serializable values + TwoFactorEnabled bool `json:"two_factor_enabled"` // Indicates if 2FA is enabled for this user + ProgramUserID int `json:"program_user_id"` + ProgramUserTable string `json:"program_user_table"` +} + +// LoginRequest contains credentials for login +type LoginRequest struct { + Username string `json:"username"` + Password string `json:"password"` + TwoFactorCode string `json:"two_factor_code,omitempty"` // TOTP or backup code + Claims map[string]any `json:"claims"` // Additional login data + Meta map[string]any `json:"meta"` // Additional metadata to be set on user context +} + +// RegisterRequest contains information for new user registration +type RegisterRequest struct { + Username string `json:"username"` + Password string `json:"password"` + Email string `json:"email"` + UserLevel int `json:"user_level"` + Roles []string `json:"roles"` + Claims map[string]any `json:"claims"` // Additional registration data + Meta map[string]any `json:"meta"` // Additional metadata +} + +// LoginResponse contains the result of a login attempt +type LoginResponse struct { + Token string `json:"token"` + RefreshToken string `json:"refresh_token"` + User *UserContext `json:"user"` + ExpiresIn int64 `json:"expires_in"` // Token expiration in seconds + Requires2FA bool `json:"requires_2fa"` // True if 2FA code is required + TwoFactorSetupData *TwoFactorSecret `json:"two_factor_setup,omitempty"` // Present when setting up 2FA + Meta map[string]any `json:"meta"` // Additional metadata to be set on user context +} + +// LogoutRequest contains information for logout +type LogoutRequest struct { + Token string `json:"token"` + UserID int `json:"user_id"` +} + +// PasswordResetRequest initiates a password reset for a user +type PasswordResetRequest struct { + Email string `json:"email,omitempty"` + Username string `json:"username,omitempty"` +} + +// PasswordResetResponse is returned when a reset is initiated +type PasswordResetResponse struct { + // Token is the reset token to be delivered out-of-band (e.g. email). + // The stored procedure may return it for delivery or leave it empty + // if the delivery is handled entirely in the database. + Token string `json:"token"` + ExpiresIn int64 `json:"expires_in"` // seconds +} + +// PasswordResetCompleteRequest completes a password reset using the token +type PasswordResetCompleteRequest struct { + Token string `json:"token"` + NewPassword string `json:"new_password"` +} diff --git a/pkg/security/sectypes/deps_test.go b/pkg/security/sectypes/deps_test.go new file mode 100644 index 0000000..a5f5ec9 --- /dev/null +++ b/pkg/security/sectypes/deps_test.go @@ -0,0 +1,20 @@ +package sectypes_test + +import ( + "os/exec" + "strings" + "testing" +) + +// sectypes is the bottom layer: it may import the standard library only. +func TestOnlyStdlibImports(t *testing.T) { + out, err := exec.Command("go", "list", "-deps", "-f", "{{if not .Standard}}{{.ImportPath}}{{end}}", ".").Output() + if err != nil { + t.Skipf("go list unavailable: %v", err) + } + for _, p := range strings.Fields(string(out)) { + if !strings.HasSuffix(p, "/sectypes") { + t.Errorf("sectypes must import only the standard library, found %s", p) + } + } +} diff --git a/pkg/security/sectypes/keys.go b/pkg/security/sectypes/keys.go new file mode 100644 index 0000000..1548fe8 --- /dev/null +++ b/pkg/security/sectypes/keys.go @@ -0,0 +1,62 @@ +package sectypes + +import ( + "crypto/sha256" + "encoding/hex" + "time" +) + +// KeyType identifies the category of an auth key. +type KeyType string + +const ( + // KeyTypeJWTSecret is a per-user JWT signing secret for token generation. + KeyTypeJWTSecret KeyType = "jwt_secret" + // KeyTypeHeaderAPI is a static API key sent via a request header. + KeyTypeHeaderAPI KeyType = "header_api" + // KeyTypeOAuth2 holds OAuth2 client credentials (client_id / client_secret). + KeyTypeOAuth2 KeyType = "oauth2" + // KeyTypeGenericAPI is a generic application API key. + KeyTypeGenericAPI KeyType = "api" +) + +// UserKey represents a single named auth key belonging to a user. +// KeyHash stores the SHA-256 hex digest of the raw key; the raw key is never persisted. +type UserKey struct { + ID int64 `json:"id"` + UserID int `json:"user_id"` + KeyType KeyType `json:"key_type"` + KeyHash string `json:"key_hash"` // SHA-256 hex; never the raw key + Name string `json:"name"` + Scopes []string `json:"scopes,omitempty"` + Meta map[string]any `json:"meta,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + CreatedAt time.Time `json:"created_at"` + LastUsedAt *time.Time `json:"last_used_at,omitempty"` + IsActive bool `json:"is_active"` +} + +// CreateKeyRequest specifies the parameters for a new key. +type CreateKeyRequest struct { + UserID int + KeyType KeyType + Name string + Scopes []string + Meta map[string]any + ExpiresAt *time.Time +} + +// CreateKeyResponse is returned exactly once when a key is created. +// The caller is responsible for persisting RawKey; it is not stored anywhere. +type CreateKeyResponse struct { + Key UserKey + RawKey string // crypto/rand 32 bytes, base64url-encoded +} + +// HashKey returns the lowercase hex SHA-256 digest of a raw key. It is the value +// stored in UserKey.KeyHash; every keystore implementation uses it to hash raw +// keys before storage or lookup. +func HashKey(raw string) string { + sum := sha256.Sum256([]byte(raw)) + return hex.EncodeToString(sum[:]) +} diff --git a/pkg/security/sectypes/oauth.go b/pkg/security/sectypes/oauth.go new file mode 100644 index 0000000..df8c6b9 --- /dev/null +++ b/pkg/security/sectypes/oauth.go @@ -0,0 +1,40 @@ +package sectypes + +import "time" + +// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client. +type OAuthServerClient struct { + ClientID string `json:"client_id"` + RedirectURIs []string `json:"redirect_uris"` + ClientName string `json:"client_name,omitempty"` + GrantTypes []string `json:"grant_types"` + AllowedScopes []string `json:"allowed_scopes,omitempty"` + ClientSecretHash string `json:"client_secret_hash,omitempty"` + TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"` +} + +// OAuthCode is a short-lived authorization code. +type OAuthCode struct { + Code string `json:"code"` + ClientID string `json:"client_id"` + RedirectURI string `json:"redirect_uri"` + ClientState string `json:"client_state,omitempty"` + CodeChallenge string `json:"code_challenge"` + CodeChallengeMethod string `json:"code_challenge_method"` + SessionToken string `json:"session_token"` + RefreshToken string `json:"refresh_token,omitempty"` + Scopes []string `json:"scopes,omitempty"` + ExpiresAt time.Time `json:"expires_at"` +} + +// OAuthTokenInfo is the RFC 7662 token introspection response. +type OAuthTokenInfo struct { + Active bool `json:"active"` + Sub string `json:"sub,omitempty"` + Username string `json:"username,omitempty"` + Email string `json:"email,omitempty"` + UserLevel int `json:"user_level,omitempty"` + Roles []string `json:"roles,omitempty"` + Exp int64 `json:"exp,omitempty"` + Iat int64 `json:"iat,omitempty"` +} diff --git a/pkg/security/sectypes/passkey.go b/pkg/security/sectypes/passkey.go new file mode 100644 index 0000000..0a9595f --- /dev/null +++ b/pkg/security/sectypes/passkey.go @@ -0,0 +1,112 @@ +package sectypes + +import "time" + +// PasskeyCredential represents a stored WebAuthn/FIDO2 credential +type PasskeyCredential struct { + ID string `json:"id"` + UserID int `json:"user_id"` + CredentialID []byte `json:"credential_id"` // Raw credential ID from authenticator + PublicKey []byte `json:"public_key"` // COSE public key + AttestationType string `json:"attestation_type"` // none, indirect, direct + AAGUID []byte `json:"aaguid"` // Authenticator AAGUID + SignCount uint32 `json:"sign_count"` // Signature counter + CloneWarning bool `json:"clone_warning"` // True if cloning detected + Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal + BackupEligible bool `json:"backup_eligible"` // Credential can be backed up + BackupState bool `json:"backup_state"` // Credential is currently backed up + Name string `json:"name,omitempty"` // User-friendly name + CreatedAt time.Time `json:"created_at"` + LastUsedAt time.Time `json:"last_used_at"` +} + +// PasskeyRegistrationOptions contains options for beginning passkey registration +type PasskeyRegistrationOptions struct { + Challenge []byte `json:"challenge"` + RelyingParty PasskeyRelyingParty `json:"rp"` + User PasskeyUser `json:"user"` + PubKeyCredParams []PasskeyCredentialParam `json:"pubKeyCredParams"` + Timeout int64 `json:"timeout,omitempty"` // Milliseconds + ExcludeCredentials []PasskeyCredentialDescriptor `json:"excludeCredentials,omitempty"` + AuthenticatorSelection *PasskeyAuthenticatorSelection `json:"authenticatorSelection,omitempty"` + Attestation string `json:"attestation,omitempty"` // none, indirect, direct, enterprise + Extensions map[string]any `json:"extensions,omitempty"` +} + +// PasskeyAuthenticationOptions contains options for beginning passkey authentication +type PasskeyAuthenticationOptions struct { + Challenge []byte `json:"challenge"` + Timeout int64 `json:"timeout,omitempty"` + RelyingPartyID string `json:"rpId,omitempty"` + AllowCredentials []PasskeyCredentialDescriptor `json:"allowCredentials,omitempty"` + UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged + Extensions map[string]any `json:"extensions,omitempty"` +} + +// PasskeyRelyingParty identifies the relying party +type PasskeyRelyingParty struct { + ID string `json:"id"` // Domain (e.g., "example.com") + Name string `json:"name"` // Display name +} + +// PasskeyUser identifies the user +type PasskeyUser struct { + ID []byte `json:"id"` // User handle (unique, persistent) + Name string `json:"name"` // Username + DisplayName string `json:"displayName"` // Display name +} + +// PasskeyCredentialParam specifies supported public key algorithm +type PasskeyCredentialParam struct { + Type string `json:"type"` // "public-key" + Alg int `json:"alg"` // COSE algorithm identifier (e.g., -7 for ES256, -257 for RS256) +} + +// PasskeyCredentialDescriptor describes a credential +type PasskeyCredentialDescriptor struct { + Type string `json:"type"` // "public-key" + ID []byte `json:"id"` // Credential ID + Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal +} + +// PasskeyAuthenticatorSelection specifies authenticator requirements +type PasskeyAuthenticatorSelection struct { + AuthenticatorAttachment string `json:"authenticatorAttachment,omitempty"` // platform, cross-platform + RequireResidentKey bool `json:"requireResidentKey,omitempty"` + ResidentKey string `json:"residentKey,omitempty"` // discouraged, preferred, required + UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged +} + +// PasskeyRegistrationResponse contains the client's registration response +type PasskeyRegistrationResponse struct { + ID string `json:"id"` // Base64URL encoded credential ID + RawID []byte `json:"rawId"` // Raw credential ID + Type string `json:"type"` // "public-key" + Response PasskeyAuthenticatorAttestationResponse `json:"response"` + ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"` + Transports []string `json:"transports,omitempty"` +} + +// PasskeyAuthenticatorAttestationResponse contains attestation data +type PasskeyAuthenticatorAttestationResponse struct { + ClientDataJSON []byte `json:"clientDataJSON"` + AttestationObject []byte `json:"attestationObject"` + Transports []string `json:"transports,omitempty"` +} + +// PasskeyAuthenticationResponse contains the client's authentication response +type PasskeyAuthenticationResponse struct { + ID string `json:"id"` // Base64URL encoded credential ID + RawID []byte `json:"rawId"` // Raw credential ID + Type string `json:"type"` // "public-key" + Response PasskeyAuthenticatorAssertionResponse `json:"response"` + ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"` +} + +// PasskeyAuthenticatorAssertionResponse contains assertion data +type PasskeyAuthenticatorAssertionResponse struct { + ClientDataJSON []byte `json:"clientDataJSON"` + AuthenticatorData []byte `json:"authenticatorData"` + Signature []byte `json:"signature"` + UserHandle []byte `json:"userHandle,omitempty"` +} diff --git a/pkg/security/sectypes/policy.go b/pkg/security/sectypes/policy.go new file mode 100644 index 0000000..d4266a3 --- /dev/null +++ b/pkg/security/sectypes/policy.go @@ -0,0 +1,98 @@ +package sectypes + +import ( + "fmt" + "reflect" + "regexp" + "strings" +) + +type ColumnSecurity struct { + Schema string `json:"schema"` + Tablename string `json:"tablename"` + Path []string `json:"path"` + ExtraFilters map[string]string `json:"extra_filters"` + UserID int `json:"user_id"` + Accesstype string `json:"accesstype"` + MaskStart int `json:"mask_start"` + MaskEnd int `json:"mask_end"` + MaskInvert bool `json:"mask_invert"` + MaskChar string `json:"mask_char"` + Control string `json:"control"` + ID int `json:"id"` +} + +type RowSecurity struct { + Schema string `json:"schema"` + Tablename string `json:"tablename"` + Template string `json:"template"` + HasBlock bool `json:"has_block"` + // UserID is the opaque user reference the security rules were loaded for. + // It may be an int, a string/UUID, or a *UserContext, depending on what the + // RowSecurityProvider/SecurityContext.GetUserRef implementation returns. + UserID any `json:"user_id"` +} + +// safeIdentRe matches an unquoted SQL identifier. Identifiers substituted into a +// row-security template must match it; anything else is rejected. +var safeIdentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) + +// userIDScalar reduces the opaque user reference to a scalar that is safe to +// bind as a query argument. A *UserContext is reduced to its UserID; other +// structured values are rejected rather than stringified into SQL. +func userIDScalar(ref any) (any, error) { + switch v := ref.(type) { + case nil: + return nil, fmt.Errorf("row security: no user reference") + case *UserContext: + if v == nil { + return nil, fmt.Errorf("row security: nil user context") + } + return v.UserID, nil + case UserContext: + return v.UserID, nil + case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64: + return v, nil + case string: + return v, nil + default: + return nil, fmt.Errorf("row security: unsupported user reference type %T", ref) + } +} + +// GetTemplate expands the row-security template into a WHERE clause and its +// bind arguments. {PrimaryKeyName}, {TableName} and {SchemaName} are validated +// identifiers substituted in place; every {UserID} becomes a `?` placeholder +// with the user reference bound as an argument, so user data never reaches the +// SQL text. +func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) (clause string, args []any, err error) { + str := m.Template + + for placeholder, ident := range map[string]string{ + "{PrimaryKeyName}": pPrimaryKeyName, + "{TableName}": m.Tablename, + "{SchemaName}": m.Schema, + } { + if !strings.Contains(str, placeholder) { + continue + } + if !safeIdentRe.MatchString(ident) { + return "", nil, fmt.Errorf("row security: invalid identifier %q for %s", ident, placeholder) + } + str = strings.ReplaceAll(str, placeholder, ident) + } + + n := strings.Count(str, "{UserID}") + if n == 0 { + return str, nil, nil + } + uid, err := userIDScalar(m.UserID) + if err != nil { + return "", nil, err + } + args = make([]any, n) + for i := range args { + args[i] = uid + } + return strings.ReplaceAll(str, "{UserID}", "?"), args, nil +} diff --git a/pkg/security/sectypes/twofactor.go b/pkg/security/sectypes/twofactor.go new file mode 100644 index 0000000..cfb2bf4 --- /dev/null +++ b/pkg/security/sectypes/twofactor.go @@ -0,0 +1,10 @@ +package sectypes + +// TwoFactorSecret contains 2FA setup information +type TwoFactorSecret struct { + Secret string `json:"secret"` // Base32 encoded secret + QRCodeURL string `json:"qr_code_url"` // URL for QR code generation + BackupCodes []string `json:"backup_codes"` // One-time backup codes + Issuer string `json:"issuer"` // Application name + AccountName string `json:"account_name"` // User identifier (email/username) +} diff --git a/pkg/security/sql_names.go b/pkg/security/sql_names.go deleted file mode 100644 index 40e5b75..0000000 --- a/pkg/security/sql_names.go +++ /dev/null @@ -1,267 +0,0 @@ -package security - -import ( - "fmt" - "reflect" - "regexp" -) - -var validSQLIdentifier = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`) - -// SQLNames defines all configurable SQL stored procedure and table names -// used by the security package. Override individual fields to remap -// to custom database objects. Use DefaultSQLNames() for baseline defaults, -// and MergeSQLNames() to apply partial overrides. -type SQLNames struct { - // Auth procedures (DatabaseAuthenticator) - Login string // default: "resolvespec_login" - Register string // default: "resolvespec_register" - Logout string // default: "resolvespec_logout" - Session string // default: "resolvespec_session" - SessionUpdate string // default: "resolvespec_session_update" - RefreshToken string // default: "resolvespec_refresh_token" - - // JWT procedures (JWTAuthenticator) - JWTLogin string // default: "resolvespec_jwt_login" - JWTLogout string // default: "resolvespec_jwt_logout" - - // Security policy procedures - ColumnSecurity string // default: "resolvespec_column_security" - RowSecurity string // default: "resolvespec_row_security" - - // TOTP procedures (DatabaseTwoFactorProvider) - TOTPEnable string // default: "resolvespec_totp_enable" - TOTPDisable string // default: "resolvespec_totp_disable" - TOTPGetStatus string // default: "resolvespec_totp_get_status" - TOTPGetSecret string // default: "resolvespec_totp_get_secret" - TOTPRegenerateBackup string // default: "resolvespec_totp_regenerate_backup_codes" - TOTPValidateBackupCode string // default: "resolvespec_totp_validate_backup_code" - - // Passkey procedures (DatabasePasskeyProvider) - PasskeyStoreCredential string // default: "resolvespec_passkey_store_credential" - PasskeyGetCredsByUsername string // default: "resolvespec_passkey_get_credentials_by_username" - PasskeyGetCredential string // default: "resolvespec_passkey_get_credential" - PasskeyUpdateCounter string // default: "resolvespec_passkey_update_counter" - PasskeyGetUserCredentials string // default: "resolvespec_passkey_get_user_credentials" - PasskeyDeleteCredential string // default: "resolvespec_passkey_delete_credential" - PasskeyUpdateName string // default: "resolvespec_passkey_update_name" - PasskeyLogin string // default: "resolvespec_passkey_login" - - // Password reset procedures (DatabaseAuthenticator) - PasswordResetRequest string // default: "resolvespec_password_reset_request" - PasswordResetComplete string // default: "resolvespec_password_reset" - - // OAuth2 procedures (DatabaseAuthenticator OAuth2 methods) - OAuthGetOrCreateUser string // default: "resolvespec_oauth_getorcreateuser" - OAuthCreateSession string // default: "resolvespec_oauth_createsession" - OAuthGetRefreshToken string // default: "resolvespec_oauth_getrefreshtoken" - OAuthUpdateRefreshToken string // default: "resolvespec_oauth_updaterefreshtoken" - OAuthGetUser string // default: "resolvespec_oauth_getuser" - - // OAuth2 server procedures (OAuthServer persistence) - OAuthRegisterClient string // default: "resolvespec_oauth_register_client" - OAuthGetClient string // default: "resolvespec_oauth_get_client" - OAuthSaveCode string // default: "resolvespec_oauth_save_code" - OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code" - OAuthIntrospect string // default: "resolvespec_oauth_introspect" - OAuthRevoke string // default: "resolvespec_oauth_revoke" -} - -// DefaultSQLNames returns an SQLNames with all default resolvespec_* values. -func DefaultSQLNames() *SQLNames { - return &SQLNames{ //nolint:gosec // G101: false positive: identifier/example, not a credential - Login: "resolvespec_login", - Register: "resolvespec_register", - Logout: "resolvespec_logout", - Session: "resolvespec_session", - SessionUpdate: "resolvespec_session_update", - RefreshToken: "resolvespec_refresh_token", - - JWTLogin: "resolvespec_jwt_login", - JWTLogout: "resolvespec_jwt_logout", - - ColumnSecurity: "resolvespec_column_security", - RowSecurity: "resolvespec_row_security", - - TOTPEnable: "resolvespec_totp_enable", - TOTPDisable: "resolvespec_totp_disable", - TOTPGetStatus: "resolvespec_totp_get_status", - TOTPGetSecret: "resolvespec_totp_get_secret", - TOTPRegenerateBackup: "resolvespec_totp_regenerate_backup_codes", - TOTPValidateBackupCode: "resolvespec_totp_validate_backup_code", - - PasskeyStoreCredential: "resolvespec_passkey_store_credential", - PasskeyGetCredsByUsername: "resolvespec_passkey_get_credentials_by_username", - PasskeyGetCredential: "resolvespec_passkey_get_credential", - PasskeyUpdateCounter: "resolvespec_passkey_update_counter", - PasskeyGetUserCredentials: "resolvespec_passkey_get_user_credentials", - PasskeyDeleteCredential: "resolvespec_passkey_delete_credential", - PasskeyUpdateName: "resolvespec_passkey_update_name", - PasskeyLogin: "resolvespec_passkey_login", - - PasswordResetRequest: "resolvespec_password_reset_request", - PasswordResetComplete: "resolvespec_password_reset", - - OAuthGetOrCreateUser: "resolvespec_oauth_getorcreateuser", - OAuthCreateSession: "resolvespec_oauth_createsession", - OAuthGetRefreshToken: "resolvespec_oauth_getrefreshtoken", - OAuthUpdateRefreshToken: "resolvespec_oauth_updaterefreshtoken", - OAuthGetUser: "resolvespec_oauth_getuser", - - OAuthRegisterClient: "resolvespec_oauth_register_client", - OAuthGetClient: "resolvespec_oauth_get_client", - OAuthSaveCode: "resolvespec_oauth_save_code", - OAuthExchangeCode: "resolvespec_oauth_exchange_code", - OAuthIntrospect: "resolvespec_oauth_introspect", - OAuthRevoke: "resolvespec_oauth_revoke", - } -} - -// MergeSQLNames returns a copy of base with any non-empty fields from override applied. -// If override is nil, a copy of base is returned. -func MergeSQLNames(base, override *SQLNames) *SQLNames { - if override == nil { - copied := *base - return &copied - } - merged := *base - if override.Login != "" { - merged.Login = override.Login - } - if override.Register != "" { - merged.Register = override.Register - } - if override.Logout != "" { - merged.Logout = override.Logout - } - if override.Session != "" { - merged.Session = override.Session - } - if override.SessionUpdate != "" { - merged.SessionUpdate = override.SessionUpdate - } - if override.RefreshToken != "" { - merged.RefreshToken = override.RefreshToken - } - if override.JWTLogin != "" { - merged.JWTLogin = override.JWTLogin - } - if override.JWTLogout != "" { - merged.JWTLogout = override.JWTLogout - } - if override.ColumnSecurity != "" { - merged.ColumnSecurity = override.ColumnSecurity - } - if override.RowSecurity != "" { - merged.RowSecurity = override.RowSecurity - } - if override.TOTPEnable != "" { - merged.TOTPEnable = override.TOTPEnable - } - if override.TOTPDisable != "" { - merged.TOTPDisable = override.TOTPDisable - } - if override.TOTPGetStatus != "" { - merged.TOTPGetStatus = override.TOTPGetStatus - } - if override.TOTPGetSecret != "" { - merged.TOTPGetSecret = override.TOTPGetSecret - } - if override.TOTPRegenerateBackup != "" { - merged.TOTPRegenerateBackup = override.TOTPRegenerateBackup - } - if override.TOTPValidateBackupCode != "" { - merged.TOTPValidateBackupCode = override.TOTPValidateBackupCode - } - if override.PasskeyStoreCredential != "" { - merged.PasskeyStoreCredential = override.PasskeyStoreCredential - } - if override.PasskeyGetCredsByUsername != "" { - merged.PasskeyGetCredsByUsername = override.PasskeyGetCredsByUsername - } - if override.PasskeyGetCredential != "" { - merged.PasskeyGetCredential = override.PasskeyGetCredential - } - if override.PasskeyUpdateCounter != "" { - merged.PasskeyUpdateCounter = override.PasskeyUpdateCounter - } - if override.PasskeyGetUserCredentials != "" { - merged.PasskeyGetUserCredentials = override.PasskeyGetUserCredentials - } - if override.PasskeyDeleteCredential != "" { - merged.PasskeyDeleteCredential = override.PasskeyDeleteCredential - } - if override.PasskeyUpdateName != "" { - merged.PasskeyUpdateName = override.PasskeyUpdateName - } - if override.PasskeyLogin != "" { - merged.PasskeyLogin = override.PasskeyLogin - } - if override.PasswordResetRequest != "" { - merged.PasswordResetRequest = override.PasswordResetRequest - } - if override.PasswordResetComplete != "" { - merged.PasswordResetComplete = override.PasswordResetComplete - } - if override.OAuthGetOrCreateUser != "" { - merged.OAuthGetOrCreateUser = override.OAuthGetOrCreateUser - } - if override.OAuthCreateSession != "" { - merged.OAuthCreateSession = override.OAuthCreateSession - } - if override.OAuthGetRefreshToken != "" { - merged.OAuthGetRefreshToken = override.OAuthGetRefreshToken - } - if override.OAuthUpdateRefreshToken != "" { - merged.OAuthUpdateRefreshToken = override.OAuthUpdateRefreshToken - } - if override.OAuthGetUser != "" { - merged.OAuthGetUser = override.OAuthGetUser - } - if override.OAuthRegisterClient != "" { - merged.OAuthRegisterClient = override.OAuthRegisterClient - } - if override.OAuthGetClient != "" { - merged.OAuthGetClient = override.OAuthGetClient - } - if override.OAuthSaveCode != "" { - merged.OAuthSaveCode = override.OAuthSaveCode - } - if override.OAuthExchangeCode != "" { - merged.OAuthExchangeCode = override.OAuthExchangeCode - } - if override.OAuthIntrospect != "" { - merged.OAuthIntrospect = override.OAuthIntrospect - } - if override.OAuthRevoke != "" { - merged.OAuthRevoke = override.OAuthRevoke - } - return &merged -} - -// ValidateSQLNames checks that all non-empty fields in names are valid SQL identifiers. -// Returns an error if any field contains invalid characters. -func ValidateSQLNames(names *SQLNames) error { - v := reflect.ValueOf(names).Elem() - typ := v.Type() - for i := 0; i < v.NumField(); i++ { - field := v.Field(i) - if field.Kind() != reflect.String { - continue - } - val := field.String() - if val != "" && !validSQLIdentifier.MatchString(val) { - return fmt.Errorf("SQLNames.%s contains invalid characters: %q", typ.Field(i).Name, val) - } - } - return nil -} - -// resolveSQLNames merges an optional override with defaults. -// Used by constructors that accept variadic *SQLNames. -func resolveSQLNames(override ...*SQLNames) *SQLNames { - if len(override) > 0 && override[0] != nil { - return MergeSQLNames(DefaultSQLNames(), override[0]) - } - return DefaultSQLNames() -} diff --git a/pkg/security/sql_names_test.go b/pkg/security/sql_names_test.go deleted file mode 100644 index be7a500..0000000 --- a/pkg/security/sql_names_test.go +++ /dev/null @@ -1,145 +0,0 @@ -package security - -import ( - "reflect" - "testing" -) - -func TestDefaultSQLNames_AllFieldsNonEmpty(t *testing.T) { - names := DefaultSQLNames() - v := reflect.ValueOf(names).Elem() - typ := v.Type() - - for i := 0; i < v.NumField(); i++ { - field := v.Field(i) - if field.Kind() != reflect.String { - continue - } - if field.String() == "" { - t.Errorf("DefaultSQLNames().%s is empty", typ.Field(i).Name) - } - } -} - -func TestMergeSQLNames_PartialOverride(t *testing.T) { - base := DefaultSQLNames() - override := &SQLNames{ - Login: "custom_login", - TOTPEnable: "custom_totp_enable", - PasskeyLogin: "custom_passkey_login", - } - - merged := MergeSQLNames(base, override) - - if merged.Login != "custom_login" { - t.Errorf("MergeSQLNames().Login = %q, want %q", merged.Login, "custom_login") - } - if merged.TOTPEnable != "custom_totp_enable" { - t.Errorf("MergeSQLNames().TOTPEnable = %q, want %q", merged.TOTPEnable, "custom_totp_enable") - } - if merged.PasskeyLogin != "custom_passkey_login" { - t.Errorf("MergeSQLNames().PasskeyLogin = %q, want %q", merged.PasskeyLogin, "custom_passkey_login") - } - // Non-overridden fields should retain defaults - if merged.Logout != "resolvespec_logout" { - t.Errorf("MergeSQLNames().Logout = %q, want %q", merged.Logout, "resolvespec_logout") - } - if merged.Session != "resolvespec_session" { - t.Errorf("MergeSQLNames().Session = %q, want %q", merged.Session, "resolvespec_session") - } -} - -func TestMergeSQLNames_NilOverride(t *testing.T) { - base := DefaultSQLNames() - merged := MergeSQLNames(base, nil) - - // Should be a copy, not the same pointer - if merged == base { - t.Error("MergeSQLNames with nil override should return a copy, not the same pointer") - } - - // All values should match - v1 := reflect.ValueOf(base).Elem() - v2 := reflect.ValueOf(merged).Elem() - typ := v1.Type() - - for i := 0; i < v1.NumField(); i++ { - f1 := v1.Field(i) - f2 := v2.Field(i) - if f1.Kind() != reflect.String { - continue - } - if f1.String() != f2.String() { - t.Errorf("MergeSQLNames(base, nil).%s = %q, want %q", typ.Field(i).Name, f2.String(), f1.String()) - } - } -} - -func TestMergeSQLNames_DoesNotMutateBase(t *testing.T) { - base := DefaultSQLNames() - originalLogin := base.Login - - override := &SQLNames{Login: "custom_login"} - _ = MergeSQLNames(base, override) - - if base.Login != originalLogin { - t.Errorf("MergeSQLNames mutated base: Login = %q, want %q", base.Login, originalLogin) - } -} - -func TestMergeSQLNames_AllFieldsMerged(t *testing.T) { - base := DefaultSQLNames() - override := &SQLNames{} - v := reflect.ValueOf(override).Elem() - for i := 0; i < v.NumField(); i++ { - if v.Field(i).Kind() == reflect.String { - v.Field(i).SetString("custom_sentinel") - } - } - - merged := MergeSQLNames(base, override) - mv := reflect.ValueOf(merged).Elem() - typ := mv.Type() - for i := 0; i < mv.NumField(); i++ { - if mv.Field(i).Kind() != reflect.String { - continue - } - if mv.Field(i).String() != "custom_sentinel" { - t.Errorf("MergeSQLNames did not merge field %s", typ.Field(i).Name) - } - } -} - -func TestValidateSQLNames_Valid(t *testing.T) { - names := DefaultSQLNames() - if err := ValidateSQLNames(names); err != nil { - t.Errorf("ValidateSQLNames(defaults) error = %v", err) - } -} - -func TestValidateSQLNames_Invalid(t *testing.T) { - names := DefaultSQLNames() - names.Login = "resolvespec_login; DROP TABLE users; --" - - err := ValidateSQLNames(names) - if err == nil { - t.Error("ValidateSQLNames should reject names with invalid characters") - } -} - -func TestResolveSQLNames_NoOverride(t *testing.T) { - names := resolveSQLNames() - if names.Login != "resolvespec_login" { - t.Errorf("resolveSQLNames().Login = %q, want default", names.Login) - } -} - -func TestResolveSQLNames_WithOverride(t *testing.T) { - names := resolveSQLNames(&SQLNames{Login: "custom_login"}) - if names.Login != "custom_login" { - t.Errorf("resolveSQLNames().Login = %q, want %q", names.Login, "custom_login") - } - if names.Logout != "resolvespec_logout" { - t.Errorf("resolveSQLNames().Logout = %q, want default", names.Logout) - } -} diff --git a/pkg/security/table_names.go b/pkg/security/table_names.go deleted file mode 100644 index fd4108b..0000000 --- a/pkg/security/table_names.go +++ /dev/null @@ -1,101 +0,0 @@ -package security - -import ( - "errors" - "fmt" - "reflect" -) - -// ErrDirectModeUnsupported is returned by Direct-mode operations that have no -// portable equivalent because they depend on an external schema this package -// does not own (e.g. core.secaccess / core.hub_link for column/row security). -// Callers needing that functionality must run against Postgres with the -// resolvespec_column_security / resolvespec_row_security stored procedures -// installed (ModeProcedure or ModeAuto with the procedures present). -var ErrDirectModeUnsupported = errors.New("direct mode does not support column/row security; requires the resolvespec_column_security/resolvespec_row_security stored procedures") - -// TableNames defines all configurable table names used by Direct-mode SQL -// in the security package. Override individual fields to remap to custom -// table names. Use DefaultTableNames() for baseline defaults, and -// MergeTableNames() to apply partial overrides. -type TableNames struct { - Users string // default: "users" - UserSessions string // default: "user_sessions" - TokenBlacklist string // default: "token_blacklist" - UserTOTPBackupCodes string // default: "user_totp_backup_codes" - UserPasskeyCredentials string // default: "user_passkey_credentials" - UserPasswordResets string // default: "user_password_resets" - OAuthClients string // default: "oauth_clients" - OAuthCodes string // default: "oauth_codes" -} - -// DefaultTableNames returns a TableNames with all default table names. -func DefaultTableNames() *TableNames { - return &TableNames{ //nolint:gosec // G101: false positive: identifier/example, not a credential - Users: "users", - UserSessions: "user_sessions", - TokenBlacklist: "token_blacklist", - UserTOTPBackupCodes: "user_totp_backup_codes", - UserPasskeyCredentials: "user_passkey_credentials", - UserPasswordResets: "user_password_resets", - OAuthClients: "oauth_clients", - OAuthCodes: "oauth_codes", - } -} - -// MergeTableNames returns a copy of base with any non-empty fields from override applied. -// If override is nil, a copy of base is returned. -func MergeTableNames(base, override *TableNames) *TableNames { - if override == nil { - copied := *base - return &copied - } - merged := *base - if override.Users != "" { - merged.Users = override.Users - } - if override.UserSessions != "" { - merged.UserSessions = override.UserSessions - } - if override.TokenBlacklist != "" { - merged.TokenBlacklist = override.TokenBlacklist - } - if override.UserTOTPBackupCodes != "" { - merged.UserTOTPBackupCodes = override.UserTOTPBackupCodes - } - if override.UserPasskeyCredentials != "" { - merged.UserPasskeyCredentials = override.UserPasskeyCredentials - } - if override.UserPasswordResets != "" { - merged.UserPasswordResets = override.UserPasswordResets - } - if override.OAuthClients != "" { - merged.OAuthClients = override.OAuthClients - } - if override.OAuthCodes != "" { - merged.OAuthCodes = override.OAuthCodes - } - return &merged -} - -// ValidateTableNames checks that all non-empty fields in names are valid SQL identifiers. -func ValidateTableNames(names *TableNames) error { - v := reflect.ValueOf(names).Elem() - typ := v.Type() - for i := 0; i < v.NumField(); i++ { - field := v.Field(i) - if field.Kind() != reflect.String { - continue - } - val := field.String() - if val != "" && !validSQLIdentifier.MatchString(val) { - return fmt.Errorf("TableNames.%s contains invalid characters: %q", typ.Field(i).Name, val) - } - } - return nil -} - -// resolveTableNames merges an optional override with defaults. -func resolveTableNames(override *TableNames) *TableNames { - return MergeTableNames(DefaultTableNames(), override) -} diff --git a/pkg/security/table_names_test.go b/pkg/security/table_names_test.go deleted file mode 100644 index 3990a06..0000000 --- a/pkg/security/table_names_test.go +++ /dev/null @@ -1,134 +0,0 @@ -package security - -import ( - "reflect" - "testing" -) - -func TestDefaultTableNames_AllFieldsNonEmpty(t *testing.T) { - names := DefaultTableNames() - v := reflect.ValueOf(names).Elem() - typ := v.Type() - - for i := 0; i < v.NumField(); i++ { - field := v.Field(i) - if field.Kind() != reflect.String { - continue - } - if field.String() == "" { - t.Errorf("DefaultTableNames().%s is empty", typ.Field(i).Name) - } - } -} - -func TestMergeTableNames_PartialOverride(t *testing.T) { - base := DefaultTableNames() - override := &TableNames{Users: "custom_users", OAuthCodes: "custom_oauth_codes"} - - merged := MergeTableNames(base, override) - - if merged.Users != "custom_users" { - t.Errorf("MergeTableNames().Users = %q, want %q", merged.Users, "custom_users") - } - if merged.OAuthCodes != "custom_oauth_codes" { - t.Errorf("MergeTableNames().OAuthCodes = %q, want %q", merged.OAuthCodes, "custom_oauth_codes") - } - if merged.UserSessions != "user_sessions" { - t.Errorf("MergeTableNames().UserSessions = %q, want default", merged.UserSessions) - } -} - -func TestMergeTableNames_NilOverride(t *testing.T) { - base := DefaultTableNames() - merged := MergeTableNames(base, nil) - - if merged == base { - t.Error("MergeTableNames with nil override should return a copy, not the same pointer") - } - if *merged != *base { - t.Errorf("MergeTableNames(base, nil) = %+v, want %+v", merged, base) - } -} - -func TestMergeTableNames_DoesNotMutateBase(t *testing.T) { - base := DefaultTableNames() - original := base.Users - - override := &TableNames{Users: "custom_users"} - _ = MergeTableNames(base, override) - - if base.Users != original { - t.Errorf("MergeTableNames mutated base: Users = %q, want %q", base.Users, original) - } -} - -func TestValidateTableNames_Valid(t *testing.T) { - if err := ValidateTableNames(DefaultTableNames()); err != nil { - t.Errorf("ValidateTableNames(defaults) error = %v", err) - } -} - -func TestValidateTableNames_Invalid(t *testing.T) { - names := DefaultTableNames() - names.Users = "users; DROP TABLE users; --" - - if err := ValidateTableNames(names); err == nil { - t.Error("ValidateTableNames should reject names with invalid characters") - } -} - -func TestResolveTableNames_NoOverride(t *testing.T) { - names := resolveTableNames(nil) - if names.Users != "users" { - t.Errorf("resolveTableNames(nil).Users = %q, want default", names.Users) - } -} - -func TestResolveTableNames_WithOverride(t *testing.T) { - names := resolveTableNames(&TableNames{Users: "custom_users"}) - if names.Users != "custom_users" { - t.Errorf("resolveTableNames().Users = %q, want %q", names.Users, "custom_users") - } - if names.UserSessions != "user_sessions" { - t.Errorf("resolveTableNames().UserSessions = %q, want default", names.UserSessions) - } -} - -func TestDefaultKeyStoreTableNames(t *testing.T) { - names := DefaultKeyStoreTableNames() - if names.UserKeys != "user_keys" { - t.Errorf("DefaultKeyStoreTableNames().UserKeys = %q, want %q", names.UserKeys, "user_keys") - } -} - -func TestMergeKeyStoreTableNames_PartialOverride(t *testing.T) { - base := DefaultKeyStoreTableNames() - merged := MergeKeyStoreTableNames(base, &KeyStoreTableNames{UserKeys: "custom_keys"}) - if merged.UserKeys != "custom_keys" { - t.Errorf("MergeKeyStoreTableNames().UserKeys = %q, want %q", merged.UserKeys, "custom_keys") - } -} - -func TestMergeKeyStoreTableNames_NilOverride(t *testing.T) { - base := DefaultKeyStoreTableNames() - merged := MergeKeyStoreTableNames(base, nil) - if merged == base { - t.Error("MergeKeyStoreTableNames with nil override should return a copy, not the same pointer") - } - if *merged != *base { - t.Errorf("MergeKeyStoreTableNames(base, nil) = %+v, want %+v", merged, base) - } -} - -func TestValidateKeyStoreTableNames_Invalid(t *testing.T) { - names := &KeyStoreTableNames{UserKeys: "bad name!"} - if err := ValidateKeyStoreTableNames(names); err == nil { - t.Error("ValidateKeyStoreTableNames should reject names with invalid characters") - } -} - -func TestValidateKeyStoreTableNames_Valid(t *testing.T) { - if err := ValidateKeyStoreTableNames(DefaultKeyStoreTableNames()); err != nil { - t.Errorf("ValidateKeyStoreTableNames(defaults) error = %v", err) - } -} diff --git a/pkg/security/totp_middleware.go b/pkg/security/totp/authenticator.go similarity index 62% rename from pkg/security/totp_middleware.go rename to pkg/security/totp/authenticator.go index 13db3ed..a61dd5e 100644 --- a/pkg/security/totp_middleware.go +++ b/pkg/security/totp/authenticator.go @@ -1,32 +1,41 @@ -package security +package totp import ( "context" "fmt" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" "net/http" ) -// TwoFactorAuthenticator wraps an Authenticator and adds 2FA support -type TwoFactorAuthenticator struct { - baseAuth Authenticator - totp *TOTPGenerator - provider TwoFactorAuthProvider +// BaseAuthenticator is the subset of security.Authenticator that Authenticator wraps. +// It is declared here so totp does not import the core security package. +type BaseAuthenticator interface { + Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) + Logout(ctx context.Context, req sectypes.LogoutRequest) error + Authenticate(r *http.Request) (*sectypes.UserContext, error) } -// NewTwoFactorAuthenticator creates a new 2FA-enabled authenticator -func NewTwoFactorAuthenticator(baseAuth Authenticator, provider TwoFactorAuthProvider, config *TwoFactorConfig) *TwoFactorAuthenticator { +// Authenticator wraps an Authenticator and adds 2FA support +type Authenticator struct { + baseAuth BaseAuthenticator + totp *Generator + provider AuthProvider +} + +// NewAuthenticator creates a new 2FA-enabled authenticator +func NewAuthenticator(baseAuth BaseAuthenticator, provider AuthProvider, config *Config) *Authenticator { if config == nil { - config = DefaultTwoFactorConfig() + config = DefaultConfig() } - return &TwoFactorAuthenticator{ + return &Authenticator{ baseAuth: baseAuth, - totp: NewTOTPGenerator(config), + totp: NewGenerator(config), provider: provider, } } // Login authenticates with 2FA support -func (t *TwoFactorAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) { +func (t *Authenticator) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) { // First, perform standard authentication resp, err := t.baseAuth.Login(ctx, req) if err != nil { @@ -87,22 +96,22 @@ func (t *TwoFactorAuthenticator) Login(ctx context.Context, req LoginRequest) (* } // Logout delegates to base authenticator -func (t *TwoFactorAuthenticator) Logout(ctx context.Context, req LogoutRequest) error { +func (t *Authenticator) Logout(ctx context.Context, req sectypes.LogoutRequest) error { return t.baseAuth.Logout(ctx, req) } // Authenticate delegates to base authenticator -func (t *TwoFactorAuthenticator) Authenticate(r *http.Request) (*UserContext, error) { +func (t *Authenticator) Authenticate(r *http.Request) (*sectypes.UserContext, error) { return t.baseAuth.Authenticate(r) } // Setup2FA initiates 2FA setup for a user -func (t *TwoFactorAuthenticator) Setup2FA(userID int, issuer, accountName string) (*TwoFactorSecret, error) { +func (t *Authenticator) Setup2FA(userID int, issuer, accountName string) (*sectypes.TwoFactorSecret, error) { return t.provider.Generate2FASecret(userID, issuer, accountName) } // Enable2FA completes 2FA setup after user confirms with a valid code -func (t *TwoFactorAuthenticator) Enable2FA(userID int, secret, verificationCode string) error { +func (t *Authenticator) Enable2FA(userID int, secret, verificationCode string) error { // Verify the code before enabling valid, err := t.totp.ValidateCode(secret, verificationCode) if err != nil { @@ -124,11 +133,11 @@ func (t *TwoFactorAuthenticator) Enable2FA(userID int, secret, verificationCode } // Disable2FA removes 2FA from a user account -func (t *TwoFactorAuthenticator) Disable2FA(userID int) error { +func (t *Authenticator) Disable2FA(userID int) error { return t.provider.Disable2FA(userID) } // RegenerateBackupCodes creates new backup codes for a user -func (t *TwoFactorAuthenticator) RegenerateBackupCodes(userID int, count int) ([]string, error) { +func (t *Authenticator) RegenerateBackupCodes(userID int, count int) ([]string, error) { return t.provider.GenerateBackupCodes(userID, count) } diff --git a/pkg/security/totp_integration_test.go b/pkg/security/totp/integration_test.go similarity index 77% rename from pkg/security/totp_integration_test.go rename to pkg/security/totp/integration_test.go index 88eb3c0..1ffd3f1 100644 --- a/pkg/security/totp_integration_test.go +++ b/pkg/security/totp/integration_test.go @@ -1,4 +1,4 @@ -package security_test +package totp_test import ( "context" @@ -7,19 +7,20 @@ import ( "testing" "time" - "github.com/bitechdev/ResolveSpec/pkg/security" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" + "github.com/bitechdev/ResolveSpec/pkg/security/totp" ) var ErrInvalidCredentials = errors.New("invalid credentials") // MockAuthenticator is a simple authenticator for testing 2FA type MockAuthenticator struct { - users map[string]*security.UserContext + users map[string]*sectypes.UserContext } func NewMockAuthenticator() *MockAuthenticator { return &MockAuthenticator{ - users: map[string]*security.UserContext{ + users: map[string]*sectypes.UserContext{ "testuser": { UserID: 1, UserName: "testuser", @@ -29,13 +30,13 @@ func NewMockAuthenticator() *MockAuthenticator { } } -func (m *MockAuthenticator) Login(ctx context.Context, req security.LoginRequest) (*security.LoginResponse, error) { +func (m *MockAuthenticator) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) { user, exists := m.users[req.Username] if !exists || req.Password != "password" { return nil, ErrInvalidCredentials } - return &security.LoginResponse{ + return §ypes.LoginResponse{ Token: "mock-token", RefreshToken: "mock-refresh-token", User: user, @@ -43,28 +44,29 @@ func (m *MockAuthenticator) Login(ctx context.Context, req security.LoginRequest }, nil } -func (m *MockAuthenticator) LoginWithCookie(ctx context.Context, req security.LoginRequest, _ http.ResponseWriter) (*security.LoginResponse, error) { +func (m *MockAuthenticator) LoginWithCookie(ctx context.Context, req sectypes.LoginRequest, _ http.ResponseWriter) (*sectypes.LoginResponse, error) { return m.Login(ctx, req) } -func (m *MockAuthenticator) Logout(ctx context.Context, req security.LogoutRequest) error { +func (m *MockAuthenticator) Logout(ctx context.Context, req sectypes.LogoutRequest) error { return nil } -func (m *MockAuthenticator) LogoutWithCookie(ctx context.Context, req security.LogoutRequest, _ http.ResponseWriter) error { +func (m *MockAuthenticator) LogoutWithCookie(ctx context.Context, req sectypes.LogoutRequest, _ http.ResponseWriter) error { return m.Logout(ctx, req) } -func (m *MockAuthenticator) Authenticate(r *http.Request) (*security.UserContext, error) { +func (m *MockAuthenticator) Authenticate(r *http.Request) (*sectypes.UserContext, error) { return m.users["testuser"], nil } -func (m *MockAuthenticator) SetAuthenticateCallback(_ func(r *http.Request) (*security.UserContext, error)) {} +func (m *MockAuthenticator) SetAuthenticateCallback(_ func(r *http.Request) (*sectypes.UserContext, error)) { +} func TestTwoFactorAuthenticator_Setup(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup 2FA secret, err := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") @@ -95,8 +97,8 @@ func TestTwoFactorAuthenticator_Setup(t *testing.T) { func TestTwoFactorAuthenticator_Enable2FA(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup 2FA secret, err := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") @@ -105,7 +107,7 @@ func TestTwoFactorAuthenticator_Enable2FA(t *testing.T) { } // Generate valid code - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, err := totp.GenerateCode(secret.Secret, time.Now()) if err != nil { t.Fatalf("GenerateCode() error = %v", err) @@ -130,8 +132,8 @@ func TestTwoFactorAuthenticator_Enable2FA(t *testing.T) { func TestTwoFactorAuthenticator_Enable2FA_InvalidCode(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup 2FA secret, err := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") @@ -154,10 +156,10 @@ func TestTwoFactorAuthenticator_Enable2FA_InvalidCode(t *testing.T) { func TestTwoFactorAuthenticator_Login_Without2FA(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", } @@ -178,17 +180,17 @@ func TestTwoFactorAuthenticator_Login_Without2FA(t *testing.T) { func TestTwoFactorAuthenticator_Login_With2FA_NoCode(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) // Try to login without 2FA code - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", } @@ -209,12 +211,12 @@ func TestTwoFactorAuthenticator_Login_With2FA_NoCode(t *testing.T) { func TestTwoFactorAuthenticator_Login_With2FA_ValidCode(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) @@ -222,7 +224,7 @@ func TestTwoFactorAuthenticator_Login_With2FA_ValidCode(t *testing.T) { newCode, _ := totp.GenerateCode(secret.Secret, time.Now()) // Login with 2FA code - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: newCode, @@ -248,17 +250,17 @@ func TestTwoFactorAuthenticator_Login_With2FA_ValidCode(t *testing.T) { func TestTwoFactorAuthenticator_Login_With2FA_InvalidCode(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) // Try to login with invalid code - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: "000000", @@ -272,12 +274,12 @@ func TestTwoFactorAuthenticator_Login_With2FA_InvalidCode(t *testing.T) { func TestTwoFactorAuthenticator_Login_WithBackupCode(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) @@ -285,7 +287,7 @@ func TestTwoFactorAuthenticator_Login_WithBackupCode(t *testing.T) { backupCodes, _ := tfaAuth.RegenerateBackupCodes(1, 10) // Login with backup code - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: backupCodes[0], @@ -301,7 +303,7 @@ func TestTwoFactorAuthenticator_Login_WithBackupCode(t *testing.T) { } // Try to use same backup code again - req2 := security.LoginRequest{ + req2 := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: backupCodes[0], @@ -315,12 +317,12 @@ func TestTwoFactorAuthenticator_Login_WithBackupCode(t *testing.T) { func TestTwoFactorAuthenticator_Disable2FA(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) @@ -337,7 +339,7 @@ func TestTwoFactorAuthenticator_Disable2FA(t *testing.T) { } // Login should not require 2FA - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", } @@ -354,12 +356,12 @@ func TestTwoFactorAuthenticator_Disable2FA(t *testing.T) { func TestTwoFactorAuthenticator_RegenerateBackupCodes(t *testing.T) { baseAuth := NewMockAuthenticator() - provider := security.NewMemoryTwoFactorProvider(nil) - tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, provider, nil) + provider := totp.NewMemoryProvider(nil) + tfaAuth := totp.NewAuthenticator(baseAuth, provider, nil) // Setup and enable 2FA secret, _ := tfaAuth.Setup2FA(1, "TestApp", "test@example.com") - totp := security.NewTOTPGenerator(nil) + totp := totp.NewGenerator(nil) code, _ := totp.GenerateCode(secret.Secret, time.Now()) tfaAuth.Enable2FA(1, secret.Secret, code) @@ -380,7 +382,7 @@ func TestTwoFactorAuthenticator_RegenerateBackupCodes(t *testing.T) { } // Old codes should not work - req := security.LoginRequest{ + req := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: codes1[0], @@ -392,7 +394,7 @@ func TestTwoFactorAuthenticator_RegenerateBackupCodes(t *testing.T) { } // New codes should work - req2 := security.LoginRequest{ + req2 := sectypes.LoginRequest{ Username: "testuser", Password: "password", TwoFactorCode: codes2[0], diff --git a/pkg/security/totp_provider_memory.go b/pkg/security/totp/memory.go similarity index 68% rename from pkg/security/totp_provider_memory.go rename to pkg/security/totp/memory.go index 3ff3799..dbbb8ed 100644 --- a/pkg/security/totp_provider_memory.go +++ b/pkg/security/totp/memory.go @@ -1,34 +1,35 @@ -package security +package totp import ( "crypto/sha256" "encoding/hex" "fmt" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" "sync" ) -// MemoryTwoFactorProvider is an in-memory implementation of TwoFactorAuthProvider for testing/examples -type MemoryTwoFactorProvider struct { +// MemoryProvider is an in-memory implementation of AuthProvider for testing/examples +type MemoryProvider struct { mu sync.RWMutex secrets map[int]string // userID -> secret backupCodes map[int]map[string]bool // userID -> backup codes (code -> used) - totpGen *TOTPGenerator + totpGen *Generator } -// NewMemoryTwoFactorProvider creates a new in-memory 2FA provider -func NewMemoryTwoFactorProvider(config *TwoFactorConfig) *MemoryTwoFactorProvider { +// NewMemoryProvider creates a new in-memory 2FA provider +func NewMemoryProvider(config *Config) *MemoryProvider { if config == nil { - config = DefaultTwoFactorConfig() + config = DefaultConfig() } - return &MemoryTwoFactorProvider{ + return &MemoryProvider{ secrets: make(map[int]string), backupCodes: make(map[int]map[string]bool), - totpGen: NewTOTPGenerator(config), + totpGen: NewGenerator(config), } } // Generate2FASecret creates a new secret for a user -func (m *MemoryTwoFactorProvider) Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error) { +func (m *MemoryProvider) Generate2FASecret(userID int, issuer, accountName string) (*sectypes.TwoFactorSecret, error) { secret, err := m.totpGen.GenerateSecret() if err != nil { return nil, err @@ -41,7 +42,7 @@ func (m *MemoryTwoFactorProvider) Generate2FASecret(userID int, issuer, accountN return nil, err } - return &TwoFactorSecret{ + return §ypes.TwoFactorSecret{ Secret: secret, QRCodeURL: qrURL, BackupCodes: backupCodes, @@ -51,12 +52,12 @@ func (m *MemoryTwoFactorProvider) Generate2FASecret(userID int, issuer, accountN } // Validate2FACode verifies a TOTP code -func (m *MemoryTwoFactorProvider) Validate2FACode(secret string, code string) (bool, error) { +func (m *MemoryProvider) Validate2FACode(secret string, code string) (bool, error) { return m.totpGen.ValidateCode(secret, code) } // Enable2FA activates 2FA for a user -func (m *MemoryTwoFactorProvider) Enable2FA(userID int, secret string, backupCodes []string) error { +func (m *MemoryProvider) Enable2FA(userID int, secret string, backupCodes []string) error { m.mu.Lock() defer m.mu.Unlock() @@ -77,7 +78,7 @@ func (m *MemoryTwoFactorProvider) Enable2FA(userID int, secret string, backupCod } // Disable2FA deactivates 2FA for a user -func (m *MemoryTwoFactorProvider) Disable2FA(userID int) error { +func (m *MemoryProvider) Disable2FA(userID int) error { m.mu.Lock() defer m.mu.Unlock() @@ -87,7 +88,7 @@ func (m *MemoryTwoFactorProvider) Disable2FA(userID int) error { } // Get2FAStatus checks if user has 2FA enabled -func (m *MemoryTwoFactorProvider) Get2FAStatus(userID int) (bool, error) { +func (m *MemoryProvider) Get2FAStatus(userID int) (bool, error) { m.mu.RLock() defer m.mu.RUnlock() @@ -96,7 +97,7 @@ func (m *MemoryTwoFactorProvider) Get2FAStatus(userID int) (bool, error) { } // Get2FASecret retrieves the user's 2FA secret -func (m *MemoryTwoFactorProvider) Get2FASecret(userID int) (string, error) { +func (m *MemoryProvider) Get2FASecret(userID int) (string, error) { m.mu.RLock() defer m.mu.RUnlock() @@ -108,7 +109,7 @@ func (m *MemoryTwoFactorProvider) Get2FASecret(userID int) (string, error) { } // GenerateBackupCodes creates backup codes for 2FA -func (m *MemoryTwoFactorProvider) GenerateBackupCodes(userID int, count int) ([]string, error) { +func (m *MemoryProvider) GenerateBackupCodes(userID int, count int) ([]string, error) { codes, err := GenerateBackupCodes(count) if err != nil { return nil, err @@ -128,7 +129,7 @@ func (m *MemoryTwoFactorProvider) GenerateBackupCodes(userID int, count int) ([] } // ValidateBackupCode checks and consumes a backup code -func (m *MemoryTwoFactorProvider) ValidateBackupCode(userID int, code string) (bool, error) { +func (m *MemoryProvider) ValidateBackupCode(userID int, code string) (bool, error) { m.mu.Lock() defer m.mu.Unlock() diff --git a/pkg/security/totp.go b/pkg/security/totp/totp.go similarity index 73% rename from pkg/security/totp.go rename to pkg/security/totp/totp.go index 5f018cf..295cccd 100644 --- a/pkg/security/totp.go +++ b/pkg/security/totp/totp.go @@ -1,4 +1,4 @@ -package security +package totp import ( "crypto/hmac" @@ -9,6 +9,7 @@ import ( "encoding/base32" "encoding/binary" "fmt" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" "hash" "math" "net/url" @@ -16,10 +17,10 @@ import ( "time" ) -// TwoFactorAuthProvider defines interface for 2FA operations -type TwoFactorAuthProvider interface { +// AuthProvider defines interface for 2FA operations +type AuthProvider interface { // Generate2FASecret creates a new secret for a user - Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error) + Generate2FASecret(userID int, issuer, accountName string) (*sectypes.TwoFactorSecret, error) // Validate2FACode verifies a TOTP code Validate2FACode(secret string, code string) (bool, error) @@ -43,26 +44,17 @@ type TwoFactorAuthProvider interface { ValidateBackupCode(userID int, code string) (bool, error) } -// TwoFactorSecret contains 2FA setup information -type TwoFactorSecret struct { - Secret string `json:"secret"` // Base32 encoded secret - QRCodeURL string `json:"qr_code_url"` // URL for QR code generation - BackupCodes []string `json:"backup_codes"` // One-time backup codes - Issuer string `json:"issuer"` // Application name - AccountName string `json:"account_name"` // User identifier (email/username) -} - -// TwoFactorConfig holds TOTP configuration -type TwoFactorConfig struct { +// Config holds TOTP configuration +type Config struct { Algorithm string // SHA1, SHA256, SHA512 Digits int // Number of digits in code (6 or 8) Period int // Time step in seconds (default 30) SkewWindow int // Number of time steps to check before/after (default 1) } -// DefaultTwoFactorConfig returns standard TOTP configuration -func DefaultTwoFactorConfig() *TwoFactorConfig { - return &TwoFactorConfig{ +// DefaultConfig returns standard TOTP configuration +func DefaultConfig() *Config { + return &Config{ Algorithm: "SHA1", Digits: 6, Period: 30, @@ -70,23 +62,23 @@ func DefaultTwoFactorConfig() *TwoFactorConfig { } } -// TOTPGenerator handles TOTP code generation and validation -type TOTPGenerator struct { - config *TwoFactorConfig +// Generator handles TOTP code generation and validation +type Generator struct { + config *Config } -// NewTOTPGenerator creates a new TOTP generator with config -func NewTOTPGenerator(config *TwoFactorConfig) *TOTPGenerator { +// NewGenerator creates a new TOTP generator with config +func NewGenerator(config *Config) *Generator { if config == nil { - config = DefaultTwoFactorConfig() + config = DefaultConfig() } - return &TOTPGenerator{ + return &Generator{ config: config, } } // GenerateSecret creates a random base32-encoded secret -func (t *TOTPGenerator) GenerateSecret() (string, error) { +func (t *Generator) GenerateSecret() (string, error) { secret := make([]byte, 20) _, err := rand.Read(secret) if err != nil { @@ -96,7 +88,7 @@ func (t *TOTPGenerator) GenerateSecret() (string, error) { } // GenerateQRCodeURL creates a URL for QR code generation -func (t *TOTPGenerator) GenerateQRCodeURL(secret, issuer, accountName string) string { +func (t *Generator) GenerateQRCodeURL(secret, issuer, accountName string) string { params := url.Values{} params.Set("secret", secret) params.Set("issuer", issuer) @@ -109,7 +101,7 @@ func (t *TOTPGenerator) GenerateQRCodeURL(secret, issuer, accountName string) st } // GenerateCode creates a TOTP code for a given time -func (t *TOTPGenerator) GenerateCode(secret string, timestamp time.Time) (string, error) { +func (t *Generator) GenerateCode(secret string, timestamp time.Time) (string, error) { // Decode secret key, err := base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(strings.ToUpper(secret)) if err != nil { @@ -142,7 +134,7 @@ func (t *TOTPGenerator) GenerateCode(secret string, timestamp time.Time) (string } // ValidateCode checks if a code is valid for the secret -func (t *TOTPGenerator) ValidateCode(secret, code string) (bool, error) { +func (t *Generator) ValidateCode(secret, code string) (bool, error) { now := time.Now() // Check current time and skew window @@ -162,7 +154,7 @@ func (t *TOTPGenerator) ValidateCode(secret, code string) (bool, error) { } // getHashFunc returns the hash function based on algorithm -func (t *TOTPGenerator) getHashFunc() func() hash.Hash { +func (t *Generator) getHashFunc() func() hash.Hash { switch strings.ToUpper(t.config.Algorithm) { case "SHA256": return sha256.New diff --git a/pkg/security/totp_test.go b/pkg/security/totp/totp_test.go similarity index 87% rename from pkg/security/totp_test.go rename to pkg/security/totp/totp_test.go index 9a216c3..c4dcbff 100644 --- a/pkg/security/totp_test.go +++ b/pkg/security/totp/totp_test.go @@ -1,4 +1,4 @@ -package security +package totp import ( "strings" @@ -7,7 +7,7 @@ import ( ) func TestTOTPGenerator_GenerateSecret(t *testing.T) { - totp := NewTOTPGenerator(nil) + totp := NewGenerator(nil) secret, err := totp.GenerateSecret() if err != nil { @@ -25,7 +25,7 @@ func TestTOTPGenerator_GenerateSecret(t *testing.T) { } func TestTOTPGenerator_GenerateQRCodeURL(t *testing.T) { - totp := NewTOTPGenerator(nil) + totp := NewGenerator(nil) secret := "JBSWY3DPEHPK3PXP" issuer := "TestApp" @@ -47,13 +47,13 @@ func TestTOTPGenerator_GenerateQRCodeURL(t *testing.T) { } func TestTOTPGenerator_GenerateCode(t *testing.T) { - config := &TwoFactorConfig{ + config := &Config{ Algorithm: "SHA1", Digits: 6, Period: 30, SkewWindow: 1, } - totp := NewTOTPGenerator(config) + totp := NewGenerator(config) secret := "JBSWY3DPEHPK3PXP" @@ -78,13 +78,13 @@ func TestTOTPGenerator_GenerateCode(t *testing.T) { } func TestTOTPGenerator_ValidateCode(t *testing.T) { - config := &TwoFactorConfig{ + config := &Config{ Algorithm: "SHA1", Digits: 6, Period: 30, SkewWindow: 1, } - totp := NewTOTPGenerator(config) + totp := NewGenerator(config) secret := "JBSWY3DPEHPK3PXP" @@ -118,13 +118,13 @@ func TestTOTPGenerator_ValidateCode(t *testing.T) { } func TestTOTPGenerator_ValidateCode_WithSkew(t *testing.T) { - config := &TwoFactorConfig{ + config := &Config{ Algorithm: "SHA1", Digits: 6, Period: 30, SkewWindow: 2, // Allow 2 periods before/after } - totp := NewTOTPGenerator(config) + totp := NewGenerator(config) secret := "JBSWY3DPEHPK3PXP" @@ -152,13 +152,13 @@ func TestTOTPGenerator_DifferentAlgorithms(t *testing.T) { for _, algo := range algorithms { t.Run(algo, func(t *testing.T) { - config := &TwoFactorConfig{ + config := &Config{ Algorithm: algo, Digits: 6, Period: 30, SkewWindow: 1, } - totp := NewTOTPGenerator(config) + totp := NewGenerator(config) code, err := totp.GenerateCode(secret, time.Now()) if err != nil { @@ -178,13 +178,13 @@ func TestTOTPGenerator_DifferentAlgorithms(t *testing.T) { } func TestTOTPGenerator_8Digits(t *testing.T) { - config := &TwoFactorConfig{ + config := &Config{ Algorithm: "SHA1", Digits: 8, Period: 30, SkewWindow: 1, } - totp := NewTOTPGenerator(config) + totp := NewGenerator(config) secret := "JBSWY3DPEHPK3PXP" @@ -234,27 +234,27 @@ func TestGenerateBackupCodes(t *testing.T) { } func TestDefaultTwoFactorConfig(t *testing.T) { - config := DefaultTwoFactorConfig() + config := DefaultConfig() if config.Algorithm != "SHA1" { - t.Errorf("DefaultTwoFactorConfig() Algorithm = %s, want SHA1", config.Algorithm) + t.Errorf("DefaultConfig() Algorithm = %s, want SHA1", config.Algorithm) } if config.Digits != 6 { - t.Errorf("DefaultTwoFactorConfig() Digits = %d, want 6", config.Digits) + t.Errorf("DefaultConfig() Digits = %d, want 6", config.Digits) } if config.Period != 30 { - t.Errorf("DefaultTwoFactorConfig() Period = %d, want 30", config.Period) + t.Errorf("DefaultConfig() Period = %d, want 30", config.Period) } if config.SkewWindow != 1 { - t.Errorf("DefaultTwoFactorConfig() SkewWindow = %d, want 1", config.SkewWindow) + t.Errorf("DefaultConfig() SkewWindow = %d, want 1", config.SkewWindow) } } func TestTOTPGenerator_InvalidSecret(t *testing.T) { - totp := NewTOTPGenerator(nil) + totp := NewGenerator(nil) // Test with invalid base32 secret _, err := totp.GenerateCode("INVALID!!!", time.Now()) @@ -270,7 +270,7 @@ func TestTOTPGenerator_InvalidSecret(t *testing.T) { // Benchmark tests func BenchmarkTOTPGenerator_GenerateCode(b *testing.B) { - totp := NewTOTPGenerator(nil) + totp := NewGenerator(nil) secret := "JBSWY3DPEHPK3PXP" now := time.Now() @@ -281,7 +281,7 @@ func BenchmarkTOTPGenerator_GenerateCode(b *testing.B) { } func BenchmarkTOTPGenerator_ValidateCode(b *testing.B) { - totp := NewTOTPGenerator(nil) + totp := NewGenerator(nil) secret := "JBSWY3DPEHPK3PXP" code, _ := totp.GenerateCode(secret, time.Now()) diff --git a/pkg/security/totp_provider_database.go b/pkg/security/totp_provider_database.go index 17c306d..2f1250c 100644 --- a/pkg/security/totp_provider_database.go +++ b/pkg/security/totp_provider_database.go @@ -5,93 +5,46 @@ import ( "crypto/sha256" "database/sql" "encoding/hex" - "encoding/json" "fmt" - "sync" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/totp" ) -// DatabaseTwoFactorProvider implements TwoFactorAuthProvider using PostgreSQL stored procedures -// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults) -// See totp_database_schema.sql for procedure definitions +// DatabaseTwoFactorProvider implements TwoFactorAuthProvider on top of the lookup package +// (stored procedures on Postgres by default, direct SQL elsewhere). +// See lookup/database_schema.sql for procedure definitions type DatabaseTwoFactorProvider struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - totpGen *TOTPGenerator - sqlNames *SQLNames - tableNames *TableNames - queryMode QueryMode - capability *dbCapability + src *lookupSource + totpGen *totp.Generator } // NewDatabaseTwoFactorProvider creates a new database-backed 2FA provider -func NewDatabaseTwoFactorProvider(db *sql.DB, config *TwoFactorConfig, names ...*SQLNames) *DatabaseTwoFactorProvider { +func NewDatabaseTwoFactorProvider(db *sql.DB, config *totp.Config) *DatabaseTwoFactorProvider { if config == nil { - config = DefaultTwoFactorConfig() - } - return &DatabaseTwoFactorProvider{ - db: db, - totpGen: NewTOTPGenerator(config), - sqlNames: resolveSQLNames(names...), - tableNames: DefaultTableNames(), - capability: newDBCapability(), + config = totp.DefaultConfig() } + return &DatabaseTwoFactorProvider{src: newLookupSource(db), totpGen: totp.NewGenerator(config)} } // WithDBFactory configures a factory used to reopen the database connection if it is closed. func (p *DatabaseTwoFactorProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseTwoFactorProvider { - p.dbFactory = factory + p.src.opts.DBFactory = factory return p } -// WithTableNames configures Direct-mode table names. If names is nil, defaults are used. -func (p *DatabaseTwoFactorProvider) WithTableNames(names *TableNames) *DatabaseTwoFactorProvider { - p.tableNames = resolveTableNames(names) +// WithLookup configures dialect, query mode and names. Call before first use. +func (p *DatabaseTwoFactorProvider) WithLookup(cfg lookup.Config) *DatabaseTwoFactorProvider { + p.src.cfg = cfg return p } -// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto). -func (p *DatabaseTwoFactorProvider) WithQueryMode(mode QueryMode) *DatabaseTwoFactorProvider { - p.queryMode = mode +// WithLookupProvider uses an existing provider instead of building one. +func (p *DatabaseTwoFactorProvider) WithLookupProvider(lp *lookup.Provider) *DatabaseTwoFactorProvider { + p.src.provider = lp return p } -func (p *DatabaseTwoFactorProvider) getDB() *sql.DB { - p.dbMu.RLock() - defer p.dbMu.RUnlock() - return p.db -} - -func (p *DatabaseTwoFactorProvider) 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 *DatabaseTwoFactorProvider) runDBOpWithReconnect(run func(*sql.DB) error) error { - db := p.getDB() - if db == nil { - return fmt.Errorf("database connection is nil") - } - err := run(db) - if isDBClosed(err) { - if reconnErr := p.reconnectDB(); reconnErr == nil { - err = run(p.getDB()) - } - } - return err -} +func (p *DatabaseTwoFactorProvider) store() lookup.TOTPStore { return p.src.get().TOTP } // Generate2FASecret creates a new secret for a user func (p *DatabaseTwoFactorProvider) Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error) { @@ -102,7 +55,7 @@ func (p *DatabaseTwoFactorProvider) Generate2FASecret(userID int, issuer, accoun qrURL := p.totpGen.GenerateQRCodeURL(secret, issuer, accountName) - backupCodes, err := GenerateBackupCodes(10) + backupCodes, err := totp.GenerateBackupCodes(10) if err != nil { return nil, fmt.Errorf("failed to generate backup codes: %w", err) } @@ -130,124 +83,31 @@ func (p *DatabaseTwoFactorProvider) Enable2FA(userID int, secret string, backupC hashedCodes[i] = hex.EncodeToString(hash[:]) } - // Convert to JSON array - codesJSON, err := json.Marshal(hashedCodes) - if err != nil { - return fmt.Errorf("failed to marshal backup codes: %w", err) - } - ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPEnable) { - return p.enable2FADirect(ctx, userID, secret, hashedCodes) - } - - // Call stored procedure - var success bool - var errorMsg sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, p.sqlNames.TOTPEnable) - err = p.getDB().QueryRow(query, userID, secret, string(codesJSON)).Scan(&success, &errorMsg) - if err != nil { - return fmt.Errorf("enable 2FA query failed: %w", err) - } - - if !success { - if errorMsg.Valid { - return fmt.Errorf("%s", errorMsg.String) - } - return fmt.Errorf("failed to enable 2FA") - } - - return nil + return p.store().Enable(ctx, userID, secret, hashedCodes) } // Disable2FA deactivates 2FA for a user func (p *DatabaseTwoFactorProvider) Disable2FA(userID int) error { ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPDisable) { - return p.disable2FADirect(ctx, userID) - } - - var success bool - var errorMsg sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, p.sqlNames.TOTPDisable) - err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg) - if err != nil { - return fmt.Errorf("disable 2FA query failed: %w", err) - } - - if !success { - if errorMsg.Valid { - return fmt.Errorf("%s", errorMsg.String) - } - return fmt.Errorf("failed to disable 2FA") - } - - return nil + return p.store().Disable(ctx, userID) } // Get2FAStatus checks if user has 2FA enabled func (p *DatabaseTwoFactorProvider) Get2FAStatus(userID int) (bool, error) { ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPGetStatus) { - return p.get2FAStatusDirect(ctx, userID) - } - - var success bool - var errorMsg sql.NullString - var enabled bool - - query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, p.sqlNames.TOTPGetStatus) - err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg, &enabled) - if err != nil { - return false, fmt.Errorf("get 2FA status query failed: %w", err) - } - - if !success { - if errorMsg.Valid { - return false, fmt.Errorf("%s", errorMsg.String) - } - return false, fmt.Errorf("failed to get 2FA status") - } - - return enabled, nil + return p.store().Status(ctx, userID) } // Get2FASecret retrieves the user's 2FA secret func (p *DatabaseTwoFactorProvider) Get2FASecret(userID int) (string, error) { ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPGetSecret) { - return p.get2FASecretDirect(ctx, userID) - } - - var success bool - var errorMsg sql.NullString - var secret sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, p.sqlNames.TOTPGetSecret) - err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg, &secret) - if err != nil { - return "", fmt.Errorf("get 2FA secret query failed: %w", err) - } - - if !success { - if errorMsg.Valid { - return "", fmt.Errorf("%s", errorMsg.String) - } - return "", fmt.Errorf("failed to get 2FA secret") - } - - if !secret.Valid { - return "", fmt.Errorf("2FA secret not found") - } - - return secret.String, nil + return p.store().Secret(ctx, userID) } // GenerateBackupCodes creates backup codes for 2FA func (p *DatabaseTwoFactorProvider) GenerateBackupCodes(userID int, count int) ([]string, error) { - codes, err := GenerateBackupCodes(count) + codes, err := totp.GenerateBackupCodes(count) if err != nil { return nil, fmt.Errorf("failed to generate backup codes: %w", err) } @@ -260,34 +120,8 @@ func (p *DatabaseTwoFactorProvider) GenerateBackupCodes(userID int, count int) ( } ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPRegenerateBackup) { - if err := p.regenerateBackupCodesDirect(ctx, userID, hashedCodes); err != nil { - return nil, err - } - return codes, nil - } - - // Convert to JSON array - codesJSON, err := json.Marshal(hashedCodes) - if err != nil { - return nil, fmt.Errorf("failed to marshal backup codes: %w", err) - } - - // Call stored procedure - var success bool - var errorMsg sql.NullString - - query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, p.sqlNames.TOTPRegenerateBackup) - err = p.getDB().QueryRow(query, userID, string(codesJSON)).Scan(&success, &errorMsg) - if err != nil { - return nil, fmt.Errorf("regenerate backup codes query failed: %w", err) - } - - if !success { - if errorMsg.Valid { - return nil, fmt.Errorf("%s", errorMsg.String) - } - return nil, fmt.Errorf("failed to regenerate backup codes") + if err := p.store().RegenerateBackupCodes(ctx, userID, hashedCodes); err != nil { + return nil, err } // Return unhashed codes to user (only time they see them) @@ -301,26 +135,5 @@ func (p *DatabaseTwoFactorProvider) ValidateBackupCode(userID int, code string) codeHash := hex.EncodeToString(hash[:]) ctx := context.Background() - if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPValidateBackupCode) { - return p.validateBackupCodeDirect(ctx, userID, codeHash) - } - - var success bool - var errorMsg sql.NullString - var valid bool - - query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, p.sqlNames.TOTPValidateBackupCode) - err := p.getDB().QueryRow(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 + return p.store().ValidateBackupCode(ctx, userID, codeHash) } diff --git a/pkg/security/totp_provider_database_direct.go b/pkg/security/totp_provider_database_direct.go deleted file mode 100644 index 500c4c2..0000000 --- a/pkg/security/totp_provider_database_direct.go +++ /dev/null @@ -1,151 +0,0 @@ -package security - -import ( - "context" - "database/sql" - "errors" - "fmt" - "time" -) - -// Direct-mode implementations mirroring the resolvespec_totp_* stored -// procedures in database_schema.sql using plain SQL against TableNames. - -func (p *DatabaseTwoFactorProvider) enable2FADirect(ctx context.Context, userID int, secret string, hashedCodes []string) error { - return p.runDBOpWithReconnect(func(db *sql.DB) error { - updQuery := rewritePlaceholders(db, fmt.Sprintf( - `UPDATE %s SET totp_secret = ?, totp_enabled = ?, totp_enabled_at = ? WHERE id = ?`, p.tableNames.Users)) - res, err := db.ExecContext(ctx, updQuery, secret, true, time.Now(), userID) - if err != nil { - return err - } - if rows, err := res.RowsAffected(); err != nil { - return err - } else if rows == 0 { - return fmt.Errorf("user not found") - } - - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes)) - if _, err := db.ExecContext(ctx, delQuery, userID); err != nil { - return err - } - - insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, code_hash) VALUES (?, ?)`, p.tableNames.UserTOTPBackupCodes)) - for _, hash := range hashedCodes { - if _, err := db.ExecContext(ctx, insQuery, userID, hash); err != nil { - return err - } - } - return nil - }) -} - -func (p *DatabaseTwoFactorProvider) disable2FADirect(ctx context.Context, userID int) error { - return p.runDBOpWithReconnect(func(db *sql.DB) error { - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET totp_secret = NULL, totp_enabled = ? WHERE id = ?`, p.tableNames.Users)) - res, err := db.ExecContext(ctx, updQuery, false, userID) - if err != nil { - return err - } - if rows, err := res.RowsAffected(); err != nil { - return err - } else if rows == 0 { - return fmt.Errorf("user not found") - } - - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes)) - _, err = db.ExecContext(ctx, delQuery, userID) - return err - }) -} - -func (p *DatabaseTwoFactorProvider) get2FAStatusDirect(ctx context.Context, userID int) (bool, error) { - var enabled sql.NullBool - err := p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT totp_enabled FROM %s WHERE id = ?`, p.tableNames.Users)) - return db.QueryRowContext(ctx, query, userID).Scan(&enabled) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return false, fmt.Errorf("user not found") - } - return false, fmt.Errorf("get 2FA status query failed: %w", err) - } - return enabled.Bool, nil -} - -func (p *DatabaseTwoFactorProvider) get2FASecretDirect(ctx context.Context, userID int) (string, error) { - var secret sql.NullString - var enabled sql.NullBool - err := p.runDBOpWithReconnect(func(db *sql.DB) error { - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT totp_secret, totp_enabled FROM %s WHERE id = ?`, p.tableNames.Users)) - return db.QueryRowContext(ctx, query, userID).Scan(&secret, &enabled) - }) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return "", fmt.Errorf("user not found") - } - return "", fmt.Errorf("get 2FA secret query failed: %w", err) - } - if !enabled.Bool { - return "", fmt.Errorf("TOTP not enabled for user") - } - return secret.String, nil -} - -func (p *DatabaseTwoFactorProvider) regenerateBackupCodesDirect(ctx context.Context, userID int, hashedCodes []string) error { - return p.runDBOpWithReconnect(func(db *sql.DB) error { - var count int - checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE id = ? AND totp_enabled = ?`, p.tableNames.Users)) - if err := db.QueryRowContext(ctx, checkQuery, userID, true).Scan(&count); err != nil { - return err - } - if count == 0 { - return fmt.Errorf("user not found or TOTP not enabled") - } - - delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes)) - if _, err := db.ExecContext(ctx, delQuery, userID); err != nil { - return err - } - - insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, code_hash) VALUES (?, ?)`, p.tableNames.UserTOTPBackupCodes)) - for _, hash := range hashedCodes { - if _, err := db.ExecContext(ctx, insQuery, userID, hash); err != nil { - return err - } - } - return nil - }) -} - -func (p *DatabaseTwoFactorProvider) validateBackupCodeDirect(ctx context.Context, userID int, codeHash string) (bool, error) { - var valid bool - err := p.runDBOpWithReconnect(func(db *sql.DB) error { - var codeID int - var used bool - query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, used FROM %s WHERE user_id = ? AND code_hash = ?`, p.tableNames.UserTOTPBackupCodes)) - err := db.QueryRowContext(ctx, query, userID, codeHash).Scan(&codeID, &used) - if errors.Is(err, sql.ErrNoRows) { - valid = false - return nil - } - if err != nil { - return err - } - if used { - return fmt.Errorf("backup code already used") - } - - updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET used = ?, used_at = ? WHERE id = ?`, p.tableNames.UserTOTPBackupCodes)) - if _, err := db.ExecContext(ctx, updQuery, true, time.Now(), codeID); err != nil { - return err - } - valid = true - return nil - }) - if err != nil { - return false, err - } - return valid, nil -} diff --git a/pkg/security/totp_provider_database_test.go b/pkg/security/totp_provider_database_test.go index 5003359..002880e 100644 --- a/pkg/security/totp_provider_database_test.go +++ b/pkg/security/totp_provider_database_test.go @@ -7,7 +7,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/security" ) -// Note: These tests require a PostgreSQL database with the schema from totp_database_schema.sql +// Note: These tests require a PostgreSQL database with the schema from lookup/database_schema.sql // Set TEST_DATABASE_URL environment variable or skip tests func setupTestDB(t *testing.T) *sql.DB { diff --git a/pkg/security/txsettings.go b/pkg/security/txsettings.go index ebec0cc..b50821b 100644 --- a/pkg/security/txsettings.go +++ b/pkg/security/txsettings.go @@ -1,12 +1,8 @@ package security import ( - "encoding/hex" - "fmt" - "regexp" - "sort" - "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" ) // TxSettingsFunc returns the transaction-local settings (e.g. RLS GUCs such as @@ -14,9 +10,6 @@ import ( // OnTxBegin, before any other SQL. Returning an error rolls the transaction back. type TxSettingsFunc func(secCtx SecurityContext) (map[string]string, error) -// settingNameRE matches a custom GUC name: two or more dot-separated identifiers. -var settingNameRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)+$`) - // SetTxSettings sets the function that provides transaction-local settings for // every transaction opened by a spec that registered its security hooks with this // list. Pass nil to disable. May be called before or after RegisterSecurityHooks. @@ -54,33 +47,7 @@ func StampTxSettings(secCtx SecurityContext, list *SecurityList, tx common.Datab // ApplyTxSettings sets each entry as a transaction-local setting on tx, in name // order. Postgres only; any other driver with a non-empty map is an error so a -// missing RLS stamp fails closed. +// missing RLS stamp fails closed. The SQL lives in lookup.ApplyTxSettings. func ApplyTxSettings(secCtx SecurityContext, tx common.Database, settings map[string]string) error { - if len(settings) == 0 { - return nil - } - if tx == nil { - return fmt.Errorf("tx settings: no transaction") - } - if drv := tx.DriverName(); drv != "postgres" && drv != "pgsql" { - return fmt.Errorf("tx settings: unsupported driver %q", drv) - } - names := make([]string, 0, len(settings)) - for name := range settings { - if !settingNameRE.MatchString(name) { - return fmt.Errorf("tx settings: invalid setting name %q", name) - } - names = append(names, name) - } - sort.Strings(names) - for _, name := range names { - // The value is hex-encoded so it needs no quoting and cannot be read as a - // bind placeholder by any adapter. - query := fmt.Sprintf("SELECT set_config('%s', convert_from(decode('%s', 'hex'), 'UTF8'), true)", - name, hex.EncodeToString([]byte(settings[name]))) - if _, err := tx.Exec(secCtx.GetContext(), query); err != nil { - return fmt.Errorf("tx settings: set %s: %w", name, err) - } - } - return nil + return lookup.ApplyTxSettings(secCtx.GetContext(), tx, settings) } diff --git a/pkg/security/txsettings_test.go b/pkg/security/txsettings_test.go index 1b1bd37..7f24c69 100644 --- a/pkg/security/txsettings_test.go +++ b/pkg/security/txsettings_test.go @@ -3,7 +3,6 @@ package security import ( "context" "errors" - "strings" "testing" "github.com/DATA-DOG/go-sqlmock" @@ -22,76 +21,6 @@ func txSettingsDB(t *testing.T) (common.Database, sqlmock.Sqlmock) { return database.NewPgSQLAdapter(db), mock } -func TestApplyTxSettingsStampsInNameOrderOnTx(t *testing.T) { - pool, mock := txSettingsDB(t) - sc := &mockSecurityContext{ctx: context.Background()} - - mock.ExpectBegin() - mock.ExpectExec(`set_config\('app\.tenant', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) - mock.ExpectExec(`set_config\('app\.user_id', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) - mock.ExpectCommit() - - err := pool.RunInTransaction(context.Background(), func(tx common.Database) error { - return ApplyTxSettings(sc, tx, map[string]string{"app.user_id": "7", "app.tenant": "o'x?"}) - }) - if err != nil { - t.Fatal(err) - } - if err := mock.ExpectationsWereMet(); err != nil { - t.Fatal(err) - } -} - -func TestApplyTxSettingsValueIsNeverInlined(t *testing.T) { - pool, mock := txSettingsDB(t) - sc := &mockSecurityContext{ctx: context.Background()} - var seen string - mock.ExpectBegin() - mock.ExpectExec(`set_config`).WillReturnResult(sqlmock.NewResult(0, 0)) - mock.ExpectCommit() - _ = pool.RunInTransaction(context.Background(), func(tx common.Database) error { - // Capture via a wrapper so the raw statement can be inspected. - return ApplyTxSettings(sc, &queryRecorder{Database: tx, got: &seen}, map[string]string{"app.v": "'; DROP TABLE x; --?"}) - }) - if strings.Contains(seen, "DROP") || strings.Contains(seen, "?") { - t.Fatalf("value leaked into SQL text: %s", seen) - } -} - -type queryRecorder struct { - common.Database - got *string -} - -func (q *queryRecorder) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) { - *q.got = query - return q.Database.Exec(ctx, query, args...) -} - -func TestApplyTxSettingsRejectsBadNameAndDriver(t *testing.T) { - pool, _ := txSettingsDB(t) - sc := &mockSecurityContext{ctx: context.Background()} - - for _, name := range []string{"user_id", "app.x'); DROP", "app..x", "app.x y"} { - if err := ApplyTxSettings(sc, pool, map[string]string{name: "1"}); err == nil { - t.Fatalf("name %q must be rejected", name) - } - } - if err := ApplyTxSettings(sc, pool, nil); err != nil { - t.Fatalf("empty settings must be a no-op: %v", err) - } - if err := ApplyTxSettings(sc, &driverStub{Database: pool, name: "sqlite"}, map[string]string{"app.x": "1"}); err == nil { - t.Fatal("non-postgres driver must fail closed") - } -} - -type driverStub struct { - common.Database - name string -} - -func (d *driverStub) DriverName() string { return d.name } - func TestStampTxSettingsUsesConfiguredFunc(t *testing.T) { pool, mock := txSettingsDB(t) sc := &mockSecurityContext{ctx: context.Background()} diff --git a/pkg/security/types.go b/pkg/security/types.go new file mode 100644 index 0000000..e31bf6d --- /dev/null +++ b/pkg/security/types.go @@ -0,0 +1,54 @@ +package security + +import ( + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" + "github.com/bitechdev/ResolveSpec/pkg/security/sectypes" +) + +// Shared data types live in sectypes so that lookup and the sub packages can use +// them without importing this package. They are aliased here permanently, so +// security.X and sectypes.X are identical types. + +type ( + UserContext = sectypes.UserContext + LoginRequest = sectypes.LoginRequest + RegisterRequest = sectypes.RegisterRequest + LoginResponse = sectypes.LoginResponse + LogoutRequest = sectypes.LogoutRequest + PasswordResetRequest = sectypes.PasswordResetRequest + PasswordResetResponse = sectypes.PasswordResetResponse + PasswordResetCompleteRequest = sectypes.PasswordResetCompleteRequest + KeyType = sectypes.KeyType + UserKey = sectypes.UserKey + CreateKeyRequest = sectypes.CreateKeyRequest + CreateKeyResponse = sectypes.CreateKeyResponse + OAuthServerClient = sectypes.OAuthServerClient + OAuthCode = sectypes.OAuthCode + OAuthTokenInfo = sectypes.OAuthTokenInfo + PasskeyCredential = sectypes.PasskeyCredential + PasskeyRegistrationOptions = sectypes.PasskeyRegistrationOptions + PasskeyAuthenticationOptions = sectypes.PasskeyAuthenticationOptions + PasskeyRelyingParty = sectypes.PasskeyRelyingParty + PasskeyUser = sectypes.PasskeyUser + PasskeyCredentialParam = sectypes.PasskeyCredentialParam + PasskeyCredentialDescriptor = sectypes.PasskeyCredentialDescriptor + PasskeyAuthenticatorSelection = sectypes.PasskeyAuthenticatorSelection + PasskeyRegistrationResponse = sectypes.PasskeyRegistrationResponse + PasskeyAuthenticatorAttestationResponse = sectypes.PasskeyAuthenticatorAttestationResponse + PasskeyAuthenticationResponse = sectypes.PasskeyAuthenticationResponse + PasskeyAuthenticatorAssertionResponse = sectypes.PasskeyAuthenticatorAssertionResponse + TwoFactorSecret = sectypes.TwoFactorSecret + ColumnSecurity = sectypes.ColumnSecurity + RowSecurity = sectypes.RowSecurity +) + +const ( + KeyTypeJWTSecret = sectypes.KeyTypeJWTSecret + KeyTypeHeaderAPI = sectypes.KeyTypeHeaderAPI + KeyTypeOAuth2 = sectypes.KeyTypeOAuth2 + KeyTypeGenericAPI = sectypes.KeyTypeGenericAPI +) + +// errInvalidAPIKey is the single error API-key login returns for unknown, expired, inactive +// and wrong-type keys. +var errInvalidAPIKey = lookup.ErrInvalidAPIKey