mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7662d5055c | ||
|
|
ea6a2e705f | ||
|
|
2516fcb13d | ||
|
|
c9fa8c60f2 | ||
|
|
60bd0a6dd3 | ||
|
|
982c90bfdd |
@@ -646,7 +646,7 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md).
|
|||||||
|
|
||||||
#### Security
|
#### 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).
|
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Direct Mode" for the SQLite/portable-SQL path).
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,341 @@
|
|||||||
|
# pkg/security lookup sub package plan
|
||||||
|
|
||||||
|
Status: plan only, no code changed. Related: `audit/mcp_plan.md` (work item 1, API key login).
|
||||||
|
|
||||||
|
## Problem
|
||||||
|
|
||||||
|
- `pkg/security` mixes two data-access styles: stored procedures (`resolvespec_*`) and ~1,800 lines of hand-written
|
||||||
|
"direct" SQL (`*_direct.go`) selected per call by `QueryMode` / `ShouldUseProcedure`.
|
||||||
|
- Direct SQL is written once with `?` placeholders and only rewritten for Postgres. It assumes one fixed schema
|
||||||
|
(table and column names, JSON stored as TEXT, bool/time handling), and is only really exercised on SQLite.
|
||||||
|
- Table names are configurable (`TableNames`), column names are not. Procedure names are configurable (`SQLNames`).
|
||||||
|
- Postgres cannot be run "tables only" in a first-class way, and other databases have no defined support.
|
||||||
|
- Rule going forward: **`pkg/security` itself contains no SQL.** All lookups go through one sub package.
|
||||||
|
|
||||||
|
## Goal
|
||||||
|
|
||||||
|
New sub package `pkg/security/lookup` that owns every database read/write the security package needs.
|
||||||
|
|
||||||
|
| Requirement | Decision |
|
||||||
|
|---|---|
|
||||||
|
| Postgres default | Existing stored procedures, existing names, unchanged behaviour out of the box |
|
||||||
|
| Postgres direct | Optional: work on tables directly with no procs installed |
|
||||||
|
| SQLite | First-class: configurable tables and columns |
|
||||||
|
| Other DBs | MySQL/MariaDB and MSSQL via dialects; adding more = adding a dialect |
|
||||||
|
| Config | Procedure names, table names, column names, mode (procedure / direct / auto), per backend |
|
||||||
|
| `pkg/security` | Calls lookup interfaces only; no `SELECT`/`INSERT`/`UPDATE`/`DELETE`, no `pg_proc` probing |
|
||||||
|
|
||||||
|
## Design
|
||||||
|
|
||||||
|
### Package layout
|
||||||
|
|
||||||
|
```
|
||||||
|
pkg/security/ # core: behaviour interfaces, SecurityList, middleware, chain, composite, hooks, write security, tx settings, type aliases
|
||||||
|
pkg/security/sectypes/ # shared data types (no deps, no SQL, no logic beyond small helpers)
|
||||||
|
pkg/security/providers/ # concrete authenticators + security providers (see Package split)
|
||||||
|
pkg/security/oauth/ # OAuth2 client login + OAuth2 authorization server
|
||||||
|
pkg/security/totp/ # two-factor: generator, providers, TwoFactorAuthenticator
|
||||||
|
pkg/security/passkey/ # WebAuthn passkey provider + passkey login flow
|
||||||
|
pkg/security/lookup/
|
||||||
|
lookup.go # store interfaces + record types + Config + New(db, cfg)
|
||||||
|
schema.go # Schema: table + column names per entity, defaults, merge, validate
|
||||||
|
mode.go # Mode (Auto/Procedure/Direct), per-operation resolution, proc probe (pg only)
|
||||||
|
dialect/ # Dialect interface + postgres, sqlite, mysql, mssql
|
||||||
|
procedure/ # procedure backend (current SQLNames, p_success/p_error/p_data contract)
|
||||||
|
direct/ # dialect-driven SQL backend (no fixed SQL strings per dialect)
|
||||||
|
ddl/ # reference schemas per dialect (replaces database_schema*.sql variants)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Shared types: `pkg/security/sectypes`
|
||||||
|
|
||||||
|
`lookup` cannot import `pkg/security` (cycle: security -> lookup -> security), so the plain data types move to a
|
||||||
|
dependency-free sub package that both import.
|
||||||
|
|
||||||
|
- Moves to `sectypes`: `UserContext`, `LoginRequest`, `LoginResponse`, `RegisterRequest`, `LogoutRequest`,
|
||||||
|
`PasswordResetRequest/Response/CompleteRequest`, `KeyType`, `UserKey`, `CreateKeyRequest/Response`,
|
||||||
|
`PasskeyCredential` (+ passkey request/option structs the stores return), `TwoFactorSecret`, OAuth server
|
||||||
|
client/code/token-info structs (`OAuthServerClient`, `OAuthCode`, `OAuthTokenInfo`), `ColumnSecurity`, `RowSecurity`.
|
||||||
|
- Stays in `pkg/security`: all behaviour interfaces (`Authenticator`, `SecurityProvider`, `Registrable`, ...),
|
||||||
|
authenticators, middleware, `OAuthServer`, TOTP generator, hooks. They reference `sectypes` types.
|
||||||
|
- Compatibility: `pkg/security` re-exports each moved type as an alias (`type UserContext = sectypes.UserContext`) and
|
||||||
|
the `KeyType*` constants, so `security.UserContext` etc. keep compiling and are identical types. In-repo users
|
||||||
|
(eventbroker, funcspec, mqttspec, resolvemcp, resolvespec, restheadspec, websocketspec, ~17 files) need no edits.
|
||||||
|
- `lookup` stores take and return `sectypes` types directly; no separate record types and no conversion layer.
|
||||||
|
- Rules: `sectypes` imports only the standard library (and `oauth2` types only if unavoidable, else a local struct);
|
||||||
|
JSON tags unchanged so wire formats and procedure `p_data` payloads stay identical.
|
||||||
|
|
||||||
|
### Package split
|
||||||
|
|
||||||
|
Dependency direction (no cycles): `sectypes` <- `lookup` <- `providers`/`oauth`/`totp`/`passkey` -> `security` (core) -> `sectypes`.
|
||||||
|
Core `security` never imports the sub packages. Sub packages do not import each other (see Dependency rules).
|
||||||
|
|
||||||
|
### Dependency rules (how cycles are avoided)
|
||||||
|
|
||||||
|
1. **Layers, imports only point down.** L0 `sectypes` (stdlib only) -> L1 `lookup`, core `security` interfaces ->
|
||||||
|
L2 `providers`, `oauth`, `totp`, `passkey`. A package may import lower layers, never its own layer or above.
|
||||||
|
2. **Define interfaces where they are consumed, not where they are implemented** (Go idiom). E.g. `oauth` declares
|
||||||
|
the small `SessionCreator` it needs; `providers.DatabaseAuthenticator` satisfies it without `oauth` importing `providers`.
|
||||||
|
3. **Shared data goes down, not sideways.** If two L2 packages need the same struct, it moves to `sectypes`
|
||||||
|
(or a tiny `internal/` package), never "A imports B for one type".
|
||||||
|
4. **No L2 -> L2 imports.** Composition happens in the application (or an optional top-level `security/setup`
|
||||||
|
package that imports everything and is imported by nobody in `pkg/security`).
|
||||||
|
5. **Dependency injection by constructor**, passing interfaces/stores (`lookup.Provider`, `security.Authenticator`);
|
||||||
|
no package-level registries that need a back-import; use functional options for optional collaborators.
|
||||||
|
6. **Core never imports concrete implementations**; where core needs behaviour it calls an interface it owns
|
||||||
|
(hooks, `SecurityContext`, `Authenticator`).
|
||||||
|
7. **Tests:** external test packages (`package foo_test`) for cross-package integration tests, so test-only
|
||||||
|
imports cannot create cycles; shared fixtures in an `internal/testutil` package.
|
||||||
|
8. **Guard in CI:** `go list -deps` / a small test that asserts the layer rules (e.g. `sectypes` imports only stdlib,
|
||||||
|
`lookup` does not import `security`, no L2 package imports another L2 package). `go build` already rejects true cycles.
|
||||||
|
|
||||||
|
| Package | Contents (from today's files) |
|
||||||
|
|---|---|
|
||||||
|
| `security` (core) | `Authenticator`, `SecurityProvider`, `Registrable`, `Refreshable`, `APIKeyLoginable`, ... interfaces; `SecurityList`; `SecurityContext`; middleware + cookie options; `ChainAuthenticator`; `CompositeSecurityProvider`; hooks; `WriteDataContext`; `TxSettings`; type aliases to `sectypes` |
|
||||||
|
| `providers` | `DatabaseAuthenticator`, `JWTAuthenticator`, `HeaderAuthenticator`, `KeyStoreAuthenticator`, `ConfigKeyStore`, `DatabaseKeyStore`, `DatabaseColumnSecurityProvider`, `DatabaseRowSecurityProvider`, `Config*SecurityProvider` |
|
||||||
|
| `oauth` | `OAuth2Config`, `OAuth2Provider`, Google/GitHub/Microsoft/Facebook/multi-provider constructors, OAuth2 refresh, `OAuthServer` + `OAuthServerConfig`, oauth server persistence (via `lookup.OAuthClientStore`) |
|
||||||
|
| `passkey` | `PasskeyProvider` impl (`DatabasePasskeyProvider`), registration/authentication flows, passkey request/option types that are not shared (shared ones stay in `sectypes`) |
|
||||||
|
| `totp` | `TwoFactorAuthProvider`, `TwoFactorConfig`, `TOTPGenerator`, `MemoryTwoFactorProvider`, `DatabaseTwoFactorProvider`, `TwoFactorAuthenticator` |
|
||||||
|
|
||||||
|
Consequences to design for:
|
||||||
|
- **Methods cannot span packages.** Today OAuth2 and passkey logic are methods on `DatabaseAuthenticator`
|
||||||
|
(`oauth2_methods*.go`, `oauth_server_db*.go`, passkey methods) and `NewOAuthServer` takes `*DatabaseAuthenticator`.
|
||||||
|
They become standalone types in `oauth` / `providers` that depend on `lookup` stores and on small interfaces
|
||||||
|
(e.g. `oauth.SessionCreator`) instead of the concrete authenticator. `NewGoogleAuthenticator(...)` etc. return a
|
||||||
|
`providers.DatabaseAuthenticator` configured with an `oauth.Provider`, or an `oauth.Authenticator` that implements
|
||||||
|
`security.Authenticator`; pick one in step 5 (see Open).
|
||||||
|
- **Constructors cannot be re-exported from core `security`** (it would import the sub packages = cycle). Types that
|
||||||
|
move to `sectypes` keep aliases; constructors and concrete types do not. This is a breaking import change.
|
||||||
|
In-repo callers affected (outside `pkg/security`): `pkg/resolvemcp` (`oauth2.go`, `oauth2_server.go`, `handler.go`),
|
||||||
|
`pkg/middleware/clientqueue.go`, docs and examples. Provide a mechanical migration table
|
||||||
|
(`security.NewDatabaseAuthenticator` -> `providers.NewDatabaseAuthenticator`, `security.OAuthServer` -> `oauth.Server`, ...).
|
||||||
|
- **Interfaces core needs from sub packages** (e.g. 2FA hook points) are defined in core or `sectypes`, implemented in
|
||||||
|
`totp`; core never imports `totp`.
|
||||||
|
- `examples*.go` / `oauth2_examples.go` / `passkey_examples.go` move next to the package they exemplify (or to
|
||||||
|
`_example_test.go` files) so core has no dependency on them.
|
||||||
|
- Tests move with their code; shared helpers (sqlite test DB, `authenticatedRequest`) go to an internal test helper package.
|
||||||
|
|
||||||
|
### Store interfaces (one per domain, mirrors current procs)
|
||||||
|
|
||||||
|
| Store | Operations (current proc in brackets) |
|
||||||
|
|---|---|
|
||||||
|
| `AuthStore` | `Login` [login], `Register` [register], `Logout` [logout], `Session` [session], `TouchSession` [session_update], `Refresh` [refresh_token], `LoginAPIKey` [login_api_key], `JWTLogin`, `JWTLogout`, `ResetRequest`, `ResetComplete` |
|
||||||
|
| `KeyStore` | `Create`, `List`, `Delete`, `Validate` [keystore_*] |
|
||||||
|
| `OAuthClientStore` | register client, get client, save code, exchange code, introspect, revoke |
|
||||||
|
| `OAuthUserStore` | get-or-create user, create session, get/update refresh token, get user |
|
||||||
|
| `PasskeyStore` | store, get, update counter, list, delete, rename, get by username, login |
|
||||||
|
| `TOTPStore` | enable, disable, status, secret, regenerate backup codes, validate backup code |
|
||||||
|
| `PolicyStore` | column security, row security: procedure backend (default) + direct backend over the `sec_*` table layout below |
|
||||||
|
|
||||||
|
Each store has a procedure implementation and a direct implementation. A `Provider` bundles them; `security`
|
||||||
|
constructors take a `lookup.Provider` (or build one from `db` + `lookup.Config`, so existing constructors keep working).
|
||||||
|
|
||||||
|
### Config
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Config struct {
|
||||||
|
Dialect string // "postgres" | "sqlite" | "mysql" | "mssql"; empty = detect from driver
|
||||||
|
Mode Mode // Auto | Procedure | Direct; default: Procedure for postgres, Direct otherwise
|
||||||
|
Overrides map[Op]Mode // optional per-operation mode, e.g. direct for Session, procedure for Login
|
||||||
|
Procs ProcNames // = today's SQLNames (+ LoginAPIKey), defaults unchanged
|
||||||
|
Schema Schema // tables + columns, see below
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
- `Schema` = per entity `{Table string; Columns map[Column]string}` with typed column keys covering every column
|
||||||
|
(`users.id`, `users.username`, `users.password`, `users.is_active`, ...). Defaults reproduce the current schema,
|
||||||
|
so zero config behaves as today. Optional `Schema` name per entity for `schema.table` qualification.
|
||||||
|
- Merge + validation as today: non-empty override wins; every identifier checked against `^[a-zA-Z_][a-zA-Z0-9_]*$`
|
||||||
|
(plus optional single `schema.` prefix). Identifiers are quoted by the dialect, never interpolated raw.
|
||||||
|
- No back-compat for config: `SQLNames`, `KeyStoreSQLNames`, `TableNames`, `KeyStoreTableNames` and `QueryMode` are
|
||||||
|
removed from `pkg/security`; `lookup.Config` replaces them (decision 8).
|
||||||
|
|
||||||
|
### Dialect interface (the per-database adaptor)
|
||||||
|
|
||||||
|
One adaptor per database type (`dialect/postgres`, `sqlite`, `mysql`, `mssql`), selected by `Config.Dialect` or
|
||||||
|
detected from the driver, registered through `dialect.Register(name, factory)` so more databases can be added later
|
||||||
|
without touching core code. Each adaptor supplies only the things that differ; the direct backend builds queries from it.
|
||||||
|
|
||||||
|
| Concern | Dialect method |
|
||||||
|
|---|---|
|
||||||
|
| Placeholders | `Placeholder(n)` (`$n`, `?`, `@pn`) |
|
||||||
|
| Identifier quoting | `Quote(ident)` (`"x"`, `` `x` ``, `[x]`) |
|
||||||
|
| Booleans | `Bool(v)` / scan helper (bool vs 0/1) |
|
||||||
|
| Time | `Now()` expr or Go-side `time.Now()`; scan helper for drivers returning strings |
|
||||||
|
| Insert returning id | `InsertReturningID(table, cols, idCol)` returns the SQL + scan strategy: postgres `... RETURNING id` (QueryRow), sqlite/mysql `LastInsertId`, mssql `... OUTPUT INSERTED.id` (QueryRow). The only dialect-specific write construct (decision 12 / confirmed) |
|
||||||
|
| Get-or-create | none; standard SQL select-then-insert inside a tx (no upsert) |
|
||||||
|
| Limit/top | only if a query needs it |
|
||||||
|
| JSON columns (scopes/meta/roles) | `EncodeJSON` / `DecodeJSON` (native jsonb vs TEXT) |
|
||||||
|
| Random / hashing | done in Go (token generation, SHA-256 key hash, bcrypt) so direct mode needs no `pgcrypto` and no DB functions |
|
||||||
|
| Driver detection | `Detect(*sql.DB)` from driver type (replaces `driverIsPostgres` / `driverIsPortableOnly`) |
|
||||||
|
|
||||||
|
Queries are assembled by a small internal builder (select/insert/update/delete with named columns from `Schema`),
|
||||||
|
not string-concatenated per dialect and not via an ORM, to keep `pkg/security` free of bun/gorm.
|
||||||
|
|
||||||
|
### PolicyStore table layout (column / row security, direct backend)
|
||||||
|
|
||||||
|
Approved layout. Both the procedure backend (`resolvespec_column_security` / `resolvespec_row_security`, rewritten
|
||||||
|
in `database_schema.sql`) and the direct backend read these tables; the former external schema is no longer
|
||||||
|
referenced anywhere in the repo. All table and column names are configurable via `Schema`, defaults shown.
|
||||||
|
|
||||||
|
`sec_group_members` (optional; omit to use direct user rules only)
|
||||||
|
|
||||||
|
| Column | Type | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `group_id` | int, not null | group a user belongs to |
|
||||||
|
| `user_id` | int, not null | FK users.id; PK (`group_id`, `user_id`) |
|
||||||
|
|
||||||
|
`sec_column_rules`
|
||||||
|
|
||||||
|
| Column | Type | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `id` | int PK | |
|
||||||
|
| `user_id` | int null | rule for one user |
|
||||||
|
| `group_id` | int null | rule for every member of the group; exactly one of `user_id` / `group_id` set (check constraint) |
|
||||||
|
| `schema_name` | text, not null | matched case-insensitively |
|
||||||
|
| `table_name` | text, not null | matched case-insensitively |
|
||||||
|
| `column_path` | text, not null | dot path under the table (`col` or `col.sub.field`) = `ColumnSecurity.Path` joined by `.` |
|
||||||
|
| `access_type` | text, not null | `ColumnSecurity.Accesstype` (e.g. `mask`, `hide`, `read`) |
|
||||||
|
| `mask_start`, `mask_end` | int null | default 0 |
|
||||||
|
| `mask_invert` | bool null | default false |
|
||||||
|
| `mask_char` | text null | default `*` |
|
||||||
|
| `extra_filters` | text/JSON null | `ExtraFilters` map, JSON-encoded via dialect `EncodeJSON` |
|
||||||
|
| `is_active` | bool, not null | default true |
|
||||||
|
|
||||||
|
`sec_row_rules`
|
||||||
|
|
||||||
|
| Column | Type | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `id` | int PK | |
|
||||||
|
| `user_id` | int null / `group_id` int null | as above, exactly one set |
|
||||||
|
| `schema_name`, `table_name` | text, not null | case-insensitive match |
|
||||||
|
| `template` | text null | SQL fragment with the existing placeholders (`RowSecurity.Template`) |
|
||||||
|
| `has_block` | bool, not null | default false; true = no rows visible (`RowSecurity.HasBlock`) |
|
||||||
|
| `is_active` | bool, not null | default true |
|
||||||
|
|
||||||
|
Resolution rules (direct backend; also the contract the conformance tests assert):
|
||||||
|
- Applicable rules = active rules where `user_id` = caller, plus rules of every group the caller belongs to.
|
||||||
|
- Column security: all applicable rules for the exact schema + table, returned as `[]ColumnSecurity` (union).
|
||||||
|
Exact table match, not a prefix match (a prefix would match `users_archive` for `users`).
|
||||||
|
- Row security: any applicable `has_block` wins; otherwise templates of all applicable rules are combined with
|
||||||
|
`AND` (each wrapped in parentheses); no rule = `RowSecurity{}` with `ErrNoRowSecurity` semantics unchanged.
|
||||||
|
- Templates are still validated/substituted by the existing safe-identifier code in core; the store only loads text.
|
||||||
|
- Both loaders keep the guarantees of today: user reference reduced to a scalar, failures are errors (fail closed),
|
||||||
|
no rule is "no rules".
|
||||||
|
|
||||||
|
### Mode resolution
|
||||||
|
|
||||||
|
- `Procedure`: always call the proc; missing proc = error (no silent fallback).
|
||||||
|
- `Direct`: always use tables via the dialect builder.
|
||||||
|
- `Auto`: Postgres probes `pg_proc` once per proc (cached, as today); other dialects resolve to `Direct`.
|
||||||
|
Probe lives in `lookup` and is the only place that queries the catalog.
|
||||||
|
- Postgres default stays `Procedure` so current installs do not change behaviour.
|
||||||
|
- Roles/user-level safety rules already enforced in direct mode (Register ignores client-supplied level/roles, bcrypt
|
||||||
|
hash, opt-in password upgrade) become backend-agnostic tests that both backends must pass.
|
||||||
|
|
||||||
|
## Work items
|
||||||
|
|
||||||
|
### 0. Extract shared types (prerequisite, no behaviour change)
|
||||||
|
- Create `pkg/security/sectypes`, move the types listed above, add aliases in `pkg/security`.
|
||||||
|
- Verify `go build ./...` and `go test ./pkg/...` unchanged; `go vet` for alias/import cycles.
|
||||||
|
- Do this first and on its own so the diff is a pure move.
|
||||||
|
|
||||||
|
### 0b. Package split (after 0, before `lookup` wiring)
|
||||||
|
- Create `providers`, `oauth`, `totp`; move files per the Package split table, one package per commit:
|
||||||
|
`totp` and `passkey` (self-contained) -> `providers` (key stores, authenticators, policy providers) -> `oauth` (needs de-methoding
|
||||||
|
from `DatabaseAuthenticator`).
|
||||||
|
- Break the method-on-`DatabaseAuthenticator` coupling for OAuth2 and passkey first (extract interfaces), then move.
|
||||||
|
- Update in-repo callers and docs; add migration table to `pkg/security/README.md`.
|
||||||
|
- Behaviour unchanged; at this point stores are still the old direct/proc code, only relocated.
|
||||||
|
|
||||||
|
### 1. Skeleton and contracts
|
||||||
|
- Create `lookup` package: records, store interfaces, `Config`, `Schema` (defaults, merge, validate), `Mode`.
|
||||||
|
- No behaviour change yet; compile-only.
|
||||||
|
|
||||||
|
### 2. Dialects
|
||||||
|
- Implement `postgres`, `sqlite`, `mysql`, `mssql` against the dialect interface; driver detection.
|
||||||
|
- Unit tests per dialect: placeholders, quoting, bool/time round trip, insert-returning-id.
|
||||||
|
|
||||||
|
### 3. Procedure backend
|
||||||
|
- Move existing proc calls out of `pkg/security` into `lookup/procedure` using `ProcNames` (current defaults).
|
||||||
|
- Keep the `p_success, p_error, p_data` contracts and reconnect-on-closed-DB helper.
|
||||||
|
- Include `resolvespec_login_api_key` (added in mcp_plan item 1) with the generic error behaviour.
|
||||||
|
|
||||||
|
### 4. Direct backend
|
||||||
|
- Port each `*_direct.go` to `lookup/direct` using `Schema` + dialect builder, one store at a time:
|
||||||
|
AuthStore -> KeyStore -> OAuth stores -> Passkey -> TOTP -> PolicyStore (column/row security tables).
|
||||||
|
- Add direct `LoginAPIKey` here (select by key hash, active, unexpired, key type in header_api/api, user active;
|
||||||
|
one generic error), since SQL is now allowed only inside `lookup`.
|
||||||
|
- Transactions: multi-step writes (login = session insert + last_login; register; reset complete) run in one tx.
|
||||||
|
|
||||||
|
### 5. Wire `pkg/security`
|
||||||
|
- Constructors accept `lookup.Provider` / `lookup.Config`; old options map onto it (deprecated).
|
||||||
|
- Replace every `*_direct.go` call and `ShouldUseProcedure` branch with a store call.
|
||||||
|
- Delete `*_direct.go`, `query_mode.go` probe/placeholder code, direct `TableNames` use; keep only aliases.
|
||||||
|
- Check no non-test code in `pkg/security` contains SQL keywords (CI grep guard).
|
||||||
|
|
||||||
|
### 5b. `pkg/security/breaking_changes.md`
|
||||||
|
- Create at step 0 and append as each step lands: moved types (aliased, no action), moved constructors/types
|
||||||
|
with rename table (old -> new import path and symbol), removed config types (`SQLNames`, `TableNames`,
|
||||||
|
`KeyStore*Names`, `QueryMode`) with the `lookup.Config` replacement, removed `database_schema_sqlite.sql`,
|
||||||
|
`ModeAuto` behaviour change, API key login procedure.
|
||||||
|
|
||||||
|
### 6. Schemas and docs
|
||||||
|
- `lookup/ddl`: reference DDL for postgres (tables only, no procs), sqlite, mysql, mssql; existing proc scripts
|
||||||
|
stay beside the procedure backend. Replace `database_schema_sqlite.sql`.
|
||||||
|
- Document: default (procs), Postgres tables-only, SQLite, custom column mapping, adding a dialect.
|
||||||
|
- Update `pkg/security` README/QUICK_REFERENCE; note API key procedure in security docs (mcp_plan item 10).
|
||||||
|
|
||||||
|
### 7. Tests
|
||||||
|
- Shared conformance suite run against every backend/dialect: login/register/logout/session/refresh, reset,
|
||||||
|
API keys (valid/expired/inactive/unknown/wrong type), keystore, OAuth server, passkey, TOTP, privilege rules.
|
||||||
|
- Backends covered: sqlite (in-memory, direct), postgres direct and postgres procedure (needs a Postgres instance,
|
||||||
|
skipped without `RESOLVESPEC_TEST_PG_DSN`), mysql/mssql behind env DSNs. Dialect unit tests need no DB.
|
||||||
|
- Procedure backend unit tests with sqlmock for the call/contract shape.
|
||||||
|
- Check for existing test data before creating any; ask before generating.
|
||||||
|
- Run with `-race`; migrate current `direct_mode_test.go` / `query_mode_test.go` cases into the suite.
|
||||||
|
|
||||||
|
## Order
|
||||||
|
|
||||||
|
0. Extract `sectypes` types + aliases (0)
|
||||||
|
0b. Package split: `totp`, `passkey` -> `providers` -> `oauth` (0b), callers + docs updated
|
||||||
|
1. Skeleton + Schema/Config (1)
|
||||||
|
2. Dialects (2)
|
||||||
|
3. Procedure backend extraction, `pkg/security` wired to it for procs only (3, part of 5) - zero behaviour change
|
||||||
|
4. Direct backend per store, then remove old `*_direct.go` as each store lands (4, 5)
|
||||||
|
5. DDL + docs + conformance suite (6, 7)
|
||||||
|
6. Resume `audit/mcp_plan.md` step 2 (guard) on top of `lookup`
|
||||||
|
|
||||||
|
## Breaking changes
|
||||||
|
|
||||||
|
- `QueryMode`, `SQLNames`, `TableNames`, `KeyStoreSQLNames`, `KeyStoreTableNames` removed; replaced by `lookup.Config`.
|
||||||
|
- `ModeAuto` no longer silently falls back from procedure to SQL on non-Postgres drivers without telling: resolution
|
||||||
|
is explicit and logged once per op.
|
||||||
|
- `database_schema_sqlite.sql` replaced by `lookup/ddl`.
|
||||||
|
- Concrete types and constructors move to `providers`, `oauth`, `totp` (e.g. `security.NewDatabaseAuthenticator`
|
||||||
|
-> `providers.NewDatabaseAuthenticator`, `security.NewOAuthServer` -> `oauth.NewServer`,
|
||||||
|
`security.NewTOTPGenerator` -> `totp.NewGenerator`). No aliases possible (import cycle); import paths must change.
|
||||||
|
- Type identity is preserved via aliases; code that used reflection on the package path of these types
|
||||||
|
(`security.UserContext` -> `sectypes.UserContext`) would see the new path (none found in-repo; re-check).
|
||||||
|
- Anything outside `lookup` that relied on SQL living in `pkg/security` (none found in-repo) must use the stores.
|
||||||
|
|
||||||
|
## Decisions
|
||||||
|
|
||||||
|
| # | Topic | Decision |
|
||||||
|
|---|---|---|
|
||||||
|
| 1 | Shared package name | `sectypes` |
|
||||||
|
| 2 | Type aliases in `security` | Kept permanently (public API) |
|
||||||
|
| 3 | `ColumnSecurity` / `RowSecurity` | Move to `sectypes` (types and their helper logic that has no outside deps) |
|
||||||
|
| 4 | Moved type names | Renamed for new paths (`oauth.Server`, `totp.Generator`, ...) |
|
||||||
|
| 5 | OAuth2 client login | Dedicated `oauth.Authenticator` type; `New{Google,GitHub,Microsoft,Facebook}Authenticator` return it |
|
||||||
|
| 6 | Migration | Rename table only; recorded in new `pkg/security/breaking_changes.md` |
|
||||||
|
| 7 | Dialects v1 | postgres, sqlite, mysql/mariadb, mssql; more later via the dialect interface |
|
||||||
|
| 8 | Deprecated config (`SQLNames`, `TableNames`, `KeyStore*Names`, `QueryMode`) | Removed, no aliases; recorded in `breaking_changes.md` |
|
||||||
|
| 9 | Column mapping | Every column of every entity configurable |
|
||||||
|
| 10 | DB handle | `*sql.DB` |
|
||||||
|
| 11 | Column / row security | Procedures (default) **and** a table layout for direct mode |
|
||||||
|
| 12 | Postgres direct SQL | Standard SQL only: no `ON CONFLICT` / `MERGE`; get-or-create = select then insert in a tx |
|
||||||
|
| 13 | Token format | Keep `sess_<hex>_<unix>`, generated in Go for direct mode |
|
||||||
|
|
||||||
|
## Open
|
||||||
|
|
||||||
|
- None.
|
||||||
+16
-6
@@ -10,7 +10,7 @@
|
|||||||
- Each un-transacted call takes its own pool connection → bursts with a small pool (see `dbtrace`).
|
- Each un-transacted call takes its own pool connection → bursts with a small pool (see `dbtrace`).
|
||||||
- Already fixed: read/create hooks in `resolvespec` + `restheadspec` (commit `47708fc`, tag >= v1.1.28). Consumers on older tags still show the bug.
|
- Already fixed: read/create hooks in `resolvespec` + `restheadspec` (commit `47708fc`, tag >= v1.1.28). Consumers on older tags still show the bug.
|
||||||
|
|
||||||
## Current state (verified by reading code; not yet by `dbtrace`)
|
## Current state — BEFORE this work (historical baseline; everything below is now fixed, see Progress and Status)
|
||||||
| Spec | Read | Create | Update | Delete |
|
| Spec | Read | Create | Update | Delete |
|
||||||
|---|---|---|---|---|
|
|---|---|---|---|---|
|
||||||
| restheadspec | tx; `AfterRead` post-commit on pool | tx; `AfterCreate` post-commit on pool | tx; re-fetch + `BeforeScan` post-commit on pool (`:1667-1674`) | **single: no tx, hook + select + delete on pool (`:1945-1994`)**; batch: tx, per-item `BeforeDelete` inside |
|
| restheadspec | tx; `AfterRead` post-commit on pool | tx; `AfterCreate` post-commit on pool | tx; re-fetch + `BeforeScan` post-commit on pool (`:1667-1674`) | **single: no tx, hook + select + delete on pool (`:1945-1994`)**; batch: tx, per-item `BeforeDelete` inside |
|
||||||
@@ -65,8 +65,18 @@
|
|||||||
- Update re-fetch is a plain SELECT in that second tx. No `RETURNING`.
|
- Update re-fetch is a plain SELECT in that second tx. No `RETURNING`.
|
||||||
- `OnTxBegin` failure aborts the whole request, rolls back, returns an error with no detail leaked to the client.
|
- `OnTxBegin` failure aborts the whole request, rolls back, returns an error with no detail leaked to the client.
|
||||||
|
|
||||||
## Open
|
## Status summary
|
||||||
- Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking.
|
**Done**
|
||||||
|
- P0-P7 all DONE (baseline, delete in one tx, `OnTxBegin` + `runInTx`, second short tx, websocketspec/mqttspec, resolvemcp, funcspec, security stamping).
|
||||||
|
- Regression tests in all six specs, plus source guard `pkg/common/tx_guard_test.go`.
|
||||||
|
- `AfterRead` decided and done (restheadspec: second short tx; websocketspec/mqttspec: inside the read tx).
|
||||||
|
- websocketspec `BeforeDisconnect`/`AfterDisconnect` wired (see Progress); `unwiredHooks` allowlist is now empty.
|
||||||
|
|
||||||
|
**Not done**
|
||||||
|
- Real-Postgres `dbtrace` measurement for websocketspec, mqttspec, resolvemcp, restheadspec, funcspec (only resolvespec measured: `pooled=0` on every op). "Done when" bullet 1 is proven for resolvespec only.
|
||||||
|
- resolvespec batch delete: per-item `BeforeDelete` (one hook per request today). Deferred on purpose: behavior change.
|
||||||
|
- Confirm the consumer's ResolveSpec version is >= v1.1.28 (read/create already in tx). Not blocking; needs the consumer.
|
||||||
|
- Known, pre-existing, not ours: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (use `-count=1`); mqttspec integration tests need a DB.
|
||||||
|
|
||||||
## Phases
|
## Phases
|
||||||
| # | Status | Change | Files | Notes |
|
| # | Status | Change | Files | Notes |
|
||||||
@@ -88,7 +98,7 @@
|
|||||||
- DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7).
|
- DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7).
|
||||||
- DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates).
|
- DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates).
|
||||||
- NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there.
|
- NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there.
|
||||||
- OPEN: restheadspec `AfterRead` still runs post-commit with `Tx = h.db` (`:1004`); decision says read has no second tx. Needs a call: run it inside the read tx, or in a short second tx.
|
- RESOLVED: restheadspec `AfterRead` question (see DONE AfterRead below).
|
||||||
- DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`.
|
- DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`.
|
||||||
- DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx").
|
- DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx").
|
||||||
- DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock).
|
- DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock).
|
||||||
@@ -103,7 +113,7 @@
|
|||||||
|
|
||||||
## Tests
|
## Tests
|
||||||
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
|
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
|
||||||
- Done: delete tx tests (`pkg/*/delete_tx_test.go`, sqlmock, 1-conn pool detects pool use). Missing: same for read/create/update, `OnTxBegin`, other specs.
|
- Done: delete tx tests (`pkg/*/delete_tx_test.go`, sqlmock, 1-conn pool detects pool use), and read/create/update/`OnTxBegin`/other-spec tests (see "DONE regression tests" in Progress). Still missing: `dbtrace` `pooled == 0` on real Postgres for all specs except resolvespec.
|
||||||
- Add per spec/op: hook `Tx` is not the pool; `OnTxBegin` fires once per tx, before other hooks; single-ID delete = 1 tx; `dbtrace` `pooled == 0` on the request path.
|
- Add per spec/op: hook `Tx` is not the pool; `OnTxBegin` fires once per tx, before other hooks; single-ID delete = 1 tx; `dbtrace` `pooled == 0` on the request path.
|
||||||
- Test data: reuse `pkg/testmodels`; **ask before generating new data** (per project rule).
|
- Test data: reuse `pkg/testmodels`; **ask before generating new data** (per project rule).
|
||||||
- Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB).
|
- Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB).
|
||||||
@@ -118,5 +128,5 @@
|
|||||||
- `dbtrace` shows `pooled=0` for every handler op on a hooked model.
|
- `dbtrace` shows `pooled=0` for every handler op on a hooked model.
|
||||||
- RLS GUC set in `OnTxBegin` is visible to read, create, update, delete queries and hooks.
|
- RLS GUC set in `OnTxBegin` is visible to read, create, update, delete queries and hooks.
|
||||||
- No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers.
|
- No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers.
|
||||||
- OPEN: websocketspec `BeforeDisconnect`/`AfterDisconnect` are defined but never executed (connection lifecycle, not DB). Allowlisted in `TestEveryDefinedHookHasACallSite`; wire them to remove the entry.
|
- DONE: websocketspec `BeforeDisconnect`/`AfterDisconnect` fire from `Connection.Close()` (single close path, `closedOnce`, so exactly once per registered connection however it closes: read error, write error, slow-consumer eviction, shutdown). Connection lifecycle, not DB: `Tx` is not set. The hook context is detached from the connection cancel (`context.WithoutCancel`) so `AfterDisconnect` still has a live context. Errors are logged and never block the close. A connection rejected by `BeforeConnect` gets no disconnect hooks. `ConnectionManager.Shutdown` now closes connections outside its lock (a hook calling `Count()` would have deadlocked). Allowlist in `TestEveryDefinedHookHasACallSite` is now empty. Tests: `pkg/websocketspec/connection_test.go`. mqttspec already fired these in `Handler.Shutdown`; its per-client disconnect is unchanged.
|
||||||
- DONE: column-level hide/mask columns are dropped from create/update payloads (`security.ApplyWriteColumnSecurity`); rules preloaded in `BeforeHandle` for create/update. resolvemcp update now runs `BeforeHandle`.
|
- DONE: column-level hide/mask columns are dropped from create/update payloads (`security.ApplyWriteColumnSecurity`); rules preloaded in `BeforeHandle` for create/update. resolvemcp update now runs `BeforeHandle`.
|
||||||
|
|||||||
@@ -275,6 +275,14 @@ func (b *BunAdapter) GetUnderlyingDB() interface{} {
|
|||||||
return b.getDB()
|
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 {
|
func (b *BunAdapter) DriverName() string {
|
||||||
// Normalize Bun's dialect name to match the project's canonical vocabulary.
|
// Normalize Bun's dialect name to match the project's canonical vocabulary.
|
||||||
// Bun returns "pg" for PostgreSQL; the rest of the project uses "postgres".
|
// Bun returns "pg" for PostgreSQL; the rest of the project uses "postgres".
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package database
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -227,6 +228,20 @@ func (g *GormAdapter) GetUnderlyingDB() interface{} {
|
|||||||
return g.getDB()
|
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 {
|
func (g *GormAdapter) DriverName() string {
|
||||||
return normalizeGormDriverName(g.getDB())
|
return normalizeGormDriverName(g.getDB())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -223,6 +223,11 @@ func (p *PgSQLAdapter) GetUnderlyingDB() interface{} {
|
|||||||
return p.db
|
return p.db
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SQLDB implements common.SQLDBProvider.
|
||||||
|
func (p *PgSQLAdapter) SQLDB() *sql.DB {
|
||||||
|
return p.db
|
||||||
|
}
|
||||||
|
|
||||||
func (p *PgSQLAdapter) DriverName() string {
|
func (p *PgSQLAdapter) DriverName() string {
|
||||||
return p.driverName
|
return p.driverName
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package common
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -38,6 +39,15 @@ type Database interface {
|
|||||||
DriverName() string
|
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)
|
// SelectQuery interface for building SELECT queries (compatible with both GORM and Bun)
|
||||||
type SelectQuery interface {
|
type SelectQuery interface {
|
||||||
Model(model interface{}) SelectQuery
|
Model(model interface{}) SelectQuery
|
||||||
|
|||||||
@@ -102,10 +102,7 @@ func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) {
|
|||||||
// Anything else defined in a spec's hooks.go must have an Execute call site: an
|
// Anything else defined in a spec's hooks.go must have an Execute call site: an
|
||||||
// unwired hook silently disables whatever is registered on it (resolvespec's
|
// unwired hook silently disables whatever is registered on it (resolvespec's
|
||||||
// AfterRead skipped column-level security masking until it was wired).
|
// AfterRead skipped column-level security masking until it was wired).
|
||||||
var unwiredHooks = map[string]string{
|
var unwiredHooks = map[string]string{}
|
||||||
"websocketspec/BeforeDisconnect": "connection close is not hooked yet",
|
|
||||||
"websocketspec/AfterDisconnect": "connection close is not hooked yet",
|
|
||||||
}
|
|
||||||
|
|
||||||
var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`)
|
var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`)
|
||||||
|
|
||||||
|
|||||||
+18
-23
@@ -19,7 +19,7 @@ In-memory store seeded from a static list. Suitable for a small, fixed set of se
|
|||||||
|
|
||||||
```go
|
```go
|
||||||
// Pre-load keys from config (KeyHash = SHA-256 hex of the raw key)
|
// 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,
|
UserID: 1,
|
||||||
KeyType: security.KeyTypeGenericAPI,
|
KeyType: security.KeyTypeGenericAPI,
|
||||||
@@ -33,7 +33,7 @@ store := security.NewConfigKeyStore([]security.UserKey{
|
|||||||
|
|
||||||
### DatabaseKeyStore
|
### 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
|
```go
|
||||||
db, _ := sql.Open("postgres", dsn)
|
db, _ := sql.Open("postgres", dsn)
|
||||||
@@ -43,8 +43,8 @@ store := security.NewDatabaseKeyStore(db)
|
|||||||
// With options
|
// With options
|
||||||
store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||||
CacheTTL: 5 * time.Minute,
|
CacheTTL: 5 * time.Minute,
|
||||||
SQLNames: &security.KeyStoreSQLNames{
|
Lookup: lookup.Config{
|
||||||
ValidateKey: "myapp_keystore_validate", // override one procedure name
|
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: <key>`
|
3. `X-API-Key: <key>`
|
||||||
|
|
||||||
```go
|
```go
|
||||||
auth := security.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
|
auth := providers.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
|
||||||
// Restrict to a specific type:
|
// Restrict to a specific type:
|
||||||
auth = security.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
|
auth = providers.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
|
||||||
```
|
```
|
||||||
|
|
||||||
Plug it into a handler:
|
Plug it into a handler:
|
||||||
@@ -109,10 +109,10 @@ On successful validation the request context receives a `UserContext` where:
|
|||||||
|
|
||||||
## Database setup
|
## 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
|
```sql
|
||||||
\i pkg/security/keystore_schema.sql
|
\i pkg/security/lookup/keystore_schema.sql
|
||||||
```
|
```
|
||||||
|
|
||||||
This creates:
|
This creates:
|
||||||
@@ -123,28 +123,23 @@ This creates:
|
|||||||
- `resolvespec_keystore_delete_key(p_user_id, p_key_id)`
|
- `resolvespec_keystore_delete_key(p_user_id, p_key_id)`
|
||||||
- `resolvespec_keystore_validate_key(p_key_hash, p_key_type)`
|
- `resolvespec_keystore_validate_key(p_key_hash, p_key_type)`
|
||||||
|
|
||||||
### Custom procedure names
|
### Custom names and modes
|
||||||
|
|
||||||
```go
|
```go
|
||||||
store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||||
SQLNames: &security.KeyStoreSQLNames{
|
Lookup: lookup.Config{
|
||||||
GetUserKeys: "myschema_get_keys",
|
Procs: lookup.ProcNames{
|
||||||
CreateKey: "myschema_create_key",
|
KeystoreGetUserKeys: "myschema_get_keys",
|
||||||
DeleteKey: "myschema_delete_key",
|
KeystoreCreateKey: "myschema_create_key",
|
||||||
ValidateKey: "myschema_validate_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
|
## Security notes
|
||||||
|
|
||||||
- Raw keys are never stored. Only the SHA-256 hex digest is persisted.
|
- Raw keys are never stored. Only the SHA-256 hex digest is persisted.
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
|
|||||||
### 1. Database Setup
|
### 1. Database Setup
|
||||||
|
|
||||||
```sql
|
```sql
|
||||||
-- Run the schema from database_schema.sql
|
-- Run the schema from lookup/database_schema.sql
|
||||||
CREATE TABLE IF NOT EXISTS users (
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
id SERIAL PRIMARY KEY,
|
id SERIAL PRIMARY KEY,
|
||||||
username VARCHAR(255) NOT NULL UNIQUE,
|
username VARCHAR(255) NOT NULL UNIQUE,
|
||||||
@@ -53,7 +53,7 @@ CREATE TABLE IF NOT EXISTS user_sessions (
|
|||||||
);
|
);
|
||||||
|
|
||||||
-- OAuth2 stored procedures (7 functions)
|
-- OAuth2 stored procedures (7 functions)
|
||||||
-- See database_schema.sql for full implementation
|
-- See lookup/database_schema.sql for full implementation
|
||||||
```
|
```
|
||||||
|
|
||||||
### 2. Google OAuth2
|
### 2. Google OAuth2
|
||||||
|
|||||||
@@ -43,16 +43,16 @@ CREATE TABLE IF NOT EXISTS user_sessions (
|
|||||||
**`resolvespec_oauth_getrefreshtoken(p_refresh_token)`**
|
**`resolvespec_oauth_getrefreshtoken(p_refresh_token)`**
|
||||||
- Gets OAuth2 session data by refresh token
|
- Gets OAuth2 session data by refresh token
|
||||||
- Returns: `{user_id, access_token, token_type, expiry}`
|
- 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)`**
|
**`resolvespec_oauth_updaterefreshtoken(p_update_data)`**
|
||||||
- Updates session with new tokens after refresh
|
- Updates session with new tokens after refresh
|
||||||
- Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}`
|
- 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)`**
|
**`resolvespec_oauth_getuser(p_user_id)`**
|
||||||
- Gets user data by ID for building UserContext
|
- Gets user data by ID for building UserContext
|
||||||
- Location: `database_schema.sql:791`
|
- Location: `lookup/database_schema.sql:791`
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthe
|
|||||||
## Setup
|
## Setup
|
||||||
|
|
||||||
### Database Schema
|
### 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
|
- Creates `user_passkey_credentials` table
|
||||||
- Adds stored procedures for passkey operations
|
- Adds stored procedures for passkey operations
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@
|
|||||||
// Step 1: Create security providers
|
// Step 1: Create security providers
|
||||||
auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended)
|
auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended)
|
||||||
// OR: auth := security.NewJWTAuthenticator("secret-key", db)
|
// OR: auth := security.NewJWTAuthenticator("secret-key", db)
|
||||||
// OR: auth := security.NewHeaderAuthenticator()
|
// OR: auth := providers.NewHeaderAuthenticator()
|
||||||
// OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2
|
// OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2
|
||||||
|
|
||||||
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
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)`
|
- Session/Login: `(p_success bool, p_error text, p_data jsonb)`
|
||||||
- Security: `(p_success bool, p_error text, p_rules 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:
|
// Requires these tables:
|
||||||
// - users (id, username, email, password, user_level, roles, is_active)
|
// - users (id, username, email, password, user_level, roles, is_active)
|
||||||
// - user_sessions (session_token, user_id, expires_at, created_at, last_activity_at)
|
// - 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:
|
// Features:
|
||||||
// - Login with username/password
|
// - Login with username/password
|
||||||
@@ -313,16 +313,15 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
|
|||||||
}
|
}
|
||||||
|
|
||||||
query := `
|
query := `
|
||||||
SELECT control, accesstype, jsonvalue
|
SELECT schema_name || '.' || table_name || '.' || column_path AS control,
|
||||||
FROM core.secaccess
|
access_type AS accesstype, COALESCE(extra_filters, '') AS jsonvalue
|
||||||
WHERE rid_hub IN (
|
FROM sec_column_rules
|
||||||
SELECT rid_hub_parent FROM core.hub_link
|
WHERE is_active = true
|
||||||
WHERE rid_hub_child = ? AND parent_hubtype = 'secgroup'
|
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 = ?))
|
||||||
AND control ILIKE ?
|
|
||||||
`
|
`
|
||||||
|
|
||||||
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -378,19 +377,19 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
|
|||||||
|
|
||||||
```go
|
```go
|
||||||
// Test Authenticator
|
// Test Authenticator
|
||||||
auth := security.NewHeaderAuthenticator()
|
auth := providers.NewHeaderAuthenticator()
|
||||||
req := httptest.NewRequest("GET", "/", nil)
|
req := httptest.NewRequest("GET", "/", nil)
|
||||||
req.Header.Set("X-User-ID", "123")
|
req.Header.Set("X-User-ID", "123")
|
||||||
userCtx, err := auth.Authenticate(req)
|
userCtx, err := auth.Authenticate(req)
|
||||||
assert.Equal(t, 123, userCtx.UserID)
|
assert.Equal(t, 123, userCtx.UserID)
|
||||||
|
|
||||||
// Test ColumnSecurityProvider
|
// Test ColumnSecurityProvider
|
||||||
colSec := security.NewConfigColumnSecurityProvider(rules)
|
colSec := providers.NewConfigColumnSecurityProvider(rules)
|
||||||
cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees")
|
cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees")
|
||||||
assert.Equal(t, "mask", cols[0].Accesstype)
|
assert.Equal(t, "mask", cols[0].Accesstype)
|
||||||
|
|
||||||
// Test RowSecurityProvider
|
// Test RowSecurityProvider
|
||||||
rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
|
||||||
row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders")
|
row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders")
|
||||||
assert.Equal(t, "user_id = {UserID}", row.Template)
|
assert.Equal(t, "user_id = {UserID}", row.Template)
|
||||||
```
|
```
|
||||||
|
|||||||
+99
-96
@@ -35,6 +35,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic
|
|||||||
| `resolvespec_login` | Session-based login | DatabaseAuthenticator |
|
| `resolvespec_login` | Session-based login | DatabaseAuthenticator |
|
||||||
| `resolvespec_logout` | Session invalidation | DatabaseAuthenticator |
|
| `resolvespec_logout` | Session invalidation | DatabaseAuthenticator |
|
||||||
| `resolvespec_session` | Session validation | 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_session_update` | Update session activity | DatabaseAuthenticator |
|
||||||
| `resolvespec_refresh_token` | Token refresh | DatabaseAuthenticator |
|
| `resolvespec_refresh_token` | Token refresh | DatabaseAuthenticator |
|
||||||
| `resolvespec_jwt_login` | JWT user validation | JWTAuthenticator |
|
| `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_request` | Create password reset token | DatabaseAuthenticator |
|
||||||
| `resolvespec_password_reset` | Validate token and set new password | 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).
|
- **procedure** (`lookup/procedure`): calls the `resolvespec_*` stored procedures (`p_success` / `p_error` / `p_data` contract). 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.
|
- **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
|
```go
|
||||||
type QueryMode int
|
type Config struct {
|
||||||
|
Dialect string // "postgres", "sqlite", "mysql", "mssql", or one you registered; empty = detect from the driver
|
||||||
const (
|
Mode lookup.Mode // default for every operation
|
||||||
ModeAuto QueryMode = iota // default
|
Overrides map[lookup.Op]lookup.Mode // per-operation mode, e.g. lookup.OpSession: lookup.ModeDirect
|
||||||
ModeProcedure
|
Procs lookup.ProcNames // procedure names, empty fields keep the default
|
||||||
ModeDirect
|
Schema lookup.Schema // table/column names, missing entries keep the default
|
||||||
)
|
|
||||||
```
|
|
||||||
|
|
||||||
- **`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
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
`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
|
```go
|
||||||
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
|
// SQLite or MySQL: nothing to configure, direct SQL is the default.
|
||||||
TableNames: &security.TableNames{Users: "app_users"}, // only override what differs
|
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
|
### 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`).
|
- Direct login, register, refresh, API-key login, password reset and passkey login run in one transaction.
|
||||||
- Session tokens generated by Direct mode use the same `sess_<hex>_<unix-timestamp>` shape as the plpgsql procedures.
|
- Passwords are stored as bcrypt; legacy cleartext values are accepted at login and only rewritten when `UpgradePasswordHash` is enabled.
|
||||||
- `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.
|
- Session tokens use the shape `sess_<hex>_<unix-timestamp>` 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
|
## Quick Start
|
||||||
|
|
||||||
@@ -297,7 +299,7 @@ type UserContext struct {
|
|||||||
|
|
||||||
**HeaderAuthenticator** - Simple header-based authentication:
|
**HeaderAuthenticator** - Simple header-based authentication:
|
||||||
```go
|
```go
|
||||||
auth := security.NewHeaderAuthenticator()
|
auth := providers.NewHeaderAuthenticator()
|
||||||
// Expects: X-User-ID, X-User-Name, X-User-Level, etc.
|
// 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
|
// Supports: Login, Logout, Session management, Token refresh
|
||||||
// All operations use stored procedures: resolvespec_login, resolvespec_logout,
|
// All operations use stored procedures: resolvespec_login, resolvespec_logout,
|
||||||
// resolvespec_session, resolvespec_session_update, resolvespec_refresh_token
|
// 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:
|
**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
|
// 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
|
```go
|
||||||
baseAuth := security.NewDatabaseAuthenticator(db)
|
baseAuth := security.NewDatabaseAuthenticator(db)
|
||||||
|
|
||||||
// Use in-memory provider (for testing)
|
// Use in-memory provider (for testing)
|
||||||
tfaProvider := security.NewMemoryTwoFactorProvider(nil)
|
tfaProvider := totp.NewMemoryProvider(nil)
|
||||||
|
|
||||||
// Or use database provider (for production)
|
// Or use database provider (for production)
|
||||||
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
||||||
// Requires: users table with totp fields, user_totp_backup_codes table
|
// 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
|
// Supports: TOTP codes, backup codes, QR code generation
|
||||||
// Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc.
|
// Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc.
|
||||||
```
|
```
|
||||||
@@ -341,7 +343,7 @@ auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
|||||||
```go
|
```go
|
||||||
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||||
// Uses stored procedure: resolvespec_column_security
|
// 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:
|
**ConfigColumnSecurityProvider** - Static configuration:
|
||||||
@@ -351,7 +353,7 @@ rules := map[string][]security.ColumnSecurity{
|
|||||||
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5},
|
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
colSec := security.NewConfigColumnSecurityProvider(rules)
|
colSec := providers.NewConfigColumnSecurityProvider(rules)
|
||||||
```
|
```
|
||||||
|
|
||||||
### Row Security Providers
|
### Row Security Providers
|
||||||
@@ -370,7 +372,7 @@ templates := map[string]string{
|
|||||||
blocked := map[string]bool{
|
blocked := map[string]bool{
|
||||||
"public.admin_logs": true,
|
"public.admin_logs": true,
|
||||||
}
|
}
|
||||||
rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
|
||||||
```
|
```
|
||||||
|
|
||||||
## Usage Examples
|
## Usage Examples
|
||||||
@@ -381,7 +383,7 @@ rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
|||||||
func main() {
|
func main() {
|
||||||
db := setupDatabase()
|
db := setupDatabase()
|
||||||
|
|
||||||
// Run migrations (see database_schema.sql)
|
// Run migrations (see lookup/database_schema.sql)
|
||||||
// db.Exec("CREATE TABLE users ...")
|
// db.Exec("CREATE TABLE users ...")
|
||||||
// db.Exec("CREATE TABLE user_sessions ...")
|
// db.Exec("CREATE TABLE user_sessions ...")
|
||||||
|
|
||||||
@@ -475,8 +477,8 @@ func handleRefresh(securityList *security.SecurityList) http.HandlerFunc {
|
|||||||
```go
|
```go
|
||||||
// 1. Wrap existing authenticator with 2FA support
|
// 1. Wrap existing authenticator with 2FA support
|
||||||
baseAuth := security.NewDatabaseAuthenticator(db)
|
baseAuth := security.NewDatabaseAuthenticator(db)
|
||||||
tfaProvider := security.NewMemoryTwoFactorProvider(nil) // Use custom DB implementation in production
|
tfaProvider := totp.NewMemoryProvider(nil) // Use custom DB implementation in production
|
||||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
|
||||||
|
|
||||||
// 2. Use as normal authenticator
|
// 2. Use as normal authenticator
|
||||||
provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec)
|
provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec)
|
||||||
@@ -548,18 +550,18 @@ has2FA, err := tfaProvider.Get2FAStatus(userID)
|
|||||||
// Uses PostgreSQL stored procedures for all operations
|
// Uses PostgreSQL stored procedures for all operations
|
||||||
db := setupDatabase()
|
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
|
// - Add totp_secret, totp_enabled, totp_enabled_at to users table
|
||||||
// - Create user_totp_backup_codes table
|
// - Create user_totp_backup_codes table
|
||||||
// - Create resolvespec_totp_* stored procedures
|
// - Create resolvespec_totp_* stored procedures
|
||||||
|
|
||||||
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
||||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
|
||||||
```
|
```
|
||||||
|
|
||||||
**Option 2: Implement Custom Provider**
|
**Option 2: Implement Custom Provider**
|
||||||
|
|
||||||
Implement `TwoFactorAuthProvider` for custom storage:
|
Implement `totp.AuthProvider` for custom storage:
|
||||||
|
|
||||||
```go
|
```go
|
||||||
type DBTwoFactorProvider struct {
|
type DBTwoFactorProvider struct {
|
||||||
@@ -585,15 +587,15 @@ func (p *DBTwoFactorProvider) Get2FASecret(userID int) (string, error) {
|
|||||||
### Configuration
|
### Configuration
|
||||||
|
|
||||||
```go
|
```go
|
||||||
config := &security.TwoFactorConfig{
|
config := &totp.Config{
|
||||||
Algorithm: "SHA256", // SHA1, SHA256, SHA512
|
Algorithm: "SHA256", // SHA1, SHA256, SHA512
|
||||||
Digits: 8, // 6 or 8
|
Digits: 8, // 6 or 8
|
||||||
Period: 30, // Seconds per code
|
Period: 30, // Seconds per code
|
||||||
SkewWindow: 2, // Accept codes ±2 periods
|
SkewWindow: 2, // Accept codes ±2 periods
|
||||||
}
|
}
|
||||||
|
|
||||||
totp := security.NewTOTPGenerator(config)
|
totp := totp.NewGenerator(config)
|
||||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, config)
|
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, config)
|
||||||
```
|
```
|
||||||
|
|
||||||
### API Response Structure
|
### API Response Structure
|
||||||
@@ -662,9 +664,9 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Create providers
|
// Create providers
|
||||||
auth := security.NewHeaderAuthenticator()
|
auth := providers.NewHeaderAuthenticator()
|
||||||
colSec := security.NewConfigColumnSecurityProvider(columnRules)
|
colSec := providers.NewConfigColumnSecurityProvider(columnRules)
|
||||||
rowSec := security.NewConfigRowSecurityProvider(rowTemplates, nil)
|
rowSec := providers.NewConfigRowSecurityProvider(rowTemplates, nil)
|
||||||
|
|
||||||
// Combine providers and register hooks
|
// Combine providers and register hooks
|
||||||
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
@@ -1013,7 +1015,7 @@ restheadspec.RegisterSecurityHooks(handler, securityList) // or funcspec/resolve
|
|||||||
|
|
||||||
### DB Requirements
|
### 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`)
|
- `user_password_resets` table (`user_id`, `token_hash` SHA-256, `expires_at`, `used`, `used_at`)
|
||||||
- `resolvespec_password_reset_request` stored procedure
|
- `resolvespec_password_reset_request` stored procedure
|
||||||
- `resolvespec_password_reset` 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
|
- `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`
|
- 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
|
```go
|
||||||
type SQLNames struct {
|
lookup.ProcNames{
|
||||||
// ...
|
PasswordResetRequest: "resolvespec_password_reset_request", // default
|
||||||
PasswordResetRequest string // default: "resolvespec_password_reset_request"
|
PasswordResetComplete: "resolvespec_password_reset", // default
|
||||||
PasswordResetComplete string // default: "resolvespec_password_reset"
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -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.
|
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
|
#### New DB Types
|
||||||
|
|
||||||
@@ -1198,10 +1201,10 @@ auth.OAuthIntrospectToken(ctx, token) // RFC 7662 — returns OAuthTokenInfo
|
|||||||
auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session
|
auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session
|
||||||
```
|
```
|
||||||
|
|
||||||
#### SQLNames Fields
|
#### Procedure names
|
||||||
|
|
||||||
```go
|
```go
|
||||||
type SQLNames struct {
|
type ProcNames struct {
|
||||||
// ... existing fields ...
|
// ... existing fields ...
|
||||||
OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
|
OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
|
||||||
OAuthGetClient string // default: "resolvespec_oauth_get_client"
|
OAuthGetClient string // default: "resolvespec_oauth_get_client"
|
||||||
|
|||||||
@@ -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`.
|
||||||
@@ -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 {
|
func (c *ChainAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
||||||
return c.authenticators[0].LogoutWithCookie(ctx, req, w)
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -88,6 +88,14 @@ func (c *CompositeSecurityProvider) RefreshToken(ctx context.Context, refreshTok
|
|||||||
return nil, fmt.Errorf("authenticator does not support token refresh")
|
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
|
// ValidateToken implements Validatable if the authenticator supports it
|
||||||
func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) {
|
func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) {
|
||||||
if validatable, ok := c.auth.(Validatable); ok {
|
if validatable, ok := c.auth.(Validatable); ok {
|
||||||
|
|||||||
@@ -5,37 +5,41 @@ import (
|
|||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"strings"
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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 {
|
func futureTime() time.Time {
|
||||||
return time.Now().Add(1 * time.Hour)
|
return time.Now().Add(1 * time.Hour)
|
||||||
}
|
}
|
||||||
|
|
||||||
// newDirectTestDB opens a fresh in-memory SQLite database and applies the
|
// 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.
|
// Direct-mode test a real, isolated database to exercise end-to-end.
|
||||||
func newDirectTestDB(t *testing.T) *sql.DB {
|
func newDirectTestDB(t *testing.T) *sql.DB {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
db, err := sql.Open("sqlite3", "file::memory:?cache=shared")
|
db, err := sql.Open("sqlite", ":memory:")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to open sqlite db: %v", err)
|
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
|
db.SetMaxOpenConns(1) // keep the shared in-memory db single-connection so state isn't lost
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
schemaPath := filepath.Join("database_schema_sqlite.sql")
|
schema, err := ddl.SQL("sqlite")
|
||||||
schema, err := os.ReadFile(schemaPath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to read schema: %v", err)
|
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)
|
t.Fatalf("failed to apply schema: %v", err)
|
||||||
}
|
}
|
||||||
return db
|
return db
|
||||||
@@ -49,7 +53,7 @@ func authenticatedRequest(token string) *http.Request {
|
|||||||
|
|
||||||
func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
regResp, err := auth.Register(ctx, RegisterRequest{
|
regResp, err := auth.Register(ctx, RegisterRequest{
|
||||||
@@ -105,7 +109,7 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
|||||||
|
|
||||||
func TestDirectMode_SessionLifecycle(t *testing.T) {
|
func TestDirectMode_SessionLifecycle(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
loginResp, err := auth.Register(ctx, RegisterRequest{Username: "bob", Password: "p", Email: "bob@example.com"})
|
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) {
|
func TestDirectMode_PasswordReset(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
if _, err := auth.Register(ctx, RegisterRequest{Username: "carol", Password: "old", Email: "carol@example.com"}); err != nil {
|
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) {
|
func TestDirectMode_JWTLoginAndLogout(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
jwtAuth := NewJWTAuthenticator("secret", db).WithQueryMode(ModeDirect)
|
jwtAuth := NewJWTAuthenticator("secret", db).WithLookup(directConfig)
|
||||||
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
if _, err := directAuth.Register(ctx, RegisterRequest{Username: "dave", Password: "p", Email: "dave@example.com"}); err != nil {
|
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) {
|
func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
regResp, err := auth.Register(ctx, RegisterRequest{Username: "erin", Password: "p", Email: "erin@example.com"})
|
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
|
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 {
|
if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil {
|
||||||
t.Fatalf("Enable2FA() error = %v", err)
|
t.Fatalf("Enable2FA() error = %v", err)
|
||||||
@@ -248,7 +252,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
|
|||||||
|
|
||||||
func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
|
func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
regResp, err := auth.Register(ctx, RegisterRequest{Username: "frank", Password: "p", Email: "frank@example.com"})
|
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
|
userID := regResp.User.UserID
|
||||||
|
|
||||||
passkeys := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{
|
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{
|
cred, err := passkeys.CompleteRegistration(ctx, userID, PasskeyRegistrationResponse{
|
||||||
@@ -317,7 +321,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
|
|||||||
|
|
||||||
func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
|
func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
userCtx := &UserContext{UserName: "gina", Email: "gina@example.com", Roles: []string{"user"}}
|
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) {
|
func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
regResp, err := auth.Register(ctx, RegisterRequest{Username: "henry", Password: "p", Email: "henry@example.com"})
|
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)
|
t.Fatalf("Register() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{QueryMode: ModeDirect})
|
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{Lookup: directConfig})
|
||||||
|
|
||||||
createResp, err := ks.CreateKey(ctx, CreateKeyRequest{
|
createResp, err := ks.CreateKey(ctx, CreateKeyRequest{
|
||||||
UserID: regResp.User.UserID,
|
UserID: regResp.User.UserID,
|
||||||
@@ -399,7 +403,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
|
|||||||
|
|
||||||
func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
client := &OAuthServerClient{
|
client := &OAuthServerClient{
|
||||||
@@ -489,7 +493,7 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
|||||||
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
||||||
for _, enabled := range []bool{false, true} {
|
for _, enabled := range []bool{false, true} {
|
||||||
db := newDirectTestDB(t)
|
db := newDirectTestDB(t)
|
||||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled})
|
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig, UpgradePasswordHash: enabled})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
if _, err := db.Exec(`DELETE FROM users`); err != nil {
|
if _, err := db.Exec(`DELETE FROM users`); err != nil {
|
||||||
@@ -518,18 +522,6 @@ func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestVerifyPasswordEdgeCases(t *testing.T) {
|
func isBcryptHash(s string) bool {
|
||||||
h, _ := hashPassword("pw")
|
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
||||||
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")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,81 +5,6 @@ import (
|
|||||||
"net/http"
|
"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
|
// Authenticator handles user authentication operations
|
||||||
type Authenticator interface {
|
type Authenticator interface {
|
||||||
// Login authenticates credentials and returns a token
|
// Login authenticates credentials and returns a token
|
||||||
@@ -144,6 +69,13 @@ type Refreshable interface {
|
|||||||
RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error)
|
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
|
// Validatable allows providers to validate tokens without full authentication
|
||||||
type Validatable interface {
|
type Validatable interface {
|
||||||
// ValidateToken checks if a token is valid without extracting full user context
|
// ValidateToken checks if a token is valid without extracting full user context
|
||||||
|
|||||||
@@ -2,64 +2,12 @@ package security
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
"time"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// hashSHA256Hex returns the lowercase hex SHA-256 digest of the given string.
|
// hashSHA256Hex is kept as a short alias for sectypes.HashKey inside this package.
|
||||||
// Used by all keystore implementations to hash raw keys before storage or lookup.
|
func hashSHA256Hex(raw string) string { return sectypes.HashKey(raw) }
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// KeyStore manages per-user auth keys with pluggable storage backends.
|
// KeyStore manages per-user auth keys with pluggable storage backends.
|
||||||
// Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures).
|
// Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures).
|
||||||
|
|||||||
@@ -5,16 +5,16 @@ import (
|
|||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/sync/singleflight"
|
"golang.org/x/sync/singleflight"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends"
|
||||||
)
|
)
|
||||||
|
|
||||||
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
||||||
@@ -24,36 +24,28 @@ type DatabaseKeyStoreOptions struct {
|
|||||||
// CacheTTL is the duration to cache ValidateKey results.
|
// CacheTTL is the duration to cache ValidateKey results.
|
||||||
// Default: 2 minutes.
|
// Default: 2 minutes.
|
||||||
CacheTTL time.Duration
|
CacheTTL time.Duration
|
||||||
// SQLNames provides custom procedure names. If nil, uses DefaultKeyStoreSQLNames().
|
// Lookup selects dialect, query mode and procedure/table/column names.
|
||||||
SQLNames *KeyStoreSQLNames
|
// The zero value uses stored procedures on Postgres and direct SQL elsewhere.
|
||||||
// TableNames provides custom table names for Direct mode. If nil, uses DefaultKeyStoreTableNames().
|
Lookup lookup.Config
|
||||||
TableNames *KeyStoreTableNames
|
// LookupProvider, when set, is used instead of building one from Lookup and the db.
|
||||||
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
|
LookupProvider *lookup.Provider
|
||||||
QueryMode QueryMode
|
|
||||||
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
||||||
// If nil, reconnection is disabled.
|
// If nil, reconnection is disabled.
|
||||||
DBFactory func() (*sql.DB, error)
|
DBFactory func() (*sql.DB, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// DatabaseKeyStore is a KeyStore backed by PostgreSQL stored procedures.
|
// DatabaseKeyStore is a KeyStore backed by the lookup package (stored procedures on
|
||||||
// All DB operations go through configurable procedure names; the raw key is
|
// Postgres by default, direct SQL elsewhere). The raw key is never passed to the database.
|
||||||
// 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
|
// 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
|
// cache TTL, a deleted key may continue to authenticate for up to CacheTTL
|
||||||
// (default 2 minutes) if the cache entry cannot be invalidated.
|
// (default 2 minutes) if the cache entry cannot be invalidated.
|
||||||
type DatabaseKeyStore struct {
|
type DatabaseKeyStore struct {
|
||||||
db *sql.DB
|
src *lookupSource
|
||||||
dbMu sync.RWMutex
|
cache *cache.Cache
|
||||||
dbFactory func() (*sql.DB, error)
|
cacheTTL time.Duration
|
||||||
sqlNames *KeyStoreSQLNames
|
|
||||||
tableNames *KeyStoreTableNames
|
|
||||||
queryMode QueryMode
|
|
||||||
capability *dbCapability
|
|
||||||
cache *cache.Cache
|
|
||||||
cacheTTL time.Duration
|
|
||||||
|
|
||||||
// validateLoads collapses concurrent key lookups for the same key
|
// validateLoads collapses concurrent key lookups for the same key
|
||||||
validateLoads singleflight.Group
|
validateLoads singleflight.Group
|
||||||
@@ -72,42 +64,14 @@ func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseK
|
|||||||
if c == nil {
|
if c == nil {
|
||||||
c = cache.GetDefaultCache()
|
c = cache.GetDefaultCache()
|
||||||
}
|
}
|
||||||
names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames)
|
src := newLookupSource(db)
|
||||||
tableNames := resolveKeyStoreTableNames(o.TableNames)
|
src.cfg = o.Lookup
|
||||||
return &DatabaseKeyStore{
|
src.provider = o.LookupProvider
|
||||||
db: db,
|
src.opts = backends.Options{DBFactory: o.DBFactory}
|
||||||
dbFactory: o.DBFactory,
|
return &DatabaseKeyStore{src: src, cache: c, cacheTTL: o.CacheTTL}
|
||||||
sqlNames: names,
|
|
||||||
tableNames: tableNames,
|
|
||||||
queryMode: o.QueryMode,
|
|
||||||
capability: newDBCapability(),
|
|
||||||
cache: c,
|
|
||||||
cacheTTL: o.CacheTTL,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ks *DatabaseKeyStore) getDB() *sql.DB {
|
func (ks *DatabaseKeyStore) keys() lookup.KeyStore { return ks.src.get().Keys }
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// CreateKey generates a raw key, stores its SHA-256 hash via the create procedure,
|
// CreateKey generates a raw key, stores its SHA-256 hash via the create procedure,
|
||||||
// and returns the raw key once.
|
// and returns the raw key once.
|
||||||
@@ -119,110 +83,29 @@ func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest)
|
|||||||
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
||||||
hash := hashSHA256Hex(rawKey)
|
hash := hashSHA256Hex(rawKey)
|
||||||
|
|
||||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.CreateKey) {
|
key, err := ks.keys().Create(ctx, req, hash)
|
||||||
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,
|
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
|
return nil, err
|
||||||
}
|
}
|
||||||
|
return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUserKeys returns all active, non-expired keys for the given user.
|
// GetUserKeys returns all active, non-expired keys for the given user.
|
||||||
// Pass an empty KeyType to return all types.
|
// Pass an empty KeyType to return all types.
|
||||||
func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) {
|
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.keys().List(ctx, userID, keyType)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteKey soft-deletes a key after verifying ownership and invalidates its cache entry.
|
// 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.
|
// 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.
|
// 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 {
|
func (ks *DatabaseKeyStore) DeleteKey(ctx context.Context, userID int, keyID int64) error {
|
||||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.DeleteKey) {
|
keyHash, err := ks.keys().Delete(ctx, userID, keyID)
|
||||||
return ks.deleteKeyDirect(ctx, userID, keyID)
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
if keyHash != "" && ks.cache != nil {
|
||||||
var success bool
|
_ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash))
|
||||||
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))
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -261,61 +144,18 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT
|
|||||||
// validateKeyLoad validates against the database and fills the cache.
|
// validateKeyLoad validates against the database and fills the cache.
|
||||||
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
|
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
|
||||||
dbtrace.Raw(ctx, "keystore.validate")
|
dbtrace.Raw(ctx, "keystore.validate")
|
||||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) {
|
key, err := ks.keys().Validate(ctx, hash, keyType)
|
||||||
key, err := ks.validateKeyDirect(ctx, hash, keyType)
|
if err != nil {
|
||||||
if err != nil {
|
return nil, err
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if ks.cache != nil {
|
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 {
|
func keystoreCacheKey(hash string) string {
|
||||||
return "keystore:validate:" + hash
|
return "keystore:validate:" + hash
|
||||||
}
|
}
|
||||||
|
|
||||||
// nullStringOr returns s.String if valid, otherwise the fallback.
|
|
||||||
func nullStringOr(s sql.NullString, fallback string) string {
|
|
||||||
if s.Valid && s.String != "" {
|
|
||||||
return s.String
|
|
||||||
}
|
|
||||||
return fallback
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<hex>_" 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,156 @@
|
|||||||
|
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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConformanceMSSQLContainer(t *testing.T) {
|
||||||
|
rt := containerRuntime(t)
|
||||||
|
port := startContainer(t, rt, "mcr.microsoft.com/mssql/server:2022-latest", "1433", map[string]string{
|
||||||
|
"ACCEPT_EULA": "Y", "MSSQL_SA_PASSWORD": containerPassword,
|
||||||
|
})
|
||||||
|
dsn := func(db string) string {
|
||||||
|
return fmt.Sprintf("sqlserver://sa:%s@127.0.0.1:%s?database=%s&encrypt=disable", containerPassword, port, db)
|
||||||
|
}
|
||||||
|
admin := waitReady(t, "sqlserver", dsn("master"), 3*time.Minute)
|
||||||
|
if _, err := admin.Exec("CREATE DATABASE cf_direct"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
runOnServer(t, "sqlserver", dsn("cf_direct"), "mssql", lookup.Config{}, true)
|
||||||
|
}
|
||||||
@@ -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) (userID int, signCount uint32, err 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)
|
||||||
|
}
|
||||||
@@ -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, ", ")) //nolint:gosec // test seeding with fixed table names
|
||||||
|
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 i := range rules {
|
||||||
|
paths[strings.Join(rules[i].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 }
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -428,30 +428,78 @@ EXCEPTION
|
|||||||
END;
|
END;
|
||||||
$$ LANGUAGE plpgsql;
|
$$ 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
|
-- 8. resolvespec_column_security - Loads column security rules for user
|
||||||
-- Input: user_id (int), schema (text), table_name (text)
|
-- Input: user_id (int), schema (text), table_name (text)
|
||||||
-- Output: p_success (bool), p_error (text), p_rules (array of security rules as jsonb)
|
-- 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)
|
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 $$
|
RETURNS TABLE(p_success boolean, p_error text, p_rules jsonb) AS $$
|
||||||
DECLARE
|
DECLARE
|
||||||
v_rules jsonb;
|
v_rules jsonb;
|
||||||
BEGIN
|
BEGIN
|
||||||
-- Query column security rules from core.secaccess
|
|
||||||
SELECT jsonb_agg(
|
SELECT jsonb_agg(
|
||||||
jsonb_build_object(
|
jsonb_build_object(
|
||||||
'control', control,
|
'control', r.schema_name || '.' || r.table_name || '.' || r.column_path,
|
||||||
'accesstype', accesstype,
|
'accesstype', r.access_type,
|
||||||
'jsonvalue', jsonvalue
|
'jsonvalue', COALESCE(r.extra_filters, '')
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
INTO v_rules
|
INTO v_rules
|
||||||
FROM core.secaccess
|
FROM sec_column_rules r
|
||||||
WHERE rid_hub IN (
|
WHERE r.is_active = true
|
||||||
SELECT rid_hub_parent
|
AND lower(r.schema_name) = lower(p_schema)
|
||||||
FROM core.hub_link
|
AND lower(r.table_name) = lower(p_table_name)
|
||||||
WHERE rid_hub_child = p_user_id AND parent_hubtype = 'secgroup'
|
AND (
|
||||||
)
|
r.user_id = p_user_id
|
||||||
AND control ILIKE (p_schema || '.' || p_table_name || '%');
|
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
|
IF v_rules IS NULL THEN
|
||||||
v_rules := '[]'::jsonb;
|
v_rules := '[]'::jsonb;
|
||||||
@@ -464,20 +512,36 @@ EXCEPTION
|
|||||||
END;
|
END;
|
||||||
$$ LANGUAGE plpgsql;
|
$$ 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)
|
-- Input: schema (text), table_name (text), user_id (int)
|
||||||
-- Output: p_template (text), p_block (bool)
|
-- 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)
|
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 $$
|
RETURNS TABLE(p_template text, p_block boolean) AS $$
|
||||||
|
DECLARE
|
||||||
|
v_block boolean;
|
||||||
|
v_template text;
|
||||||
BEGIN
|
BEGIN
|
||||||
-- Call the existing core function if it exists, or implement your own logic
|
SELECT COALESCE(bool_or(r.has_block), false),
|
||||||
-- This is a placeholder that you should customize based on your core.api_sec_rowtemplate logic
|
COALESCE(string_agg('(' || r.template || ')', ' AND ' ORDER BY r.id)
|
||||||
RETURN QUERY SELECT ''::text, false;
|
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:
|
IF v_block THEN
|
||||||
-- RETURN QUERY SELECT template, has_block
|
v_template := '';
|
||||||
-- FROM core.row_security_config
|
END IF;
|
||||||
-- WHERE schema_name = p_schema AND table_name = p_table_name AND user_id = p_user_id;
|
|
||||||
|
RETURN QUERY SELECT v_template, v_block;
|
||||||
END;
|
END;
|
||||||
$$ LANGUAGE plpgsql;
|
$$ LANGUAGE plpgsql;
|
||||||
|
|
||||||
@@ -650,7 +714,7 @@ BEGIN
|
|||||||
v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2');
|
v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2');
|
||||||
|
|
||||||
-- Convert roles array to comma-separated string
|
-- 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;
|
INTO v_roles;
|
||||||
|
|
||||||
-- Try to find existing user by email
|
-- Try to find existing user by email
|
||||||
@@ -701,7 +765,7 @@ BEGIN
|
|||||||
v_access_token := p_session_data->>'access_token';
|
v_access_token := p_session_data->>'access_token';
|
||||||
v_refresh_token := p_session_data->>'refresh_token';
|
v_refresh_token := p_session_data->>'refresh_token';
|
||||||
v_token_type := COALESCE(p_session_data->>'token_type', 'Bearer');
|
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');
|
v_auth_provider := COALESCE(p_session_data->>'auth_provider', 'oauth2');
|
||||||
|
|
||||||
-- Insert or update session
|
-- Insert or update session
|
||||||
@@ -857,7 +921,7 @@ BEGIN
|
|||||||
v_new_session_token := p_update_data->>'new_session_token';
|
v_new_session_token := p_update_data->>'new_session_token';
|
||||||
v_new_access_token := p_update_data->>'new_access_token';
|
v_new_access_token := p_update_data->>'new_access_token';
|
||||||
v_new_refresh_token := p_update_data->>'new_refresh_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 session in user_sessions table
|
||||||
UPDATE user_sessions
|
UPDATE user_sessions
|
||||||
@@ -1214,7 +1278,7 @@ BEGIN
|
|||||||
|
|
||||||
-- Convert transports array
|
-- Convert transports array
|
||||||
IF p_credential->'transports' IS NOT NULL THEN
|
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;
|
INTO v_transports;
|
||||||
END IF;
|
END IF;
|
||||||
|
|
||||||
@@ -1304,12 +1368,11 @@ BEGIN
|
|||||||
'name', name,
|
'name', name,
|
||||||
'created_at', created_at,
|
'created_at', created_at,
|
||||||
'last_used_at', last_used_at
|
'last_used_at', last_used_at
|
||||||
)
|
) ORDER BY created_at DESC
|
||||||
), '[]'::jsonb)
|
), '[]'::jsonb)
|
||||||
INTO v_credentials
|
INTO v_credentials
|
||||||
FROM user_passkey_credentials
|
FROM user_passkey_credentials
|
||||||
WHERE user_id = p_user_id
|
WHERE user_id = p_user_id;
|
||||||
ORDER BY created_at DESC;
|
|
||||||
|
|
||||||
RETURN QUERY SELECT true, NULL::text, v_credentials;
|
RETURN QUERY SELECT true, NULL::text, v_credentials;
|
||||||
EXCEPTION
|
EXCEPTION
|
||||||
@@ -1451,6 +1514,64 @@ EXCEPTION
|
|||||||
END;
|
END;
|
||||||
$$ LANGUAGE plpgsql;
|
$$ 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
|
-- 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
|
-- 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)
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
LANGUAGE plpgsql AS $$
|
LANGUAGE plpgsql AS $$
|
||||||
DECLARE
|
DECLARE
|
||||||
v_client_id text;
|
v_client_id text;
|
||||||
v_row jsonb;
|
v_row jsonb;
|
||||||
BEGIN
|
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)
|
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
|
||||||
VALUES (
|
VALUES (
|
||||||
v_client_id,
|
v_client_id,
|
||||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')),
|
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
||||||
p_data->>'client_name',
|
p_request->>'client_name',
|
||||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'grant_types')), ARRAY['authorization_code']),
|
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,
|
||||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email']),
|
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_data->>'client_secret_hash', ''),
|
NULLIF(p_request->>'client_secret_hash', ''),
|
||||||
COALESCE(NULLIF(p_data->>'token_endpoint_auth_method', ''), 'none')
|
COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none')
|
||||||
)
|
)
|
||||||
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
|
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
|
||||||
|
|
||||||
@@ -1717,22 +1838,22 @@ BEGIN
|
|||||||
END;
|
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)
|
RETURNS TABLE(p_success bool, p_error text)
|
||||||
LANGUAGE plpgsql AS $$
|
LANGUAGE plpgsql AS $$
|
||||||
BEGIN
|
BEGIN
|
||||||
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at)
|
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at)
|
||||||
VALUES (
|
VALUES (
|
||||||
p_data->>'code',
|
p_request->>'code',
|
||||||
p_data->>'client_id',
|
p_request->>'client_id',
|
||||||
p_data->>'redirect_uri',
|
p_request->>'redirect_uri',
|
||||||
p_data->>'client_state',
|
p_request->>'client_state',
|
||||||
p_data->>'code_challenge',
|
p_request->>'code_challenge',
|
||||||
COALESCE(p_data->>'code_challenge_method', 'S256'),
|
COALESCE(p_request->>'code_challenge_method', 'S256'),
|
||||||
p_data->>'session_token',
|
p_request->>'session_token',
|
||||||
p_data->>'refresh_token',
|
p_request->>'refresh_token',
|
||||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'scopes')),
|
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)),
|
||||||
(p_data->>'expires_at')::timestamp
|
(p_request->>'expires_at')::timestamptz::timestamp
|
||||||
);
|
);
|
||||||
|
|
||||||
RETURN QUERY SELECT true, null::text;
|
RETURN QUERY SELECT true, null::text;
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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);
|
||||||
|
|
||||||
@@ -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)
|
||||||
|
);
|
||||||
|
|
||||||
@@ -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);
|
||||||
|
|
||||||
@@ -1,13 +1,15 @@
|
|||||||
-- Portable schema for Direct-mode (non-stored-procedure) operation.
|
-- Reference schema for the lookup direct backend: SQLite. Direct backend only.
|
||||||
-- Plain CREATE TABLE statements only, no functions/triggers, using types
|
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
|
||||||
-- understood by SQLite (and portable to MySQL). Used by Direct-mode tests
|
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
|
||||||
-- and as a reference for deployments without Postgres.
|
|
||||||
|
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
|
||||||
|
-- roles: comma-separated roles
|
||||||
|
|
||||||
CREATE TABLE IF NOT EXISTS users (
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
username VARCHAR(255) NOT NULL UNIQUE,
|
username VARCHAR(255) NOT NULL UNIQUE,
|
||||||
email 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,
|
user_level INTEGER DEFAULT 0,
|
||||||
roles VARCHAR(500),
|
roles VARCHAR(500),
|
||||||
is_active BOOLEAN DEFAULT 1,
|
is_active BOOLEAN DEFAULT 1,
|
||||||
@@ -23,6 +25,7 @@ CREATE TABLE IF NOT EXISTS users (
|
|||||||
totp_enabled_at TIMESTAMP
|
totp_enabled_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
|
|
||||||
CREATE TABLE IF NOT EXISTS user_sessions (
|
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
session_token VARCHAR(500) NOT NULL UNIQUE,
|
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||||
@@ -38,10 +41,12 @@ CREATE TABLE IF NOT EXISTS user_sessions (
|
|||||||
auth_provider VARCHAR(50)
|
auth_provider VARCHAR(50)
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_session_token ON user_sessions(session_token);
|
CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id);
|
||||||
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_user_sessions_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_refresh_token ON user_sessions(refresh_token);
|
||||||
|
|
||||||
|
|
||||||
CREATE TABLE IF NOT EXISTS token_blacklist (
|
CREATE TABLE IF NOT EXISTS token_blacklist (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
@@ -51,6 +56,9 @@ CREATE TABLE IF NOT EXISTS token_blacklist (
|
|||||||
created_at TIMESTAMP
|
created_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
|
|
||||||
|
-- code_hash: SHA-256 hex of the backup code
|
||||||
|
|
||||||
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
|
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
user_id INTEGER NOT NULL,
|
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_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);
|
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 (
|
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
user_id INTEGER NOT NULL,
|
user_id INTEGER NOT NULL,
|
||||||
credential_id TEXT NOT NULL UNIQUE, -- base64 text (Direct mode), not native bytea
|
credential_id TEXT NOT NULL UNIQUE,
|
||||||
public_key TEXT NOT NULL, -- base64 text
|
public_key TEXT NOT NULL,
|
||||||
attestation_type VARCHAR(50) DEFAULT 'none',
|
attestation_type VARCHAR(50) DEFAULT 'none',
|
||||||
aaguid TEXT, -- base64 text
|
aaguid TEXT,
|
||||||
sign_count INTEGER DEFAULT 0,
|
sign_count INTEGER DEFAULT 0,
|
||||||
clone_warning BOOLEAN DEFAULT 0,
|
clone_warning BOOLEAN DEFAULT 0,
|
||||||
transports TEXT, -- JSON-encoded []string
|
transports TEXT,
|
||||||
backup_eligible BOOLEAN DEFAULT 0,
|
backup_eligible BOOLEAN DEFAULT 0,
|
||||||
backup_state BOOLEAN DEFAULT 0,
|
backup_state BOOLEAN DEFAULT 0,
|
||||||
name VARCHAR(255),
|
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_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 (
|
CREATE TABLE IF NOT EXISTS user_password_resets (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
@@ -93,19 +110,32 @@ CREATE TABLE IF NOT EXISTS user_password_resets (
|
|||||||
used_at TIMESTAMP
|
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 (
|
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
client_id VARCHAR(255) NOT NULL UNIQUE,
|
client_id VARCHAR(255) NOT NULL UNIQUE,
|
||||||
redirect_uris TEXT NOT NULL, -- JSON-encoded []string
|
redirect_uris TEXT NOT NULL,
|
||||||
client_name VARCHAR(255),
|
client_name VARCHAR(255),
|
||||||
grant_types TEXT, -- JSON-encoded []string
|
grant_types TEXT,
|
||||||
allowed_scopes TEXT, -- JSON-encoded []string
|
allowed_scopes TEXT,
|
||||||
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
|
client_secret_hash TEXT,
|
||||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||||
is_active BOOLEAN DEFAULT 1,
|
is_active BOOLEAN DEFAULT 1,
|
||||||
created_at TIMESTAMP
|
created_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
|
|
||||||
|
-- scopes: JSON-encoded array
|
||||||
|
|
||||||
CREATE TABLE IF NOT EXISTS oauth_codes (
|
CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
code VARCHAR(255) NOT NULL UNIQUE,
|
code VARCHAR(255) NOT NULL UNIQUE,
|
||||||
@@ -116,28 +146,83 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
|||||||
code_challenge_method VARCHAR(10) DEFAULT 'S256',
|
code_challenge_method VARCHAR(10) DEFAULT 'S256',
|
||||||
session_token TEXT NOT NULL,
|
session_token TEXT NOT NULL,
|
||||||
refresh_token TEXT,
|
refresh_token TEXT,
|
||||||
scopes TEXT, -- JSON-encoded []string
|
scopes TEXT,
|
||||||
expires_at TIMESTAMP NOT NULL,
|
expires_at TIMESTAMP NOT NULL,
|
||||||
created_at TIMESTAMP
|
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);
|
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 (
|
CREATE TABLE IF NOT EXISTS user_keys (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
user_id INTEGER NOT NULL,
|
user_id INTEGER NOT NULL,
|
||||||
key_type VARCHAR(50) NOT NULL,
|
key_type VARCHAR(50) NOT NULL,
|
||||||
key_hash VARCHAR(64) NOT NULL UNIQUE,
|
key_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||||
name VARCHAR(255) NOT NULL DEFAULT '',
|
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||||
scopes TEXT, -- JSON-encoded []string
|
scopes TEXT,
|
||||||
meta TEXT, -- JSON-encoded map
|
meta TEXT,
|
||||||
expires_at TIMESTAMP,
|
expires_at TIMESTAMP,
|
||||||
created_at TIMESTAMP,
|
created_at TIMESTAMP,
|
||||||
last_used_at TIMESTAMP,
|
last_used_at TIMESTAMP,
|
||||||
is_active BOOLEAN DEFAULT 1
|
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_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_key_type ON user_keys(key_type);
|
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);
|
||||||
|
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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 "<pkgpath>.<Type>")
|
||||||
|
// 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 "<pkgpath>.<Type>" 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) <mid> VALUES (...) <tail>" 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()
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -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>_<unix>".
|
||||||
|
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
|
||||||
|
}
|
||||||
@@ -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().UTC().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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,473 @@
|
|||||||
|
// 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)
|
||||||
|
}
|
||||||
|
// Timestamp columns carry no zone: bind every instant as UTC so drivers that send an
|
||||||
|
// offset (SQL Server) and ones that drop it agree on the stored wall clock.
|
||||||
|
if tv, ok := v.(time.Time); ok {
|
||||||
|
return tv.UTC()
|
||||||
|
}
|
||||||
|
if tp, ok := v.(*time.Time); ok {
|
||||||
|
if tp == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return tp.UTC()
|
||||||
|
}
|
||||||
|
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() (cols []string, args []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...)
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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 }
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,293 @@
|
|||||||
|
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) (userID int, signCount uint32, err error) {
|
||||||
|
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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
package security
|
package direct
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/subtle"
|
"crypto/subtle"
|
||||||
@@ -15,7 +15,8 @@ const maxPasswordBytes = 72
|
|||||||
|
|
||||||
var errPasswordTooLong = errors.New("password must be at most 72 bytes")
|
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 {
|
if len(password) > maxPasswordBytes {
|
||||||
return "", errPasswordTooLong
|
return "", errPasswordTooLong
|
||||||
}
|
}
|
||||||
@@ -30,12 +31,11 @@ func isBcryptHash(s string) bool {
|
|||||||
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
||||||
}
|
}
|
||||||
|
|
||||||
// verifyPassword checks supplied against the stored value. A stored bcrypt hash
|
// VerifyPassword checks supplied against the stored value. A stored bcrypt hash is compared
|
||||||
// is compared with bcrypt. A legacy cleartext value (written before hashing was
|
// with bcrypt. A legacy cleartext value (written before hashing was implemented) is compared
|
||||||
// implemented) is compared in constant time and, on a match, needsRehash is true
|
// in constant time and, on a match, needsRehash is true so the caller can upgrade the row.
|
||||||
// so the caller can upgrade the row to a bcrypt hash. An empty stored value
|
// An empty stored value (e.g. an OAuth2-only user) never matches.
|
||||||
// (e.g. an OAuth2-only user) never matches.
|
func VerifyPassword(stored, supplied string) (ok, needsRehash bool) {
|
||||||
func verifyPassword(stored, supplied string) (ok, needsRehash bool) {
|
|
||||||
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
|
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
|
||||||
return false, false
|
return false, false
|
||||||
}
|
}
|
||||||
@@ -53,9 +53,9 @@ var (
|
|||||||
dummyHash string
|
dummyHash string
|
||||||
)
|
)
|
||||||
|
|
||||||
// burnPasswordCheck spends roughly one bcrypt comparison so an unknown username
|
// BurnPasswordCheck spends roughly one bcrypt comparison so an unknown username costs about
|
||||||
// costs about the same as a wrong password.
|
// the same as a wrong password.
|
||||||
func burnPasswordCheck(supplied string) {
|
func BurnPasswordCheck(supplied string) {
|
||||||
dummyHashOnce.Do(func() {
|
dummyHashOnce.Do(func() {
|
||||||
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
|
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
|
||||||
dummyHash = string(h)
|
dummyHash = string(h)
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -83,7 +83,7 @@ BEGIN
|
|||||||
p_request->>'scopes',
|
p_request->>'scopes',
|
||||||
p_request->'meta',
|
p_request->'meta',
|
||||||
CASE WHEN p_request->>'expires_at' IS NOT NULL
|
CASE WHEN p_request->>'expires_at' IS NOT NULL
|
||||||
THEN (p_request->>'expires_at')::TIMESTAMP
|
THEN (p_request->>'expires_at')::timestamptz::timestamp
|
||||||
ELSE NULL
|
ELSE NULL
|
||||||
END
|
END
|
||||||
)
|
)
|
||||||
@@ -185,3 +185,81 @@ EXCEPTION WHEN OTHERS THEN
|
|||||||
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||||
END;
|
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;
|
||||||
@@ -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")
|
||||||
|
)
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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" //nolint:gosec // operation name, not a credential
|
||||||
|
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" //nolint:gosec // operation name, not a credential
|
||||||
|
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" //nolint:gosec // operation name, not a credential
|
||||||
|
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" //nolint:gosec // operation name, not a credential
|
||||||
|
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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,318 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Auth implements lookup.AuthStore with stored procedures.
|
||||||
|
type Auth struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.AuthStore = (*Auth)(nil)
|
||||||
|
|
||||||
|
// NewAuth creates the procedure-backed AuthStore.
|
||||||
|
func NewAuth(run Runner, procs lookup.ProcNames) *Auth { return &Auth{run: run, procs: procs} }
|
||||||
|
|
||||||
|
// callData runs "SELECT p_success, p_error, p_data::text FROM proc($1::jsonb)".
|
||||||
|
func (a *Auth) callData(ctx context.Context, proc, queryErrOp string, arg any) (sql.NullString, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, dataJSON sql.NullString
|
||||||
|
err := a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, proc) //nolint:gosec // G201: identifier comes from validated config
|
||||||
|
return db.QueryRowContext(ctx, query, arg).Scan(&success, &errorMsg, &dataJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return sql.NullString{}, fmt.Errorf("%s query failed: %w", queryErrOp, err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return sql.NullString{}, failure(errorMsg, queryErrOp+" failed")
|
||||||
|
}
|
||||||
|
return dataJSON, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Login implements lookup.AuthStore.
|
||||||
|
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||||
|
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal login request: %w", err)
|
||||||
|
}
|
||||||
|
data, err := a.callData(ctx, a.procs.Login, "login", string(reqJSON))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var response sectypes.LoginResponse
|
||||||
|
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse login response: %w", err)
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register implements lookup.AuthStore.
|
||||||
|
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
|
||||||
|
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal register request: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg, dataJSON sql.NullString
|
||||||
|
err = a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.Register)
|
||||||
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("register query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "registration failed")
|
||||||
|
}
|
||||||
|
var response sectypes.LoginResponse
|
||||||
|
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse register response: %w", err)
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logout implements lookup.AuthStore.
|
||||||
|
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||||
|
reqJSON, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal logout request: %w", err)
|
||||||
|
}
|
||||||
|
_, err = a.callData(ctx, a.procs.Logout, "logout", string(reqJSON))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Session implements lookup.AuthStore.
|
||||||
|
func (a *Auth) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, userJSON sql.NullString
|
||||||
|
err := a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
||||||
|
return db.QueryRowContext(ctx, query, token, reference).Scan(&success, &errorMsg, &userJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("session query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "invalid or expired session")
|
||||||
|
}
|
||||||
|
if !userJSON.Valid {
|
||||||
|
return nil, fmt.Errorf("no user data in session")
|
||||||
|
}
|
||||||
|
var user sectypes.UserContext
|
||||||
|
if err := json.Unmarshal([]byte(userJSON.String), &user); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||||
|
}
|
||||||
|
return &user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TouchSession implements lookup.AuthStore.
|
||||||
|
func (a *Auth) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error {
|
||||||
|
userJSON, err := json.Marshal(user)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg, updatedUserJSON sql.NullString
|
||||||
|
return a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.SessionUpdate)
|
||||||
|
return db.QueryRowContext(ctx, query, token, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh implements lookup.AuthStore.
|
||||||
|
func (a *Auth) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) {
|
||||||
|
// Get the current session to pass to refresh.
|
||||||
|
var success bool
|
||||||
|
var errorMsg, userJSON sql.NullString
|
||||||
|
err := a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session)
|
||||||
|
return db.QueryRowContext(ctx, query, refreshToken, "refresh").Scan(&success, &errorMsg, &userJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("refresh token query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "invalid refresh token")
|
||||||
|
}
|
||||||
|
|
||||||
|
var newSuccess bool
|
||||||
|
var newErrorMsg, newUserJSON sql.NullString
|
||||||
|
err = a.run.Run(func(db *sql.DB) error {
|
||||||
|
refreshQuery := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.RefreshToken)
|
||||||
|
return db.QueryRowContext(ctx, refreshQuery, refreshToken, userJSON).Scan(&newSuccess, &newErrorMsg, &newUserJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("refresh token generation failed: %w", err)
|
||||||
|
}
|
||||||
|
if !newSuccess {
|
||||||
|
return nil, failure(newErrorMsg, "failed to refresh token")
|
||||||
|
}
|
||||||
|
|
||||||
|
var userCtx sectypes.UserContext
|
||||||
|
if err := json.Unmarshal([]byte(newUserJSON.String), &userCtx); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A resolvespec_refresh_token implementation that issues its own rotating
|
||||||
|
// refresh token (independent of the access/session token) returns it
|
||||||
|
// under claims.refresh_token, since UserContext has no dedicated field
|
||||||
|
// for it. Surface that into LoginResponse.RefreshToken so callers don't
|
||||||
|
// need to reach into User.Claims themselves. claims.expires_in
|
||||||
|
// (seconds) similarly overrides the default access-token ExpiresIn when
|
||||||
|
// the procedure provides a real value. Implementations that don't set
|
||||||
|
// these claims keep today's behavior unchanged (empty RefreshToken,
|
||||||
|
// 24h ExpiresIn default).
|
||||||
|
resp := §ypes.LoginResponse{
|
||||||
|
Token: userCtx.SessionID, // New session token from stored procedure
|
||||||
|
User: &userCtx,
|
||||||
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||||
|
}
|
||||||
|
if rt, ok := userCtx.Claims["refresh_token"].(string); ok && rt != "" {
|
||||||
|
resp.RefreshToken = rt
|
||||||
|
}
|
||||||
|
if expiresIn, ok := userCtx.Claims["expires_in"].(float64); ok && expiresIn > 0 {
|
||||||
|
resp.ExpiresIn = int64(expiresIn)
|
||||||
|
}
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginAPIKey implements lookup.AuthStore. Unknown, expired and inactive keys all return
|
||||||
|
// lookup.ErrInvalidAPIKey; the raw key is never logged.
|
||||||
|
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||||
|
if rawKey == "" {
|
||||||
|
return nil, lookup.ErrInvalidAPIKey
|
||||||
|
}
|
||||||
|
reqJSON, err := json.Marshal(map[string]any{"api_key": rawKey, "claims": claims})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal api key login request: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg, dataJSON sql.NullString
|
||||||
|
err = a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.LoginAPIKey)
|
||||||
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("api key login query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, lookup.ErrInvalidAPIKey
|
||||||
|
}
|
||||||
|
var response sectypes.LoginResponse
|
||||||
|
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse api key login response: %w", err)
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// JWTLogin implements lookup.AuthStore. The password is verified inside the procedure;
|
||||||
|
// the hash is never returned. The token is a placeholder until JWT signing is wired in.
|
||||||
|
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var userJSON []byte
|
||||||
|
err := a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user FROM %s($1, $2)`, a.procs.JWTLogin)
|
||||||
|
return db.QueryRowContext(ctx, query, req.Username, req.Password).Scan(&success, &errorMsg, &userJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("login query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "invalid credentials")
|
||||||
|
}
|
||||||
|
var user struct {
|
||||||
|
ID int `json:"id"`
|
||||||
|
Username string `json:"username"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
UserLevel int `json:"user_level"`
|
||||||
|
Roles string `json:"roles"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(userJSON, &user); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user data: %w", err)
|
||||||
|
}
|
||||||
|
roles := []string{}
|
||||||
|
if user.Roles != "" {
|
||||||
|
roles = strings.Split(user.Roles, ",")
|
||||||
|
}
|
||||||
|
expiresAt := time.Now().Add(24 * time.Hour)
|
||||||
|
return §ypes.LoginResponse{
|
||||||
|
Token: fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix()),
|
||||||
|
User: §ypes.UserContext{
|
||||||
|
UserID: user.ID,
|
||||||
|
UserName: user.Username,
|
||||||
|
Email: user.Email,
|
||||||
|
UserLevel: user.UserLevel,
|
||||||
|
Roles: roles,
|
||||||
|
},
|
||||||
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// JWTLogout implements lookup.AuthStore.
|
||||||
|
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err := a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, a.procs.JWTLogout)
|
||||||
|
return db.QueryRowContext(ctx, query, req.Token, req.UserID).Scan(&success, &errorMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("logout query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errorMsg, "logout failed")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetRequest implements lookup.AuthStore.
|
||||||
|
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
|
||||||
|
reqJSON, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal password reset request: %w", err)
|
||||||
|
}
|
||||||
|
data, err := a.callData(ctx, a.procs.PasswordResetRequest, "password reset request", string(reqJSON))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var response sectypes.PasswordResetResponse
|
||||||
|
if data.Valid && data.String != "" {
|
||||||
|
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse password reset response: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetComplete implements lookup.AuthStore.
|
||||||
|
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
|
||||||
|
reqJSON, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal password reset complete request: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err = a.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1::jsonb)`, a.procs.PasswordResetComplete)
|
||||||
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("password reset complete query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errorMsg, "password reset failed")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// normalizeTimes rewrites the zone-less timestamps the procedures emit (Postgres `timestamp`
|
||||||
|
// columns serialise as "2026-01-02T03:04:05.123456") as UTC RFC 3339, so the standard
|
||||||
|
// time.Time decoder accepts them. It handles a JSON object or an array of objects and only
|
||||||
|
// touches string fields whose name ends in "_at" or is "expiry". Anything else, including
|
||||||
|
// input that is not valid JSON, is returned unchanged.
|
||||||
|
func normalizeTimes(raw []byte) []byte {
|
||||||
|
var v any
|
||||||
|
if err := json.Unmarshal(raw, &v); err != nil {
|
||||||
|
return raw
|
||||||
|
}
|
||||||
|
switch x := v.(type) {
|
||||||
|
case map[string]any:
|
||||||
|
fixTimes(x)
|
||||||
|
case []any:
|
||||||
|
for _, e := range x {
|
||||||
|
if m, ok := e.(map[string]any); ok {
|
||||||
|
fixTimes(m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return raw
|
||||||
|
}
|
||||||
|
out, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
return raw
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func fixTimes(m map[string]any) {
|
||||||
|
for k, v := range m {
|
||||||
|
s, ok := v.(string)
|
||||||
|
if !ok || (!strings.HasSuffix(k, "_at") && k != "expiry") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := time.Parse(time.RFC3339Nano, s); err == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, layout := range []string{"2006-01-02T15:04:05.999999999", "2006-01-02 15:04:05.999999999"} {
|
||||||
|
if t, err := time.Parse(layout, s); err == nil {
|
||||||
|
m[k] = t.UTC().Format(time.RFC3339Nano)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,156 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Keys implements lookup.KeyStore with the resolvespec_keystore_* procedures.
|
||||||
|
type Keys struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.KeyStore = (*Keys)(nil)
|
||||||
|
|
||||||
|
// NewKeys creates the procedure-backed KeyStore.
|
||||||
|
func NewKeys(run Runner, procs lookup.ProcNames) *Keys { return &Keys{run: run, procs: procs} }
|
||||||
|
|
||||||
|
// orDefault returns the procedure's error message when it is non-empty, otherwise def.
|
||||||
|
func orDefault(s sql.NullString, def string) string {
|
||||||
|
if s.Valid && s.String != "" {
|
||||||
|
return s.String
|
||||||
|
}
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create implements lookup.KeyStore.
|
||||||
|
func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
|
||||||
|
type createRequest struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
KeyType sectypes.KeyType `json:"key_type"`
|
||||||
|
KeyHash string `json:"key_hash"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Scopes []string `json:"scopes,omitempty"`
|
||||||
|
Meta map[string]any `json:"meta,omitempty"`
|
||||||
|
ExpiresAt *time.Time `json:"expires_at,omitempty"`
|
||||||
|
}
|
||||||
|
reqJSON, err := json.Marshal(createRequest{
|
||||||
|
UserID: req.UserID,
|
||||||
|
KeyType: req.KeyType,
|
||||||
|
KeyHash: keyHash,
|
||||||
|
Name: req.Name,
|
||||||
|
Scopes: req.Scopes,
|
||||||
|
Meta: req.Meta,
|
||||||
|
ExpiresAt: req.ExpiresAt,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg, keyJSON sql.NullString
|
||||||
|
err = k.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, k.procs.KeystoreCreateKey)
|
||||||
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(orDefault(errorMsg, "create key failed"))
|
||||||
|
}
|
||||||
|
key, err := decodeKey([]byte(keyJSON.String))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse created key: %w", err)
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// List implements lookup.KeyStore.
|
||||||
|
func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, keysJSON sql.NullString
|
||||||
|
err := k.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, k.procs.KeystoreGetUserKeys)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get user keys procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(orDefault(errorMsg, "get user keys failed"))
|
||||||
|
}
|
||||||
|
var keys []sectypes.UserKey
|
||||||
|
if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" {
|
||||||
|
var raw []json.RawMessage
|
||||||
|
if err := json.Unmarshal([]byte(keysJSON.String), &raw); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user keys: %w", err)
|
||||||
|
}
|
||||||
|
for _, r := range raw {
|
||||||
|
k, err := decodeKey(r)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user keys: %w", err)
|
||||||
|
}
|
||||||
|
keys = append(keys, *k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if keys == nil {
|
||||||
|
keys = []sectypes.UserKey{}
|
||||||
|
}
|
||||||
|
return keys, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete implements lookup.KeyStore. The procedure returns the key hash.
|
||||||
|
func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, keyHash sql.NullString
|
||||||
|
err := k.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, k.procs.KeystoreDeleteKey)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("delete key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return "", errors.New(orDefault(errorMsg, "delete key failed"))
|
||||||
|
}
|
||||||
|
return keyHash.String, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate implements lookup.KeyStore.
|
||||||
|
func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, keyJSON sql.NullString
|
||||||
|
err := k.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, k.procs.KeystoreValidateKey)
|
||||||
|
return db.QueryRowContext(ctx, query, keyHash, string(keyType)).Scan(&success, &errorMsg, &keyJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("validate key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(orDefault(errorMsg, "invalid or expired key"))
|
||||||
|
}
|
||||||
|
key, err := decodeKey([]byte(keyJSON.String))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse validated key: %w", err)
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeKey reads one key record from a key procedure.
|
||||||
|
func decodeKey(raw []byte) (*sectypes.UserKey, error) {
|
||||||
|
var k sectypes.UserKey
|
||||||
|
if err := json.Unmarshal(normalizeTimes(raw), &k); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &k, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,315 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OAuthUsers implements lookup.OAuthUserStore with the resolvespec_oauth_* procedures.
|
||||||
|
type OAuthUsers struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.OAuthUserStore = (*OAuthUsers)(nil)
|
||||||
|
|
||||||
|
// NewOAuthUsers creates the procedure-backed OAuthUserStore.
|
||||||
|
func NewOAuthUsers(run Runner, procs lookup.ProcNames) *OAuthUsers {
|
||||||
|
return &OAuthUsers{run: run, procs: procs}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetOrCreateUser implements lookup.OAuthUserStore.
|
||||||
|
func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
|
||||||
|
userJSON, err := json.Marshal(map[string]any{
|
||||||
|
"username": user.UserName,
|
||||||
|
"email": user.Email,
|
||||||
|
"remote_id": user.RemoteID,
|
||||||
|
"user_level": user.UserLevel,
|
||||||
|
"roles": user.Roles,
|
||||||
|
"auth_provider": provider,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to marshal user data: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
var userID sql.NullInt64
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_user_id
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, o.procs.OAuthGetOrCreateUser), userJSON).Scan(&success, &errMsg, &userID)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to get or create user: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return 0, failure(errMsg, "failed to get or create user")
|
||||||
|
}
|
||||||
|
if !userID.Valid {
|
||||||
|
return 0, fmt.Errorf("user ID not returned")
|
||||||
|
}
|
||||||
|
return int(userID.Int64), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateSession implements lookup.OAuthUserStore.
|
||||||
|
func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error {
|
||||||
|
sessionJSON, err := json.Marshal(map[string]any{
|
||||||
|
"session_token": s.SessionToken,
|
||||||
|
"user_id": s.UserID,
|
||||||
|
"access_token": s.AccessToken,
|
||||||
|
"refresh_token": s.RefreshToken,
|
||||||
|
"token_type": s.TokenType,
|
||||||
|
"expires_at": s.ExpiresAt,
|
||||||
|
"auth_provider": s.Provider,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal session data: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, o.procs.OAuthCreateSession), sessionJSON).Scan(&success, &errMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to create session: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errMsg, "failed to create session")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetByRefreshToken implements lookup.OAuthUserStore.
|
||||||
|
func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
var data []byte
|
||||||
|
err := o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, o.procs.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &data)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errMsg, "invalid or expired refresh token")
|
||||||
|
}
|
||||||
|
var session lookup.OAuthRefreshSession
|
||||||
|
if err := json.Unmarshal(normalizeTimes(data), &session); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse session data: %w", err)
|
||||||
|
}
|
||||||
|
return &session, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateRefreshToken implements lookup.OAuthUserStore.
|
||||||
|
func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
|
||||||
|
updateJSON, err := json.Marshal(map[string]any{
|
||||||
|
"user_id": userID,
|
||||||
|
"old_refresh_token": oldRefreshToken,
|
||||||
|
"new_session_token": newSessionToken,
|
||||||
|
"new_access_token": newAccessToken,
|
||||||
|
"new_refresh_token": newRefreshToken,
|
||||||
|
"expires_at": expiresAt,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal update data: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, o.procs.OAuthUpdateRefreshToken), updateJSON).Scan(&success, &errMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to update session: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errMsg, "failed to update session")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUser implements lookup.OAuthUserStore.
|
||||||
|
func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
var data []byte
|
||||||
|
err := o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, o.procs.OAuthGetUser), userID).Scan(&success, &errMsg, &data)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get user data: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errMsg, "failed to get user data")
|
||||||
|
}
|
||||||
|
var userCtx sectypes.UserContext
|
||||||
|
if err := json.Unmarshal(data, &userCtx); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||||
|
}
|
||||||
|
return &userCtx, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthClients implements lookup.OAuthClientStore with the resolvespec_oauth_* server procedures.
|
||||||
|
type OAuthClients struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.OAuthClientStore = (*OAuthClients)(nil)
|
||||||
|
|
||||||
|
// NewOAuthClients creates the procedure-backed OAuthClientStore.
|
||||||
|
func NewOAuthClients(run Runner, procs lookup.ProcNames) *OAuthClients {
|
||||||
|
return &OAuthClients{run: run, procs: procs}
|
||||||
|
}
|
||||||
|
|
||||||
|
// callData runs a `(p_success, p_error, p_data)` procedure with one argument.
|
||||||
|
func (o *OAuthClients) callData(ctx context.Context, proc string, arg any) (data []byte, ok bool, errMsg sql.NullString, err error) {
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, proc), arg).Scan(&ok, &errMsg, &data)
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// callNoData runs a `(p_success, p_error)` procedure with one argument.
|
||||||
|
func (o *OAuthClients) callNoData(ctx context.Context, proc string, arg any) (ok bool, errMsg sql.NullString, err error) {
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1)
|
||||||
|
`, proc), arg).Scan(&ok, &errMsg)
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterClient implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
|
||||||
|
input, err := json.Marshal(client)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal client: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
var data []byte
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, o.procs.OAuthRegisterClient), input).Scan(&success, &errMsg, &data)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errMsg, "failed to register client")
|
||||||
|
}
|
||||||
|
var result sectypes.OAuthServerClient
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse registered client: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetClient implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||||
|
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthGetClient, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return nil, failure(errMsg, "client not found")
|
||||||
|
}
|
||||||
|
var result sectypes.OAuthServerClient
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse client: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveCode implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
|
||||||
|
input, err := json.Marshal(code) //nolint:gosec // G117: intentional: field must be serialized
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal code: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errMsg sql.NullString
|
||||||
|
err = o.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, o.procs.OAuthSaveCode), input).Scan(&success, &errMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to save code: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errMsg, "failed to save code")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExchangeCode implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||||
|
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthExchangeCode, code)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return nil, failure(errMsg, "invalid or expired code")
|
||||||
|
}
|
||||||
|
var result sectypes.OAuthCode
|
||||||
|
if err := json.Unmarshal(normalizeTimes(data), &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse code data: %w", err)
|
||||||
|
}
|
||||||
|
result.Code = code
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Introspect implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
|
||||||
|
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthIntrospect, token)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to introspect token: %w", err)
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return nil, failure(errMsg, "introspection failed")
|
||||||
|
}
|
||||||
|
var result sectypes.OAuthTokenInfo
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse token info: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Revoke implements lookup.OAuthClientStore.
|
||||||
|
func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
|
||||||
|
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthRevoke, token)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to revoke token: %w", err)
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return failure(errMsg, "failed to revoke token")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,282 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passkey implements lookup.PasskeyStore with the resolvespec_passkey_* procedures.
|
||||||
|
// Credential ids cross the lookup interface as base64 text; the procedures that take a
|
||||||
|
// bytea credential id receive the decoded bytes.
|
||||||
|
type Passkey struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.PasskeyStore = (*Passkey)(nil)
|
||||||
|
|
||||||
|
// NewPasskey creates the procedure-backed PasskeyStore.
|
||||||
|
func NewPasskey(run Runner, procs lookup.ProcNames) *Passkey {
|
||||||
|
return &Passkey{run: run, procs: procs}
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeCredentialID(b64 string) ([]byte, error) {
|
||||||
|
id, err := base64.StdEncoding.DecodeString(b64)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid credential ID: %w", err)
|
||||||
|
}
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
|
||||||
|
credJSON, err := json.Marshal(map[string]any{
|
||||||
|
"user_id": rec.UserID,
|
||||||
|
"credential_id": rec.CredentialID,
|
||||||
|
"public_key": rec.PublicKey,
|
||||||
|
"attestation_type": rec.AttestationType,
|
||||||
|
"sign_count": rec.SignCount,
|
||||||
|
"transports": rec.Transports,
|
||||||
|
"backup_eligible": rec.BackupEligible,
|
||||||
|
"backup_state": rec.BackupState,
|
||||||
|
"name": rec.Name,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to marshal credential data: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var credentialID sql.NullInt64
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential_id FROM %s($1::jsonb)`, p.procs.PasskeyStoreCredential)
|
||||||
|
return db.QueryRowContext(ctx, query, string(credJSON)).Scan(&success, &errorMsg, &credentialID)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to store credential: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return 0, failure(errorMsg, "failed to store credential")
|
||||||
|
}
|
||||||
|
return credentialID.Int64, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) Get(ctx context.Context, credentialID string) (userID int, signCount uint32, err error) {
|
||||||
|
raw, err := decodeCredentialID(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, err
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg, credentialJSON sql.NullString
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential::text FROM %s($1)`, p.procs.PasskeyGetCredential)
|
||||||
|
return db.QueryRowContext(ctx, query, raw).Scan(&success, &errorMsg, &credentialJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return 0, 0, failure(errorMsg, "credential not found")
|
||||||
|
}
|
||||||
|
var cred struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
SignCount uint32 `json:"sign_count"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(normalizeTimes([]byte(credentialJSON.String)), &cred); err != nil {
|
||||||
|
return 0, 0, fmt.Errorf("failed to parse credential: %w", err)
|
||||||
|
}
|
||||||
|
return cred.UserID, cred.SignCount, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateCounter implements lookup.PasskeyStore. Like the code it replaces, it only reports
|
||||||
|
// an error when the query itself fails; the procedure's success flag is not checked.
|
||||||
|
func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
|
||||||
|
raw, err := decodeCredentialID(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var cloneWarning sql.NullBool
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_clone_warning FROM %s($1, $2)`, p.procs.PasskeyUpdateCounter)
|
||||||
|
return db.QueryRowContext(ctx, query, raw, newCounter).Scan(&success, &errorMsg, &cloneWarning)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return cloneWarning.Valid && cloneWarning.Bool, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// List implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, credentialsJSON sql.NullString
|
||||||
|
err := p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetUserCredentials)
|
||||||
|
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &credentialsJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "failed to get credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
var rawCreds []struct {
|
||||||
|
ID int `json:"id"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
CredentialID string `json:"credential_id"`
|
||||||
|
PublicKey string `json:"public_key"`
|
||||||
|
AttestationType string `json:"attestation_type"`
|
||||||
|
AAGUID string `json:"aaguid"`
|
||||||
|
SignCount uint32 `json:"sign_count"`
|
||||||
|
CloneWarning bool `json:"clone_warning"`
|
||||||
|
Transports []string `json:"transports"`
|
||||||
|
BackupEligible bool `json:"backup_eligible"`
|
||||||
|
BackupState bool `json:"backup_state"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
LastUsedAt time.Time `json:"last_used_at"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(normalizeTimes([]byte(credentialsJSON.String)), &rawCreds); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
credentials := make([]sectypes.PasskeyCredential, 0, len(rawCreds))
|
||||||
|
for i := range rawCreds {
|
||||||
|
raw := rawCreds[i]
|
||||||
|
credID, err := base64.StdEncoding.DecodeString(raw.CredentialID)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
pubKey, err := base64.StdEncoding.DecodeString(raw.PublicKey)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
aaguid, _ := base64.StdEncoding.DecodeString(raw.AAGUID)
|
||||||
|
credentials = append(credentials, sectypes.PasskeyCredential{
|
||||||
|
ID: fmt.Sprintf("%d", raw.ID),
|
||||||
|
UserID: raw.UserID,
|
||||||
|
CredentialID: credID,
|
||||||
|
PublicKey: pubKey,
|
||||||
|
AttestationType: raw.AttestationType,
|
||||||
|
AAGUID: aaguid,
|
||||||
|
SignCount: raw.SignCount,
|
||||||
|
CloneWarning: raw.CloneWarning,
|
||||||
|
Transports: raw.Transports,
|
||||||
|
BackupEligible: raw.BackupEligible,
|
||||||
|
BackupState: raw.BackupState,
|
||||||
|
Name: raw.Name,
|
||||||
|
CreatedAt: raw.CreatedAt,
|
||||||
|
LastUsedAt: raw.LastUsedAt,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return credentials, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error {
|
||||||
|
raw, err := decodeCredentialID(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, p.procs.PasskeyDeleteCredential)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, raw).Scan(&success, &errorMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to delete credential: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errorMsg, "failed to delete credential")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Rename implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error {
|
||||||
|
raw, err := decodeCredentialID(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3)`, p.procs.PasskeyUpdateName)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, raw, name).Scan(&success, &errorMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to update credential name: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errorMsg, "failed to update credential name")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ByUsername implements lookup.PasskeyStore.
|
||||||
|
func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, credentialsJSON sql.NullString
|
||||||
|
var userID sql.NullInt64
|
||||||
|
err := p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetCredsByUsername)
|
||||||
|
return db.QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return 0, nil, failure(errorMsg, "failed to get credentials")
|
||||||
|
}
|
||||||
|
var creds []lookup.PasskeyCredentialRef
|
||||||
|
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
|
||||||
|
return 0, nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||||
|
}
|
||||||
|
return int(userID.Int64), creds, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Login implements lookup.PasskeyStore: it creates the session for a user whose passkey
|
||||||
|
// assertion was already verified.
|
||||||
|
func (p *Passkey) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||||
|
reqData := map[string]any{"user_id": userID}
|
||||||
|
if claims != nil {
|
||||||
|
if ip, ok := claims["ip_address"].(string); ok {
|
||||||
|
reqData["ip_address"] = ip
|
||||||
|
}
|
||||||
|
if ua, ok := claims["user_agent"].(string); ok {
|
||||||
|
reqData["user_agent"] = ua
|
||||||
|
}
|
||||||
|
}
|
||||||
|
reqJSON, err := json.Marshal(reqData)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal passkey login request: %w", err)
|
||||||
|
}
|
||||||
|
var success bool
|
||||||
|
var errorMsg, dataJSON sql.NullString
|
||||||
|
err = p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, p.procs.PasskeyLogin)
|
||||||
|
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("passkey login query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "passkey login failed")
|
||||||
|
}
|
||||||
|
var response sectypes.LoginResponse
|
||||||
|
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse passkey login response: %w", err)
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Policy implements lookup.PolicyStore with the column and row security procedures.
|
||||||
|
type Policy struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.PolicyStore = (*Policy)(nil)
|
||||||
|
|
||||||
|
// NewPolicy creates the procedure-backed PolicyStore.
|
||||||
|
func NewPolicy(run Runner, procs lookup.ProcNames) *Policy { return &Policy{run: run, procs: procs} }
|
||||||
|
|
||||||
|
// ColumnSecurity implements lookup.PolicyStore.
|
||||||
|
func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var rulesJSON []byte
|
||||||
|
err := p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_rules FROM %s($1, $2, $3)`, p.procs.ColumnSecurity)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, schema, table).Scan(&success, &errorMsg, &rulesJSON)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to load column security: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, failure(errorMsg, "failed to load column security")
|
||||||
|
}
|
||||||
|
|
||||||
|
type securityRecord struct {
|
||||||
|
Control string `json:"control"`
|
||||||
|
Accesstype string `json:"accesstype"`
|
||||||
|
JSONValue string `json:"jsonvalue"`
|
||||||
|
}
|
||||||
|
var records []securityRecord
|
||||||
|
if err := json.Unmarshal(rulesJSON, &records); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse security rules: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var rules []sectypes.ColumnSecurity
|
||||||
|
for _, rec := range records {
|
||||||
|
parts := strings.Split(rec.Control, ".")
|
||||||
|
if len(parts) < 3 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rules = append(rules, sectypes.ColumnSecurity{
|
||||||
|
Schema: schema,
|
||||||
|
Tablename: table,
|
||||||
|
Path: parts[2:],
|
||||||
|
Accesstype: rec.Accesstype,
|
||||||
|
UserID: userID,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RowSecurity implements lookup.PolicyStore. userRef is unwrapped to a scalar user id
|
||||||
|
// because the procedure's p_user_id is an integer.
|
||||||
|
func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
|
||||||
|
switch v := userRef.(type) {
|
||||||
|
case *sectypes.UserContext:
|
||||||
|
if v != nil {
|
||||||
|
userRef = v.UserID
|
||||||
|
}
|
||||||
|
case sectypes.UserContext:
|
||||||
|
userRef = v.UserID
|
||||||
|
}
|
||||||
|
|
||||||
|
var template sql.NullString
|
||||||
|
var hasBlock sql.NullBool
|
||||||
|
err := p.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.procs.RowSecurity)
|
||||||
|
return db.QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
|
||||||
|
}
|
||||||
|
return sectypes.RowSecurity{
|
||||||
|
Schema: schema,
|
||||||
|
Tablename: table,
|
||||||
|
UserID: userRef,
|
||||||
|
Template: template.String,
|
||||||
|
HasBlock: hasBlock.Bool,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,92 @@
|
|||||||
|
// Package procedure is the stored-procedure backend of lookup: it calls the
|
||||||
|
// resolvespec_* functions (names from lookup.ProcNames) and keeps their
|
||||||
|
// p_success / p_error / p_data contract. Error texts match the ones the security
|
||||||
|
// package returned before the extraction.
|
||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
|
||||||
|
type Runner interface {
|
||||||
|
Run(run func(*sql.DB) error) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// RunFunc adapts a function to Runner. The security package passes its own
|
||||||
|
// reconnecting helper this way.
|
||||||
|
type RunFunc func(run func(*sql.DB) error) error
|
||||||
|
|
||||||
|
// Run implements Runner.
|
||||||
|
func (f RunFunc) Run(run func(*sql.DB) error) error { return f(run) }
|
||||||
|
|
||||||
|
// DB is a standalone Runner over a *sql.DB with an optional reconnect factory.
|
||||||
|
type DB struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
db *sql.DB
|
||||||
|
factory func() (*sql.DB, error)
|
||||||
|
onReconnect func()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDB wraps db. factory (optional) is called to obtain a fresh handle when the current
|
||||||
|
// one is closed; onReconnect (optional) runs after a successful reconnect, e.g. to reset
|
||||||
|
// cached procedure probes.
|
||||||
|
func NewDB(db *sql.DB, factory func() (*sql.DB, error), onReconnect func()) *DB {
|
||||||
|
return &DB{db: db, factory: factory, onReconnect: onReconnect}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns the current handle.
|
||||||
|
func (d *DB) Get() *sql.DB {
|
||||||
|
d.mu.RLock()
|
||||||
|
defer d.mu.RUnlock()
|
||||||
|
return d.db
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *DB) reconnect() error {
|
||||||
|
if d.factory == nil {
|
||||||
|
return fmt.Errorf("no db factory configured for reconnect")
|
||||||
|
}
|
||||||
|
newDB, err := d.factory()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
d.mu.Lock()
|
||||||
|
d.db = newDB
|
||||||
|
d.mu.Unlock()
|
||||||
|
if d.onReconnect != nil {
|
||||||
|
d.onReconnect()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run implements Runner.
|
||||||
|
func (d *DB) Run(run func(*sql.DB) error) error {
|
||||||
|
db := d.Get()
|
||||||
|
if db == nil {
|
||||||
|
return fmt.Errorf("database connection is nil")
|
||||||
|
}
|
||||||
|
err := run(db)
|
||||||
|
if IsClosed(err) {
|
||||||
|
if reconnErr := d.reconnect(); reconnErr == nil {
|
||||||
|
err = run(d.Get())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsClosed reports whether err indicates the *sql.DB has been closed.
|
||||||
|
func IsClosed(err error) bool {
|
||||||
|
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// failure builds the error for p_success = false: the procedure's own message when it
|
||||||
|
// returned one, otherwise def.
|
||||||
|
func failure(errMsg sql.NullString, def string) error {
|
||||||
|
if errMsg.Valid {
|
||||||
|
return fmt.Errorf("%s", errMsg.String)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("%s", def)
|
||||||
|
}
|
||||||
@@ -0,0 +1,226 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/base64"
|
||||||
|
"errors"
|
||||||
|
"regexp"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/DATA-DOG/go-sqlmock"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newMock(t *testing.T) (*DB, sqlmock.Sqlmock) {
|
||||||
|
t.Helper()
|
||||||
|
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
return NewDB(db, nil, nil), mock
|
||||||
|
}
|
||||||
|
|
||||||
|
func q(s string) string { return regexp.QuoteMeta(s) }
|
||||||
|
|
||||||
|
func TestAuthLoginCallsProcedureWithJSON(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
procs := lookup.DefaultProcNames()
|
||||||
|
procs.Login = "custom_login"
|
||||||
|
a := NewAuth(run, procs)
|
||||||
|
|
||||||
|
mock.ExpectQuery(q("SELECT p_success, p_error, p_data::text FROM custom_login($1::jsonb)")).
|
||||||
|
WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
|
||||||
|
AddRow(true, nil, `{"token":"sess_1","user":{"user_id":7,"user_name":"bob"}}`))
|
||||||
|
|
||||||
|
resp, err := a.Login(context.Background(), sectypes.LoginRequest{Username: "bob", Password: "x"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if resp.Token != "sess_1" || resp.User == nil || resp.User.UserID != 7 {
|
||||||
|
t.Fatalf("unexpected response: %+v", resp)
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthLoginFailureUsesProcedureMessage(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
a := NewAuth(run, lookup.DefaultProcNames())
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, "bad credentials", nil))
|
||||||
|
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "bad credentials" {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
|
||||||
|
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() == "" {
|
||||||
|
t.Fatalf("expected default error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunnerReconnectsOnClosedDB(t *testing.T) {
|
||||||
|
first, _, err := sqlmock.New()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = first.Close()
|
||||||
|
|
||||||
|
second, mock, err := sqlmock.New()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer second.Close()
|
||||||
|
|
||||||
|
reconnected := false
|
||||||
|
run := NewDB(first, func() (*sql.DB, error) { return second, nil }, func() { reconnected = true })
|
||||||
|
mock.ExpectQuery("SELECT 1").WillReturnRows(sqlmock.NewRows([]string{"x"}).AddRow(1))
|
||||||
|
|
||||||
|
err = run.Run(func(db *sql.DB) error {
|
||||||
|
var x int
|
||||||
|
return db.QueryRow("SELECT 1").Scan(&x)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !reconnected || run.Get() != second {
|
||||||
|
t.Fatal("expected reconnect to the new handle")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunnerNoFactoryReturnsClosedError(t *testing.T) {
|
||||||
|
db, _, _ := sqlmock.New()
|
||||||
|
_ = db.Close()
|
||||||
|
err := NewDB(db, nil, nil).Run(func(db *sql.DB) error { return db.QueryRow("SELECT 1").Scan(new(int)) })
|
||||||
|
if !IsClosed(err) {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPasskeyGetDecodesCredentialID(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||||
|
raw := []byte{1, 2, 3, 4}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_passkey_get_credential").WithArgs(raw).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_credential"}).
|
||||||
|
AddRow(true, nil, `{"user_id":9,"sign_count":4}`))
|
||||||
|
uid, count, err := p.Get(context.Background(), base64.StdEncoding.EncodeToString(raw))
|
||||||
|
if err != nil || uid != 9 || count != 4 {
|
||||||
|
t.Fatalf("got %d %d %v", uid, count, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPasskeyInvalidBase64(t *testing.T) {
|
||||||
|
run, _ := newMock(t)
|
||||||
|
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||||
|
if _, _, err := p.Get(context.Background(), "***"); err == nil {
|
||||||
|
t.Fatal("expected error")
|
||||||
|
}
|
||||||
|
if err := p.Delete(context.Background(), 1, "***"); err == nil {
|
||||||
|
t.Fatal("expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPasskeyUpdateCounterReportsCloneWarning(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||||
|
id := base64.StdEncoding.EncodeToString([]byte("abc"))
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_passkey_update_counter").WithArgs([]byte("abc"), uint32(5)).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_clone_warning"}).AddRow(true, nil, true))
|
||||||
|
warn, err := p.UpdateCounter(context.Background(), id, 5)
|
||||||
|
if err != nil || !warn {
|
||||||
|
t.Fatalf("got %v %v", warn, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_passkey_update_counter").WillReturnError(errors.New("boom"))
|
||||||
|
if _, err := p.UpdateCounter(context.Background(), id, 6); err == nil {
|
||||||
|
t.Fatal("expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPasskeyByUsername(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||||
|
mock.ExpectQuery("resolvespec_passkey_get_credentials_by_username").WithArgs("bob").
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id", "p_credentials"}).
|
||||||
|
AddRow(true, nil, 3, `[{"credential_id":"YWJj","transports":["usb"]}]`))
|
||||||
|
uid, refs, err := p.ByUsername(context.Background(), "bob")
|
||||||
|
if err != nil || uid != 3 || len(refs) != 1 || refs[0].CredentialID != "YWJj" || refs[0].Transports[0] != "usb" {
|
||||||
|
t.Fatalf("got %d %+v %v", uid, refs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthUsersGetOrCreateUser(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
o := NewOAuthUsers(run, lookup.DefaultProcNames())
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, 11))
|
||||||
|
id, err := o.GetOrCreateUser(context.Background(), §ypes.UserContext{UserName: "u", Email: "e"}, "github")
|
||||||
|
if err != nil || id != 11 {
|
||||||
|
t.Fatalf("got %d %v", id, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, nil))
|
||||||
|
if _, err := o.GetOrCreateUser(context.Background(), §ypes.UserContext{}, "github"); err == nil || err.Error() != "user ID not returned" {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthUsersRefreshRoundTrip(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
o := NewOAuthUsers(run, lookup.DefaultProcNames())
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_getrefreshtoken").WithArgs("r1").
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
|
||||||
|
AddRow(true, nil, `{"user_id":2,"access_token":"a","token_type":"Bearer","expiry":"2030-01-01T00:00:00Z"}`))
|
||||||
|
s, err := o.GetByRefreshToken(context.Background(), "r1")
|
||||||
|
if err != nil || s.UserID != 2 || s.AccessToken != "a" {
|
||||||
|
t.Fatalf("got %+v %v", s, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_updaterefreshtoken").WithArgs(sqlmock.AnyArg()).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(false, "session not found"))
|
||||||
|
err = o.UpdateRefreshToken(context.Background(), 2, "r1", "s2", "a2", "r2", time.Now())
|
||||||
|
if err == nil || err.Error() != "session not found" {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthClientsExchangeCodeSetsCode(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
c := NewOAuthClients(run, lookup.DefaultProcNames())
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("abc").
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, `{"client_id":"cid"}`))
|
||||||
|
code, err := c.ExchangeCode(context.Background(), "abc")
|
||||||
|
if err != nil || code.Code != "abc" || code.ClientID != "cid" {
|
||||||
|
t.Fatalf("got %+v %v", code, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("zzz").
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
|
||||||
|
if _, err := c.ExchangeCode(context.Background(), "zzz"); err == nil || err.Error() != "invalid or expired code" {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthClientsRevoke(t *testing.T) {
|
||||||
|
run, mock := newMock(t)
|
||||||
|
c := NewOAuthClients(run, lookup.DefaultProcNames())
|
||||||
|
mock.ExpectQuery("resolvespec_oauth_revoke").WithArgs("t").
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(true, nil))
|
||||||
|
if err := c.Revoke(context.Background(), "t"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
package procedure
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TOTP implements lookup.TOTPStore with the resolvespec_totp_* procedures.
|
||||||
|
type TOTP struct {
|
||||||
|
run Runner
|
||||||
|
procs lookup.ProcNames
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ lookup.TOTPStore = (*TOTP)(nil)
|
||||||
|
|
||||||
|
// NewTOTP creates the procedure-backed TOTPStore.
|
||||||
|
func NewTOTP(run Runner, procs lookup.ProcNames) *TOTP { return &TOTP{run: run, procs: procs} }
|
||||||
|
|
||||||
|
// exec runs a "p_success, p_error" procedure and maps failure to an error.
|
||||||
|
func (t *TOTP) exec(ctx context.Context, query, op, def string, args ...any) error {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err := t.run.Run(func(db *sql.DB) error {
|
||||||
|
return db.QueryRowContext(ctx, query, args...).Scan(&success, &errorMsg)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s query failed: %w", op, err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return failure(errorMsg, def)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable implements lookup.TOTPStore.
|
||||||
|
func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
|
||||||
|
codesJSON, err := json.Marshal(hashedCodes)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||||
|
}
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, t.procs.TOTPEnable)
|
||||||
|
return t.exec(ctx, query, "enable 2FA", "failed to enable 2FA", userID, secret, string(codesJSON))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Disable implements lookup.TOTPStore.
|
||||||
|
func (t *TOTP) Disable(ctx context.Context, userID int) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, t.procs.TOTPDisable)
|
||||||
|
return t.exec(ctx, query, "disable 2FA", "failed to disable 2FA", userID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Status implements lookup.TOTPStore.
|
||||||
|
func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) {
|
||||||
|
var success, enabled bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err := t.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, t.procs.TOTPGetStatus)
|
||||||
|
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &enabled)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("get 2FA status query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return false, failure(errorMsg, "failed to get 2FA status")
|
||||||
|
}
|
||||||
|
return enabled, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Secret implements lookup.TOTPStore.
|
||||||
|
func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg, secret sql.NullString
|
||||||
|
err := t.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, t.procs.TOTPGetSecret)
|
||||||
|
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &secret)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return "", failure(errorMsg, "failed to get 2FA secret")
|
||||||
|
}
|
||||||
|
if !secret.Valid {
|
||||||
|
return "", fmt.Errorf("2FA secret not found")
|
||||||
|
}
|
||||||
|
return secret.String, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegenerateBackupCodes implements lookup.TOTPStore.
|
||||||
|
func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
|
||||||
|
codesJSON, err := json.Marshal(hashedCodes)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||||
|
}
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, t.procs.TOTPRegenerateBackup)
|
||||||
|
return t.exec(ctx, query, "regenerate backup codes", "failed to regenerate backup codes", userID, string(codesJSON))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateBackupCode implements lookup.TOTPStore. A failure without a message means
|
||||||
|
// "not valid", not an error.
|
||||||
|
func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
|
||||||
|
var success, valid bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
err := t.run.Run(func(db *sql.DB) error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, t.procs.TOTPValidateBackupCode)
|
||||||
|
return db.QueryRowContext(ctx, query, userID, codeHash).Scan(&success, &errorMsg, &valid)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("validate backup code query failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return false, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return valid, nil
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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" //nolint:gosec // table name, not a credential
|
||||||
|
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) }
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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 }
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -441,7 +441,7 @@ func ExampleOAuth2Complete() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func setupOAuth2Tables(db *sql.DB) {
|
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
|
// This is a helper function - in production, use migrations
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
|
||||||
"golang.org/x/oauth2"
|
"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
|
// 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) {
|
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.src.get().OAuthUser.GetOrCreateUser(ctx, userCtx, providerName)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// oauth2CreateSession creates a new OAuth2 session using stored procedure
|
// 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 {
|
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.src.get().OAuthUser.CreateSession(ctx, lookup.OAuthSession{
|
||||||
return a.oauth2CreateSessionDirect(ctx, sessionToken, userID, token, expiresAt, providerName)
|
SessionToken: sessionToken,
|
||||||
}
|
UserID: userID,
|
||||||
|
AccessToken: token.AccessToken,
|
||||||
sessionData := map[string]interface{}{
|
RefreshToken: token.RefreshToken,
|
||||||
"session_token": sessionToken,
|
TokenType: token.TokenType,
|
||||||
"user_id": userID,
|
ExpiresAt: expiresAt,
|
||||||
"access_token": token.AccessToken,
|
Provider: providerName,
|
||||||
"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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// validateState validates state using in-memory storage
|
// 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
|
// 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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -447,12 +376,12 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Update session in database with new tokens
|
// 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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get user data
|
// Get user data
|
||||||
userCtx, err := a.oauthGetUserByID(ctx, session.UserID)
|
userCtx, err := a.src.get().OAuthUser.GetUser(ctx, session.UserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -2,229 +2,34 @@ package security
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"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.
|
// OAuthRegisterClient persists an OAuth2 client registration.
|
||||||
func (a *DatabaseAuthenticator) OAuthRegisterClient(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
|
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.src.get().OAuthClient.RegisterClient(ctx, client)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OAuthGetClient retrieves a registered client by ID.
|
// OAuthGetClient retrieves a registered client by ID.
|
||||||
func (a *DatabaseAuthenticator) OAuthGetClient(ctx context.Context, clientID string) (*OAuthServerClient, error) {
|
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.src.get().OAuthClient.GetClient(ctx, clientID)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OAuthSaveCode persists an authorization code.
|
// OAuthSaveCode persists an authorization code.
|
||||||
func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCode) error {
|
func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCode) error {
|
||||||
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthSaveCode) {
|
return a.src.get().OAuthClient.SaveCode(ctx, code)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OAuthExchangeCode retrieves and deletes an authorization code (single use).
|
// OAuthExchangeCode retrieves and deletes an authorization code (single use).
|
||||||
func (a *DatabaseAuthenticator) OAuthExchangeCode(ctx context.Context, code string) (*OAuthCode, error) {
|
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.src.get().OAuthClient.ExchangeCode(ctx, code)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OAuthIntrospectToken validates a token and returns its metadata (RFC 7662).
|
// OAuthIntrospectToken validates a token and returns its metadata (RFC 7662).
|
||||||
func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token string) (*OAuthTokenInfo, error) {
|
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.src.get().OAuthClient.Introspect(ctx, token)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OAuthRevokeToken revokes a token by deleting the session (RFC 7009).
|
// OAuthRevokeToken revokes a token by deleting the session (RFC 7009).
|
||||||
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
|
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
|
||||||
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthRevoke) {
|
return a.src.get().OAuthClient.Revoke(ctx, token)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -19,7 +19,7 @@ import (
|
|||||||
func newTestOAuthServer(t *testing.T) (*OAuthServer, *DatabaseAuthenticator) {
|
func newTestOAuthServer(t *testing.T) (*OAuthServer, *DatabaseAuthenticator) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
db := newDirectTestDB(t)
|
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)
|
srv := NewOAuthServer(OAuthServerConfig{Issuer: "https://auth.example.com", PersistCodes: true}, auth)
|
||||||
t.Cleanup(srv.Close)
|
t.Cleanup(srv.Close)
|
||||||
return srv, auth
|
return srv, auth
|
||||||
|
|||||||
@@ -3,118 +3,8 @@ package security
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"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
|
// PasskeyProvider handles passkey registration and authentication
|
||||||
type PasskeyProvider interface {
|
type PasskeyProvider interface {
|
||||||
// BeginRegistration creates registration options for a new passkey
|
// BeginRegistration creates registration options for a new passkey
|
||||||
|
|||||||
@@ -5,26 +5,21 @@ import (
|
|||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends"
|
||||||
)
|
)
|
||||||
|
|
||||||
// DatabasePasskeyProvider implements PasskeyProvider using database storage
|
// DatabasePasskeyProvider implements PasskeyProvider on top of the lookup package
|
||||||
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
// (stored procedures on Postgres by default, direct SQL elsewhere).
|
||||||
type DatabasePasskeyProvider struct {
|
type DatabasePasskeyProvider struct {
|
||||||
db *sql.DB
|
src *lookupSource
|
||||||
dbMu sync.RWMutex
|
rpID string // Relying Party ID (domain)
|
||||||
dbFactory func() (*sql.DB, error)
|
rpName string // Relying Party display name
|
||||||
rpID string // Relying Party ID (domain)
|
rpOrigin string // Expected origin for WebAuthn
|
||||||
rpName string // Relying Party display name
|
timeout int64 // Timeout in milliseconds (default: 60000)
|
||||||
rpOrigin string // Expected origin for WebAuthn
|
|
||||||
timeout int64 // Timeout in milliseconds (default: 60000)
|
|
||||||
sqlNames *SQLNames
|
|
||||||
tableNames *TableNames
|
|
||||||
queryMode QueryMode
|
|
||||||
capability *dbCapability
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DatabasePasskeyProviderOptions configures the passkey provider
|
// DatabasePasskeyProviderOptions configures the passkey provider
|
||||||
@@ -37,12 +32,10 @@ type DatabasePasskeyProviderOptions struct {
|
|||||||
RPOrigin string
|
RPOrigin string
|
||||||
// Timeout is the timeout for operations in milliseconds (default: 60000)
|
// Timeout is the timeout for operations in milliseconds (default: 60000)
|
||||||
Timeout int64
|
Timeout int64
|
||||||
// SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames().
|
// Lookup selects dialect, query mode and procedure/table/column names.
|
||||||
SQLNames *SQLNames
|
Lookup lookup.Config
|
||||||
// TableNames provides custom table names for Direct mode. If nil, uses DefaultTableNames().
|
// LookupProvider, when set, is used instead of building one from Lookup and the db.
|
||||||
TableNames *TableNames
|
LookupProvider *lookup.Provider
|
||||||
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
|
|
||||||
QueryMode QueryMode
|
|
||||||
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
||||||
// If nil, reconnection is disabled.
|
// If nil, reconnection is disabled.
|
||||||
DBFactory func() (*sql.DB, error)
|
DBFactory func() (*sql.DB, error)
|
||||||
@@ -53,60 +46,20 @@ func NewDatabasePasskeyProvider(db *sql.DB, opts DatabasePasskeyProviderOptions)
|
|||||||
if opts.Timeout == 0 {
|
if opts.Timeout == 0 {
|
||||||
opts.Timeout = 60000 // 60 seconds default
|
opts.Timeout = 60000 // 60 seconds default
|
||||||
}
|
}
|
||||||
|
src := newLookupSource(db)
|
||||||
sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames)
|
src.cfg = opts.Lookup
|
||||||
tableNames := resolveTableNames(opts.TableNames)
|
src.provider = opts.LookupProvider
|
||||||
|
src.opts = backends.Options{DBFactory: opts.DBFactory}
|
||||||
return &DatabasePasskeyProvider{
|
return &DatabasePasskeyProvider{
|
||||||
db: db,
|
src: src,
|
||||||
dbFactory: opts.DBFactory,
|
rpID: opts.RPID,
|
||||||
rpID: opts.RPID,
|
rpName: opts.RPName,
|
||||||
rpName: opts.RPName,
|
rpOrigin: opts.RPOrigin,
|
||||||
rpOrigin: opts.RPOrigin,
|
timeout: opts.Timeout,
|
||||||
timeout: opts.Timeout,
|
|
||||||
sqlNames: sqlNames,
|
|
||||||
tableNames: tableNames,
|
|
||||||
queryMode: opts.QueryMode,
|
|
||||||
capability: newDBCapability(),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *DatabasePasskeyProvider) getDB() *sql.DB {
|
func (p *DatabasePasskeyProvider) store() lookup.PasskeyStore { return p.src.get().Passkey }
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// BeginRegistration creates registration options for a new passkey
|
// BeginRegistration creates registration options for a new passkey
|
||||||
func (p *DatabasePasskeyProvider) BeginRegistration(ctx context.Context, userID int, username, displayName string) (*PasskeyRegistrationOptions, error) {
|
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)
|
credIDB64 := base64.StdEncoding.EncodeToString(response.RawID)
|
||||||
pubKeyB64 := base64.StdEncoding.EncodeToString(response.Response.AttestationObject)
|
pubKeyB64 := base64.StdEncoding.EncodeToString(response.Response.AttestationObject)
|
||||||
|
|
||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyStoreCredential) {
|
credentialID, err := p.store().Store(ctx, lookup.PasskeyCredentialRecord{
|
||||||
credentialID, err := p.storeCredentialDirect(ctx, storeCredentialParams{
|
UserID: userID,
|
||||||
UserID: userID,
|
CredentialID: credIDB64,
|
||||||
CredentialID: credIDB64,
|
PublicKey: pubKeyB64,
|
||||||
PublicKey: pubKeyB64,
|
AttestationType: "none",
|
||||||
AttestationType: "none",
|
Transports: response.Transports,
|
||||||
SignCount: 0,
|
Name: "Passkey",
|
||||||
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)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to marshal credential data: %w", err)
|
return nil, 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 &PasskeyCredential{
|
return &PasskeyCredential{
|
||||||
ID: fmt.Sprintf("%d", credentialID.Int64),
|
ID: fmt.Sprintf("%d", credentialID),
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
CredentialID: response.RawID,
|
CredentialID: response.RawID,
|
||||||
PublicKey: response.Response.AttestationObject,
|
PublicKey: response.Response.AttestationObject,
|
||||||
@@ -260,41 +164,15 @@ func (p *DatabasePasskeyProvider) BeginAuthentication(ctx context.Context, usern
|
|||||||
// If username is provided, get user's credentials
|
// If username is provided, get user's credentials
|
||||||
var allowCredentials []PasskeyCredentialDescriptor
|
var allowCredentials []PasskeyCredentialDescriptor
|
||||||
if username != "" {
|
if username != "" {
|
||||||
var creds []passkeyCredential
|
_, refs, err := p.store().ByUsername(ctx, username)
|
||||||
|
if err != nil {
|
||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) {
|
return nil, err
|
||||||
_, 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)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
creds := refs
|
||||||
|
|
||||||
allowCredentials = make([]PasskeyCredentialDescriptor, 0, len(creds))
|
allowCredentials = make([]PasskeyCredentialDescriptor, 0, len(creds))
|
||||||
for _, cred := range creds {
|
for _, cred := range creds {
|
||||||
credID, err := base64.StdEncoding.DecodeString(cred.ID)
|
credID, err := base64.StdEncoding.DecodeString(cred.CredentialID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -327,214 +205,47 @@ func (p *DatabasePasskeyProvider) CompleteAuthentication(ctx context.Context, re
|
|||||||
|
|
||||||
credIDB64 := base64.StdEncoding.EncodeToString(response.RawID)
|
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
|
// TODO: Verify signature here
|
||||||
// For now, we'll just update the counter as a placeholder
|
// 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)
|
// Update counter (in production, this should be done after successful verification)
|
||||||
newCounter := cred.SignCount + 1
|
cloneWarning, err := store.UpdateCounter(ctx, credIDB64, 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)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to update counter: %w", err)
|
return 0, fmt.Errorf("failed to update counter: %w", err)
|
||||||
}
|
}
|
||||||
|
if cloneWarning {
|
||||||
if cloneWarning.Valid && cloneWarning.Bool {
|
|
||||||
return 0, fmt.Errorf("credential cloning detected")
|
return 0, fmt.Errorf("credential cloning detected")
|
||||||
}
|
}
|
||||||
|
|
||||||
return cred.UserID, nil
|
return userID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetCredentials returns all passkey credentials for a user
|
// GetCredentials returns all passkey credentials for a user
|
||||||
func (p *DatabasePasskeyProvider) GetCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error) {
|
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.store().List(ctx, userID)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteCredential removes a passkey credential
|
// DeleteCredential removes a passkey credential
|
||||||
func (p *DatabasePasskeyProvider) DeleteCredential(ctx context.Context, userID int, credentialID string) error {
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid credential ID: %w", err)
|
return fmt.Errorf("invalid credential ID: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyDeleteCredential) {
|
return p.store().Delete(ctx, userID, credentialID)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateCredentialName updates the friendly name of a credential
|
// UpdateCredentialName updates the friendly name of a credential
|
||||||
func (p *DatabasePasskeyProvider) UpdateCredentialName(ctx context.Context, userID int, credentialID string, name string) error {
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid credential ID: %w", err)
|
return fmt.Errorf("invalid credential ID: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyUpdateName) {
|
return p.store().Rename(ctx, userID, credentialID, name)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -5,7 +5,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
"regexp"
|
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -18,36 +17,6 @@ import (
|
|||||||
"golang.org/x/sync/singleflight"
|
"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
|
// 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
|
// entry is loaded for the user and table. It means "no rules", as opposed to a
|
||||||
// failure, which callers must treat as fatal.
|
// 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.
|
// ErrNoColumnSecurity is the column-security equivalent of ErrNoRowSecurity.
|
||||||
var ErrNoColumnSecurity = errors.New("no column security data")
|
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
|
// SecurityList manages security state and caching
|
||||||
// It wraps a SecurityProvider and provides caching and utility methods
|
// It wraps a SecurityProvider and provides caching and utility methods
|
||||||
type SecurityList struct {
|
type SecurityList struct {
|
||||||
|
|||||||
+96
-812
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||||
|
}
|
||||||
+16
-13
@@ -1,10 +1,13 @@
|
|||||||
package security
|
package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
// KeyStoreAuthenticator implements the Authenticator interface using a KeyStore.
|
// KeyStoreAuthenticator implements the Authenticator interface using a KeyStore.
|
||||||
@@ -17,45 +20,45 @@ import (
|
|||||||
// 2. Authorization: ApiKey <key>
|
// 2. Authorization: ApiKey <key>
|
||||||
// 3. X-API-Key header
|
// 3. X-API-Key header
|
||||||
type KeyStoreAuthenticator struct {
|
type KeyStoreAuthenticator struct {
|
||||||
keyStore KeyStore
|
keyStore security.KeyStore
|
||||||
keyType KeyType // empty = accept any type
|
keyType sectypes.KeyType // empty = accept any type
|
||||||
authenticateCallback func(r *http.Request) (*UserContext, error)
|
authenticateCallback func(r *http.Request) (*sectypes.UserContext, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewKeyStoreAuthenticator creates a KeyStoreAuthenticator.
|
// NewKeyStoreAuthenticator creates a KeyStoreAuthenticator.
|
||||||
// Pass an empty keyType to accept keys of any type.
|
// 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}
|
return &KeyStoreAuthenticator{keyStore: ks, keyType: keyType}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Login is not supported for keystore authentication.
|
// 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")
|
return nil, fmt.Errorf("keystore authenticator does not support login")
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoginWithCookie is not supported for keystore authentication.
|
// 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")
|
return nil, fmt.Errorf("keystore authenticator does not support login")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Logout is not supported for keystore authentication.
|
// 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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// LogoutWithCookie is not supported for keystore authentication.
|
// 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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetAuthenticateCallback registers a fallback called when key authentication fails.
|
// 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
|
a.authenticateCallback = fn
|
||||||
}
|
}
|
||||||
|
|
||||||
// Authenticate extracts an API key from the request and validates it against the KeyStore.
|
// Authenticate extracts an API key from the request and validates it against the KeyStore.
|
||||||
// Returns a UserContext built from the matching UserKey on success.
|
// 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)
|
rawKey := extractAPIKey(r)
|
||||||
if rawKey == "" {
|
if rawKey == "" {
|
||||||
if a.authenticateCallback != nil {
|
if a.authenticateCallback != nil {
|
||||||
@@ -93,7 +96,7 @@ func extractAPIKey(r *http.Request) string {
|
|||||||
|
|
||||||
// userKeyToUserContext converts a UserKey into a UserContext.
|
// userKeyToUserContext converts a UserKey into a UserContext.
|
||||||
// Scopes are mapped to Roles. Key type and name are stored in Claims.
|
// 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{
|
claims := map[string]any{
|
||||||
"key_type": string(k.KeyType),
|
"key_type": string(k.KeyType),
|
||||||
"key_name": k.Name,
|
"key_name": k.Name,
|
||||||
@@ -109,7 +112,7 @@ func userKeyToUserContext(k *UserKey) *UserContext {
|
|||||||
roles = []string{}
|
roles = []string{}
|
||||||
}
|
}
|
||||||
|
|
||||||
return &UserContext{
|
return §ypes.UserContext{
|
||||||
UserID: k.UserID,
|
UserID: k.UserID,
|
||||||
SessionID: fmt.Sprintf("key:%d", k.ID),
|
SessionID: fmt.Sprintf("key:%d", k.ID),
|
||||||
Roles: roles,
|
Roles: roles,
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
package security
|
package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
@@ -10,6 +10,8 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConfigKeyStore is an in-memory keystore backed by a static slice of UserKey values.
|
// ConfigKeyStore is an in-memory keystore backed by a static slice of UserKey values.
|
||||||
@@ -20,16 +22,16 @@ import (
|
|||||||
// Keys created at runtime via CreateKey are held in memory only and lost on restart.
|
// Keys created at runtime via CreateKey are held in memory only and lost on restart.
|
||||||
type ConfigKeyStore struct {
|
type ConfigKeyStore struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
keys []UserKey
|
keys []sectypes.UserKey
|
||||||
next int64 // monotonic ID counter for runtime-created keys (atomic)
|
next int64 // monotonic ID counter for runtime-created keys (atomic)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewConfigKeyStore creates a ConfigKeyStore seeded with the provided keys.
|
// NewConfigKeyStore creates a ConfigKeyStore seeded with the provided keys.
|
||||||
// Pass nil or an empty slice to start with no pre-loaded 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.
|
// 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
|
var maxID int64
|
||||||
copied := make([]UserKey, len(keys))
|
copied := make([]sectypes.UserKey, len(keys))
|
||||||
copy(copied, keys)
|
copy(copied, keys)
|
||||||
for i := range copied {
|
for i := range copied {
|
||||||
if copied[i].CreatedAt.IsZero() {
|
if copied[i].CreatedAt.IsZero() {
|
||||||
@@ -44,16 +46,16 @@ func NewConfigKeyStore(keys []UserKey) *ConfigKeyStore {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// CreateKey generates a new raw key, stores its SHA-256 hash, and returns the raw key once.
|
// 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)
|
rawBytes := make([]byte, 32)
|
||||||
if _, err := rand.Read(rawBytes); err != nil {
|
if _, err := rand.Read(rawBytes); err != nil {
|
||||||
return nil, fmt.Errorf("failed to generate key material: %w", err)
|
return nil, fmt.Errorf("failed to generate key material: %w", err)
|
||||||
}
|
}
|
||||||
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
||||||
hash := hashSHA256Hex(rawKey)
|
hash := sectypes.HashKey(rawKey)
|
||||||
|
|
||||||
id := atomic.AddInt64(&s.next, 1)
|
id := atomic.AddInt64(&s.next, 1)
|
||||||
key := UserKey{
|
key := sectypes.UserKey{
|
||||||
ID: id,
|
ID: id,
|
||||||
UserID: req.UserID,
|
UserID: req.UserID,
|
||||||
KeyType: req.KeyType,
|
KeyType: req.KeyType,
|
||||||
@@ -70,17 +72,17 @@ func (s *ConfigKeyStore) CreateKey(_ context.Context, req CreateKeyRequest) (*Cr
|
|||||||
s.keys = append(s.keys, key)
|
s.keys = append(s.keys, key)
|
||||||
s.mu.Unlock()
|
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.
|
// GetUserKeys returns all active, non-expired keys for the given user.
|
||||||
// Pass an empty KeyType to return all types.
|
// 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()
|
now := time.Now()
|
||||||
s.mu.RLock()
|
s.mu.RLock()
|
||||||
defer s.mu.RUnlock()
|
defer s.mu.RUnlock()
|
||||||
|
|
||||||
var result []UserKey
|
var result []sectypes.UserKey
|
||||||
for i := range s.keys {
|
for i := range s.keys {
|
||||||
k := &s.keys[i]
|
k := &s.keys[i]
|
||||||
if k.UserID != userID || !k.IsActive {
|
if k.UserID != userID || !k.IsActive {
|
||||||
@@ -117,8 +119,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.
|
// ValidateKey hashes the raw key and finds a matching, active, non-expired entry.
|
||||||
// Uses constant-time comparison to prevent timing side-channels.
|
// Uses constant-time comparison to prevent timing side-channels.
|
||||||
// Pass an empty KeyType to accept any type.
|
// Pass an empty KeyType to accept any type.
|
||||||
func (s *ConfigKeyStore) ValidateKey(_ context.Context, rawKey string, keyType KeyType) (*UserKey, error) {
|
func (s *ConfigKeyStore) ValidateKey(_ context.Context, rawKey string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
|
||||||
hash := hashSHA256Hex(rawKey)
|
hash := sectypes.HashKey(rawKey)
|
||||||
hashBytes, _ := hex.DecodeString(hash)
|
hashBytes, _ := hex.DecodeString(hash)
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
|
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -13,86 +13,6 @@ import (
|
|||||||
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
"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
|
// Test parseRoles helper
|
||||||
func TestParseRoles(t *testing.T) {
|
func TestParseRoles(t *testing.T) {
|
||||||
tests := []struct {
|
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
|
// authenticateSync authenticates and waits for the asynchronous session
|
||||||
// activity update so sqlmock expectations are never touched concurrently.
|
// activity update so sqlmock expectations are never touched concurrently.
|
||||||
func authenticateSync(auth *DatabaseAuthenticator, req *http.Request) (*UserContext, error) {
|
func authenticateSync(auth *DatabaseAuthenticator, req *http.Request) (*UserContext, error) {
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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"`
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user