From ab3d2b5b04a9d8fd8863fbddd2048fc1d29ac813 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:18:01 +0200 Subject: [PATCH 01/23] docs(audit): add funcspec server-side audit --- audit/pkg/functionspec.audit.md | 228 ++++++++++++++++++++++++++++++++ 1 file changed, 228 insertions(+) create mode 100644 audit/pkg/functionspec.audit.md diff --git a/audit/pkg/functionspec.audit.md b/audit/pkg/functionspec.audit.md new file mode 100644 index 0000000..522b6e7 --- /dev/null +++ b/audit/pkg/functionspec.audit.md @@ -0,0 +1,228 @@ +# Audit: `pkg/funcspec` + +| | | +|---|---| +| **Package** | `github.com/bitechdev/ResolveSpec/pkg/funcspec` | +| **Files** | `function_api.go` (1251), `parameters.go` (411), `hooks.go` (179), `hooks_example.go`, `security_adapter.go` (117) | +| **Tests** | `function_api_test.go` (1278), `hooks_test.go` (589), `parameters_test.go` (549) — 2 416 lines; `go test` and `go test -race` pass, 76.1 % statement coverage | +| **Audit date** | 2026-09-30 | +| **Axes** | thread locking/waiting, slowness, security, panic handling & logging | +| **Threat model** | hostile internet client; query string, headers and body are attacker-controlled | +| **Depth** | targeted (server-side request path; verified against source) | + +## Summary + +`funcspec` exposes app-defined SQL templates as endpoints. The template is +trusted; everything the client adds to it is not. The package builds SQL by +string manipulation and has two kinds of client-controlled SQL fragments +(`X-Custom-SQL-W`, `X-Custom-SQL-Or`, `sort`) that are guarded only by a keyword +denylist (`ValidSQL(..., "select")`, `function_api.go:951-980`). That is not an +injection boundary: the fragment lands inside a query that may already carry +tenant or auth predicates, and the OR path produces wrong precedence that +widens results (findings 1-3). + +The auth integration is weaker than it looks. `RegisterSecurityHooks` is opt-in, +the anonymous default is `UserID 0`, and the auth hooks return an error *and* +set `Abort`, so `Execute` returns the error first and the handler answers +**400 `hook_error`**, not 401 (finding 6). + +Error handling leaks: `sendError` returns the DB error text and the full SQL to +the client, and the panic recovery writes the panic value into the 500 body +(finding 5). + +Resource limits are absent: default limit 100 000, no cap on `X-Limit`, no +LIMIT at all when counting is skipped, unbounded `[post_body]` read, a 15-minute +timeout, and a `COUNT(1)` over the full query on every list request (finding 8). + +Positives: there are no data races under the existing tests; the `[variable]` +substitution is quote-context aware; `Content-Type` and transaction handling are +consistent; hooks run inside the transaction. + +## Findings + +| # | Severity | Axis | Finding | +|---|---|---|---| +| 1 | **High** | security | `X-Custom-SQL-W`, `X-Custom-SQL-Or` and `sort` are appended as raw SQL, protected only by a keyword denylist (`ValidSQL "select"`) | +| 2 | **High** | security | `sqlQryWhereOr` emits `a AND b OR (c)`; OR conditions escape the AND-ed predicates (auth/tenant filters) | +| 3 | **Medium** | security | `sqlQryWhere`/`sqlQryWhereOr` locate WHERE/GROUP BY/ORDER BY/LIMIT by substring on the lower-cased query, including inside literals, subqueries and CTEs | +| 4 | **Medium** | security | Unquoted string `X-FieldFilter` value in `ApplyFilters`; `SearchOps` keyed per column (one op per column, random order) | +| 5 | **High** | security / logging | `sendError` returns `err.Error()` and the full SQL; panic recovery writes the panic value to the 500 body | +| 6 | **High** | security / correctness | Auth hooks return an error plus `Abort`; handler replies 400 `hook_error` instead of 401; hooks are opt-in; anonymous = `UserID 0` | +| 7 | **Medium** | security | `DecodeParam` (`ZIP_`/`__`) is applied to every header/param, ignores errors and recurses without a depth limit; headers matched with `HasPrefix` | +| 8 | **High** | slowness | Default limit 100 000, no `X-Limit` cap, no LIMIT with `NoCount`/`skipcount`, unbounded `io.ReadAll` of `[post_body]`, 15-minute timeout, full `COUNT(1)` per list request | +| 9 | **Medium** | locking | `HookRegistry` map and `variablesCallback` are unsynchronized; `Register`/`Clear*` race with `Execute` | +| 10 | **Medium** | security | Dollar-quote substitution (`[post_body]`, `[user]`, `[method]`, …) skips backslash escaping; `[id_session]` substituted unquoted | +| 11 | **Low** | correctness | `Content-Range` offset comes from the `offset` query param only; header offset ignored | +| 12 | **Low** | correctness | `X-Select-Fields`/`X-Not-Select-Fields` accepted but no-ops; `sort` `-col` negates instead of DESC | +| 13 | **Low** | panic / logging | Recovery is handler-level only; `Serving: Records` logged at Info per request; hook/filter strings logged unscrubbed (X8) | +| 14 | **Low** | correctness | `BeforeResponse` runs post-commit on the pool, not the tx (see `audit/single_tran.md`) | +| 15 | **Low** | security | Security adapter hard-codes schema `public` and entity `sql_query`; per-entity rules cannot be applied | +| 16 | **Info** | testing | Regexes compiled per call (`ValidSQL`, `sqlStripStringLiterals`); no `-race` in CI (X1); no hostile-input tests for findings 1-4 | + +## 1. Raw SQL fragments behind a keyword denylist — High + +`ApplyFilters` (`parameters.go:283-297`) passes `X-Custom-SQL-W` and +`X-Custom-SQL-Or` through `ValidSQL(..., "select")` and splices the result into +the query. `sort` goes the same way into `ORDER BY` (`function_api.go:~226`). +The denylist (`function_api.go:964-979`) removes `;`, `--`, `/*`, `*/`, `xp_`, +`sp_` and a few keywords **followed by a space**. It is not a parser: +- Subqueries, function calls (`pg_sleep`, `pg_read_*` where permitted), + `SELECT` itself and `)` are not blocked; a `)` can close the + `COUNT(1) FROM (%s) cnts` wrapper (`function_api.go:~241`). +- Keywords are removed rather than rejected, so input can be shaped so that + removal assembles a different token. +- Whitespace variants (tab, newline) bypass the `keyword␠` patterns. + +Whether the raw fragments are reachable is decided by the handler; they are +parsed whenever `ParseParameters` runs, i.e. always. Fix: drop the two headers +from the wire contract, or accept only a column/operator/value structure built +by the server; validate `sort` against `^[A-Za-z0-9_.]+( (ASC|DESC))?(,…)*$` +and ideally an allowlist of columns. + +## 2. OR precedence widens results — High + +`sqlQryWhereOr` (`parameters.go:381-411`) rewrites `WHERE a AND b` into +`WHERE a AND b OR (c)`. SQL evaluates `AND` first, so the result is +`(a AND b) OR c`: any row satisfying `c` is returned regardless of `a`/`b`. +Where `a` is a tenant or ownership predicate in the template, a client-supplied +OR condition (`X-SearchOr`, `X-Custom-SQL-Or`, search operator with logic OR) +returns other tenants' rows. Verified with `ParseParameters` + `ApplyFilters` +on generated headers. Fix: wrap the existing WHERE body in parentheses before +appending `OR (...)`, or build a predicate tree. + +## 3. Substring-based clause location — Medium + +Both helpers use `strings.Index` on `" where "`, `" group by"`, `" order by"`, +`" limit "` over the whole lower-cased query. A match inside a string literal, +a subquery, a CTE or a column alias selects the wrong insertion point, and +`wherePos > 0` decides AND-append vs. new WHERE on the first match anywhere. +`ApplyDistinct` (`parameters.go:363-378`) similarly inserts after the first +`SELECT` substring, and the ORDER BY test (`function_api.go:~224`) compares the +first `order by` to the first `from `. `sqlStripStringLiterals` exists +(`function_api.go:858`) but is not used by these helpers. + +## 4. Filter handling inconsistencies — Medium + +- `ApplyFilters` builds `col = value` for `X-FieldFilter` without quoting the + value (`parameters.go:248-250`), so a string value becomes a column reference + (`status = active`). `mergeHeaderParams` quotes the same filter, so the + `SqlQuery` path applies it twice with different semantics. +- `RequestParameters.SearchOps` is a map keyed by column; two operators on one + column overwrite each other and map iteration order makes the generated WHERE + non-deterministic. + +## 5. Information disclosure in errors — High + +- `sendError` (`function_api.go:1150-1172`) sets `Detail = err.Error()` and, + for `*common.SQLError`, `SQL` = the final statement, including the template, + substituted values and any injected fragment. Used by every failure path + (`query_failed`, `count_failed`, `hook_error`). +- Panic recovery in `SqlQueryList` (`:80-86`) and `SqlQuery` (`:433-439`) calls + `http.Error(w, fmt.Sprintf("Internal server error: %v", err), 500)`; the + panic value reaches the client. Same class as `middleware` finding 4. +Fix: log server-side, return a generic message plus a request id. + +## 6. Auth hook abort returns 400, hooks opt-in — High + +`RegisterSecurityHooks` (`security_adapter.go:14-55`) sets `Abort`, +`AbortCode=401` **and returns an error**. `HookRegistry.Execute` +(`hooks.go:113-137`) returns the error before it evaluates `Abort`, and the +handler maps that to `sendError(400, "hook_error", …)` +(`function_api.go:~202`). The 401 branch in the handler is only reachable for +hooks that set `Abort` without returning an error. Clients therefore see 400 +with `Detail: "hook execution failed: authentication required"`. + +Also: without `RegisterSecurityHooks` there is no authentication at all; a +missing user context is replaced with `UserID 0, "anonymous"` +(`function_api.go:~103`) and the request proceeds. Fix: return nil after +setting `Abort` in the auth hooks, or have the handler honour `AbortCode` when +the error wraps an abort; consider fail-closed by default. + +## 7. Header/param decoding — Medium + +`decodeValue` (`parameters.go:203`) calls `restheadspec.DecodeParam` and drops +the error. `DecodeParam` replaces all `ZIP_`/`__` occurrences and decodes +recursively with no depth limit, so one value can force repeated base64/gzip +work (decompression amplification, since size is not capped). Header keys are +matched with `HasPrefix`, so `X-SearchOp-` variants and unrelated +headers with the same prefix are interpreted. + +## 8. Unbounded resource use — High + +- `parameters.go:54` default `Limit: 100000`; `X-Limit` and `limit` accept any + positive integer. +- In `SqlQueryList` the `LIMIT`/`OFFSET` clause is added **only inside + `if !options.NoCount`** (`function_api.go:~232-251`); `NoCount` or + `X-SkipCount` returns the whole result set. +- `COUNT(1) FROM ()` runs on every list request (double execution + cost). +- `[post_body]` uses `io.ReadAll(r.Body)` (`function_api.go:913`) with no + `http.MaxBytesReader`; the body is also embedded into the SQL text. +- `context.WithTimeout(…, 15*time.Minute)` (`:91`, `:444`) holds a transaction + and pooled connection for up to 15 minutes per request. +- `ValidSQL` and `sqlStripStringLiterals` compile regexes on each call. +Fix: hard cap on limit, always apply LIMIT, cap body size, configurable timeout. + +## 9. Unsynchronized registry — Medium + +`HookRegistry.hooks` (`hooks.go`) is a plain map; `Register`, `Clear`, +`ClearAll` mutate it while `Execute` reads it from request goroutines. Safe +only if all registration completes before serving. `Handler.variablesCallback` +(`function_api.go:60-68`) has the same property. Fix: `sync.RWMutex` and copy-on- +read, or document and enforce "register before serve". + +## 10. Dollar-quote substitution — Medium + +`safeSubstituteVar` returns the raw value when the placeholder is adjacent to +`$` (`function_api.go:1044-1049`), so neither backslash nor quote escaping +applies. The tag is neutralised only for `$M$`, `$PBODY$` and the equivalents +in `replaceMetaVariables`; a caller-supplied value in a template that uses a +different tag (or `$$`) is not. `isInsideDollarQuote` inspects only the first +occurrence of the placeholder. `[id_session]` is replaced without any quoting +(`function_api.go:~900`); its source is the auth layer, but it becomes an +injection point if a session token format allows quotes. + +## 11-12. Behavioural defects — Low + +- `Content-Range` offset uses only `r.URL.Query().Get("offset")` + (`function_api.go:~319`) while the applied offset can come from + `X-Offset`; the reported range is wrong for header-driven paging. +- `ApplyFieldSelection` (`parameters.go:226-241`) only logs; the headers have + no effect. `sort=-col` is not converted to DESC; it is emitted as `ORDER BY + -col`, which negates the column value. + +## 13. Panic handling and logging — Low + +Recovery exists per handler only (no middleware-level recovery for hooks run +outside), and the stack is logged via `logger.Error`, which forwards to Sentry +unscrubbed (X8). `logger.Info("Serving: Records …")` runs on every list request. +`logger.Debug` lines include the generated filter SQL and attacker-supplied +values. Hook failures log `err` with attacker-influenced text. + +## 14. `BeforeResponse` outside the transaction — Low + +`BeforeResponse` executes after `RunInTransaction` returns, with +`hookCtx.Tx = h.db` (`function_api.go:~336-343`, `:~640`). A hook that writes +cannot be rolled back with the query, and a failure returns 500 after the work +committed. Tracked in `audit/single_tran.md`. + +## 15. Security adapter — Low + +`funcSpecSecurityContext.GetSchema()` returns `"public"` and `GetEntity()` +returns `"sql_query"` for every endpoint (`security_adapter.go:84-92`), so +column/row security rules keyed by entity cannot distinguish funcspec endpoints. +`GetModel`, `GetQuery`, `SetQuery` are stubs. + +## 16. Testing — Info + +Tests cover handler flow, hooks and parameter parsing. No test exercises the +hostile inputs of findings 1-4 or the 401-vs-400 outcome. There is no `-race` +job in CI (X1). The earlier note about a failing +`TestReplaceMetaVariables/Replace_[user]` no longer reproduces: the package +passes today. + +## Cross-references + +X1 (no `-race`), X7 (inconsistent panic handling), X8 (logger forwards to +Sentry unscrubbed), `middleware` finding 4 (panic value in body), +`audit/single_tran.md` (post-commit hooks). From 54e6a3b17c0eaf0036e903e5fd68923326201460 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:18:01 +0200 Subject: [PATCH 02/23] feat(resolvespec-python): add Python client for ResolveSpec, HeaderSpec, FunctionSpec and WebSocketSpec --- resolvespec-python/.gitignore | 6 + resolvespec-python/README.md | 142 ++++++++ resolvespec-python/pyproject.toml | 24 ++ .../src/resolvespec/__init__.py | 43 +++ .../src/resolvespec/funcspec.py | 197 ++++++++++ .../src/resolvespec/headerspec.py | 336 ++++++++++++++++++ resolvespec-python/src/resolvespec/http.py | 76 ++++ .../src/resolvespec/resolvespec.py | 137 +++++++ resolvespec-python/src/resolvespec/types.py | 166 +++++++++ .../src/resolvespec/websocket.py | 335 +++++++++++++++++ resolvespec-python/tests/test_funcspec.py | 132 +++++++ resolvespec-python/tests/test_headerspec.py | 236 ++++++++++++ resolvespec-python/tests/test_resolvespec.py | 210 +++++++++++ resolvespec-python/tests/test_websocket.py | 152 ++++++++ resolvespec-python/todo.md | 28 +- 15 files changed, 2206 insertions(+), 14 deletions(-) create mode 100644 resolvespec-python/.gitignore create mode 100644 resolvespec-python/README.md create mode 100644 resolvespec-python/pyproject.toml create mode 100644 resolvespec-python/src/resolvespec/__init__.py create mode 100644 resolvespec-python/src/resolvespec/funcspec.py create mode 100644 resolvespec-python/src/resolvespec/headerspec.py create mode 100644 resolvespec-python/src/resolvespec/http.py create mode 100644 resolvespec-python/src/resolvespec/resolvespec.py create mode 100644 resolvespec-python/src/resolvespec/types.py create mode 100644 resolvespec-python/src/resolvespec/websocket.py create mode 100644 resolvespec-python/tests/test_funcspec.py create mode 100644 resolvespec-python/tests/test_headerspec.py create mode 100644 resolvespec-python/tests/test_resolvespec.py create mode 100644 resolvespec-python/tests/test_websocket.py diff --git a/resolvespec-python/.gitignore b/resolvespec-python/.gitignore new file mode 100644 index 0000000..27e6129 --- /dev/null +++ b/resolvespec-python/.gitignore @@ -0,0 +1,6 @@ +__pycache__/ +*.egg-info/ +.venv/ +dist/ +.pytest_cache/ +.coverage diff --git a/resolvespec-python/README.md b/resolvespec-python/README.md new file mode 100644 index 0000000..ff60cd2 --- /dev/null +++ b/resolvespec-python/README.md @@ -0,0 +1,142 @@ +# resolvespec (Python) + +Python client for ResolveSpec REST, HeaderSpec (restheadspec), FunctionSpec and WebSocketSpec. Port of `resolvespec-js`. + +- Python >= 3.11, `httpx` (REST, sync + async), `websockets` (WS, async) +- Options/filters/sorts are plain dicts using the wire key names (`TypedDict` hints in `resolvespec.types`) + +``` +pip install resolvespec +``` + +## Clients + +| Protocol | Sync | Async | Transport | +|---|---|---|---| +| ResolveSpec | `ResolveSpecClient` | `AsyncResolveSpecClient` | POST + JSON body `{operation, id, data, options}` | +| HeaderSpec | `HeaderSpecClient` | `AsyncHeaderSpecClient` | GET/POST/PUT/DELETE, options as `X-*` headers | +| FunctionSpec | `FuncSpecClient` | `AsyncFuncSpecClient` | user-defined SQL endpoints; params via query string + `X-*` headers | +| WebSocketSpec | - | `WebSocketClient` | WebSocket JSON messages | + +Constructor (REST): `Client(base_url, token=None, headers=None, timeout=30.0)` + +- `token` -> `Authorization: Bearer`; wins over `headers` +- `headers`: custom headers, merged case-insensitively; snapshot at construction +- Sync: context manager / `close()`. Async: `async with` / `await aclose()` +- Cached sync factories: `get_resolvespec_client()`, `get_headerspec_client()` (same args -> same instance) + +## ResolveSpec + +URL: `{base}/{schema}/{entity}[/{id}]` + +| Method | Signature | +|---|---| +| `get_metadata` | `(schema, entity)` (GET) | +| `read` | `(schema, entity, id=None, options=None)` | +| `create` | `(schema, entity, data, options=None)` | +| `update` | `(schema, entity, data, id=None, options=None)` | +| `delete` | `(schema, entity, id)` | + +`id`: int/str -> URL path; `list[str]` -> body `id`. +Returns `{"success", "data", "metadata"?, "error"?}`. + +## HeaderSpec + +| Method | HTTP | Signature | +|---|---|---| +| `read` | GET | `(schema, entity, id=None, options=None)` | +| `create` | POST | `(schema, entity, data, options=None)` | +| `update` | PUT | `(schema, entity, id, data, options=None)` | +| `delete` | DELETE | `(schema, entity, id)` | + +Response metadata derived from `Content-Range` (`offset-end/total`) and `X-Limit`. +`build_headers(options)`, `encode_header_value()` / `decode_header_value()` (`ZIP_` / `__` base64) are exported. + +### Option -> header + +| Option | Header | +|---|---| +| `columns` / `omit_columns` | `X-Select-Fields` / `X-Not-Select-Fields` | +| filter `eq` + AND | `X-FieldFilter-{col}` | +| filter AND / OR | `X-SearchOp-{op}-{col}` / `X-SearchOr-{op}-{col}` | +| spatial (`st_*`, `bbox`) / vector (`*_within`) filter | `X-SpatialFilter-{col}` / `X-VectorFilter-{col}` (JSON) | +| `sort` | `X-Sort` (`+col,-col`) | +| `limit` / `offset` | `X-Limit` / `X-Offset` | +| `cursor_forward` / `cursor_backward` | `X-Cursor-Forward` / `X-Cursor-Backward` | +| `preload` | `X-Preload` (`Rel:c1,c2\|Rel2`), `X-Preload-Where`, `X-Preload-{n}[-Where]` | +| `expand` | `X-Expand` | +| `custom_sql_joins` / `custom_sql_or` | `X-Custom-SQL-Join` / `X-Custom-SQL-Or` | +| `search_columns` | `X-SearchCols` | +| `advanced_sql` | `X-AdvSQL-{col}` | +| `computedColumns` | `X-CQL-SEL-{name}` | +| `customOperators` | `X-Custom-SQL-W` (AND-joined) | +| `vector_search` | `X-Vector-Search-{col}`, `-Vector`, `-As`, `-Dir` | +| `fetch_row_number` | `X-Fetch-RowNumber` | +| `clean_json`, `distinct`, `skip_count`, `skip_cache`, `atomic_transaction`, `single_record_as_object` | `X-Clean-JSON`, `X-Distinct`, `X-SkipCount`, `X-SkipCache`, `X-Transaction-Atomic`, `X-Single-Record-As-Object` | +| `pk_row` | `X-PKRow` | +| `response_format` (`simple`/`detail`/`syncfusion`) | `X-SimpleApi` / `X-DetailApi` / `X-Syncfusion` | +| `xfiles` | `X-Files` (`ZIP_` base64 JSON) | + +Filter operator -> header op: `eq equals`, `neq notequals`, `gt greaterthan`, `gte greaterthanorequal`, `lt lessthan`, `lte lessthanorequal`, `like/ilike/contains contains`, `startswith beginswith`, `endswith`, `in`, `between`, `between_inclusive betweeninclusive`, `is_null empty`, `is_not_null notempty`. + +## FunctionSpec + +Routes are defined by the server app, so calls take a `path`. The server never reads a request body. + +| Method | Server handler | Result | +|---|---|---| +| `query(path, params=None, options=None, *, method="GET")` | `SqlQuery` (single record) | `{success, data}` | +| `query_list(path, params=None, options=None, *, method="GET")` | `SqlQueryList` | `{success, data, metadata}` (from `Content-Range: items a-b/total`) | + +- `params` -> query string. `bool` -> `true/false`, `None` skipped, `list` -> repeated key (server: `IN` filter). `p-` prefixed names are substituted into the SQL. +- `options` -> `X-*` headers. Query values override headers of the same name. +- 206 Partial Content (more rows than returned) is treated as success. + +| Option | Header | +|---|---| +| `filters` (`eq`+AND) | `X-FieldFilter-{col}` | +| `filters` (other) | `X-SearchOp-{op}-{col}` / `X-SearchOr-{op}-{col}` | +| `search_filters` `{col: text}` | `X-SearchFilter-{col}` (ILIKE) | +| `custom_sql_where` / `custom_sql_or` | `X-Custom-SQL-W` / `X-Custom-SQL-Or` | +| `sort` | `X-Sort` as SQL terms: `col ASC,col DESC` | +| `limit` / `offset` | `X-Limit` / `X-Offset` | +| `distinct`, `skip_count`, `skip_cache` | `X-Distinct`, `X-SkipCount`, `X-SkipCache` | +| `response_format` | `X-SimpleApi` / `X-DetailApi` / `X-Syncfusion` (`data` shape changes: array / `{items,...}` / `{result,count}`) | + +Server limits: +- `sort` goes verbatim into `ORDER BY`; `-col` (restheadspec style) does **not** mean DESC. +- `X-Select-Fields` / `X-Not-Select-Fields` are no-ops server-side, so not exposed. +- One search operator per column; same column twice keeps the last. +- Values starting with `ZIP_` / `__` are base64-decoded by the server; such plaintext cannot be sent. +- Non-ASCII / control-char values are sent `ZIP_`-encoded automatically. + +## WebSocketSpec + +`WebSocketClient(url, *, reconnect=True, reconnect_interval=3.0, max_reconnect_attempts=10, heartbeat_interval=30.0, request_timeout=30.0, subscribe_timeout=10.0, headers=None)` + +| Method | Notes | +|---|---| +| `connect()` / `close()` | also `async with` | +| `request(operation, entity, *, schema, record_id, data, options)` | returns response `data` | +| `read(entity, *, schema, record_id, filters, columns, sort, preload, limit, offset)` | | +| `create(entity, data, *, schema)` | | +| `update(entity, id, data, *, schema)` | | +| `delete(entity, id, *, schema)` | | +| `meta(entity, *, schema)` | | +| `subscribe(entity, callback, *, schema, filters)` | returns subscription id; callback gets notification dict (sync or async) | +| `unsubscribe(subscription_id)` | | +| `on(event, cb)` / `off(event)` | events: `connect`, `disconnect`, `error`, `message`, `state_change` | +| `state`, `is_connected()`, `get_subscriptions()` | | + +Auto-reconnect does not restore subscriptions; re-subscribe on `connect`. + +## Errors + +`ResolveSpecError(message, status_code, code, details)` on non-2xx (REST) or failed response / timeout / not connected (WS). + +## Dev + +``` +pip install -e '.[dev]' +pytest +``` diff --git a/resolvespec-python/pyproject.toml b/resolvespec-python/pyproject.toml new file mode 100644 index 0000000..b5aa80b --- /dev/null +++ b/resolvespec-python/pyproject.toml @@ -0,0 +1,24 @@ +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[project] +name = "resolvespec" +version = "1.0.0" +description = "Python client for ResolveSpec REST, HeaderSpec and WebSocket APIs" +readme = "README.md" +requires-python = ">=3.11" +license = { text = "MIT" } +authors = [{ name = "Hein (Warkanum) Puth" }] +keywords = ["resolvespec", "headerspec", "websocket", "rest-client", "api-client"] +dependencies = ["httpx>=0.27", "websockets>=13"] + +[project.optional-dependencies] +dev = ["pytest>=8", "pytest-asyncio>=0.23", "pytest-cov"] + +[tool.hatch.build.targets.wheel] +packages = ["src/resolvespec"] + +[tool.pytest.ini_options] +testpaths = ["tests"] +asyncio_mode = "auto" diff --git a/resolvespec-python/src/resolvespec/__init__.py b/resolvespec-python/src/resolvespec/__init__.py new file mode 100644 index 0000000..f273fe4 --- /dev/null +++ b/resolvespec-python/src/resolvespec/__init__.py @@ -0,0 +1,43 @@ +"""ResolveSpec Python client: REST (ResolveSpec), HeaderSpec and WebSocketSpec.""" +from typing import Mapping, Optional + +from .headerspec import ( + AsyncHeaderSpecClient, + HeaderSpecClient, + build_headers, + decode_header_value, + encode_header_value, +) +from .funcspec import AsyncFuncSpecClient, FuncSpecClient +from .http import ResolveSpecError, merge_headers +from .resolvespec import AsyncResolveSpecClient, ResolveSpecClient +from .types import * # noqa: F401,F403 +from .websocket import Subscription, WebSocketClient + + +def _cache_key(base_url: str, token: Optional[str], headers: Optional[Mapping[str, str]]): + return ( + base_url, + token, + tuple(sorted((k.lower(), v) for k, v in (headers or {}).items())), + ) + + +_resolvespec: dict = {} +_headerspec: dict = {} + + +def get_resolvespec_client(base_url: str, token: Optional[str] = None, headers: Optional[Mapping[str, str]] = None) -> ResolveSpecClient: + """Cached sync client, keyed by base_url + token + headers (case-insensitive names).""" + key = _cache_key(base_url, token, headers) + if key not in _resolvespec: + _resolvespec[key] = ResolveSpecClient(base_url, token, headers) + return _resolvespec[key] + + +def get_headerspec_client(base_url: str, token: Optional[str] = None, headers: Optional[Mapping[str, str]] = None) -> HeaderSpecClient: + """Cached sync client, keyed by base_url + token + headers (case-insensitive names).""" + key = _cache_key(base_url, token, headers) + if key not in _headerspec: + _headerspec[key] = HeaderSpecClient(base_url, token, headers) + return _headerspec[key] diff --git a/resolvespec-python/src/resolvespec/funcspec.py b/resolvespec-python/src/resolvespec/funcspec.py new file mode 100644 index 0000000..35131fd --- /dev/null +++ b/resolvespec-python/src/resolvespec/funcspec.py @@ -0,0 +1,197 @@ +"""FunctionSpec client: calls user-defined SQL endpoints (Go pkg/funcspec). + +Routes are defined by the server application, so calls take a `path`. +Parameters are sent as query string values and/or `X-*` headers; the server never +reads a request body. Query-string values override headers of the same name. + +Server behaviour worth knowing (pkg/funcspec): + - `sort` is inserted raw into ORDER BY, so it must be SQL (`col DESC`), not `-col`. + - Field selection (`X-Select-Fields`) is a no-op server-side, so it is not exposed. + - Only one search operator per column is kept. + - Values starting with `ZIP_` or `__` are base64-decoded by the server (even after our + own encoding), so such plaintext values cannot be sent faithfully. +""" +from __future__ import annotations + +import re +from typing import Any, Dict, List, Mapping, Optional + +import httpx + +from .headerspec import _OPERATOR_MAP, _bool, _filter_value, encode_header_value +from .http import client_headers, error_from, merge_headers, parse_json +from .types import APIResponse, FuncSpecOptions + +Params = Mapping[str, Any] + +_CONTENT_RANGE = re.compile(r"(\d+)-(\d+)/(\d+)") + + +def _safe(value: str) -> str: + """Encode values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces).""" + if not value.isascii() or not value.isprintable() or value != value.strip(): + return encode_header_value(value) + return value + + +def build_headers(options: Mapping[str, Any]) -> Dict[str, str]: + """Build the X-* headers understood by funcspec.ParseParameters.""" + h: Dict[str, str] = {} + o = options + + for f in o.get("filters") or []: + operator = f["operator"] + logic = f.get("logic_operator") or "AND" + value = _safe(_filter_value(f)) + if operator == "eq" and logic == "AND": + h[f"X-FieldFilter-{f['column']}"] = value + else: + kind = "X-SearchOr" if logic == "OR" else "X-SearchOp" + h[f"{kind}-{_OPERATOR_MAP.get(operator, operator)}-{f['column']}"] = value + + for col, text in (o.get("search_filters") or {}).items(): + h[f"X-SearchFilter-{col}"] = _safe(str(text)) # CAST(col AS TEXT) ILIKE %text% + + if o.get("custom_sql_where"): + h["X-Custom-SQL-W"] = _safe(o["custom_sql_where"]) + if o.get("custom_sql_or"): + h["X-Custom-SQL-Or"] = _safe(o["custom_sql_or"]) + + if o.get("sort"): + h["X-Sort"] = _safe(",".join(_sort_term(s) for s in o["sort"])) + if o.get("limit") is not None: + h["X-Limit"] = str(o["limit"]) + if o.get("offset") is not None: + h["X-Offset"] = str(o["offset"]) + + for name, key in (("X-Distinct", "distinct"), ("X-SkipCount", "skip_count"), ("X-SkipCache", "skip_cache")): + if o.get(key) is not None: + h[name] = _bool(o[key]) + + fmt = o.get("response_format") + if fmt: + h[{"simple": "X-SimpleApi", "detail": "X-DetailApi", "syncfusion": "X-Syncfusion"}[fmt]] = "true" + return h + + +def _sort_term(s: Mapping[str, str]) -> str: + # funcspec puts this verbatim into ORDER BY + return f"{s['column']} {'DESC' if s.get('direction', 'asc').upper() == 'DESC' else 'ASC'}" + + +def build_query(params: Optional[Params]) -> Dict[str, Any]: + """Query-string values: bools -> true/false, lists -> repeated keys (server: IN filter).""" + out: Dict[str, Any] = {} + for k, v in (params or {}).items(): + if v is None: + continue + if isinstance(v, (list, tuple)): + out[k] = [_safe(_q(x)) for x in v] + else: + out[k] = _safe(_q(v)) + return out + + +def _q(v: Any) -> str: + return _bool(v) if isinstance(v, bool) else str(v) + + +def _metadata(response: httpx.Response, options: Optional[Mapping[str, Any]]) -> Dict[str, int]: + """Content-Range is `items {offset}-{offset+len}/{total}`.""" + m = _CONTENT_RANGE.search(response.headers.get("content-range", "")) + start, end, total = (int(x) for x in m.groups()) if m else (0, 0, 0) + return { + "total": total, + "count": end - start, + "filtered": total, + "offset": start, + "limit": int((options or {}).get("limit") or 0), + } + + +def _wrap(response: httpx.Response, options: Optional[Mapping[str, Any]], with_metadata: bool) -> APIResponse: + data = parse_json(response) + if not response.is_success: # 206 Partial Content is success + raise error_from(response, data) + result: APIResponse = {"success": True, "data": data} + if with_metadata: + result["metadata"] = _metadata(response, options) + return result + + +class _Base: + def __init__( + self, + base_url: str, + token: Optional[str] = None, + headers: Optional[Mapping[str, str]] = None, + timeout: Optional[float] = 30.0, + ): + self.base_url = base_url + self.token = token + self.headers = dict(headers or {}) # snapshot + self.timeout = timeout + + def _req(self, method: str, path: str, params: Optional[Params], options: Optional[Mapping[str, Any]]): + url = f"{self.base_url.rstrip('/')}/{path.lstrip('/')}" + headers = merge_headers( + client_headers(self.token, self.headers), + build_headers(options) if options else {}, + ) + return method.upper(), url, headers, build_query(params) + + +class FuncSpecClient(_Base): + """Synchronous client. Use as a context manager or call close().""" + + def __init__(self, *args: Any, transport: Optional[httpx.BaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.Client(timeout=self.timeout, transport=transport) + + def close(self) -> None: + self._http.close() + + def __enter__(self) -> "FuncSpecClient": + return self + + def __exit__(self, *exc: Any) -> None: + self.close() + + def _send(self, req, options, with_metadata) -> APIResponse: + method, url, headers, query = req + return _wrap(self._http.request(method, url, headers=headers, params=query), options, with_metadata) + + def query(self, path: str, params: Optional[Params] = None, options: Optional[FuncSpecOptions] = None, *, method: str = "GET") -> APIResponse: + """Single-record endpoint (Handler.SqlQuery). `data` is the row object.""" + return self._send(self._req(method, path, params, options), options, False) + + def query_list(self, path: str, params: Optional[Params] = None, options: Optional[FuncSpecOptions] = None, *, method: str = "GET") -> APIResponse: + """List endpoint (Handler.SqlQueryList). Adds `metadata` from Content-Range.""" + return self._send(self._req(method, path, params, options), options, True) + + +class AsyncFuncSpecClient(_Base): + """Asyncio client. Use as an async context manager or await aclose().""" + + def __init__(self, *args: Any, transport: Optional[httpx.AsyncBaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.AsyncClient(timeout=self.timeout, transport=transport) + + async def aclose(self) -> None: + await self._http.aclose() + + async def __aenter__(self) -> "AsyncFuncSpecClient": + return self + + async def __aexit__(self, *exc: Any) -> None: + await self.aclose() + + async def _send(self, req, options, with_metadata) -> APIResponse: + method, url, headers, query = req + return _wrap(await self._http.request(method, url, headers=headers, params=query), options, with_metadata) + + async def query(self, path: str, params: Optional[Params] = None, options: Optional[FuncSpecOptions] = None, *, method: str = "GET") -> APIResponse: + return await self._send(self._req(method, path, params, options), options, False) + + async def query_list(self, path: str, params: Optional[Params] = None, options: Optional[FuncSpecOptions] = None, *, method: str = "GET") -> APIResponse: + return await self._send(self._req(method, path, params, options), options, True) diff --git a/resolvespec-python/src/resolvespec/headerspec.py b/resolvespec-python/src/resolvespec/headerspec.py new file mode 100644 index 0000000..831c772 --- /dev/null +++ b/resolvespec-python/src/resolvespec/headerspec.py @@ -0,0 +1,336 @@ +"""HeaderSpec client: query options sent as HTTP headers (Go restheadspec). + +Methods: GET=read, POST=create, PUT=update, DELETE=delete. +""" +from __future__ import annotations + +import base64 +import json +import re +from typing import Any, Dict, Mapping, Optional + +import httpx + +from .http import build_url, client_headers, error_from, merge_headers, parse_json +from .types import APIResponse, FilterOption, HeaderSpecOptions + +_PREFIXES = ("ZIP_", "__") + +_OPERATOR_MAP = { + "eq": "equals", + "neq": "notequals", + "gt": "greaterthan", + "gte": "greaterthanorequal", + "lt": "lessthan", + "lte": "lessthanorequal", + "like": "contains", + "ilike": "contains", + "contains": "contains", + "startswith": "beginswith", + "endswith": "endswith", + "in": "in", + "between": "between", + "between_inclusive": "betweeninclusive", + "is_null": "empty", + "is_not_null": "notempty", +} + + +def encode_header_value(value: str) -> str: + """Base64 (UTF-8) with ZIP_ prefix, for complex header values.""" + return "ZIP_" + base64.b64encode(value.encode("utf-8")).decode("ascii") + + +def decode_header_value(value: str) -> str: + """Decode a value that may carry a ZIP_ or __ base64 prefix (nested allowed).""" + code = value + for prefix in _PREFIXES: + if code.startswith(prefix): + b64 = re.sub(r"[\n\r ]", "", code[len(prefix):]) + b64 += "=" * (-len(b64) % 4) + code = base64.b64decode(b64).decode("utf-8") + break + if code.startswith(_PREFIXES): + code = decode_header_value(code) + return code + + +def _geo_header(operator: str) -> Optional[str]: + op = operator.lower() + if op.endswith("_within"): + return "X-VectorFilter-" + if op.startswith("st_") or op in ("bbox", "&&"): + return "X-SpatialFilter-" + return None + + +def _filter_value(f: FilterOption) -> str: + v = f.get("value") + if v is None: + return "" + if isinstance(v, (list, tuple)): + return ",".join(_scalar(x) for x in v) + return _scalar(v) + + +def _scalar(v: Any) -> str: + if isinstance(v, bool): # match JS String(true) + return "true" if v else "false" + return str(v) + + +def _bool(v: bool) -> str: + return "true" if v else "false" + + +def _preload_spec(p: Mapping[str, Any]) -> str: + cols = p.get("columns") + return f"{p['relation']}:{','.join(cols)}" if cols else p["relation"] + + +def build_headers(options: HeaderSpecOptions) -> Dict[str, str]: + """Build restheadspec HTTP headers from options. See README for the mapping.""" + h: Dict[str, str] = {} + o = options + + if o.get("columns"): + h["X-Select-Fields"] = ",".join(o["columns"]) + if o.get("omit_columns"): + h["X-Not-Select-Fields"] = ",".join(o["omit_columns"]) + + for f in o.get("filters") or []: + logic = f.get("logic_operator") or "AND" + operator = f["operator"] + op = _OPERATOR_MAP.get(operator, operator) + value = _filter_value(f) + geo = _geo_header(operator) + if geo: + payload: Dict[str, Any] = {"op": operator, "value": f.get("value")} + if logic == "OR": + payload["logic"] = "or" + h[f"{geo}{f['column']}"] = json.dumps(payload, separators=(",", ":")) + elif operator == "eq" and logic == "AND": + h[f"X-FieldFilter-{f['column']}"] = value + elif logic == "OR": + h[f"X-SearchOr-{op}-{f['column']}"] = value + else: + h[f"X-SearchOp-{op}-{f['column']}"] = value + + if o.get("sort"): + h["X-Sort"] = ",".join( + ("-" if s["direction"].upper() == "DESC" else "+") + s["column"] for s in o["sort"] + ) + + if o.get("limit") is not None: + h["X-Limit"] = str(o["limit"]) + if o.get("offset") is not None: + h["X-Offset"] = str(o["offset"]) + if o.get("cursor_forward"): + h["X-Cursor-Forward"] = o["cursor_forward"] + if o.get("cursor_backward"): + h["X-Cursor-Backward"] = o["cursor_backward"] + + if o.get("preload"): + # Go applies X-Preload-Where to every preload in the matching X-Preload header, + # so preloads are grouped by where clause. + groups: Dict[str, list] = {} + for p in o["preload"]: + groups.setdefault(p.get("where") or "", []).append(_preload_spec(p)) + n = 0 + for where, specs in groups.items(): + if not where: + h["X-Preload"] = "|".join(specs) + elif "" not in groups and n == 0: + # X-Preload-Where would also apply to a where-less X-Preload, so only use it alone + h["X-Preload"] = "|".join(specs) + h["X-Preload-Where"] = where + n += 1 + else: + n += 1 + h[f"X-Preload-{n}"] = "|".join(specs) + h[f"X-Preload-{n}-Where"] = where + + if o.get("expand"): + h["X-Expand"] = "|".join(_preload_spec(e) for e in o["expand"]) + if o.get("custom_sql_joins"): + h["X-Custom-SQL-Join"] = "|".join(o["custom_sql_joins"]) + if o.get("custom_sql_or"): + h["X-Custom-SQL-Or"] = " OR ".join(o["custom_sql_or"]) + if o.get("search_columns"): + h["X-SearchCols"] = ",".join(o["search_columns"]) + for col, sql in (o.get("advanced_sql") or {}).items(): + h[f"X-AdvSQL-{col}"] = sql + + vs = o.get("vector_search") + if vs: + h[f"X-Vector-Search-{vs['column']}"] = vs.get("metric") or "l2" + h["X-Vector-Search-Vector"] = json.dumps(vs["vector"], separators=(",", ":")) + if vs.get("as"): + h["X-Vector-Search-As"] = vs["as"] + if vs.get("direction"): + h["X-Vector-Search-Dir"] = vs["direction"] + + for name, key in ( + ("X-Clean-JSON", "clean_json"), + ("X-Distinct", "distinct"), + ("X-SkipCount", "skip_count"), + ("X-SkipCache", "skip_cache"), + ("X-Transaction-Atomic", "atomic_transaction"), + ("X-Single-Record-As-Object", "single_record_as_object"), + ): + if o.get(key) is not None: + h[name] = _bool(o[key]) + + if o.get("pk_row"): + h["X-PKRow"] = o["pk_row"] + + fmt = o.get("response_format") + if fmt: + h[{"simple": "X-SimpleApi", "detail": "X-DetailApi", "syncfusion": "X-Syncfusion"}[fmt]] = "true" + + if o.get("xfiles"): + h["X-Files"] = encode_header_value(json.dumps(o["xfiles"], separators=(",", ":"))) + + if o.get("fetch_row_number"): + h["X-Fetch-RowNumber"] = o["fetch_row_number"] + + for cc in o.get("computedColumns") or []: + h[f"X-CQL-SEL-{cc['name']}"] = cc["expression"] + + if o.get("customOperators"): + h["X-Custom-SQL-W"] = " AND ".join(co["sql"] for co in o["customOperators"]) + + return h + + +def _int(s: Optional[str]) -> int: + try: + return int(s) # type: ignore[arg-type] + except (TypeError, ValueError): + return 0 + + +def _wrap(response: httpx.Response) -> APIResponse: + """Wrap a raw restheadspec body, deriving metadata from Content-Range / X-Limit.""" + data = parse_json(response) + if not response.is_success: + raise error_from(response, data) + cr = response.headers.get("content-range") + total = _int(cr.split("/")[-1]) if cr else 0 + offset = _int(cr.split("/")[0].split("-")[0].split(" ")[-1]) if cr else 0 + return { + "data": data, + "success": True, + "error": data.get("error") if isinstance(data, dict) else None, + "metadata": { + "count": total, + "total": total, + "filtered": total, + "offset": offset, + "limit": _int(response.headers.get("x-limit")), + }, + } + + +class _Base: + def __init__( + self, + base_url: str, + token: Optional[str] = None, + headers: Optional[Mapping[str, str]] = None, + timeout: Optional[float] = 30.0, + ): + self.base_url = base_url + self.token = token + self.headers = dict(headers or {}) # snapshot + self.timeout = timeout + + def _base_headers(self) -> Dict[str, str]: + return client_headers(self.token, self.headers) + + def _req(self, method, schema, entity, id, options=None, body=None): + opt = build_headers(options) if options else {} + return ( + method, + build_url(self.base_url, schema, entity, id), + merge_headers(self._base_headers(), opt), + body, + ) + + def _read_req(self, schema, entity, id, options): + return self._req("GET", schema, entity, id, options) + + def _create_req(self, schema, entity, data, options): + return self._req("POST", schema, entity, None, options, data) + + def _update_req(self, schema, entity, id, data, options): + return self._req("PUT", schema, entity, id, options, data) + + def _delete_req(self, schema, entity, id): + return self._req("DELETE", schema, entity, id) + + +class HeaderSpecClient(_Base): + """Synchronous client. Use as a context manager or call close().""" + + def __init__(self, *args: Any, transport: Optional[httpx.BaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.Client(timeout=self.timeout, transport=transport) + + def close(self) -> None: + self._http.close() + + def __enter__(self) -> "HeaderSpecClient": + return self + + def __exit__(self, *exc: Any) -> None: + self.close() + + def _send(self, req) -> APIResponse: + method, url, headers, body = req + return _wrap(self._http.request(method, url, headers=headers, json=body)) + + def read(self, schema: str, entity: str, id: Optional[str] = None, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return self._send(self._read_req(schema, entity, id, options)) + + def create(self, schema: str, entity: str, data: Any, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return self._send(self._create_req(schema, entity, data, options)) + + def update(self, schema: str, entity: str, id: str, data: Any, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return self._send(self._update_req(schema, entity, id, data, options)) + + def delete(self, schema: str, entity: str, id: str) -> APIResponse: + return self._send(self._delete_req(schema, entity, id)) + + +class AsyncHeaderSpecClient(_Base): + """Asyncio client. Use as an async context manager or await aclose().""" + + def __init__(self, *args: Any, transport: Optional[httpx.AsyncBaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.AsyncClient(timeout=self.timeout, transport=transport) + + async def aclose(self) -> None: + await self._http.aclose() + + async def __aenter__(self) -> "AsyncHeaderSpecClient": + return self + + async def __aexit__(self, *exc: Any) -> None: + await self.aclose() + + async def _send(self, req) -> APIResponse: + method, url, headers, body = req + return _wrap(await self._http.request(method, url, headers=headers, json=body)) + + async def read(self, schema: str, entity: str, id: Optional[str] = None, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return await self._send(self._read_req(schema, entity, id, options)) + + async def create(self, schema: str, entity: str, data: Any, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return await self._send(self._create_req(schema, entity, data, options)) + + async def update(self, schema: str, entity: str, id: str, data: Any, options: Optional[HeaderSpecOptions] = None) -> APIResponse: + return await self._send(self._update_req(schema, entity, id, data, options)) + + async def delete(self, schema: str, entity: str, id: str) -> APIResponse: + return await self._send(self._delete_req(schema, entity, id)) diff --git a/resolvespec-python/src/resolvespec/http.py b/resolvespec-python/src/resolvespec/http.py new file mode 100644 index 0000000..8029bed --- /dev/null +++ b/resolvespec-python/src/resolvespec/http.py @@ -0,0 +1,76 @@ +"""Shared HTTP helpers for the REST clients.""" +from __future__ import annotations + +from typing import Any, Dict, Mapping, Optional +from urllib.parse import quote + + +class ResolveSpecError(Exception): + """Raised on a non-2xx response or an unsuccessful API result.""" + + def __init__( + self, + message: str, + status_code: Optional[int] = None, + code: Optional[str] = None, + details: Any = None, + detail: Optional[str] = None, + ): + super().__init__(message) + self.message = message + self.status_code = status_code + self.code = code + self.details = details + self.detail = detail # server-side reason (funcspec / restheadspec errors) + + +def merge_headers(*sources: Mapping[str, str]) -> Dict[str, str]: + """Merge HTTP headers case-insensitively; the last source wins and keeps its spelling.""" + result: Dict[str, str] = {} + for source in sources: + for name, value in source.items(): + for existing in [k for k in result if k.lower() == name.lower()]: + del result[existing] + result[name] = value + return result + + +def client_headers(token: Optional[str], headers: Optional[Mapping[str, str]]) -> Dict[str, str]: + """Content-Type < custom headers < bearer token.""" + return merge_headers( + {"Content-Type": "application/json"}, + headers or {}, + {"Authorization": f"Bearer {token}"} if token else {}, + ) + + +def build_url(base_url: str, schema: str, entity: str, id: Optional[Any] = None) -> str: + url = f"{base_url.rstrip('/')}/{quote(schema, safe='')}/{quote(entity, safe='')}" + if id is not None and id != "": + url += f"/{quote(str(id), safe='')}" + return url + + +def drop_none(d: Mapping[str, Any]) -> Dict[str, Any]: + return {k: v for k, v in d.items() if v is not None} + + +def parse_json(response: Any) -> Any: + try: + return response.json() + except ValueError: + return None + + +def error_from(response: Any, data: Any) -> ResolveSpecError: + err = data.get("error") if isinstance(data, dict) else None + err = err if isinstance(err, dict) else {} + text = (response.text or "").strip() if data is None else "" + fallback = text[:200] or f"{response.reason_phrase} ({response.status_code})" + return ResolveSpecError( + err.get("message") or fallback, + status_code=response.status_code, + code=err.get("code"), + details=err.get("details"), + detail=err.get("detail"), + ) diff --git a/resolvespec-python/src/resolvespec/resolvespec.py b/resolvespec-python/src/resolvespec/resolvespec.py new file mode 100644 index 0000000..070721b --- /dev/null +++ b/resolvespec-python/src/resolvespec/resolvespec.py @@ -0,0 +1,137 @@ +"""ResolveSpec client: JSON body protocol (POST {operation, data, options}).""" +from __future__ import annotations + +from typing import Any, Dict, List, Mapping, Optional, Tuple + +import httpx + +from .http import build_url, client_headers, drop_none, error_from, parse_json +from .types import APIResponse, Options, RecordId + + +def _url_id(id: Optional[RecordId]) -> Optional[str]: + return str(id) if isinstance(id, (int, str)) else None + + +def _body_id(id: Optional[RecordId]) -> Optional[List[str]]: + return id if isinstance(id, list) else None + + +class _Base: + def __init__( + self, + base_url: str, + token: Optional[str] = None, + headers: Optional[Mapping[str, str]] = None, + timeout: Optional[float] = 30.0, + ): + self.base_url = base_url + self.token = token + self.headers = dict(headers or {}) # snapshot + self.timeout = timeout + + def _headers(self) -> Dict[str, str]: + return client_headers(self.token, self.headers) + + def _request( + self, method: str, schema: str, entity: str, id: Optional[str], body: Optional[Dict[str, Any]] + ) -> Tuple[str, str, Dict[str, str], Optional[Dict[str, Any]]]: + return method, build_url(self.base_url, schema, entity, id), self._headers(), body + + @staticmethod + def _result(response: httpx.Response) -> APIResponse: + data = parse_json(response) + if not response.is_success: + raise error_from(response, data) + return data + + # request builders (shared by sync and async) + def _metadata_req(self, schema, entity): + return self._request("GET", schema, entity, None, None) + + def _read_req(self, schema, entity, id, options): + body = drop_none({"operation": "read", "id": _body_id(id), "options": options}) + return self._request("POST", schema, entity, _url_id(id), body) + + def _create_req(self, schema, entity, data, options): + body = drop_none({"operation": "create", "data": data, "options": options}) + return self._request("POST", schema, entity, None, body) + + def _update_req(self, schema, entity, data, id, options): + body = drop_none({"operation": "update", "id": _body_id(id), "data": data, "options": options}) + return self._request("POST", schema, entity, _url_id(id), body) + + def _delete_req(self, schema, entity, id): + return self._request("POST", schema, entity, str(id), {"operation": "delete"}) + + +class ResolveSpecClient(_Base): + """Synchronous client. Use as a context manager or call close().""" + + def __init__(self, *args: Any, transport: Optional[httpx.BaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.Client(timeout=self.timeout, transport=transport) + + def close(self) -> None: + self._http.close() + + def __enter__(self) -> "ResolveSpecClient": + return self + + def __exit__(self, *exc: Any) -> None: + self.close() + + def _send(self, req) -> APIResponse: + method, url, headers, body = req + return self._result(self._http.request(method, url, headers=headers, json=body)) + + def get_metadata(self, schema: str, entity: str) -> APIResponse: + return self._send(self._metadata_req(schema, entity)) + + def read(self, schema: str, entity: str, id: Optional[RecordId] = None, options: Optional[Options] = None) -> APIResponse: + return self._send(self._read_req(schema, entity, id, options)) + + def create(self, schema: str, entity: str, data: Any, options: Optional[Options] = None) -> APIResponse: + return self._send(self._create_req(schema, entity, data, options)) + + def update(self, schema: str, entity: str, data: Any, id: Optional[RecordId] = None, options: Optional[Options] = None) -> APIResponse: + return self._send(self._update_req(schema, entity, data, id, options)) + + def delete(self, schema: str, entity: str, id: Any) -> APIResponse: + return self._send(self._delete_req(schema, entity, id)) + + +class AsyncResolveSpecClient(_Base): + """Asyncio client. Use as an async context manager or await aclose().""" + + def __init__(self, *args: Any, transport: Optional[httpx.AsyncBaseTransport] = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._http = httpx.AsyncClient(timeout=self.timeout, transport=transport) + + async def aclose(self) -> None: + await self._http.aclose() + + async def __aenter__(self) -> "AsyncResolveSpecClient": + return self + + async def __aexit__(self, *exc: Any) -> None: + await self.aclose() + + async def _send(self, req) -> APIResponse: + method, url, headers, body = req + return self._result(await self._http.request(method, url, headers=headers, json=body)) + + async def get_metadata(self, schema: str, entity: str) -> APIResponse: + return await self._send(self._metadata_req(schema, entity)) + + async def read(self, schema: str, entity: str, id: Optional[RecordId] = None, options: Optional[Options] = None) -> APIResponse: + return await self._send(self._read_req(schema, entity, id, options)) + + async def create(self, schema: str, entity: str, data: Any, options: Optional[Options] = None) -> APIResponse: + return await self._send(self._create_req(schema, entity, data, options)) + + async def update(self, schema: str, entity: str, data: Any, id: Optional[RecordId] = None, options: Optional[Options] = None) -> APIResponse: + return await self._send(self._update_req(schema, entity, data, id, options)) + + async def delete(self, schema: str, entity: str, id: Any) -> APIResponse: + return await self._send(self._delete_req(schema, entity, id)) diff --git a/resolvespec-python/src/resolvespec/types.py b/resolvespec-python/src/resolvespec/types.py new file mode 100644 index 0000000..085d5b0 --- /dev/null +++ b/resolvespec-python/src/resolvespec/types.py @@ -0,0 +1,166 @@ +"""Types aligned with Go pkg/common/types.go. Dict keys are the wire names.""" +from __future__ import annotations + +from typing import Any, Dict, List, NotRequired, TypedDict, Union + +Operator = str # eq neq gt gte lt lte like ilike in contains startswith endswith +# between between_inclusive is_null is_not_null +# st_dwithin bbox (spatial) | l2_within cosine_within ip_within (vector) +Operation = str # read | create | update | delete +SortDirection = str # asc | desc | ASC | DESC +VectorMetric = str # l2 | cosine | ip +ResponseFormat = str # simple | detail | syncfusion + +RecordId = Union[int, str, List[str]] + + +class Parameter(TypedDict): + name: str + value: str + sequence: NotRequired[int] + + +class FilterOption(TypedDict): + column: str + operator: str + value: Any + logic_operator: NotRequired[str] # "AND" | "OR" + + +class SortOption(TypedDict): + column: str + direction: str + + +class CustomOperator(TypedDict): + name: str + sql: str + + +class ComputedColumn(TypedDict): + name: str + expression: str + + +class PreloadOption(TypedDict, total=False): + relation: str + table_name: str + columns: List[str] + omit_columns: List[str] + sort: List[SortOption] + filters: List[FilterOption] + where: str + limit: int + offset: int + updateable: bool + computed_ql: Dict[str, str] + recursive: bool + primary_key: str + related_key: str + foreign_key: str + recursive_child_key: str + sql_joins: List[str] + join_aliases: List[str] + + +# `as` is a keyword, so the functional syntax is required. +VectorSearchOption = TypedDict( + "VectorSearchOption", + { + "column": str, + "vector": List[float], + "metric": str, # l2 (default) | cosine | ip + "as": str, # distance column alias, default _distance + "direction": str, # asc (default) | desc + }, + total=False, +) + + +class ExpandOption(TypedDict, total=False): + relation: str + columns: List[str] + + +class XFiles(TypedDict, total=False): + tablename: str + schema: str + primarykey: str + foreignkey: str + relatedkey: str + sort: List[str] + prefix: str + editable: bool + recursive: bool + expand: bool + rownumber: bool + skipcount: bool + offset: int + limit: int + columns: List[str] + omit_columns: List[str] + cql_columns: List[str] + sql_joins: List[str] + sql_or: List[str] + sql_and: List[str] + parenttables: List["XFiles"] + childtables: List["XFiles"] + filter_fields: List[Dict[str, str]] + cursor_forward: str + cursor_backward: str + + +class Options(TypedDict, total=False): + preload: List[PreloadOption] + columns: List[str] + omit_columns: List[str] + filters: List[FilterOption] + sort: List[SortOption] + limit: int + offset: int + customOperators: List[CustomOperator] + computedColumns: List[ComputedColumn] + parameters: List[Parameter] + cursor_forward: str + cursor_backward: str + fetch_row_number: str + vector_search: VectorSearchOption + + +class HeaderSpecOptions(Options, total=False): + """Options only available to the header-based (restheadspec) protocol.""" + + expand: List[ExpandOption] # X-Expand + custom_sql_joins: List[str] # X-Custom-SQL-Join + custom_sql_or: List[str] # X-Custom-SQL-Or + search_columns: List[str] # X-SearchCols + advanced_sql: Dict[str, str] # X-AdvSQL-{col} + clean_json: bool # X-Clean-JSON + distinct: bool # X-Distinct + skip_count: bool # X-SkipCount + skip_cache: bool # X-SkipCache + pk_row: str # X-PKRow + response_format: str # X-SimpleApi / X-DetailApi / X-Syncfusion + single_record_as_object: bool # X-Single-Record-As-Object + atomic_transaction: bool # X-Transaction-Atomic + xfiles: XFiles # X-Files + + +class FuncSpecOptions(TypedDict, total=False): + """Options understood by funcspec endpoints (sent as X-* headers).""" + + filters: List[FilterOption] # eq+AND -> X-FieldFilter; others X-SearchOp / X-SearchOr (one per column) + search_filters: Dict[str, str] # X-SearchFilter-{col}: text ILIKE + custom_sql_where: str # X-Custom-SQL-W + custom_sql_or: str # X-Custom-SQL-Or + sort: List[SortOption] # sent as SQL ORDER BY terms ("col DESC") + limit: int + offset: int + distinct: bool + skip_count: bool + skip_cache: bool + response_format: str # simple | detail | syncfusion + + +# Responses are plain dicts: {"success", "data", "metadata"?, "error"?} +APIResponse = Dict[str, Any] diff --git a/resolvespec-python/src/resolvespec/websocket.py b/resolvespec-python/src/resolvespec/websocket.py new file mode 100644 index 0000000..b14bdbc --- /dev/null +++ b/resolvespec-python/src/resolvespec/websocket.py @@ -0,0 +1,335 @@ +"""WebSocketSpec client (asyncio). Mirrors the Go websocketspec message protocol.""" +from __future__ import annotations + +import asyncio +import json +import logging +import uuid +from dataclasses import dataclass, field +from typing import Any, Awaitable, Callable, Dict, List, Optional, Union + +from websockets.asyncio.client import ClientConnection, connect + +from .http import ResolveSpecError +from .types import FilterOption, PreloadOption, SortOption + +log = logging.getLogger("resolvespec.websocket") + +# Connection states +DISCONNECTED = "disconnected" +CONNECTING = "connecting" +CONNECTED = "connected" +DISCONNECTING = "disconnecting" +RECONNECTING = "reconnecting" + +Notification = Dict[str, Any] +Callback = Callable[[Any], Union[None, Awaitable[None]]] +EVENTS = ("connect", "disconnect", "error", "message", "state_change") + + +@dataclass +class Subscription: + id: str + entity: str + schema: Optional[str] = None + options: Optional[Dict[str, Any]] = None + callback: Optional[Callback] = field(default=None, repr=False) + + +def _drop_none(d: Dict[str, Any]) -> Dict[str, Any]: + return {k: v for k, v in d.items() if v is not None} + + +class WebSocketClient: + """ + Usage: + async with WebSocketClient("ws://localhost:8080/ws") as ws: + rows = await ws.read("users", schema="public", limit=10) + + Events (`on(event, callback)`): connect, disconnect, error, message, state_change. + Callbacks may be sync or async. + """ + + def __init__( + self, + url: str, + *, + reconnect: bool = True, + reconnect_interval: float = 3.0, + max_reconnect_attempts: int = 10, + heartbeat_interval: float = 30.0, + request_timeout: float = 30.0, + subscribe_timeout: float = 10.0, + headers: Optional[Dict[str, str]] = None, + ): + self.url = url + self.reconnect = reconnect + self.reconnect_interval = reconnect_interval + self.max_reconnect_attempts = max_reconnect_attempts + self.heartbeat_interval = heartbeat_interval + self.request_timeout = request_timeout + self.subscribe_timeout = subscribe_timeout + self.headers = dict(headers or {}) + + self._ws: Optional[ClientConnection] = None + self._state = DISCONNECTED + self._pending: Dict[str, "asyncio.Future[Dict[str, Any]]"] = {} + self._subscriptions: Dict[str, Subscription] = {} + self._listeners: Dict[str, Callback] = {} + self._tasks: List["asyncio.Task[Any]"] = [] + self._reader: Optional["asyncio.Task[Any]"] = None + self._manual_close = False + + # ---- lifecycle ------------------------------------------------------- + + async def __aenter__(self) -> "WebSocketClient": + await self.connect() + return self + + async def __aexit__(self, *exc: Any) -> None: + await self.close() + + async def connect(self) -> None: + if self.is_connected(): + return + self._manual_close = False + self._set_state(CONNECTING) + try: + self._ws = await connect(self.url, additional_headers=self.headers or None) + except Exception as e: + self._set_state(DISCONNECTED) + await self._emit("error", e) + raise + self._set_state(CONNECTED) + self._reader = asyncio.create_task(self._read_loop(self._ws)) + self._heartbeat = asyncio.create_task(self._heartbeat_loop()) + await self._emit("connect") + + async def close(self) -> None: + self._manual_close = True + self._set_state(DISCONNECTING) + for t in (self._reader, getattr(self, "_heartbeat", None), getattr(self, "_reconnect_task", None)): + if t and t is not asyncio.current_task(): + t.cancel() + if self._ws: + await self._ws.close() + self._ws = None + self._fail_pending(ResolveSpecError("WebSocket closed")) + self._set_state(DISCONNECTED) + + def is_connected(self) -> bool: + return self._ws is not None and self._state == CONNECTED + + @property + def state(self) -> str: + return self._state + + def on(self, event: str, callback: Callback) -> None: + if event not in EVENTS: + raise ValueError(f"unknown event {event!r}; expected one of {EVENTS}") + self._listeners[event] = callback + + def off(self, event: str) -> None: + self._listeners.pop(event, None) + + def get_subscriptions(self) -> List[Subscription]: + return list(self._subscriptions.values()) + + # ---- operations ------------------------------------------------------ + + async def request( + self, + operation: str, + entity: str, + *, + schema: Optional[str] = None, + record_id: Optional[str] = None, + data: Any = None, + options: Optional[Dict[str, Any]] = None, + ) -> Any: + message = _drop_none({ + "type": "request", + "operation": operation, + "entity": entity, + "schema": schema, + "record_id": record_id, + "data": data, + "options": options, + }) + response = await self._call(message, self.request_timeout, "Request") + return response.get("data") + + async def read( + self, + entity: str, + *, + schema: Optional[str] = None, + record_id: Optional[str] = None, + filters: Optional[List[FilterOption]] = None, + columns: Optional[List[str]] = None, + sort: Optional[List[SortOption]] = None, + preload: Optional[List[PreloadOption]] = None, + limit: Optional[int] = None, + offset: Optional[int] = None, + ) -> Any: + options = _drop_none({ + "filters": filters, "columns": columns, "sort": sort, + "preload": preload, "limit": limit, "offset": offset, + }) + return await self.request("read", entity, schema=schema, record_id=record_id, options=options) + + async def create(self, entity: str, data: Any, *, schema: Optional[str] = None) -> Any: + return await self.request("create", entity, schema=schema, data=data) + + async def update(self, entity: str, id: str, data: Any, *, schema: Optional[str] = None) -> Any: + return await self.request("update", entity, schema=schema, record_id=id, data=data) + + async def delete(self, entity: str, id: str, *, schema: Optional[str] = None) -> None: + await self.request("delete", entity, schema=schema, record_id=id) + + async def meta(self, entity: str, *, schema: Optional[str] = None) -> Any: + return await self.request("meta", entity, schema=schema) + + async def subscribe( + self, + entity: str, + callback: Callback, + *, + schema: Optional[str] = None, + filters: Optional[List[FilterOption]] = None, + ) -> str: + message = _drop_none({ + "type": "subscription", + "operation": "subscribe", + "entity": entity, + "schema": schema, + "options": _drop_none({"filters": filters}), + }) + response = await self._call(message, self.subscribe_timeout, "Subscription") + sub_id = (response.get("data") or {}).get("subscription_id") + if not sub_id: + raise ResolveSpecError("Subscription failed") + self._subscriptions[sub_id] = Subscription( + sub_id, entity, schema, _drop_none({"filters": filters}) or None, callback + ) + return sub_id + + async def unsubscribe(self, subscription_id: str) -> None: + message = {"type": "subscription", "operation": "unsubscribe", "subscription_id": subscription_id} + await self._call(message, self.subscribe_timeout, "Unsubscribe") + self._subscriptions.pop(subscription_id, None) + + # ---- internals ------------------------------------------------------- + + async def _call(self, message: Dict[str, Any], timeout: float, what: str) -> Dict[str, Any]: + self._ensure_connected() + mid = str(uuid.uuid4()) + message["id"] = mid + fut: "asyncio.Future[Dict[str, Any]]" = asyncio.get_running_loop().create_future() + self._pending[mid] = fut + try: + await self._ws.send(json.dumps(message)) # type: ignore[union-attr] + response = await asyncio.wait_for(fut, timeout) + except asyncio.TimeoutError: + raise ResolveSpecError(f"{what} timeout") from None + finally: + self._pending.pop(mid, None) + if not response.get("success"): + err = response.get("error") or {} + raise ResolveSpecError( + err.get("message") or f"{what} failed", code=err.get("code"), details=err.get("details") + ) + return response + + def _ensure_connected(self) -> None: + if not self.is_connected(): + raise ResolveSpecError("WebSocket is not connected. Call connect() first.") + + def _fail_pending(self, exc: Exception) -> None: + for fut in self._pending.values(): + if not fut.done(): + fut.set_exception(exc) + self._pending.clear() + + async def _read_loop(self, ws: ClientConnection) -> None: + try: + async for raw in ws: + await self._handle_message(raw) + except asyncio.CancelledError: + raise + except Exception as e: # connection error + await self._emit("error", e) + # connection ended + if ws is not self._ws: + return + self._ws = None + if hb := getattr(self, "_heartbeat", None): + hb.cancel() + self._fail_pending(ResolveSpecError("WebSocket disconnected")) + self._set_state(DISCONNECTED) + await self._emit("disconnect", ws.close_code, ws.close_reason) + if self.reconnect and not self._manual_close: + self._reconnect_task = asyncio.create_task(self._reconnect()) + + async def _reconnect(self) -> None: + for attempt in range(1, self.max_reconnect_attempts + 1): + if self._manual_close: + return + log.debug("Reconnection attempt %d/%d", attempt, self.max_reconnect_attempts) + self._set_state(RECONNECTING) + await asyncio.sleep(self.reconnect_interval) + try: + await self.connect() + return + except Exception as e: + log.debug("Reconnection failed: %s", e) + self._set_state(DISCONNECTED) + + async def _handle_message(self, raw: Union[str, bytes]) -> None: + try: + message = json.loads(raw) + except ValueError as e: + log.debug("Error parsing message: %s", e) + return + await self._emit("message", message) + kind = message.get("type") + if kind == "response": + fut = self._pending.get(message.get("id")) + if fut and not fut.done(): + fut.set_result(message) + elif kind == "notification": + sub = self._subscriptions.get(message.get("subscription_id")) + if sub and sub.callback: + await _maybe_await(sub.callback(message)) + elif kind != "pong": + log.debug("Unknown message type: %s", kind) + + async def _heartbeat_loop(self) -> None: + try: + while True: + await asyncio.sleep(self.heartbeat_interval) + if self.is_connected(): + await self._ws.send(json.dumps({"id": str(uuid.uuid4()), "type": "ping"})) # type: ignore[union-attr] + except asyncio.CancelledError: + raise + except Exception as e: + log.debug("Heartbeat failed: %s", e) + + def _set_state(self, state: str) -> None: + if self._state != state: + self._state = state + cb = self._listeners.get("state_change") + if cb: + res = cb(state) + if asyncio.iscoroutine(res): + asyncio.ensure_future(res) + + async def _emit(self, event: str, *args: Any) -> None: + cb = self._listeners.get(event) + if cb: + await _maybe_await(cb(*args)) + + +async def _maybe_await(result: Any) -> None: + if asyncio.iscoroutine(result) or isinstance(result, asyncio.Future): + await result diff --git a/resolvespec-python/tests/test_funcspec.py b/resolvespec-python/tests/test_funcspec.py new file mode 100644 index 0000000..3d39625 --- /dev/null +++ b/resolvespec-python/tests/test_funcspec.py @@ -0,0 +1,132 @@ +import httpx +import pytest + +from resolvespec import AsyncFuncSpecClient, FuncSpecClient, ResolveSpecError +from resolvespec.funcspec import build_headers, build_query +from resolvespec.headerspec import decode_header_value + + +def make(handler, **kw): + return FuncSpecClient("http://localhost:3000", "tok", transport=httpx.MockTransport(handler), **kw) + + +def capture(status=200, body=None, headers=None): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(status, json=body if body is not None else [], headers=headers) + + return seen, handler + + +def test_filters(): + h = build_headers({"filters": [ + {"column": "status", "operator": "eq", "value": "active"}, + {"column": "age", "operator": "gte", "value": 18}, + {"column": "name", "operator": "contains", "value": "x", "logic_operator": "OR"}, + {"column": "deleted", "operator": "is_null", "value": None}, + {"column": "id", "operator": "in", "value": [1, 2]}, + {"column": "p", "operator": "between_inclusive", "value": [1, 5]}, + ]}) + assert h == { + "X-FieldFilter-status": "active", + "X-SearchOp-greaterthanorequal-age": "18", + "X-SearchOr-contains-name": "x", + "X-SearchOp-empty-deleted": "", + "X-SearchOp-in-id": "1,2", + "X-SearchOp-betweeninclusive-p": "1,5", + } + + +def test_sort_is_sql_not_prefixed(): + # server inserts sort verbatim into ORDER BY; "-col" would negate the column + h = build_headers({"sort": [{"column": "name", "direction": "asc"}, {"column": "created_at", "direction": "DESC"}]}) + assert h["X-Sort"] == "name ASC,created_at DESC" + + +def test_misc_options(): + h = build_headers({ + "search_filters": {"name": "bob"}, "custom_sql_where": "a = 1", "custom_sql_or": "b = 2", + "limit": 5, "offset": 10, "distinct": True, "skip_count": True, "skip_cache": False, + "response_format": "syncfusion", + }) + assert h == { + "X-SearchFilter-name": "bob", "X-Custom-SQL-W": "a = 1", "X-Custom-SQL-Or": "b = 2", + "X-Limit": "5", "X-Offset": "10", "X-Distinct": "true", "X-SkipCount": "true", + "X-SkipCache": "false", "X-Syncfusion": "true", + } + + +def test_ambiguous_values_are_encoded(): + h = build_headers({"custom_sql_where": "name = 'café'", "filters": [{"column": "c", "operator": "eq", "value": " pad "}]}) + assert h["X-Custom-SQL-W"].startswith("ZIP_") + assert decode_header_value(h["X-Custom-SQL-W"]) == "name = 'café'" + assert decode_header_value(h["X-FieldFilter-c"]) == " pad " + + +def test_build_query(): + q = build_query({"p-id": 5, "flag": True, "ids": [1, 2], "skip": None, "m": "match=ab"}) + assert q == {"p-id": "5", "flag": "true", "ids": ["1", "2"], "m": "match=ab"} + + +def test_query_list_request_and_metadata(): + seen, h = capture(206, [{"id": 1}, {"id": 2}], {"content-range": "items 10-12/50"}) + with make(h) as c: + res = c.query_list("/api/orders", {"p-status": "open", "id": [1, 2]}, {"limit": 2, "offset": 10}) + r = seen[0] + assert r.method == "GET" + assert r.url.path == "/api/orders" + assert r.url.params.multi_items() == [("p-status", "open"), ("id", "1"), ("id", "2")] + assert r.headers["x-limit"] == "2" and r.headers["authorization"] == "Bearer tok" + assert res == { + "success": True, + "data": [{"id": 1}, {"id": 2}], + "metadata": {"total": 50, "count": 2, "filtered": 50, "offset": 10, "limit": 2}, + } + + +def test_query_list_empty_result(): + seen, h = capture(200, [], {"content-range": "items 0-0/0"}) + with make(h) as c: + assert c.query_list("orders")["metadata"]["total"] == 0 + assert seen[0].url.path == "/orders" + + +def test_query_single_has_no_metadata_and_method(): + seen, h = capture(200, {"id": 1}) + with make(h) as c: + res = c.query("api/order", method="post") + assert seen[0].method == "POST" + assert res == {"success": True, "data": {"id": 1}} + + +def test_detail_format_data_passthrough(): + body = {"items": [{"a": 1}], "count": "1", "total": "1", "tablename": "/x", "tableprefix": "gsql"} + _, h = capture(200, body, {"content-range": "items 0-1/1"}) + with make(h) as c: + assert c.query_list("x", options={"response_format": "detail"})["data"] == body + + +def test_server_error_shape(): + err = {"success": False, "error": {"code": "query_failed", "message": "Failed to retrieve records", "detail": "no such column", "sql": "SELECT"}} + _, h = capture(400, err) + with make(h) as c: + with pytest.raises(ResolveSpecError, match="Failed to retrieve") as ei: + c.query_list("x") + assert ei.value.code == "query_failed" and ei.value.detail == "no such column" and ei.value.status_code == 400 + + +def test_plain_text_panic_error(): + with make(lambda r: httpx.Response(500, text="Internal server error: boom")) as c: + with pytest.raises(ResolveSpecError, match="boom"): + c.query("x") + + +async def test_async(): + async def handler(req): + return httpx.Response(200, json=[{"id": 1}], headers={"content-range": "items 0-1/1"}) + + async with AsyncFuncSpecClient("http://localhost:3000", transport=httpx.MockTransport(handler)) as c: + assert (await c.query_list("x"))["metadata"]["total"] == 1 + assert (await c.query("x"))["data"] == [{"id": 1}] diff --git a/resolvespec-python/tests/test_headerspec.py b/resolvespec-python/tests/test_headerspec.py new file mode 100644 index 0000000..276e4e6 --- /dev/null +++ b/resolvespec-python/tests/test_headerspec.py @@ -0,0 +1,236 @@ +import json + +import httpx +import pytest + +from resolvespec import ( + AsyncHeaderSpecClient, + HeaderSpecClient, + ResolveSpecError, + build_headers, + decode_header_value, + encode_header_value, + get_headerspec_client, +) +import base64 + +CFG = dict(base_url="http://localhost:3000", token="tok") + + +# ---- build_headers (ported from headerspec.test.ts) ---- + +def test_preload_shared_where(): + h = build_headers({"preload": [ + {"relation": "Items", "columns": ["id"], "where": "active = true"}, + {"relation": "Tags", "where": "active = true"}, + ]}) + assert h["X-Preload"] == "Items:id|Tags" + assert h["X-Preload-Where"] == "active = true" + + +def test_preload_mixed_where_numbered(): + h = build_headers({"preload": [ + {"relation": "Items", "where": "a = 1"}, + {"relation": "Category"}, + {"relation": "Tags", "where": "b = 2"}, + ]}) + assert h["X-Preload"] == "Category" + assert "X-Preload-Where" not in h + assert h["X-Preload-1"] == "Items" and h["X-Preload-1-Where"] == "a = 1" + assert h["X-Preload-2"] == "Tags" and h["X-Preload-2-Where"] == "b = 2" + + +def test_expand_joins_or_searchcols_advsql(): + h = build_headers({ + "expand": [{"relation": "Dept", "columns": ["id", "name"]}, {"relation": "Role"}], + "custom_sql_joins": ["LEFT JOIN a ON a.id = b.id", "INNER JOIN c ON c.id = b.cid"], + "custom_sql_or": ["x = 1", "y = 2"], + "search_columns": ["name", "email"], + "advanced_sql": {"total": "a + b"}, + }) + assert h["X-Expand"] == "Dept:id,name|Role" + assert h["X-Custom-SQL-Join"] == "LEFT JOIN a ON a.id = b.id|INNER JOIN c ON c.id = b.cid" + assert h["X-Custom-SQL-Or"] == "x = 1 OR y = 2" + assert h["X-SearchCols"] == "name,email" + assert h["X-AdvSQL-total"] == "a + b" + + +def test_flags_pkrow_format(): + h = build_headers({ + "clean_json": True, "distinct": True, "skip_count": True, "skip_cache": False, + "atomic_transaction": True, "single_record_as_object": False, + "pk_row": "42", "response_format": "detail", + }) + assert h["X-Clean-JSON"] == "true" + assert h["X-Distinct"] == "true" + assert h["X-SkipCount"] == "true" + assert h["X-SkipCache"] == "false" + assert h["X-Transaction-Atomic"] == "true" + assert h["X-Single-Record-As-Object"] == "false" + assert h["X-PKRow"] == "42" + assert h["X-DetailApi"] == "true" + + +def test_spatial_and_vector_filters(): + h = build_headers({"filters": [ + {"column": "geom", "operator": "st_dwithin", "value": {"geom": "POINT(0 0)", "distance": 5}, "logic_operator": "OR"}, + {"column": "emb", "operator": "cosine_within", "value": {"vector": [1, 2], "distance": 0.3}}, + ]}) + assert json.loads(h["X-SpatialFilter-geom"]) == { + "op": "st_dwithin", "value": {"geom": "POINT(0 0)", "distance": 5}, "logic": "or"} + assert json.loads(h["X-VectorFilter-emb"])["op"] == "cosine_within" + + +def test_vector_search(): + h = build_headers({"vector_search": {"column": "emb", "vector": [0.1, 0.2], "metric": "cosine", "as": "dist", "direction": "desc"}}) + assert h["X-Vector-Search-emb"] == "cosine" + assert h["X-Vector-Search-Vector"] == "[0.1,0.2]" + assert h["X-Vector-Search-As"] == "dist" + assert h["X-Vector-Search-Dir"] == "desc" + + +def test_xfiles_zip(): + xf = {"tablename": "users", "prefix": "USR", "limit": 10} + h = build_headers({"xfiles": xf}) + assert h["X-Files"].startswith("ZIP_") + assert json.loads(decode_header_value(h["X-Files"])) == xf + + +def test_columns_and_omit(): + assert build_headers({"columns": ["id", "name", "email"]})["X-Select-Fields"] == "id,name,email" + assert build_headers({"omit_columns": ["secret", "internal"]})["X-Not-Select-Fields"] == "secret,internal" + + +def test_filters(): + assert build_headers({"filters": [{"column": "status", "operator": "eq", "value": "active"}]})["X-FieldFilter-status"] == "active" + assert build_headers({"filters": [{"column": "age", "operator": "gte", "value": 18}]})["X-SearchOp-greaterthanorequal-age"] == "18" + assert build_headers({"filters": [{"column": "name", "operator": "contains", "value": "test", "logic_operator": "OR"}]})["X-SearchOr-contains-name"] == "test" + assert build_headers({"filters": [{"column": "price", "operator": "between", "value": [10, 100]}]})["X-SearchOp-between-price"] == "10,100" + assert build_headers({"filters": [{"column": "deleted_at", "operator": "is_null", "value": None}]})["X-SearchOp-empty-deleted_at"] == "" + assert build_headers({"filters": [{"column": "id", "operator": "in", "value": [1, 2, 3]}]})["X-SearchOp-in-id"] == "1,2,3" + assert build_headers({"filters": [{"column": "a", "operator": "eq", "value": True}]})["X-FieldFilter-a"] == "true" + + +def test_sort_pagination_cursor(): + h = build_headers({ + "sort": [{"column": "name", "direction": "asc"}, {"column": "created_at", "direction": "DESC"}], + "limit": 25, "offset": 0, "cursor_forward": "abc", "cursor_backward": "xyz", + }) + assert h["X-Sort"] == "+name,-created_at" + assert h["X-Limit"] == "25" and h["X-Offset"] == "0" + assert h["X-Cursor-Forward"] == "abc" and h["X-Cursor-Backward"] == "xyz" + + +def test_preload_basic_rownumber_computed_custom(): + h = build_headers({ + "preload": [{"relation": "Items", "columns": ["id", "name"]}, {"relation": "Category"}], + "fetch_row_number": "42", + "computedColumns": [{"name": "total", "expression": "price * qty"}], + "customOperators": [{"name": "a", "sql": "status = 'active'"}, {"name": "v", "sql": "verified = true"}], + }) + assert h["X-Preload"] == "Items:id,name|Category" + assert h["X-Fetch-RowNumber"] == "42" + assert h["X-CQL-SEL-total"] == "price * qty" + assert h["X-Custom-SQL-W"] == "status = 'active' AND verified = true" + + +def test_empty_options(): + assert build_headers({}) == {} + + +# ---- encode / decode ---- + +def test_roundtrip(): + for s in ("some complex value with spaces & symbols!", "café ☕ 你好"): + enc = encode_header_value(s) + assert enc.startswith("ZIP_") + assert decode_header_value(enc) == s + + +def test_decode_double_underscore_and_plain(): + assert decode_header_value("__" + base64.b64encode(b"hello").decode()) == "hello" + assert decode_header_value("__" + base64.b64encode("café ☕".encode()).decode()) == "café ☕" + assert decode_header_value("plain") == "plain" + + +def test_decode_nested(): + assert decode_header_value(encode_header_value(encode_header_value("x"))) == "x" + + +# ---- client ---- + +def make(handler, cls=HeaderSpecClient, **kw): + return cls(**{**CFG, **kw}, transport=httpx.MockTransport(handler)) + + +def test_read_sends_get_with_headers(): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(200, json=[{"id": 1}], headers={"content-range": "0-9/100", "x-limit": "10"}) + + with make(handler) as c: + res = c.read("public", "users", options={"columns": ["id", "name"], "limit": 10}) + r = seen[0] + assert str(r.url) == "http://localhost:3000/public/users" + assert r.method == "GET" + assert r.headers["x-select-fields"] == "id,name" + assert r.headers["x-limit"] == "10" + assert r.headers["authorization"] == "Bearer tok" + assert res["success"] is True + assert res["data"] == [{"id": 1}] + assert res["metadata"] == {"count": 100, "total": 100, "filtered": 100, "offset": 0, "limit": 10} + + +def test_metadata_defaults_without_content_range(): + with make(lambda r: httpx.Response(200, json=[])) as c: + assert c.read("public", "users")["metadata"]["total"] == 0 + + +def test_read_with_id_create_update_delete(): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(200, json={}) + + with make(handler) as c: + c.read("public", "users", "42") + c.create("public", "users", {"name": "Test"}) + c.update("public", "users", "1", {"name": "Updated"}, {"filters": [{"column": "active", "operator": "eq", "value": True}]}) + c.delete("public", "users", "1") + assert str(seen[0].url) == "http://localhost:3000/public/users/42" + assert seen[1].method == "POST" and json.loads(seen[1].content) == {"name": "Test"} + assert seen[2].method == "PUT" and str(seen[2].url).endswith("/public/users/1") + assert seen[2].headers["x-fieldfilter-active"] == "true" + assert seen[3].method == "DELETE" + + +def test_error_response(): + with make(lambda r: httpx.Response(400, json={"error": {"code": "err", "message": "fail"}})) as c: + with pytest.raises(ResolveSpecError, match="fail") as ei: + c.read("public", "users") + assert ei.value.status_code == 400 and ei.value.code == "err" + + +def test_error_non_json(): + with make(lambda r: httpx.Response(502, text="bad gateway")) as c: + with pytest.raises(ResolveSpecError, match="bad gateway") as ei: + c.read("public", "users") + assert ei.value.status_code == 502 + + +async def test_async_client(): + async def handler(req): + return httpx.Response(200, json=[{"id": 1}]) + + async with AsyncHeaderSpecClient(**CFG, transport=httpx.MockTransport(handler)) as c: + res = await c.read("public", "users", options={"limit": 1}) + assert res["data"] == [{"id": 1}] + + +def test_singleton(): + a = get_headerspec_client("http://hs-singleton:3000") + assert a is get_headerspec_client("http://hs-singleton:3000") + assert a is not get_headerspec_client("http://hs-singleton-b:3000") diff --git a/resolvespec-python/tests/test_resolvespec.py b/resolvespec-python/tests/test_resolvespec.py new file mode 100644 index 0000000..90686ab --- /dev/null +++ b/resolvespec-python/tests/test_resolvespec.py @@ -0,0 +1,210 @@ +import json + +import httpx +import pytest + +from resolvespec import ( + AsyncHeaderSpecClient, + AsyncResolveSpecClient, + HeaderSpecClient, + ResolveSpecClient, + ResolveSpecError, + get_headerspec_client, + get_resolvespec_client, +) + +CFG = dict(base_url="http://localhost:3000", token="test-token") + + +def make(handler, **kw): + return ResolveSpecClient(**{**CFG, **kw}, transport=httpx.MockTransport(handler)) + + +def ok(_req): + return httpx.Response(200, json={"success": True, "data": [{"id": 1}]}) + + +def capture(): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(200, json={"success": True, "data": {"id": 1, "name": "Test"}}) + + return seen, handler + + +def body(req): + return json.loads(req.content) + + +def test_read_with_numeric_id(): + seen, h = capture() + with make(h) as c: + assert c.read("public", "users", 1)["success"] is True + r = seen[0] + assert str(r.url) == "http://localhost:3000/public/users/1" + assert r.method == "POST" + assert r.headers["authorization"] == "Bearer test-token" + assert r.headers["content-type"] == "application/json" + assert body(r) == {"operation": "read"} + + +def test_read_array_id_goes_in_body(): + seen, h = capture() + with make(h) as c: + c.read("public", "users", ["1", "2"]) + assert str(seen[0].url) == "http://localhost:3000/public/users" + assert body(seen[0])["id"] == ["1", "2"] + + +def test_read_options_passthrough(): + seen, h = capture() + opts = { + "columns": ["id", "name"], "omit_columns": ["secret"], + "filters": [{"column": "active", "operator": "eq", "value": True}], + "sort": [{"column": "name", "direction": "asc"}], + "limit": 10, "offset": 0, "cursor_forward": "cursor1", "fetch_row_number": "5", + "customOperators": [{"name": "x", "sql": "a = 1"}], + } + with make(h) as c: + c.read("public", "users", options=opts) + assert body(seen[0])["options"] == opts + + +def test_create(): + seen, h = capture() + with make(h) as c: + res = c.create("public", "users", {"name": "Test"}) + assert res["data"]["name"] == "Test" + assert body(seen[0]) == {"operation": "create", "data": {"name": "Test"}} + + +def test_create_batch(): + seen, h = capture() + with make(h) as c: + c.create("public", "users", [{"a": 1}, {"a": 2}]) + assert body(seen[0])["data"] == [{"a": 1}, {"a": 2}] + + +def test_update_with_id_in_url_and_array(): + seen, h = capture() + with make(h) as c: + c.update("public", "users", {"name": "X"}, 5) + c.update("public", "users", {"name": "X"}, ["1", "2"]) + assert str(seen[0].url).endswith("/public/users/5") + assert body(seen[0]) == {"operation": "update", "data": {"name": "X"}} + assert str(seen[1].url).endswith("/public/users") + assert body(seen[1])["id"] == ["1", "2"] + + +def test_update_preserves_empty_string_and_null(): + seen, h = capture() + with make(h) as c: + c.update("public", "users", {"a": "", "b": None}, 1) + assert body(seen[0])["data"] == {"a": "", "b": None} + + +def test_delete(): + seen, h = capture() + with make(h) as c: + c.delete("public", "users", 1) + assert str(seen[0].url).endswith("/public/users/1") + assert body(seen[0]) == {"operation": "delete"} + + +def test_get_metadata(): + seen, h = capture() + with make(h) as c: + c.get_metadata("public", "users") + assert seen[0].method == "GET" + assert str(seen[0].url) == "http://localhost:3000/public/users" + assert not seen[0].content + + +def test_error_uses_server_message(): + with make(lambda r: httpx.Response(404, json={"success": False, "error": {"code": "not_found", "message": "nope"}})) as c: + with pytest.raises(ResolveSpecError, match="nope") as ei: + c.read("public", "users", 1) + assert ei.value.status_code == 404 and ei.value.code == "not_found" + + +def test_id_is_url_quoted(): + seen, h = capture() + with make(h) as c: + c.read("public", "users", "a/b") + assert str(seen[0].url).endswith("/public/users/a%2Fb") + + +def test_trailing_slash_base_url(): + seen, h = capture() + with make(h, base_url="http://localhost:3000/") as c: + c.read("public", "users") + assert str(seen[0].url) == "http://localhost:3000/public/users" + + +async def test_async_client(): + async def handler(req): + return httpx.Response(200, json={"success": True, "data": [1]}) + + async with AsyncResolveSpecClient(**CFG, transport=httpx.MockTransport(handler)) as c: + assert (await c.read("public", "users"))["data"] == [1] + assert (await c.create("public", "users", {}))["success"] + assert (await c.update("public", "users", {}, 1))["success"] + assert (await c.delete("public", "users", 1))["success"] + assert (await c.get_metadata("public", "users"))["success"] + + +# ---- custom headers (ported from custom-headers.test.ts) ---- + +@pytest.mark.parametrize("cls", [ResolveSpecClient, HeaderSpecClient]) +def test_custom_headers_on_every_op_case_insensitive(cls): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(200, json={"success": True, "data": []}) + + headers = {"X-Tenant": "acme", "authorization": "Basic ignored", + "content-type": "application/custom+json", "x-limit": "99"} + with cls("http://localhost:3000", "tok", headers, transport=httpx.MockTransport(handler)) as c: + c.read("public", "users", options={"limit": 10}) + c.create("public", "users", {}) + if cls is ResolveSpecClient: + c.update("public", "users", {}, "1") + c.get_metadata("public", "users") + else: + c.update("public", "users", "1", {}) + c.delete("public", "users", "1") + for r in seen: + assert r.headers["x-tenant"] == "acme" + assert r.headers["authorization"] == "Bearer tok" + assert r.headers["content-type"] == "application/custom+json" + if cls is HeaderSpecClient: + assert seen[0].headers["x-limit"] == "10" + assert headers["authorization"] == "Basic ignored" + assert headers["x-limit"] == "99" + + +@pytest.mark.parametrize("cls", [ResolveSpecClient, HeaderSpecClient]) +def test_custom_auth_without_token(cls): + seen = [] + + def handler(req): + seen.append(req) + return httpx.Response(200, json={"success": True, "data": []}) + + with cls("http://localhost:3000", headers={"Authorization": "Basic custom"}, transport=httpx.MockTransport(handler)) as c: + c.read("public", "users") + assert seen[0].headers["authorization"] == "Basic custom" + + +@pytest.mark.parametrize("factory", [get_resolvespec_client, get_headerspec_client]) +def test_cache_isolation_and_snapshot(factory): + headers = {"X-Tenant": "acme", "X-App": "grid"} + first = factory("http://tenant-cache", "one", headers) + assert factory("http://tenant-cache", "one", {"x-app": "grid", "x-tenant": "acme"}) is first + assert factory("http://tenant-cache", "two", headers) is not first + headers["X-Tenant"] = "other" + assert factory("http://tenant-cache", "one", headers) is not first + assert first.headers["X-Tenant"] == "acme" diff --git a/resolvespec-python/tests/test_websocket.py b/resolvespec-python/tests/test_websocket.py new file mode 100644 index 0000000..0c4f350 --- /dev/null +++ b/resolvespec-python/tests/test_websocket.py @@ -0,0 +1,152 @@ +import asyncio +import json + +import pytest +from websockets.asyncio.server import serve + +from resolvespec import ResolveSpecError, WebSocketClient + + +class Server: + """Minimal in-process WebSocketSpec server.""" + + def __init__(self): + self.received = [] + self.conns = set() + self.respond = True + + async def handler(self, ws): + self.conns.add(ws) + try: + async for raw in ws: + msg = json.loads(raw) + self.received.append(msg) + if msg["type"] == "ping": + await ws.send(json.dumps({"type": "pong"})) + continue + if not self.respond: + continue + await ws.send(json.dumps(self.reply(msg))) + finally: + self.conns.discard(ws) + + def reply(self, msg): + base = {"id": msg["id"], "type": "response", "success": True, "timestamp": "t"} + if msg["type"] == "subscription" and msg["operation"] == "subscribe": + return {**base, "data": {"subscription_id": "sub-1"}} + if msg.get("entity") == "fail": + return {**base, "success": False, "error": {"code": "bad", "message": "boom"}} + return {**base, "data": {"echo": msg.get("operation"), "record_id": msg.get("record_id")}} + + +@pytest.fixture +async def server(): + s = Server() + async with serve(s.handler, "127.0.0.1", 0) as srv: + s.url = "ws://127.0.0.1:%d" % srv.sockets[0].getsockname()[1] + yield s + + +async def test_operations_and_message_shape(server): + async with WebSocketClient(server.url, reconnect=False) as c: + assert c.state == "connected" + assert await c.read("users", schema="public", record_id="1", limit=5, filters=[{"column": "a", "operator": "eq", "value": 1}]) == {"echo": "read", "record_id": "1"} + await c.create("users", {"n": 1}, schema="public") + await c.update("users", "2", {"n": 2}) + await c.delete("users", "3") + await c.meta("users") + m = server.received + assert m[0]["type"] == "request" and m[0]["operation"] == "read" + assert m[0]["schema"] == "public" and m[0]["record_id"] == "1" + assert m[0]["options"] == {"filters": [{"column": "a", "operator": "eq", "value": 1}], "limit": 5} + assert m[1]["data"] == {"n": 1} + assert m[2]["record_id"] == "2" + assert [x["operation"] for x in m] == ["read", "create", "update", "delete", "meta"] + assert "schema" not in m[2] + assert len({x["id"] for x in m}) == 5 + + +async def test_error_response_raises(server): + async with WebSocketClient(server.url, reconnect=False) as c: + with pytest.raises(ResolveSpecError, match="boom") as ei: + await c.read("fail") + assert ei.value.code == "bad" + + +async def test_request_timeout(server): + server.respond = False + async with WebSocketClient(server.url, reconnect=False, request_timeout=0.1) as c: + with pytest.raises(ResolveSpecError, match="timeout"): + await c.read("users") + assert not c._pending + + +async def test_not_connected_raises(): + c = WebSocketClient("ws://127.0.0.1:1") + with pytest.raises(ResolveSpecError, match="not connected"): + await c.read("users") + + +async def test_subscribe_notify_unsubscribe(server): + got = asyncio.Queue() + async with WebSocketClient(server.url, reconnect=False) as c: + sid = await c.subscribe("users", got.put, schema="public", filters=[{"column": "a", "operator": "eq", "value": 1}]) + assert sid == "sub-1" + assert [s.id for s in c.get_subscriptions()] == ["sub-1"] + assert server.received[0]["operation"] == "subscribe" + assert server.received[0]["options"] == {"filters": [{"column": "a", "operator": "eq", "value": 1}]} + for ws in server.conns: + await ws.send(json.dumps({"type": "notification", "operation": "create", "subscription_id": "sub-1", + "entity": "users", "data": {"id": 9}, "timestamp": "t"})) + n = await asyncio.wait_for(got.get(), 2) + assert n["data"] == {"id": 9} + await c.unsubscribe("sub-1") + assert c.get_subscriptions() == [] + assert server.received[-1] == {**server.received[-1], "operation": "unsubscribe", "subscription_id": "sub-1"} + + +async def test_events_and_heartbeat(server): + events = [] + c = WebSocketClient(server.url, reconnect=False, heartbeat_interval=0.05) + c.on("connect", lambda: events.append("connect")) + c.on("state_change", lambda s: events.append(s)) + c.on("message", lambda m: events.append(("msg", m["type"]))) + await c.connect() + await asyncio.sleep(0.2) + await c.close() + assert events[:3] == ["connecting", "connected", "connect"] + assert ("msg", "pong") in events + assert events[-1] == "disconnected" + assert any(m["type"] == "ping" for m in server.received) + with pytest.raises(ValueError): + c.on("bogus", lambda: None) + + +async def test_reconnect_after_server_drop(server): + states = [] + c = WebSocketClient(server.url, reconnect=True, reconnect_interval=0.05) + c.on("state_change", states.append) + await c.connect() + for ws in list(server.conns): + await ws.close() + for _ in range(100): + if states.count("connected") >= 2: + break + await asyncio.sleep(0.05) + assert "reconnecting" in states + assert c.is_connected() + assert (await c.read("users"))["echo"] == "read" + await c.close() + + +async def test_pending_requests_fail_on_disconnect(server): + server.respond = False + c = WebSocketClient(server.url, reconnect=False) + await c.connect() + task = asyncio.create_task(c.read("users")) + await asyncio.sleep(0.05) + for ws in list(server.conns): + await ws.close() + with pytest.raises(ResolveSpecError, match="disconnected"): + await asyncio.wait_for(task, 2) + await c.close() diff --git a/resolvespec-python/todo.md b/resolvespec-python/todo.md index f6f8d3e..e44b1de 100644 --- a/resolvespec-python/todo.md +++ b/resolvespec-python/todo.md @@ -4,34 +4,34 @@ ### 1. ResolveSpec Client API -- [ ] Core API implementation (read, create, update, delete, get_metadata) -- [ ] Unit tests for API functions +- [x] Core API implementation (read, create, update, delete, get_metadata) +- [x] Unit tests for API functions - [ ] Integration tests with server -- [ ] Error handling and edge cases +- [x] Error handling and edge cases ### 2. HeaderSpec Client API -- [ ] Client API implementation -- [ ] Unit tests +- [x] Client API implementation +- [x] Unit tests - [ ] Integration tests with server ### 3. FunctionSpec Client API -- [ ] Client API implementation -- [ ] Unit tests +- [x] Client API implementation +- [x] Unit tests - [ ] Integration tests with server ### 4. WebSocketSpec Client API -- [ ] WebSocketClient class implementation (read, create, update, delete, meta, subscribe, unsubscribe) -- [ ] Unit tests for WebSocketClient -- [ ] Connection handling tests -- [ ] Subscription tests +- [x] WebSocketClient class implementation (read, create, update, delete, meta, subscribe, unsubscribe) +- [x] Unit tests for WebSocketClient +- [x] Connection handling tests +- [x] Subscription tests - [ ] Integration tests with server ### 5. Testing Infrastructure -- [ ] Set up test framework (pytest) +- [x] Set up test framework (pytest) - [ ] Configure test coverage reporting (pytest-cov) - [ ] Add test utilities and fixtures - [ ] Create test documentation @@ -43,8 +43,8 @@ - [ ] Usage examples for each client API - [ ] Installation guide - [ ] Contributing guidelines -- [ ] README with quick start +- [x] README (cheatsheet) --- -**Last Updated:** 2026-02-07 \ No newline at end of file +**Last Updated:** 2026-09-30 \ No newline at end of file From f6a9daa89e9c5771789daf478e4e676dd49abda9 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:18:24 +0200 Subject: [PATCH 03/23] docs(audit): add plan for single transaction per request --- audit/single_tran.md | 99 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 99 insertions(+) create mode 100644 audit/single_tran.md diff --git a/audit/single_tran.md b/audit/single_tran.md new file mode 100644 index 0000000..0463c33 --- /dev/null +++ b/audit/single_tran.md @@ -0,0 +1,99 @@ +# Single transaction per request — plan + +## Goal +- Every DB statement and every hook that touches the DB in one request runs on **one transaction / one connection**. +- Hooks never receive the raw pool (`h.db`). +- Fixes: RLS GUCs (`set_config(..., true)`) lost on reads/creates/deletes; extra pool connections; select-then-write races. + +## Why +- `set_config(..., true)` is transaction-local. A hook on the pool, or a query on another pool connection, never sees it → RLS returns 0 rows / 42501. +- 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. + +## Current state (verified by reading code; not yet by `dbtrace`) +| 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 | +| resolvespec | tx | tx | tx; re-fetch on pool (`:1297, 1449, 1602`) | **single: hook + select + delete on pool (`:1654, 1794, 1806`)**; batch: one `BeforeDelete` before tx, none per item | +| websocketspec | **pool** (`:563-672`) | **pool** (`:708`) | **pool** (`:744`) | **pool** (`:757`) | +| mqttspec | **pool** (`:674-789`) | **pool** (`:838`) | **pool** (`:875`) | **pool** (`:889`) | +| resolvemcp | **pool** (`:253`) | single: **pool** (`:445`); batch: tx | tx | tx | +| funcspec | tx; `BeforeResponse` post-commit on pool (`:337, 640`) | — | — | — | + +- Correction to earlier note: "no transactions" in mqttspec/websocketspec/resolvemcp-read is a gap for this problem, not a non-issue. +- `BeforeHandle` runs before any tx by design (auth + model checks, `PreloadSecurityRules`). Keep it DB-free except security preload (own connection, cached). + +## In-tx hook coverage today (verified) +- Already in tx with `Tx: tx`: `BeforeRead`, `BeforeCreate`, `BeforeUpdate`, `BeforeScan` (read/create/update, both specs); restheadspec batch delete `BeforeDelete` + `AfterDelete`. +- **Not in tx:** single `BeforeDelete`/`AfterDelete` (both specs), resolvespec batch delete (no per-item hook), all `After*` post-commit, all websocketspec/mqttspec hooks, resolvemcp read/single create. +- Gap beyond coverage: no single guaranteed "tx opened" point. User/RLS stamping would have to be repeated in each `Before*` hook and is missed by any path without one (e.g. resolvespec batch delete). `OnTxBegin` closes this: fires once per tx, first, for read/insert/update/delete and for the second short tx. + +## Scope +- `OnTxBegin` + `runInTx` apply to **all six**: resolvespec, restheadspec, websocketspec, mqttspec, resolvemcp, funcspec. +- Each spec has its own `HookType` (resolvespec, restheadspec, websocketspec, resolvemcp, funcspec); mqttspec aliases websocketspec, so it inherits the constant but needs its own handler wiring. +- Same semantics everywhere: fires once per tx, first, for read/insert/update/delete and the second short tx; failure aborts + rolls back, nothing leaked. +- Each spec's `security_hooks.go` registers the user/RLS stamping on `OnTxBegin`. +- Shared helper preferred over six copies: one small function in `pkg/common` (begin tx, set `Tx`, call spec-supplied begin callback), each spec passes its own hook executor. + +## funcspec (different) +- Custom SQL handlers (`SqlQuery`, `SqlQueryList`), no CRUD, no model registry; one tx per request already (`:195`, `:561`). +- Already in tx: `BeforeQuery`/`BeforeQueryList`, `BeforeSQLExec`, `AfterSQLExec`, `AfterQuery`/`AfterQueryList`, plus `BeforeOp`. +- `BeforeOp` = generic pre-hook via `ExecuteBeforeOp`, fires before every `Before*` in the tx; but it fires **twice** per tx (query hook + `BeforeSQLExec`), so it is not a once-per-tx point. +- Only gap: `BeforeResponse` runs post-commit with `Tx = h.db` (`:337`, `:640`). +- Applies from this plan: once-per-tx `OnTxBegin` (stamping user/RLS before any SQL, incl. hook-mutated SQL), `BeforeResponse` on a tx, fail-closed abort. +- Does not apply: delete/insert/update phases, second re-fetch tx (no re-fetch; SQL is user-defined), `BeforeHandle` preload. +- Decided: add a real `OnTxBegin`. `BeforeOp` unchanged. + +## Common interface (`pkg/common`, new `txhook.go`) +- Precedent: `security.SecurityContext` + per-spec `newSecurityContext(hookCtx)` adapter. Same pattern here. +- `common.TxHookName` = `"on_tx_begin"`: one shared string; each spec declares `OnTxBegin HookType = common.TxHookName` (HookTypes are per-spec types, so the constant value is shared, not the type). +- `common.TxContext` interface, implemented by each spec's `HookContext` via a small adapter: `GetContext()`, `GetTx()`, `SetTx(common.Database)`, plus `Abort` accessors for the abort path. +- `common.RunRequestTx(ctx, db, tc TxContext, onBegin func() error, body func(tx common.Database) error) error`: `RunInTransaction` -> `tc.SetTx(tx)` -> `onBegin()` (spec passes `registry.Execute(OnTxBegin, hookCtx)`) -> `body(tx)`. `onBegin` error or abort = return error = rollback. +- Shared stamping: one function in `pkg/security` taking `SecurityContext` + `common.Database` (sets tx-local user/RLS); each spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. No per-spec copies of the logic. +- Each spec keeps its own registry/`HookContext`; only the tx lifecycle and stamping are shared. +- Out of scope: unifying the six `HookContext` / registry types. + +## Design +1. **`OnTxBegin` hook** (new `HookType`, all specs). Runs first inside every tx the handler opens; gets `tx` in `hookCtx.Tx`. RLS stamping is registered once there; reads owner/tenant from request context. +2. **`runInTx` helper per handler**: wraps `RunInTransaction`, sets `hookCtx.Tx = tx`, fires `OnTxBegin`, runs the body. All handler paths use it; no path passes `h.db` to a hook. +3. **Post-commit work** (`After*`, update re-fetch, `BeforeResponse`): run in a second short `runInTx` (so `OnTxBegin` re-applies). Not inside the main tx. +4. **Delete**: hook → select → delete in one tx; 404 on no row; cache invalidation after commit. +5. **Backward compat**: hooks keep the same names/order; only `hookCtx.Tx` changes from pool to tx. `OnTxBegin` is additive. + +## Decisions (settled) +- Insert/update: re-fetch + `AfterCreate`/`AfterUpdate`-style post-commit work run in a **second short tx** (must see trigger changes). Only insert/update; read and delete have no second tx. +- 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. + +## Open +- Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking. + +## Phases +| # | Change | Files | Notes | +|---|---|---|---| +| 0 | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above | +| 1 | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | +| 2 | `OnTxBegin` hook type + `runInTx` helper | `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec first | +| 3 | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | +| 4 | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | +| 5 | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | +| 6 | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | +| 7 | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | + +## Tests +- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. +- Missing: any test asserting hook `Tx` is a tx, or counting connections per op. +- 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). +- Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB). + +## Risks +- Long tx if a hook does slow work inside it → hold connection longer; keep hooks fast. +- Pool of 1: nothing inside a tx may take a second pool connection (auth/security loads are outside; keep it so). +- Behavior change: After hooks no longer get the pool handle; hooks that relied on an independent connection break. +- websocket/mqtt long-lived connections: tx must be per message, never per connection. + +## Done when +- `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. +- No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers. From b2b815552fad13ed42efcd8bf7f691a74dcf8acb Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:31:43 +0200 Subject: [PATCH 04/23] refactor: move JS and Python clients under clients/ --- .gitignore | 2 +- README.md | 2 +- .../resolvespec-js}/.changeset/README.md | 0 .../resolvespec-js}/.changeset/config.json | 0 .../resolvespec-js}/CHANGELOG.md | 0 {resolvespec-js => clients/resolvespec-js}/PLAN.md | 0 {resolvespec-js => clients/resolvespec-js}/README.md | 0 .../resolvespec-js}/dist/index.cjs | 0 .../resolvespec-js}/dist/index.d.ts | 0 .../resolvespec-js}/dist/index.js | 0 .../resolvespec-js}/package.json | 0 .../resolvespec-js}/pnpm-lock.yaml | 0 .../resolvespec-js}/pnpm-workspace.yaml | 0 .../resolvespec-js}/src/__tests__/common.test.ts | 0 .../src/__tests__/custom-headers.test.ts | 0 .../resolvespec-js}/src/__tests__/headerspec.test.ts | 0 .../resolvespec-js}/src/__tests__/resolvespec.test.ts | 0 .../src/__tests__/websocketspec.test.ts | 0 .../resolvespec-js}/src/common/http.ts | 0 .../resolvespec-js}/src/common/index.ts | 0 .../resolvespec-js}/src/common/types.ts | 0 .../resolvespec-js}/src/headerspec/client.ts | 0 .../resolvespec-js}/src/headerspec/index.ts | 0 .../resolvespec-js}/src/index.ts | 0 .../resolvespec-js}/src/resolvespec/client.ts | 0 .../resolvespec-js}/src/resolvespec/index.ts | 0 .../resolvespec-js}/src/websocketspec/client.ts | 0 .../resolvespec-js}/src/websocketspec/index.ts | 0 .../resolvespec-js}/src/websocketspec/types.ts | 0 .../resolvespec-js}/tsconfig.json | 0 .../resolvespec-js}/vite.config.ts | 0 .../resolvespec-python}/.gitignore | 0 .../resolvespec-python}/README.md | 0 .../resolvespec-python}/pyproject.toml | 0 .../resolvespec-python}/src/resolvespec/__init__.py | 0 .../resolvespec-python}/src/resolvespec/funcspec.py | 0 .../resolvespec-python}/src/resolvespec/headerspec.py | 0 .../resolvespec-python}/src/resolvespec/http.py | 0 .../resolvespec-python}/src/resolvespec/resolvespec.py | 0 .../resolvespec-python}/src/resolvespec/types.py | 0 .../resolvespec-python}/src/resolvespec/websocket.py | 0 .../resolvespec-python}/tests/test_funcspec.py | 0 .../resolvespec-python}/tests/test_headerspec.py | 0 .../resolvespec-python}/tests/test_resolvespec.py | 0 .../resolvespec-python}/tests/test_websocket.py | 0 .../resolvespec-python}/todo.md | 0 todo.md | 10 +++++----- 47 files changed, 7 insertions(+), 7 deletions(-) rename {resolvespec-js => clients/resolvespec-js}/.changeset/README.md (100%) rename {resolvespec-js => clients/resolvespec-js}/.changeset/config.json (100%) rename {resolvespec-js => clients/resolvespec-js}/CHANGELOG.md (100%) rename {resolvespec-js => clients/resolvespec-js}/PLAN.md (100%) rename {resolvespec-js => clients/resolvespec-js}/README.md (100%) rename {resolvespec-js => clients/resolvespec-js}/dist/index.cjs (100%) rename {resolvespec-js => clients/resolvespec-js}/dist/index.d.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/dist/index.js (100%) rename {resolvespec-js => clients/resolvespec-js}/package.json (100%) rename {resolvespec-js => clients/resolvespec-js}/pnpm-lock.yaml (100%) rename {resolvespec-js => clients/resolvespec-js}/pnpm-workspace.yaml (100%) rename {resolvespec-js => clients/resolvespec-js}/src/__tests__/common.test.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/__tests__/custom-headers.test.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/__tests__/headerspec.test.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/__tests__/resolvespec.test.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/__tests__/websocketspec.test.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/common/http.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/common/index.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/common/types.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/headerspec/client.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/headerspec/index.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/index.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/resolvespec/client.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/resolvespec/index.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/websocketspec/client.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/websocketspec/index.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/src/websocketspec/types.ts (100%) rename {resolvespec-js => clients/resolvespec-js}/tsconfig.json (100%) rename {resolvespec-js => clients/resolvespec-js}/vite.config.ts (100%) rename {resolvespec-python => clients/resolvespec-python}/.gitignore (100%) rename {resolvespec-python => clients/resolvespec-python}/README.md (100%) rename {resolvespec-python => clients/resolvespec-python}/pyproject.toml (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/__init__.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/funcspec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/headerspec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/http.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/resolvespec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/types.py (100%) rename {resolvespec-python => clients/resolvespec-python}/src/resolvespec/websocket.py (100%) rename {resolvespec-python => clients/resolvespec-python}/tests/test_funcspec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/tests/test_headerspec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/tests/test_resolvespec.py (100%) rename {resolvespec-python => clients/resolvespec-python}/tests/test_websocket.py (100%) rename {resolvespec-python => clients/resolvespec-python}/todo.md (100%) diff --git a/.gitignore b/.gitignore index c1a1ec1..49f296e 100644 --- a/.gitignore +++ b/.gitignore @@ -28,5 +28,5 @@ test.db /testserver tests/data/ node_modules/ -resolvespec-js/dist/ +clients/resolvespec-js/dist/ .codex diff --git a/README.md b/README.md index 24dd9b5..4256c31 100644 --- a/README.md +++ b/README.md @@ -491,7 +491,7 @@ TypeScript/JavaScript client library supporting all three REST and WebSocket pro - Header-based REST client (`HeaderSpecClient`) - WebSocket client (`WebSocketClient`) with CRUD, subscriptions, heartbeat, reconnect -For complete documentation, see [resolvespec-js/README.md](resolvespec-js/README.md). +For complete documentation, see [clients/resolvespec-js/README.md](clients/resolvespec-js/README.md). ### Real-Time Communication diff --git a/resolvespec-js/.changeset/README.md b/clients/resolvespec-js/.changeset/README.md similarity index 100% rename from resolvespec-js/.changeset/README.md rename to clients/resolvespec-js/.changeset/README.md diff --git a/resolvespec-js/.changeset/config.json b/clients/resolvespec-js/.changeset/config.json similarity index 100% rename from resolvespec-js/.changeset/config.json rename to clients/resolvespec-js/.changeset/config.json diff --git a/resolvespec-js/CHANGELOG.md b/clients/resolvespec-js/CHANGELOG.md similarity index 100% rename from resolvespec-js/CHANGELOG.md rename to clients/resolvespec-js/CHANGELOG.md diff --git a/resolvespec-js/PLAN.md b/clients/resolvespec-js/PLAN.md similarity index 100% rename from resolvespec-js/PLAN.md rename to clients/resolvespec-js/PLAN.md diff --git a/resolvespec-js/README.md b/clients/resolvespec-js/README.md similarity index 100% rename from resolvespec-js/README.md rename to clients/resolvespec-js/README.md diff --git a/resolvespec-js/dist/index.cjs b/clients/resolvespec-js/dist/index.cjs similarity index 100% rename from resolvespec-js/dist/index.cjs rename to clients/resolvespec-js/dist/index.cjs diff --git a/resolvespec-js/dist/index.d.ts b/clients/resolvespec-js/dist/index.d.ts similarity index 100% rename from resolvespec-js/dist/index.d.ts rename to clients/resolvespec-js/dist/index.d.ts diff --git a/resolvespec-js/dist/index.js b/clients/resolvespec-js/dist/index.js similarity index 100% rename from resolvespec-js/dist/index.js rename to clients/resolvespec-js/dist/index.js diff --git a/resolvespec-js/package.json b/clients/resolvespec-js/package.json similarity index 100% rename from resolvespec-js/package.json rename to clients/resolvespec-js/package.json diff --git a/resolvespec-js/pnpm-lock.yaml b/clients/resolvespec-js/pnpm-lock.yaml similarity index 100% rename from resolvespec-js/pnpm-lock.yaml rename to clients/resolvespec-js/pnpm-lock.yaml diff --git a/resolvespec-js/pnpm-workspace.yaml b/clients/resolvespec-js/pnpm-workspace.yaml similarity index 100% rename from resolvespec-js/pnpm-workspace.yaml rename to clients/resolvespec-js/pnpm-workspace.yaml diff --git a/resolvespec-js/src/__tests__/common.test.ts b/clients/resolvespec-js/src/__tests__/common.test.ts similarity index 100% rename from resolvespec-js/src/__tests__/common.test.ts rename to clients/resolvespec-js/src/__tests__/common.test.ts diff --git a/resolvespec-js/src/__tests__/custom-headers.test.ts b/clients/resolvespec-js/src/__tests__/custom-headers.test.ts similarity index 100% rename from resolvespec-js/src/__tests__/custom-headers.test.ts rename to clients/resolvespec-js/src/__tests__/custom-headers.test.ts diff --git a/resolvespec-js/src/__tests__/headerspec.test.ts b/clients/resolvespec-js/src/__tests__/headerspec.test.ts similarity index 100% rename from resolvespec-js/src/__tests__/headerspec.test.ts rename to clients/resolvespec-js/src/__tests__/headerspec.test.ts diff --git a/resolvespec-js/src/__tests__/resolvespec.test.ts b/clients/resolvespec-js/src/__tests__/resolvespec.test.ts similarity index 100% rename from resolvespec-js/src/__tests__/resolvespec.test.ts rename to clients/resolvespec-js/src/__tests__/resolvespec.test.ts diff --git a/resolvespec-js/src/__tests__/websocketspec.test.ts b/clients/resolvespec-js/src/__tests__/websocketspec.test.ts similarity index 100% rename from resolvespec-js/src/__tests__/websocketspec.test.ts rename to clients/resolvespec-js/src/__tests__/websocketspec.test.ts diff --git a/resolvespec-js/src/common/http.ts b/clients/resolvespec-js/src/common/http.ts similarity index 100% rename from resolvespec-js/src/common/http.ts rename to clients/resolvespec-js/src/common/http.ts diff --git a/resolvespec-js/src/common/index.ts b/clients/resolvespec-js/src/common/index.ts similarity index 100% rename from resolvespec-js/src/common/index.ts rename to clients/resolvespec-js/src/common/index.ts diff --git a/resolvespec-js/src/common/types.ts b/clients/resolvespec-js/src/common/types.ts similarity index 100% rename from resolvespec-js/src/common/types.ts rename to clients/resolvespec-js/src/common/types.ts diff --git a/resolvespec-js/src/headerspec/client.ts b/clients/resolvespec-js/src/headerspec/client.ts similarity index 100% rename from resolvespec-js/src/headerspec/client.ts rename to clients/resolvespec-js/src/headerspec/client.ts diff --git a/resolvespec-js/src/headerspec/index.ts b/clients/resolvespec-js/src/headerspec/index.ts similarity index 100% rename from resolvespec-js/src/headerspec/index.ts rename to clients/resolvespec-js/src/headerspec/index.ts diff --git a/resolvespec-js/src/index.ts b/clients/resolvespec-js/src/index.ts similarity index 100% rename from resolvespec-js/src/index.ts rename to clients/resolvespec-js/src/index.ts diff --git a/resolvespec-js/src/resolvespec/client.ts b/clients/resolvespec-js/src/resolvespec/client.ts similarity index 100% rename from resolvespec-js/src/resolvespec/client.ts rename to clients/resolvespec-js/src/resolvespec/client.ts diff --git a/resolvespec-js/src/resolvespec/index.ts b/clients/resolvespec-js/src/resolvespec/index.ts similarity index 100% rename from resolvespec-js/src/resolvespec/index.ts rename to clients/resolvespec-js/src/resolvespec/index.ts diff --git a/resolvespec-js/src/websocketspec/client.ts b/clients/resolvespec-js/src/websocketspec/client.ts similarity index 100% rename from resolvespec-js/src/websocketspec/client.ts rename to clients/resolvespec-js/src/websocketspec/client.ts diff --git a/resolvespec-js/src/websocketspec/index.ts b/clients/resolvespec-js/src/websocketspec/index.ts similarity index 100% rename from resolvespec-js/src/websocketspec/index.ts rename to clients/resolvespec-js/src/websocketspec/index.ts diff --git a/resolvespec-js/src/websocketspec/types.ts b/clients/resolvespec-js/src/websocketspec/types.ts similarity index 100% rename from resolvespec-js/src/websocketspec/types.ts rename to clients/resolvespec-js/src/websocketspec/types.ts diff --git a/resolvespec-js/tsconfig.json b/clients/resolvespec-js/tsconfig.json similarity index 100% rename from resolvespec-js/tsconfig.json rename to clients/resolvespec-js/tsconfig.json diff --git a/resolvespec-js/vite.config.ts b/clients/resolvespec-js/vite.config.ts similarity index 100% rename from resolvespec-js/vite.config.ts rename to clients/resolvespec-js/vite.config.ts diff --git a/resolvespec-python/.gitignore b/clients/resolvespec-python/.gitignore similarity index 100% rename from resolvespec-python/.gitignore rename to clients/resolvespec-python/.gitignore diff --git a/resolvespec-python/README.md b/clients/resolvespec-python/README.md similarity index 100% rename from resolvespec-python/README.md rename to clients/resolvespec-python/README.md diff --git a/resolvespec-python/pyproject.toml b/clients/resolvespec-python/pyproject.toml similarity index 100% rename from resolvespec-python/pyproject.toml rename to clients/resolvespec-python/pyproject.toml diff --git a/resolvespec-python/src/resolvespec/__init__.py b/clients/resolvespec-python/src/resolvespec/__init__.py similarity index 100% rename from resolvespec-python/src/resolvespec/__init__.py rename to clients/resolvespec-python/src/resolvespec/__init__.py diff --git a/resolvespec-python/src/resolvespec/funcspec.py b/clients/resolvespec-python/src/resolvespec/funcspec.py similarity index 100% rename from resolvespec-python/src/resolvespec/funcspec.py rename to clients/resolvespec-python/src/resolvespec/funcspec.py diff --git a/resolvespec-python/src/resolvespec/headerspec.py b/clients/resolvespec-python/src/resolvespec/headerspec.py similarity index 100% rename from resolvespec-python/src/resolvespec/headerspec.py rename to clients/resolvespec-python/src/resolvespec/headerspec.py diff --git a/resolvespec-python/src/resolvespec/http.py b/clients/resolvespec-python/src/resolvespec/http.py similarity index 100% rename from resolvespec-python/src/resolvespec/http.py rename to clients/resolvespec-python/src/resolvespec/http.py diff --git a/resolvespec-python/src/resolvespec/resolvespec.py b/clients/resolvespec-python/src/resolvespec/resolvespec.py similarity index 100% rename from resolvespec-python/src/resolvespec/resolvespec.py rename to clients/resolvespec-python/src/resolvespec/resolvespec.py diff --git a/resolvespec-python/src/resolvespec/types.py b/clients/resolvespec-python/src/resolvespec/types.py similarity index 100% rename from resolvespec-python/src/resolvespec/types.py rename to clients/resolvespec-python/src/resolvespec/types.py diff --git a/resolvespec-python/src/resolvespec/websocket.py b/clients/resolvespec-python/src/resolvespec/websocket.py similarity index 100% rename from resolvespec-python/src/resolvespec/websocket.py rename to clients/resolvespec-python/src/resolvespec/websocket.py diff --git a/resolvespec-python/tests/test_funcspec.py b/clients/resolvespec-python/tests/test_funcspec.py similarity index 100% rename from resolvespec-python/tests/test_funcspec.py rename to clients/resolvespec-python/tests/test_funcspec.py diff --git a/resolvespec-python/tests/test_headerspec.py b/clients/resolvespec-python/tests/test_headerspec.py similarity index 100% rename from resolvespec-python/tests/test_headerspec.py rename to clients/resolvespec-python/tests/test_headerspec.py diff --git a/resolvespec-python/tests/test_resolvespec.py b/clients/resolvespec-python/tests/test_resolvespec.py similarity index 100% rename from resolvespec-python/tests/test_resolvespec.py rename to clients/resolvespec-python/tests/test_resolvespec.py diff --git a/resolvespec-python/tests/test_websocket.py b/clients/resolvespec-python/tests/test_websocket.py similarity index 100% rename from resolvespec-python/tests/test_websocket.py rename to clients/resolvespec-python/tests/test_websocket.py diff --git a/resolvespec-python/todo.md b/clients/resolvespec-python/todo.md similarity index 100% rename from resolvespec-python/todo.md rename to clients/resolvespec-python/todo.md diff --git a/todo.md b/todo.md index d83fb3c..a17f3d2 100644 --- a/todo.md +++ b/todo.md @@ -23,23 +23,23 @@ This document tracks incomplete features and improvements for the ResolveSpec pr ### ResolveSpec JS Client Implementation & Testing -1. **ResolveSpec Client API (resolvespec-js)** +1. **ResolveSpec Client API (clients/resolvespec-js)** - [x] Core API implementation (read, create, update, delete, getMetadata) - [ ] Unit tests for API functions - [ ] Integration tests with server - [ ] Error handling and edge cases -2. **HeaderSpec Client API (resolvespec-js)** +2. **HeaderSpec Client API (clients/resolvespec-js)** - [ ] Client API implementation - [ ] Unit tests - [ ] Integration tests with server -3. **FunctionSpec Client API (resolvespec-js)** +3. **FunctionSpec Client API (clients/resolvespec-js)** - [ ] Client API implementation - [ ] Unit tests - [ ] Integration tests with server -4. **WebSocketSpec Client API (resolvespec-js)** +4. **WebSocketSpec Client API (clients/resolvespec-js)** - [x] WebSocketClient class implementation (read, create, update, delete, meta, subscribe, unsubscribe) - [ ] Unit tests for WebSocketClient - [ ] Connection handling tests @@ -54,7 +54,7 @@ This document tracks incomplete features and improvements for the ResolveSpec pr ### ResolveSpec Python Client Implementation & Testing -See [`resolvespec-python/todo.md`](./resolvespec-python/todo.md) for detailed Python client implementation tasks. +See [`clients/resolvespec-python/todo.md`](./clients/resolvespec-python/todo.md) for detailed Python client implementation tasks. ### Core Functionality From cd96404cdda22fe1aef4bcdc7ca159c6c7e96385 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:33:25 +0200 Subject: [PATCH 05/23] fix(delete): run delete hooks and queries in one transaction - resolvespec/restheadspec: single and batch delete use one transaction - add sqlmock tests for delete transaction behaviour - testmodels: serial integer ids; update tests accordingly - add compose testserver, smoke script, podman-first Makefile targets --- Makefile | 22 ++- docker-compose.yml | 17 +++ docker/Dockerfile.testserver | 13 ++ docker/testserver.config.yaml | 95 +++++++++++++ pkg/resolvespec/delete_tx_test.go | 168 ++++++++++++++++++++++ pkg/resolvespec/handler.go | 216 ++++++++++++++--------------- pkg/restheadspec/delete_tx_test.go | 171 +++++++++++++++++++++++ pkg/restheadspec/handler.go | 82 +++++++---- pkg/testmodels/business.go | 28 ++-- scripts/testserver-smoke.sh | 35 +++++ tests/crud_test.go | 96 +++++++------ tests/integration_test.go | 71 +++++++--- 12 files changed, 788 insertions(+), 226 deletions(-) create mode 100644 docker/Dockerfile.testserver create mode 100644 docker/testserver.config.yaml create mode 100644 pkg/resolvespec/delete_tx_test.go create mode 100644 pkg/restheadspec/delete_tx_test.go create mode 100755 scripts/testserver-smoke.sh diff --git a/Makefile b/Makefile index 311b5c4..9ba8eb5 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,7 @@ -.PHONY: test test-unit test-race test-integration docker-up docker-down clean +# Container compose command: podman if installed, else docker +COMPOSE ?= $(shell command -v podman >/dev/null 2>&1 && echo "podman compose" || echo "docker compose") + +.PHONY: testserver-up testserver-down testserver-smoke test test-unit test-race test-integration docker-up docker-down clean GOLANGCI_LINT := $(shell go env GOPATH)/bin/golangci-lint @@ -82,7 +85,7 @@ lintfix: ## Run linter # Start PostgreSQL for integration tests docker-up: @echo "Starting PostgreSQL container..." - @podman compose up -d postgres-test + @$(COMPOSE) up -d postgres-test @echo "Waiting for PostgreSQL to be ready..." @sleep 5 @echo "PostgreSQL is ready!" @@ -90,12 +93,23 @@ docker-up: # Stop PostgreSQL container docker-down: @echo "Stopping PostgreSQL container..." - @podman compose down + @$(COMPOSE) down + +# Test server + PostgreSQL in containers (dbtrace enabled) + +testserver-up: + @$(COMPOSE) up -d --build postgres-test testserver + +testserver-down: + @$(COMPOSE) down + +testserver-smoke: + @COMPOSE="$(COMPOSE)" scripts/testserver-smoke.sh # Clean up Docker volumes and test data clean: @echo "Cleaning up..." - @podman compose down -v + @$(COMPOSE) down -v @echo "Cleanup complete!" # Run integration tests with Docker (full workflow) diff --git a/docker-compose.yml b/docker-compose.yml index 47e7983..4f5c819 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -18,6 +18,23 @@ services: networks: - resolvespec-test + testserver: + build: + context: . + dockerfile: docker/Dockerfile.testserver + container_name: resolvespec-testserver + environment: + RESOLVESPEC_DB_TRACE_ENABLED: "true" + RESOLVESPEC_DB_TRACE_MIN_CALLS: "1" + RESOLVESPEC_DB_TRACE_POOL_LOG: "true" + ports: + - "8080:8080" + depends_on: + postgres-test: + condition: service_healthy + networks: + - resolvespec-test + volumes: postgres-test-data: driver: local diff --git a/docker/Dockerfile.testserver b/docker/Dockerfile.testserver new file mode 100644 index 0000000..f937a3b --- /dev/null +++ b/docker/Dockerfile.testserver @@ -0,0 +1,13 @@ +FROM golang:1.25-alpine AS build +WORKDIR /src +COPY go.mod go.sum ./ +RUN go mod download +COPY . . +RUN CGO_ENABLED=0 go build -o /out/testserver ./cmd/testserver + +FROM alpine:3.20 +RUN apk add --no-cache ca-certificates +COPY --from=build /out/testserver /usr/local/bin/testserver +COPY docker/testserver.config.yaml /etc/resolvespec/config.yaml +EXPOSE 8080 +ENTRYPOINT ["testserver"] diff --git a/docker/testserver.config.yaml b/docker/testserver.config.yaml new file mode 100644 index 0000000..5a8fa98 --- /dev/null +++ b/docker/testserver.config.yaml @@ -0,0 +1,95 @@ +# ResolveSpec Test Server Configuration (docker compose, PostgreSQL) +# This is a minimal configuration for the test server + +servers: + default_server: "main" + shutdown_timeout: 30s + drain_timeout: 25s + read_timeout: 10s + write_timeout: 10s + idle_timeout: 120s + instances: + main: + name: "main" + host: "0.0.0.0" + port: 8080 + description: "Main server instance" + gzip: true + tags: + env: "test" + +logger: + dev: true + path: "" + +cache: + provider: "memory" + +middleware: + rate_limit_rps: 100.0 + rate_limit_burst: 200 + max_request_size: 10485760 + +cors: + allowed_origins: + - "*" + allowed_methods: + - "GET" + - "POST" + - "PUT" + - "DELETE" + - "OPTIONS" + allowed_headers: + - "*" + max_age: 3600 + +tracing: + enabled: false + service_name: "resolvespec" + service_version: "1.0.0" + endpoint: "" + +error_tracking: + enabled: false + provider: "noop" + environment: "development" + sample_rate: 1.0 + traces_sample_rate: 0.1 + +event_broker: + enabled: false + provider: "memory" + mode: "sync" + worker_count: 1 + buffer_size: 100 + instance_id: "" + +dbmanager: + default_connection: "default" + max_open_conns: 25 + max_idle_conns: 5 + conn_max_lifetime: 30m + conn_max_idle_time: 5m + retry_attempts: 3 + retry_delay: 1s + health_check_interval: 30s + enable_auto_reconnect: true + connections: + # "default" overrides the built-in default connection (all connections are connected at start) + default: + name: "default" + type: "postgres" + host: "postgres-test" + port: 5432 + user: "postgres" + password: "postgres" + database: "postgres" + sslmode: "disable" + application_name: "resolvespec-testserver" + default_orm: "gorm" + enable_logging: true + enable_metrics: false + connect_timeout: 10s + query_timeout: 30s + +paths: {} diff --git a/pkg/resolvespec/delete_tx_test.go b/pkg/resolvespec/delete_tx_test.go new file mode 100644 index 0000000..dd5a91d --- /dev/null +++ b/pkg/resolvespec/delete_tx_test.go @@ -0,0 +1,168 @@ +package resolvespec + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "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/modelregistry" +) + +type delItem struct { + ID int `json:"id" bun:"id,pk"` + Name string `json:"name" bun:"name"` +} + +func newDeleteHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, *sql.DB) { + t.Helper() + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + // One connection: any statement that bypasses the transaction while it is + // open cannot get a connection and fails on the request context timeout. + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + return NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry()), mock, db +} + +func runDelete(h *Handler, id string, data interface{}) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleDelete(ctx, w, id, data) + return rec +} + +// recordDeleteHook registers a BeforeDelete hook that captures hookCtx.Tx. +func recordDeleteHook(h *Handler, hookErr error) *[]common.Database { + var seen []common.Database + h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error { + seen = append(seen, ctx.Tx) + return hookErr + }) + return &seen +} + +func TestDeleteSingleUsesOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + seen := recordDeleteHook(h, nil) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := runDelete(h, "7", nil) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*seen) != 1 || (*seen)[0] == h.db { + t.Fatalf("BeforeDelete must run once on the transaction, got %v", *seen) + } +} + +func TestDeleteSingleNotFoundRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"})) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusNotFound { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteHookErrorRollsBackWithoutQueries(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + recordDeleteHook(h, errors.New("denied")) + + mock.ExpectBegin() + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusForbidden { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteSingleExecErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnError(errors.New("boom")) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteBatchUsesOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + seen := recordDeleteHook(h, nil) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := runDelete(h, "", []interface{}{"1", "2"}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*seen) != 1 || (*seen)[0] == h.db { + t.Fatalf("BeforeDelete must run once on the transaction, got %v", *seen) + } + var resp struct { + Data map[string]float64 `json:"data"` + } + if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil || resp.Data["deleted"] != 2 { + t.Fatalf("unexpected body %s (%v)", rec.Body, err) + } +} + +func TestDeleteBatchFailureRollsBackAll(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`DELETE FROM`).WillReturnError(errors.New("boom")) + mock.ExpectRollback() + + if rec := runDelete(h, "", []string{"1", "2"}); rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index f582a2f..cdbabb2 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -1210,7 +1210,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Now read the existing record from the database existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...) + selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...) // Apply conditions to select, based on the resolved target ID // (URL ID, request ID, or the "id" field embedded in the data payload). @@ -1375,7 +1375,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // First, read the existing record existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { continue // Skip if record not found @@ -1524,7 +1524,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // First, read the existing record existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { continue // Skip if record not found @@ -1640,7 +1640,6 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id logger.Info("Deleting records from %s.%s", schema, entity) - // Execute BeforeDelete hooks (covers model-rule checks before any deletion) hookCtx := &HookContext{ Context: ctx, Handler: h, @@ -1653,118 +1652,123 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id Writer: w, Tx: h.db, } + + // Hook, lookup and delete(s) share one transaction so transaction-local + // state set by hooks (e.g. RLS settings) applies to every statement. + var payload interface{} + var failure *deleteFailure + txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { + hookCtx.Tx = tx + payload, failure = h.executeDelete(ctx, tx, hookCtx, schema, tableName, model, id, data) + if failure != nil { + return failure + } + return nil + }) + if failure != nil { + h.sendError(w, failure.status, failure.code, failure.message, failure.err) + return + } + if txErr != nil { + logger.Error("Error in delete transaction: %v", txErr) + h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting records", txErr) + return + } + + // Invalidate cache for this table after commit + cacheTags := buildCacheTags(schema, tableName) + if err := invalidateCacheForTags(ctx, cacheTags); err != nil { + logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) + } + h.sendResponse(w, payload, nil) +} + +// deleteFailure describes an error response for a delete; returning it from the +// transaction closure rolls the transaction back. +type deleteFailure struct { + status int + code string + message string + err error +} + +func (f *deleteFailure) Error() string { return f.message } + +// executeDelete runs the BeforeDelete hook and the delete(s) on tx and returns +// the response payload. +func (h *Handler) executeDelete(ctx context.Context, tx common.Database, hookCtx *HookContext, schema, tableName string, model interface{}, id string, data interface{}) (interface{}, *deleteFailure) { + // Execute BeforeDelete hooks (covers model-rule checks before any deletion) if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { logger.Error("BeforeDelete hook failed: %v", err) - h.sendError(w, http.StatusForbidden, "delete_forbidden", "Delete operation not allowed", err) - return + return nil, &deleteFailure{http.StatusForbidden, "delete_forbidden", "Delete operation not allowed", err} + } + + pkName := reflection.GetPrimaryKeyName(model) + deleteByID := func(itemID interface{}) (int, error) { + result, err := tx.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID).Exec(ctx) + if err != nil { + return 0, fmt.Errorf("failed to delete record %v: %w", itemID, err) + } + return int(result.RowsAffected()), nil + } + batchFailure := func(err error) *deleteFailure { + logger.Error("Error in batch delete: %v", err) + return &deleteFailure{http.StatusInternalServerError, "delete_error", "Error deleting records", err} } // Handle batch delete from request data if data != nil { switch v := data.(type) { case []string: - // Array of IDs as strings logger.Info("Batch delete with %d IDs ([]string)", len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { - for _, itemID := range v { - - query := tx.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(reflection.GetPrimaryKeyName(model))), itemID) - if _, err := query.Exec(ctx); err != nil { - return fmt.Errorf("failed to delete record %s: %w", itemID, err) - } + for _, itemID := range v { + if _, err := deleteByID(itemID); err != nil { + return nil, batchFailure(err) } - return nil - }) - if err != nil { - logger.Error("Error in batch delete: %v", err) - h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting records", err) - return } logger.Info("Successfully deleted %d records", len(v)) - // Invalidate cache for this table - cacheTags := buildCacheTags(schema, tableName) - if err := invalidateCacheForTags(ctx, cacheTags); err != nil { - logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) - } - h.sendResponse(w, map[string]interface{}{"deleted": len(v)}, nil) - return + return map[string]interface{}{"deleted": len(v)}, nil case []interface{}: // Array of IDs or objects with ID field logger.Info("Batch delete with %d items ([]interface{})", len(v)) deletedCount := 0 - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { - for _, item := range v { - var itemID interface{} - - // Check if item is a string ID or object with id field - switch v := item.(type) { - case string: - itemID = v - case map[string]interface{}: - itemID = v["id"] - default: - // Try to use the item directly as ID - itemID = item - } - - if itemID == nil { - continue // Skip items without ID - } - - query := tx.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(reflection.GetPrimaryKeyName(model))), itemID) - result, err := query.Exec(ctx) - if err != nil { - return fmt.Errorf("failed to delete record %v: %w", itemID, err) - } - deletedCount += int(result.RowsAffected()) + for _, item := range v { + var itemID interface{} + switch iv := item.(type) { + case string: + itemID = iv + case map[string]interface{}: + itemID = iv["id"] + default: + itemID = item } - return nil - }) - if err != nil { - logger.Error("Error in batch delete: %v", err) - h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting records", err) - return + if itemID == nil { + continue // Skip items without ID + } + n, err := deleteByID(itemID) + if err != nil { + return nil, batchFailure(err) + } + deletedCount += n } logger.Info("Successfully deleted %d records", deletedCount) - // Invalidate cache for this table - cacheTags := buildCacheTags(schema, tableName) - if err := invalidateCacheForTags(ctx, cacheTags); err != nil { - logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) - } - h.sendResponse(w, map[string]interface{}{"deleted": deletedCount}, nil) - return + return map[string]interface{}{"deleted": deletedCount}, nil case []map[string]interface{}: - // Array of objects with id field logger.Info("Batch delete with %d items ([]map[string]interface{})", len(v)) deletedCount := 0 - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { - for _, item := range v { - if itemID, ok := item["id"]; ok && itemID != nil { - query := tx.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(reflection.GetPrimaryKeyName(model))), itemID) - result, err := query.Exec(ctx) - if err != nil { - return fmt.Errorf("failed to delete record %v: %w", itemID, err) - } - deletedCount += int(result.RowsAffected()) + for _, item := range v { + if itemID, ok := item["id"]; ok && itemID != nil { + n, err := deleteByID(itemID) + if err != nil { + return nil, batchFailure(err) } + deletedCount += n } - return nil - }) - if err != nil { - logger.Error("Error in batch delete: %v", err) - h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting records", err) - return } logger.Info("Successfully deleted %d records", deletedCount) - // Invalidate cache for this table - cacheTags := buildCacheTags(schema, tableName) - if err := invalidateCacheForTags(ctx, cacheTags); err != nil { - logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) - } - h.sendResponse(w, map[string]interface{}{"deleted": deletedCount}, nil) - return + return map[string]interface{}{"deleted": deletedCount}, nil case map[string]interface{}: // Single object with id field @@ -1777,13 +1781,9 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id // Single delete with URL ID if id == "" { logger.Error("Delete operation requires an ID") - h.sendError(w, http.StatusBadRequest, "missing_id", "Delete operation requires an ID", nil) - return + return nil, &deleteFailure{http.StatusBadRequest, "missing_id", "Delete operation requires an ID", nil} } - // Get primary key name - pkName := reflection.GetPrimaryKeyName(model) - // First, fetch the record that will be deleted modelType := reflect.TypeOf(model) if modelType.Kind() == reflect.Pointer { @@ -1791,42 +1791,28 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id } recordToDelete := reflect.New(modelType).Interface() - selectQuery := h.db.NewSelect().Model(recordToDelete).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) + selectQuery := tx.NewSelect().Model(recordToDelete).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { logger.Warn("Record not found for delete: %s = %s", pkName, id) - h.sendError(w, http.StatusNotFound, "not_found", "Record not found", err) - return + return nil, &deleteFailure{http.StatusNotFound, "not_found", "Record not found", err} } logger.Error("Error fetching record for delete: %v", err) - h.sendError(w, http.StatusInternalServerError, "fetch_error", "Error fetching record", err) - return + return nil, &deleteFailure{http.StatusInternalServerError, "fetch_error", "Error fetching record", err} } - query := h.db.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) - - result, err := query.Exec(ctx) + n, err := deleteByID(id) if err != nil { logger.Error("Error deleting record: %v", err) - h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting record", err) - return + return nil, &deleteFailure{http.StatusInternalServerError, "delete_error", "Error deleting record", err} } - - // Check if the record was actually deleted - if result.RowsAffected() == 0 { + if n == 0 { logger.Warn("No rows deleted for ID: %s", id) - h.sendError(w, http.StatusNotFound, "not_found", "Record not found or already deleted", nil) - return + return nil, &deleteFailure{http.StatusNotFound, "not_found", "Record not found or already deleted", nil} } logger.Info("Successfully deleted record with ID: %s", id) - // Return the deleted record data - // Invalidate cache for this table - cacheTags := buildCacheTags(schema, tableName) - if err := invalidateCacheForTags(ctx, cacheTags); err != nil { - logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) - } - h.sendResponse(w, recordToDelete, nil) + return recordToDelete, nil } // applyFilters applies all filters with proper grouping for OR logic diff --git a/pkg/restheadspec/delete_tx_test.go b/pkg/restheadspec/delete_tx_test.go new file mode 100644 index 0000000..1d8396d --- /dev/null +++ b/pkg/restheadspec/delete_tx_test.go @@ -0,0 +1,171 @@ +package restheadspec + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "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/modelregistry" +) + +type delItem struct { + ID int `json:"id" bun:"id,pk"` + Name string `json:"name" bun:"name"` +} + +func newDeleteHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, *sql.DB) { + t.Helper() + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + // One connection: any statement that bypasses the transaction while it is + // open cannot get a connection and fails on the request context timeout. + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + return NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry()), mock, db +} + +func runDelete(h *Handler, id string, data interface{}) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + ctx = WithModel(ctx, &delItem{}) + h.handleDelete(ctx, w, id, data) + return rec +} + +// recordDeleteHook registers a BeforeDelete hook that captures hookCtx.Tx. +func recordDeleteHook(h *Handler, hookErr error) *[]common.Database { + var seen []common.Database + h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error { + seen = append(seen, ctx.Tx) + return hookErr + }) + return &seen +} + +func TestDeleteSingleUsesOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + seen := recordDeleteHook(h, nil) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := runDelete(h, "7", nil) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*seen) != 1 || (*seen)[0] == h.db { + t.Fatalf("BeforeDelete must run once on the transaction, got %v", *seen) + } +} + +func TestDeleteSingleNotFoundRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"})) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusNotFound { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteHookErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + recordDeleteHook(h, errors.New("denied")) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusBadRequest { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteSingleExecErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnError(errors.New("boom")) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestDeleteBatchUsesOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + seen := recordDeleteHook(h, nil) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := runDelete(h, "", []interface{}{"1", "2"}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + // restheadspec fires the hook per item + if len(*seen) != 2 || (*seen)[0] == h.db || (*seen)[1] == h.db { + t.Fatalf("BeforeDelete must run per item on the transaction, got %v", *seen) + } + var resp map[string]float64 + if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil || resp["deleted"] != 2 { + t.Fatalf("unexpected body %s (%v)", rec.Body, err) + } +} + +func TestDeleteBatchFailureRollsBackAll(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`DELETE FROM`).WillReturnError(errors.New("boom")) + mock.ExpectRollback() + + if rec := runDelete(h, "", []string{"1", "2"}); rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index d532108..39a8349 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -1576,7 +1576,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // Now read the existing record from the database existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := tx.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) + selectQuery := h.db.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { return fmt.Errorf("record not found with ID: %v", targetID) @@ -1934,24 +1934,62 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id return } - // Get primary key name pkName := reflection.GetPrimaryKeyName(model) - // First, fetch the record that will be deleted modelType := reflect.TypeOf(model) modelType = reflection.GetPointerElement(modelType) recordToDelete := reflect.New(modelType).Interface() - selectQuery := h.db.NewSelect().Model(recordToDelete).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) + // Lookup, hooks and delete share one transaction so transaction-local + // state set by hooks (e.g. RLS settings) applies to every statement. + var failure *deleteFailure + txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { + failure = h.deleteSingleInTx(ctx, tx, w, schema, entity, tableName, model, pkName, id, recordToDelete) + if failure != nil { + return failure + } + return nil + }) + if failure != nil { + h.sendError(w, failure.status, failure.code, failure.message, failure.err) + return + } + if txErr != nil { + logger.Error("Error in delete transaction: %v", txErr) + h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting record", txErr) + return + } + + // Invalidate cache for this table after commit + cacheTags := buildCacheTags(schema, tableName) + if err := invalidateCacheForTags(ctx, cacheTags); err != nil { + logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) + } + h.sendResponse(w, recordToDelete, nil) +} + +// deleteFailure describes an error response for a delete; returning it from the +// transaction closure rolls the transaction back. +type deleteFailure struct { + status int + code string + message string + err error +} + +func (f *deleteFailure) Error() string { return f.message } + +// deleteSingleInTx fetches the record, runs the delete hooks and deletes it, all on tx. +func (h *Handler) deleteSingleInTx(ctx context.Context, tx common.Database, w common.ResponseWriter, schema, entity, tableName string, model interface{}, pkName, id string, recordToDelete interface{}) *deleteFailure { + // First, fetch the record that will be deleted + selectQuery := tx.NewSelect().Model(recordToDelete).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { logger.Warn("Record not found for delete: %s = %s", pkName, id) - h.sendError(w, http.StatusNotFound, "not_found", "Record not found", err) - return + return &deleteFailure{http.StatusNotFound, "not_found", "Record not found", err} } logger.Error("Error fetching record for delete: %v", err) - h.sendError(w, http.StatusInternalServerError, "fetch_error", "Error fetching record", err) - return + return &deleteFailure{http.StatusInternalServerError, "fetch_error", "Error fetching record", err} } // Execute BeforeDelete hooks with the record data @@ -1965,25 +2003,23 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id Operation: "delete", ID: id, Writer: w, - Tx: h.db, + Tx: tx, Data: recordToDelete, } if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { logger.Error("BeforeDelete hook failed: %v", err) - h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return + return &deleteFailure{http.StatusBadRequest, "hook_error", "Hook execution failed", err} } - query := h.db.NewDelete().Table(tableName) + query := tx.NewDelete().Table(tableName) query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) // Execute BeforeScan hooks - pass query chain so hooks can modify it hookCtx.Query = query if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil { logger.Error("BeforeScan hook failed: %v", err) - h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return + return &deleteFailure{http.StatusBadRequest, "hook_error", "Hook execution failed", err} } // Use potentially modified query from hook context @@ -1994,15 +2030,13 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id result, err := query.Exec(ctx) if err != nil { logger.Error("Error deleting record: %v", err) - h.sendError(w, http.StatusInternalServerError, "delete_error", "Error deleting record", err) - return + return &deleteFailure{http.StatusInternalServerError, "delete_error", "Error deleting record", err} } // Check if the record was actually deleted if result.RowsAffected() == 0 { logger.Warn("No rows deleted for ID: %s", id) - h.sendError(w, http.StatusNotFound, "not_found", "Record not found or already deleted", nil) - return + return &deleteFailure{http.StatusNotFound, "not_found", "Record not found or already deleted", nil} } // Execute AfterDelete hooks with the deleted record data @@ -2011,17 +2045,9 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { logger.Error("AfterDelete hook failed: %v", err) - h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) - return + return &deleteFailure{http.StatusInternalServerError, "hook_error", "Hook execution failed", err} } - - // Return the deleted record data - // Invalidate cache for this table - cacheTags := buildCacheTags(schema, tableName) - if err := invalidateCacheForTags(ctx, cacheTags); err != nil { - logger.Warn("Failed to invalidate cache for table %s: %v", tableName, err) - } - h.sendResponse(w, recordToDelete, nil) + return nil } // mergeRecordWithRequest merges a database record with the original request data diff --git a/pkg/testmodels/business.go b/pkg/testmodels/business.go index 9539626..aa711af 100644 --- a/pkg/testmodels/business.go +++ b/pkg/testmodels/business.go @@ -8,7 +8,7 @@ import ( // Department represents a company department type Department struct { - ID string `json:"id" gorm:"primaryKey;type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` Name string `json:"name"` Code string `json:"code" gorm:"uniqueIndex"` Description string `json:"description"` @@ -26,13 +26,13 @@ func (Department) TableName() string { // Employee represents a company employee type Employee struct { - ID string `json:"id" gorm:"primaryKey;type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` FirstName string `json:"first_name"` LastName string `json:"last_name"` Email string `json:"email" gorm:"uniqueIndex"` Title string `json:"title"` - DepartmentID string `json:"department_id" gorm:"type:string"` - ManagerID *string `json:"manager_id" gorm:"type:string"` + DepartmentID int32 `json:"department_id"` + ManagerID *int32 `json:"manager_id"` HireDate time.Time `json:"hire_date"` Status string `json:"status"` CreatedAt time.Time `json:"created_at"` @@ -52,7 +52,7 @@ func (Employee) TableName() string { // Project represents a company project type Project struct { - ID string `json:"id" gorm:"primaryKey;type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` Name string `json:"name"` Code string `json:"code" gorm:"uniqueIndex"` Description string `json:"description"` @@ -76,9 +76,9 @@ func (Project) TableName() string { // ProjectTask represents a task within a project type ProjectTask struct { - ID string `json:"id" gorm:"primaryKey;type:string"` - ProjectID string `json:"project_id" gorm:"type:string"` - AssigneeID string `json:"assignee_id" gorm:"type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` + ProjectID int32 `json:"project_id"` + AssigneeID int32 `json:"assignee_id"` Title string `json:"title"` Description string `json:"description"` Status string `json:"status"` @@ -99,14 +99,14 @@ func (ProjectTask) TableName() string { // Document represents any document in the system type Document struct { - ID string `json:"id" gorm:"primaryKey;type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` Name string `json:"name"` Type string `json:"type"` ContentType string `json:"content_type"` Size int64 `json:"size"` Path string `json:"path"` - OwnerID string `json:"owner_id" gorm:"type:string"` - ProjectID *string `json:"project_id" gorm:"type:string"` + OwnerID int32 `json:"owner_id"` + ProjectID *int32 `json:"project_id"` Status string `json:"status"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` @@ -122,9 +122,9 @@ func (Document) TableName() string { // Comment represents a comment on a task type Comment struct { - ID string `json:"id" gorm:"primaryKey;type:string"` - TaskID string `json:"task_id" gorm:"type:string"` - AuthorID string `json:"author_id" gorm:"type:string"` + ID int32 `json:"id" gorm:"primaryKey;autoIncrement"` + TaskID int32 `json:"task_id"` + AuthorID int32 `json:"author_id"` Content string `json:"content"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` diff --git a/scripts/testserver-smoke.sh b/scripts/testserver-smoke.sh new file mode 100755 index 0000000..de72346 --- /dev/null +++ b/scripts/testserver-smoke.sh @@ -0,0 +1,35 @@ +#!/usr/bin/env bash +# Delete-path smoke test against the compose test server. +# Usage: scripts/testserver-smoke.sh [base_url] (COMPOSE overrides the compose command) +set -euo pipefail + +BASE="${1:-http://localhost:8080}" +if [ -z "${COMPOSE:-}" ]; then + if command -v podman >/dev/null 2>&1; then COMPOSE="podman compose"; else COMPOSE="docker compose"; fi +fi +TS="$(date +%s)" +BODY="$(mktemp)" +trap 'rm -f "$BODY"' EXIT + +# call [path-suffix] -> prints HTTP status, body in $BODY +call() { + curl -s -o "$BODY" -w '%{http_code}' -X POST "$BASE/public/departments${2:-}" \ + -H 'Content-Type: application/json' -d "$1" +} +expect() { # name want got + if [ "$2" != "$3" ]; then echo "FAIL $1: want $2 got $3: $(cat "$BODY")"; exit 1; fi + echo "ok $1 ($3)" +} +ids() { grep -o '"id":[0-9]*' "$BODY" | cut -d: -f2; } + +expect create 200 "$(call "{\"operation\":\"create\",\"data\":{\"name\":\"Smoke\",\"code\":\"S$TS\"}}")" +ID="$(ids | head -1)" +expect delete 200 "$(call '{"operation":"delete"}' "/$ID")" +expect delete-again 404 "$(call '{"operation":"delete"}' "/$ID")" + +expect batch-create 200 "$(call "{\"operation\":\"create\",\"data\":[{\"name\":\"B\",\"code\":\"B1$TS\"},{\"name\":\"B\",\"code\":\"B2$TS\"}]}")" +B1="$(ids | sed -n 1p)"; B2="$(ids | sed -n 2p)" +expect batch-delete 200 "$(call "{\"operation\":\"delete\",\"data\":[\"$B1\",\"$B2\"]}")" + +echo "--- dbtrace" +$COMPOSE logs testserver 2>&1 | grep 'dbtrace' | tail -20 || true diff --git a/tests/crud_test.go b/tests/crud_test.go index fdc2508..f268d09 100644 --- a/tests/crud_test.go +++ b/tests/crud_test.go @@ -168,15 +168,14 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { // Generate unique IDs for this test run timestamp := time.Now().Unix() - deptID := fmt.Sprintf("dept_rs_%d", timestamp) - empID := fmt.Sprintf("emp_rs_%d", timestamp) + // IDs are assigned by the database (serial) and captured on create + var deptID, empID int64 // Test CREATE operation t.Run("Create_Department", func(t *testing.T) { payload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": deptID, "name": "Engineering Department", "code": fmt.Sprintf("ENG_%d", timestamp), "description": "Software Engineering", @@ -188,15 +187,15 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) + deptID = createdID(result) assert.True(t, result["success"].(bool), "Create department should succeed") - logger.Info("Department created successfully: %s", deptID) + logger.Info("Department created successfully: %d", deptID) }) t.Run("Create_Employee", func(t *testing.T) { payload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": empID, "first_name": "John", "last_name": "Doe", "email": fmt.Sprintf("john.doe.rs.%d@example.com", timestamp), @@ -212,8 +211,9 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) + empID = createdID(result) assert.True(t, result["success"].(bool), "Create employee should succeed") - logger.Info("Employee created successfully: %s", empID) + logger.Info("Employee created successfully: %d", empID) }) // Test READ operation @@ -222,7 +222,7 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { "operation": "read", } - resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%s", deptID), payload) + resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%d", deptID), payload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -230,9 +230,9 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { assert.True(t, result["success"].(bool), "Read department should succeed") data := result["data"].(map[string]interface{}) - assert.Equal(t, deptID, data["id"]) + assert.EqualValues(t, deptID, data["id"]) assert.Equal(t, "Engineering Department", data["name"]) - logger.Info("Department read successfully: %s", deptID) + logger.Info("Department read successfully: %d", deptID) }) t.Run("Read_Employees_With_Filters", func(t *testing.T) { @@ -270,17 +270,17 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { }, } - resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%s", deptID), payload) + resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%d", deptID), payload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) assert.True(t, result["success"].(bool), "Update department should succeed") - logger.Info("Department updated successfully: %s", deptID) + logger.Info("Department updated successfully: %d", deptID) // Verify update readPayload := map[string]interface{}{"operation": "read"} - resp = makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%s", deptID), readPayload) + resp = makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%d", deptID), readPayload) json.NewDecoder(resp.Body).Decode(&result) data := result["data"].(map[string]interface{}) assert.Equal(t, "Updated Software Engineering Department", data["description"]) @@ -294,13 +294,13 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { }, } - resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%s", empID), payload) + resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%d", empID), payload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) assert.True(t, result["success"].(bool), "Update employee should succeed") - logger.Info("Employee updated successfully: %s", empID) + logger.Info("Employee updated successfully: %d", empID) }) // Test DELETE operation @@ -309,17 +309,17 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { "operation": "delete", } - resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%s", empID), payload) + resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%d", empID), payload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) assert.True(t, result["success"].(bool), "Delete employee should succeed") - logger.Info("Employee deleted successfully: %s", empID) + logger.Info("Employee deleted successfully: %d", empID) // Verify deletion - after delete, reading should return empty/zero-value record or error readPayload := map[string]interface{}{"operation": "read"} - resp = makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%s", empID), readPayload) + resp = makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/employees/%d", empID), readPayload) json.NewDecoder(resp.Body).Decode(&result) // After deletion, the record should either not exist or have empty/zero ID if result["success"] != nil && result["success"].(bool) { @@ -337,13 +337,13 @@ func testResolveSpecCRUD(t *testing.T, serverURL string) { "operation": "delete", } - resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%s", deptID), payload) + resp := makeResolveSpecRequest(t, serverURL, fmt.Sprintf("/resolvespec/departments/%d", deptID), payload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) assert.True(t, result["success"].(bool), "Delete department should succeed") - logger.Info("Department deleted successfully: %s", deptID) + logger.Info("Department deleted successfully: %d", deptID) }) logger.Info("ResolveSpec API CRUD tests completed") @@ -355,13 +355,12 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Generate unique IDs for this test run timestamp := time.Now().Unix() - deptID := fmt.Sprintf("dept_rhs_%d", timestamp) - empID := fmt.Sprintf("emp_rhs_%d", timestamp) + // IDs are assigned by the database (serial) and captured on create + var deptID, empID int64 // Test CREATE operation (POST) t.Run("Create_Department", func(t *testing.T) { data := map[string]interface{}{ - "id": deptID, "name": "Marketing Department", "code": fmt.Sprintf("MKT_%d", timestamp), "description": "Marketing and Communications", @@ -372,20 +371,20 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) + deptID = createdID(result) // Check if response has "success" field (wrapped format) or direct data (unwrapped format) if success, ok := result["success"]; ok && success != nil { assert.True(t, success.(bool), "Create department should succeed") } else { // Unwrapped format - verify we got the created data back assert.NotEmpty(t, result, "Create department should return data") - assert.Equal(t, deptID, result["id"], "Created department should have correct ID") + assert.EqualValues(t, deptID, result["id"], "Created department should have correct ID") } - logger.Info("Department created successfully: %s", deptID) + logger.Info("Department created successfully: %d", deptID) }) t.Run("Create_Employee", func(t *testing.T) { data := map[string]interface{}{ - "id": empID, "first_name": "Jane", "last_name": "Smith", "email": fmt.Sprintf("jane.smith.rhs.%d@example.com", timestamp), @@ -400,20 +399,21 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { var result map[string]interface{} json.NewDecoder(resp.Body).Decode(&result) + empID = createdID(result) // Check if response has "success" field (wrapped format) or direct data (unwrapped format) if success, ok := result["success"]; ok && success != nil { assert.True(t, success.(bool), "Create employee should succeed") } else { // Unwrapped format - verify we got the created data back assert.NotEmpty(t, result, "Create employee should return data") - assert.Equal(t, empID, result["id"], "Created employee should have correct ID") + assert.EqualValues(t, empID, result["id"], "Created employee should have correct ID") } - logger.Info("Employee created successfully: %s", empID) + logger.Info("Employee created successfully: %d", empID) }) // Test READ operation (GET) t.Run("Read_Department", func(t *testing.T) { - resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%s", deptID), "GET", nil, nil) + resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%d", deptID), "GET", nil, nil) assert.Equal(t, http.StatusOK, resp.StatusCode) // RestHeadSpec may return data directly as array/object or wrapped in response object @@ -424,7 +424,7 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { var dataArray []interface{} if err := json.Unmarshal(body, &dataArray); err == nil { assert.GreaterOrEqual(t, len(dataArray), 1, "Should find department") - logger.Info("Department read successfully (simple format - array): %s", deptID) + logger.Info("Department read successfully (simple format - array): %d", deptID) return } @@ -435,7 +435,7 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { if _, hasSuccess := singleObj["success"]; !hasSuccess { // This is a direct data object (simple format, single record) assert.NotEmpty(t, singleObj, "Should find department") - logger.Info("Department read successfully (simple format - single object): %s", deptID) + logger.Info("Department read successfully (simple format - single object): %d", deptID) return } @@ -444,13 +444,13 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Check if data is an array if data, ok := singleObj["data"].([]interface{}); ok { assert.GreaterOrEqual(t, len(data), 1, "Should find department") - logger.Info("Department read successfully (detail format - array): %s", deptID) + logger.Info("Department read successfully (detail format - array): %d", deptID) return } // Check if data is a single object (SingleRecordAsObject feature in detail format) if data, ok := singleObj["data"].(map[string]interface{}); ok { assert.NotEmpty(t, data, "Should find department") - logger.Info("Department read successfully (detail format - single object): %s", deptID) + logger.Info("Department read successfully (detail format - single object): %d", deptID) return } } @@ -549,7 +549,7 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { "description": "Updated Marketing and Sales Department", } - resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%s", deptID), "PUT", data, nil) + resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%d", deptID), "PUT", data, nil) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -561,11 +561,11 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Unwrapped format - verify we got the updated data back assert.NotEmpty(t, result, "Update department should return data") } - logger.Info("Department updated successfully: %s", deptID) + logger.Info("Department updated successfully: %d", deptID) // Verify update by reading the department again // For simplicity, just verify the update succeeded, skip verification read - logger.Info("Department update verified: %s", deptID) + logger.Info("Department update verified: %d", deptID) }) t.Run("Update_Employee_With_PATCH", func(t *testing.T) { @@ -573,7 +573,7 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { "title": "Senior Marketing Manager", } - resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/employees/%s", empID), "PATCH", data, nil) + resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/employees/%d", empID), "PATCH", data, nil) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -585,12 +585,12 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Unwrapped format - verify we got the updated data back assert.NotEmpty(t, result, "Update employee should return data") } - logger.Info("Employee updated successfully: %s", empID) + logger.Info("Employee updated successfully: %d", empID) }) // Test DELETE operation (DELETE) t.Run("Delete_Employee", func(t *testing.T) { - resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/employees/%s", empID), "DELETE", nil, nil) + resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/employees/%d", empID), "DELETE", nil, nil) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -602,14 +602,14 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Unwrapped format - verify we got a response (typically {"deleted": count}) assert.NotEmpty(t, result, "Delete employee should return data") } - logger.Info("Employee deleted successfully: %s", empID) + logger.Info("Employee deleted successfully: %d", empID) // Verify deletion - just log that delete succeeded - logger.Info("Employee deletion verified: %s", empID) + logger.Info("Employee deletion verified: %d", empID) }) t.Run("Delete_Department", func(t *testing.T) { - resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%s", deptID), "DELETE", nil, nil) + resp := makeRestHeadSpecRequest(t, serverURL, fmt.Sprintf("/restheadspec/departments/%d", deptID), "DELETE", nil, nil) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -621,7 +621,7 @@ func testRestHeadSpecCRUD(t *testing.T, serverURL string) { // Unwrapped format - verify we got a response (typically {"deleted": count}) assert.NotEmpty(t, result, "Delete department should return data") } - logger.Info("Department deleted successfully: %s", deptID) + logger.Info("Department deleted successfully: %d", deptID) }) logger.Info("RestHeadSpec API CRUD tests completed") @@ -687,3 +687,15 @@ func makeRestHeadSpecRequest(t *testing.T, serverURL, path, method string, data return resp } + +// createdID extracts the database-assigned id from a create response, in either the +// wrapped ({"data": {...}}) or unwrapped format. +func createdID(result map[string]interface{}) int64 { + if data, ok := result["data"].(map[string]interface{}); ok { + result = data + } + if id, ok := result["id"].(float64); ok { + return int64(id) + } + return 0 +} diff --git a/tests/integration_test.go b/tests/integration_test.go index e20f6fb..b062e0a 100644 --- a/tests/integration_test.go +++ b/tests/integration_test.go @@ -2,6 +2,7 @@ package test import ( "encoding/json" + "fmt" "net/http" "testing" "time" @@ -9,6 +10,32 @@ import ( "github.com/stretchr/testify/assert" ) +// Database-assigned (serial) ids, captured on create; later tests build on earlier ones. +var deptID, emp1ID, emp2ID, mgrID, projID, task1ID int64 + +// createdIDs returns the ids of the records in a create response (single object or array). +func createdIDs(resp *http.Response) []int64 { + var result struct { + Data interface{} `json:"data"` + } + _ = json.NewDecoder(resp.Body).Decode(&result) + items, ok := result.Data.([]interface{}) + if !ok { + items = []interface{}{result.Data} + } + ids := make([]int64, 0, len(items)) + for _, item := range items { + if m, ok := item.(map[string]interface{}); ok { + if id, ok := m["id"].(float64); ok { + ids = append(ids, int64(id)) + continue + } + } + ids = append(ids, 0) + } + return ids +} + // TestMain sets up the test environment func TestMain(m *testing.M) { TestSetup(m) @@ -19,7 +46,6 @@ func TestDepartmentEmployees(t *testing.T) { deptPayload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": "dept1", "name": "Engineering", "code": "ENG", "description": "Engineering Department", @@ -28,25 +54,24 @@ func TestDepartmentEmployees(t *testing.T) { resp := makeRequest(t, "/departments", deptPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) + deptID = createdIDs(resp)[0] // Create employees in department empPayload := map[string]interface{}{ "operation": "create", "data": []map[string]interface{}{ { - "id": "emp1", "first_name": "John", "last_name": "Doe", "email": "john@example.com", - "department_id": "dept1", + "department_id": deptID, "title": "Senior Engineer", }, { - "id": "emp2", "first_name": "Jane", "last_name": "Smith", "email": "jane@example.com", - "department_id": "dept1", + "department_id": deptID, "title": "Engineer", }, }, @@ -54,6 +79,8 @@ func TestDepartmentEmployees(t *testing.T) { resp = makeRequest(t, "/employees", empPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) + emps := createdIDs(resp) + emp1ID, emp2ID = emps[0], emps[1] // Read department with employees readPayload := map[string]interface{}{ @@ -68,7 +95,7 @@ func TestDepartmentEmployees(t *testing.T) { }, } - resp = makeRequest(t, "/departments/dept1", readPayload) + resp = makeRequest(t, fmt.Sprintf("/departments/%d", deptID), readPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -83,29 +110,29 @@ func TestEmployeeHierarchy(t *testing.T) { mgrPayload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": "mgr1", "first_name": "Alice", "last_name": "Manager", "email": "alice@example.com", "title": "Engineering Manager", - "department_id": "dept1", + "department_id": deptID, }, } resp := makeRequest(t, "/employees", mgrPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) + mgrID = createdIDs(resp)[0] // Update employees to set manager updatePayload := map[string]interface{}{ "operation": "update", "data": map[string]interface{}{ - "manager_id": "mgr1", + "manager_id": mgrID, }, } - resp = makeRequest(t, "/employees/emp1", updatePayload) + resp = makeRequest(t, fmt.Sprintf("/employees/%d", emp1ID), updatePayload) assert.Equal(t, http.StatusOK, resp.StatusCode) - resp = makeRequest(t, "/employees/emp2", updatePayload) + resp = makeRequest(t, fmt.Sprintf("/employees/%d", emp2ID), updatePayload) assert.Equal(t, http.StatusOK, resp.StatusCode) // Read manager with reports @@ -121,7 +148,7 @@ func TestEmployeeHierarchy(t *testing.T) { }, } - resp = makeRequest(t, "/employees/mgr1", readPayload) + resp = makeRequest(t, fmt.Sprintf("/employees/%d", mgrID), readPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} @@ -136,7 +163,6 @@ func TestProjectStructure(t *testing.T) { projectPayload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": "proj1", "name": "New Website", "code": "WEB", "description": "Company website redesign", @@ -149,15 +175,15 @@ func TestProjectStructure(t *testing.T) { resp := makeRequest(t, "/projects", projectPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) + projID = createdIDs(resp)[0] // Create project tasks taskPayload := map[string]interface{}{ "operation": "create", "data": []map[string]interface{}{ { - "id": "task1", - "project_id": "proj1", - "assignee_id": "emp1", + "project_id": projID, + "assignee_id": emp1ID, "title": "Design Homepage", "description": "Create homepage design", "status": "in_progress", @@ -165,9 +191,8 @@ func TestProjectStructure(t *testing.T) { "due_date": time.Now().AddDate(0, 1, 0).Format(time.RFC3339), }, { - "id": "task2", - "project_id": "proj1", - "assignee_id": "emp2", + "project_id": projID, + "assignee_id": emp2ID, "title": "Implement Backend", "description": "Implement backend APIs", "status": "planned", @@ -179,14 +204,14 @@ func TestProjectStructure(t *testing.T) { resp = makeRequest(t, "/project_tasks", taskPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) + task1ID = createdIDs(resp)[0] // Create task comments commentPayload := map[string]interface{}{ "operation": "create", "data": map[string]interface{}{ - "id": "comment1", - "task_id": "task1", - "author_id": "mgr1", + "task_id": task1ID, + "author_id": mgrID, "content": "Looking good! Please add more animations.", }, } @@ -223,7 +248,7 @@ func TestProjectStructure(t *testing.T) { }, } - resp = makeRequest(t, "/projects/proj1", readPayload) + resp = makeRequest(t, fmt.Sprintf("/projects/%d", projID), readPayload) assert.Equal(t, http.StatusOK, resp.StatusCode) var result map[string]interface{} From f2dbe2561ce3a4fdaebda45f49547d9fa4e08ac5 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:40:06 +0200 Subject: [PATCH 06/23] feat(hooks): add OnTxBegin and runInTx for resolvespec and restheadspec Every transaction the handlers open now fires OnTxBegin first, with the transaction in hookCtx.Tx, via common.RunRequestTx. --- audit/single_tran.md | 31 ++++++---- pkg/common/txhook.go | 27 +++++++++ pkg/resolvespec/handler.go | 50 ++++++++++++----- pkg/resolvespec/hooks.go | 9 +++ pkg/resolvespec/on_tx_begin_test.go | 84 ++++++++++++++++++++++++++++ pkg/restheadspec/handler.go | 70 ++++++++++++++--------- pkg/restheadspec/hooks.go | 9 +++ pkg/restheadspec/on_tx_begin_test.go | 84 ++++++++++++++++++++++++++++ 8 files changed, 313 insertions(+), 51 deletions(-) create mode 100644 pkg/common/txhook.go create mode 100644 pkg/resolvespec/on_tx_begin_test.go create mode 100644 pkg/restheadspec/on_tx_begin_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 0463c33..aaecd87 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -69,20 +69,29 @@ - Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking. ## Phases -| # | Change | Files | Notes | -|---|---|---|---| -| 0 | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above | -| 1 | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | -| 2 | `OnTxBegin` hook type + `runInTx` helper | `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec first | -| 3 | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | -| 4 | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | -| 5 | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | -| 6 | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | -| 7 | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | +| # | Status | Change | Files | Notes | +|---|---|---|---|---| +| 0 | DONE | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above | +| 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | +| 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 | +| 3 | TODO | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | +| 4 | TODO | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | +| 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | +| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | +| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | + +## Progress +- DONE P0: baseline via `dbtrace` on real Postgres (commit `cd96404`): create/read/delete `pooled=0`; update `pooled=1` (re-fetch) = P3 target. websocketspec/mqttspec/resolvemcp not measured. +- DONE P1: single + batch delete in one tx (resolvespec, restheadspec). Not done: per-item `BeforeDelete` in resolvespec batch (behavior change, deferred). +- DONE infra: `sqlmock` delete tx tests (both specs); compose test server + `scripts/testserver-smoke.sh` (podman first); testmodels ids now serial. +- NOTE: restheadspec single delete still does the lookup before `BeforeDelete`; safe once `OnTxBegin` (P2) exists. An `AfterDelete` failure now rolls the delete back. +- 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). +- FOUND (P3 scope): resolvespec batch update (`handler.go` ~`:1377`, `:1529`) reads existing record via `h.db.NewSelect()` inside the tx = pool connection; should be `tx`. +- NEXT: P3. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. -- Missing: any test asserting hook `Tx` is a tx, or counting connections per op. +- 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. - 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). - Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB). diff --git a/pkg/common/txhook.go b/pkg/common/txhook.go new file mode 100644 index 0000000..4d0023d --- /dev/null +++ b/pkg/common/txhook.go @@ -0,0 +1,27 @@ +package common + +import "context" + +// TxHookName is the shared value of every spec's OnTxBegin HookType. +const TxHookName = "on_tx_begin" + +// TxContext is implemented by a spec's HookContext so RunRequestTx can point +// it at the transaction it opens. +type TxContext interface { + SetTx(tx Database) +} + +// RunRequestTx opens a transaction on db, points tc at it, runs onBegin (the +// spec's OnTxBegin hooks) and then body. An error from onBegin or body rolls +// the transaction back; body is not run when onBegin fails. +func RunRequestTx(ctx context.Context, db Database, tc TxContext, onBegin func() error, body func(tx Database) error) error { + return db.RunInTransaction(ctx, func(tx Database) error { + tc.SetTx(tx) + if onBegin != nil { + if err := onBegin(); err != nil { + return err + } + } + return body(tx) + }) +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index cdbabb2..9c7e9f7 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -313,7 +313,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st errMsg string ) - txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { + txErr := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "read", options, w), func(tx common.Database) error { hookCtx := &HookContext{ Context: ctx, Handler: h, @@ -734,7 +734,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if h.shouldUseNestedProcessor(v, model) { logger.Info("Using nested CUD processor for create operation") var nestedResult *common.ProcessResult - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { hookCtx := &HookContext{ Context: ctx, Handler: h, @@ -782,7 +782,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat // Standard processing without nested relations pkName := reflection.GetPrimaryKeyName(model) var responseData interface{} = v - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { hookCtx := &HookContext{ Context: ctx, Handler: h, @@ -857,7 +857,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if hasNestedData { logger.Info("Using nested CUD processor for batch create with nested data") results := make([]map[string]interface{}, 0, len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { // Temporarily swap the database to use transaction originalDB := h.nestedProcessor h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h) @@ -912,7 +912,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat pkName := reflection.GetPrimaryKeyName(model) modelElemType := reflection.GetPointerElement(reflect.TypeOf(model)) responseItems := make([]interface{}, 0, len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { for _, item := range v { hookCtx := &HookContext{ Context: ctx, @@ -989,7 +989,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if hasNestedData { logger.Info("Using nested CUD processor for batch create with nested data ([]interface{})") results := make([]interface{}, 0, len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { // Temporarily swap the database to use transaction originalDB := h.nestedProcessor h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h) @@ -1046,7 +1046,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat pkName := reflection.GetPrimaryKeyName(model) modelElemType := reflection.GetPointerElement(reflect.TypeOf(model)) responseItems := make([]interface{}, 0, len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error { for _, item := range v { itemMap, ok := item.(map[string]interface{}) if !ok { @@ -1180,7 +1180,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url } // Wrap in transaction to ensure BeforeUpdate hook is inside transaction - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { // Execute BeforeUpdate hooks inside transaction, before any queries run. // BeforeUpdate hooks may set session-scoped RLS GUCs (via SET LOCAL); // they must run before the existence-check select so that select is @@ -1334,7 +1334,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url if hasNestedData { logger.Info("Using nested CUD processor for batch update with nested data") results := make([]map[string]interface{}, 0, len(updates)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { // Temporarily swap the database to use transaction originalDB := h.nestedProcessor h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h) @@ -1368,7 +1368,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Standard batch update without nested relations pkName := reflection.GetPrimaryKeyName(model) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { for _, item := range updates { if itemID, ok := item["id"]; ok { itemIDStr := fmt.Sprintf("%v", itemID) @@ -1479,7 +1479,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url if hasNestedData { logger.Info("Using nested CUD processor for batch update with nested data ([]interface{})") results := make([]interface{}, 0, len(updates)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { // Temporarily swap the database to use transaction originalDB := h.nestedProcessor h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h) @@ -1516,7 +1516,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Standard batch update without nested relations pkName := reflection.GetPrimaryKeyName(model) list := make([]interface{}, 0) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { for _, item := range updates { if itemMap, ok := item.(map[string]interface{}); ok { if itemID, ok := itemMap["id"]; ok { @@ -1657,8 +1657,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id // state set by hooks (e.g. RLS settings) applies to every statement. var payload interface{} var failure *deleteFailure - txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { - hookCtx.Tx = tx + txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error { payload, failure = h.executeDelete(ctx, tx, hookCtx, schema, tableName, model, id, data) if failure != nil { return failure @@ -2561,3 +2560,26 @@ func mergeWithInput(dbRecord interface{}, input map[string]interface{}) map[stri } return result } + +// newTxHookContext builds the context OnTxBegin hooks receive for paths that +// create their per-item hook contexts inside the transaction. +func (h *Handler) newTxHookContext(ctx context.Context, schema, entity string, model interface{}, operation string, options common.RequestOptions, w common.ResponseWriter) *HookContext { + return &HookContext{ + Context: ctx, + Handler: h, + Schema: schema, + Entity: entity, + Model: model, + Operation: operation, + Options: options, + Writer: w, + } +} + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin +// fired first. Every transaction the handler opens goes through here. +func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(ctx, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} diff --git a/pkg/resolvespec/hooks.go b/pkg/resolvespec/hooks.go index 5651e7a..184dd31 100644 --- a/pkg/resolvespec/hooks.go +++ b/pkg/resolvespec/hooks.go @@ -41,6 +41,12 @@ const ( // Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each // individual SQL-operation hook point, so it runs once per statement executed. BeforeOp HookType = "before_op" + + // OnTxBegin fires once, first, inside every transaction the handler opens + // (including the second short transaction for post-commit work). hookCtx.Tx + // is the transaction; use it to stamp transaction-local state such as RLS + // settings. An error or abort rolls the transaction back. + OnTxBegin HookType = common.TxHookName ) // HookContext contains all the data available to a hook @@ -76,6 +82,9 @@ type HookContext struct { Tx common.Database } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is the signature for hook functions // It receives a HookContext and can modify it or return an error // If an error is returned, the operation will be aborted diff --git a/pkg/resolvespec/on_tx_begin_test.go b/pkg/resolvespec/on_tx_begin_test.go new file mode 100644 index 0000000..df0096f --- /dev/null +++ b/pkg/resolvespec/on_tx_begin_test.go @@ -0,0 +1,84 @@ +package resolvespec + +import ( + "errors" + "net/http" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// recordTxOrder records the order hooks fire in and the Tx each one saw. +func recordTxOrder(h *Handler, beginErr error) (*[]string, *[]common.Database) { + var order []string + var txs []common.Database + h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error { + order = append(order, "begin") + txs = append(txs, ctx.Tx) + return beginErr + }) + h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error { + order = append(order, "before_delete") + return nil + }) + return &order, &txs +} + +func TestOnTxBeginFiresOnceFirstOnTx(t *testing.T) { + cases := map[string]struct { + id string + data interface{} + exec int + }{ + "single": {id: "7", exec: 1}, + "batch": {data: []interface{}{"1", "2"}, exec: 2}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + order, txs := recordTxOrder(h, nil) + + mock.ExpectBegin() + if tc.id != "" { + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + } + for i := 0; i < tc.exec; i++ { + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + } + mock.ExpectCommit() + + if rec := runDelete(h, tc.id, tc.data); rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*txs) != 1 || (*txs)[0] == nil || (*txs)[0] == h.db { + t.Fatalf("OnTxBegin must fire once on the transaction, got %v", *txs) + } + if (*order)[0] != "begin" { + t.Fatalf("OnTxBegin must fire before other hooks, got %v", *order) + } + }) + } +} + +func TestOnTxBeginErrorRollsBackWithoutQueries(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + order, _ := recordTxOrder(h, errors.New("no user")) + + mock.ExpectBegin() + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code < http.StatusBadRequest { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*order) != 1 { + t.Fatalf("no hook may run after a failed OnTxBegin, got %v", *order) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 39a8349..3daf210 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -460,8 +460,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st errMsg string ) - txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { - hookCtx.Tx = tx + txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error { if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed" @@ -1322,8 +1321,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat // Process all items in a transaction results := make([]interface{}, 0) - txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { - hookCtx.Tx = tx + txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error { if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed" return err @@ -1538,11 +1536,23 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // Variable to store the updated record var updatedRecord interface{} - // Declare hook context to be used inside and outside transaction - var hookCtx *HookContext + // Hook context used inside and outside transaction + hookCtx := &HookContext{ + Context: ctx, + Handler: h, + Schema: schema, + Entity: entity, + TableName: tableName, + Model: model, + Operation: "update", + Options: options, + ID: id, + Data: dataMap, + Writer: w, + } // Process nested relations if present - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, hookCtx, func(tx common.Database) error { // Create temporary nested processor with transaction txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h) @@ -1550,21 +1560,6 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // BeforeUpdate hooks may set session-scoped RLS GUCs (via SET LOCAL); // they must run before the existence-check select so that select is // also subject to RLS on this connection/transaction. - hookCtx = &HookContext{ - Context: ctx, - Handler: h, - Schema: schema, - Entity: entity, - TableName: tableName, - Tx: tx, - Model: model, - Operation: "update", - Options: options, - ID: id, - Data: dataMap, - Writer: w, - } - if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { return fmt.Errorf("BeforeUpdate hook failed: %w", err) } @@ -1733,7 +1728,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id // Array of IDs as strings logger.Info("Batch delete with %d IDs ([]string)", len(v)) deletedCount := 0 - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error { for _, itemID := range v { // Execute hooks for each item hookCtx := &HookContext{ @@ -1790,7 +1785,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id logger.Info("Batch delete with %d items ([]interface{})", len(v)) deletedCount := 0 pkName := reflection.GetPrimaryKeyName(model) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error { for _, item := range v { var itemID interface{} @@ -1864,7 +1859,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id logger.Info("Batch delete with %d items ([]map[string]interface{})", len(v)) deletedCount := 0 pkName := reflection.GetPrimaryKeyName(model) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error { for _, item := range v { if itemID, ok := item[pkName]; ok && itemID != nil { itemIDStr := fmt.Sprintf("%v", itemID) @@ -1943,7 +1938,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id // Lookup, hooks and delete share one transaction so transaction-local // state set by hooks (e.g. RLS settings) applies to every statement. var failure *deleteFailure - txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error { + txErr := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error { failure = h.deleteSingleInTx(ctx, tx, w, schema, entity, tableName, model, pkName, id, recordToDelete) if failure != nil { return failure @@ -3532,3 +3527,26 @@ func (h *Handler) HandleOpenAPI(w common.ResponseWriter, r common.Request) { func (h *Handler) SetOpenAPIGenerator(generator func() (string, error)) { h.openAPIGenerator = generator } + +// newTxHookContext builds the context OnTxBegin hooks receive for paths that +// create their per-item hook contexts inside the transaction. +func (h *Handler) newTxHookContext(ctx context.Context, schema, entity, tableName string, model interface{}, operation string, w common.ResponseWriter) *HookContext { + return &HookContext{ + Context: ctx, + Handler: h, + Schema: schema, + Entity: entity, + TableName: tableName, + Model: model, + Operation: operation, + Writer: w, + } +} + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin +// fired first. Every transaction the handler opens goes through here. +func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(ctx, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} diff --git a/pkg/restheadspec/hooks.go b/pkg/restheadspec/hooks.go index 57b2320..a0c628b 100644 --- a/pkg/restheadspec/hooks.go +++ b/pkg/restheadspec/hooks.go @@ -41,6 +41,12 @@ const ( // Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each // individual SQL-operation hook point, so it runs once per statement executed. BeforeOp HookType = "before_op" + + // OnTxBegin fires once, first, inside every transaction the handler opens + // (including the second short transaction for post-commit work). hookCtx.Tx + // is the transaction; use it to stamp transaction-local state such as RLS + // settings. An error or abort rolls the transaction back. + OnTxBegin HookType = common.TxHookName ) // HookContext contains all the data available to a hook @@ -83,6 +89,9 @@ type HookContext struct { Tx common.Database } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is the signature for hook functions // It receives a HookContext and can modify it or return an error // If an error is returned, the operation will be aborted diff --git a/pkg/restheadspec/on_tx_begin_test.go b/pkg/restheadspec/on_tx_begin_test.go new file mode 100644 index 0000000..2b473e7 --- /dev/null +++ b/pkg/restheadspec/on_tx_begin_test.go @@ -0,0 +1,84 @@ +package restheadspec + +import ( + "errors" + "net/http" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// recordTxOrder records the order hooks fire in and the Tx each one saw. +func recordTxOrder(h *Handler, beginErr error) (*[]string, *[]common.Database) { + var order []string + var txs []common.Database + h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error { + order = append(order, "begin") + txs = append(txs, ctx.Tx) + return beginErr + }) + h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error { + order = append(order, "before_delete") + return nil + }) + return &order, &txs +} + +func TestOnTxBeginFiresOnceFirstOnTx(t *testing.T) { + cases := map[string]struct { + id string + data interface{} + exec int + }{ + "single": {id: "7", exec: 1}, + "batch": {data: []interface{}{"1", "2"}, exec: 2}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + order, txs := recordTxOrder(h, nil) + + mock.ExpectBegin() + if tc.id != "" { + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + } + for i := 0; i < tc.exec; i++ { + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + } + mock.ExpectCommit() + + if rec := runDelete(h, tc.id, tc.data); rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*txs) != 1 || (*txs)[0] == nil || (*txs)[0] == h.db { + t.Fatalf("OnTxBegin must fire once on the transaction, got %v", *txs) + } + if (*order)[0] != "begin" { + t.Fatalf("OnTxBegin must fire before other hooks, got %v", *order) + } + }) + } +} + +func TestOnTxBeginErrorRollsBackWithoutQueries(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + order, _ := recordTxOrder(h, errors.New("no user")) + + mock.ExpectBegin() + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code < http.StatusBadRequest { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(*order) != 1 { + t.Fatalf("no hook may run after a failed OnTxBegin, got %v", *order) + } +} From eb492d52aa0f9f719b23051946cb2c44df61d3e7 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:40:33 +0200 Subject: [PATCH 07/23] feat(clients): add Go, Rust, C# and Dart clients for ResolveSpec and FunctionSpec --- clients/README.md | 12 + clients/resolvespec-cs/.gitignore | 2 + clients/resolvespec-cs/README.md | 43 +++ clients/resolvespec-cs/src/FuncSpecClient.cs | 185 ++++++++++++ clients/resolvespec-cs/src/Http.cs | 78 +++++ clients/resolvespec-cs/src/ResolveSpec.csproj | 11 + .../resolvespec-cs/src/ResolveSpecClient.cs | 65 +++++ clients/resolvespec-cs/src/Types.cs | 137 +++++++++ clients/resolvespec-cs/tests/ClientTests.cs | 178 ++++++++++++ .../tests/ResolveSpec.Tests.csproj | 16 + clients/resolvespec-dart/.gitignore | 3 + clients/resolvespec-dart/README.md | 43 +++ .../resolvespec-dart/analysis_options.yaml | 1 + clients/resolvespec-dart/lib/resolvespec.dart | 7 + clients/resolvespec-dart/lib/src/client.dart | 84 ++++++ .../resolvespec-dart/lib/src/funcspec.dart | 186 ++++++++++++ .../resolvespec-dart/lib/src/resolvespec.dart | 63 ++++ clients/resolvespec-dart/lib/src/types.dart | 268 +++++++++++++++++ clients/resolvespec-dart/pubspec.yaml | 14 + .../resolvespec-dart/test/client_test.dart | 179 ++++++++++++ clients/resolvespec-go/README.md | 40 +++ clients/resolvespec-go/client.go | 121 ++++++++ clients/resolvespec-go/funcspec.go | 274 ++++++++++++++++++ clients/resolvespec-go/funcspec_test.go | 98 +++++++ clients/resolvespec-go/go.mod | 3 + clients/resolvespec-go/resolvespec.go | 94 ++++++ clients/resolvespec-go/resolvespec_test.go | 99 +++++++ clients/resolvespec-go/types.go | 107 +++++++ clients/resolvespec-rs/.gitignore | 2 + clients/resolvespec-rs/Cargo.toml | 18 ++ clients/resolvespec-rs/README.md | 40 +++ clients/resolvespec-rs/src/client.rs | 90 ++++++ clients/resolvespec-rs/src/error.rs | 34 +++ clients/resolvespec-rs/src/funcspec.rs | 263 +++++++++++++++++ clients/resolvespec-rs/src/lib.rs | 12 + clients/resolvespec-rs/src/resolvespec.rs | 130 +++++++++ clients/resolvespec-rs/src/types.rs | 192 ++++++++++++ clients/resolvespec-rs/tests/client.rs | 164 +++++++++++ 38 files changed, 3356 insertions(+) create mode 100644 clients/README.md create mode 100644 clients/resolvespec-cs/.gitignore create mode 100644 clients/resolvespec-cs/README.md create mode 100644 clients/resolvespec-cs/src/FuncSpecClient.cs create mode 100644 clients/resolvespec-cs/src/Http.cs create mode 100644 clients/resolvespec-cs/src/ResolveSpec.csproj create mode 100644 clients/resolvespec-cs/src/ResolveSpecClient.cs create mode 100644 clients/resolvespec-cs/src/Types.cs create mode 100644 clients/resolvespec-cs/tests/ClientTests.cs create mode 100644 clients/resolvespec-cs/tests/ResolveSpec.Tests.csproj create mode 100644 clients/resolvespec-dart/.gitignore create mode 100644 clients/resolvespec-dart/README.md create mode 100644 clients/resolvespec-dart/analysis_options.yaml create mode 100644 clients/resolvespec-dart/lib/resolvespec.dart create mode 100644 clients/resolvespec-dart/lib/src/client.dart create mode 100644 clients/resolvespec-dart/lib/src/funcspec.dart create mode 100644 clients/resolvespec-dart/lib/src/resolvespec.dart create mode 100644 clients/resolvespec-dart/lib/src/types.dart create mode 100644 clients/resolvespec-dart/pubspec.yaml create mode 100644 clients/resolvespec-dart/test/client_test.dart create mode 100644 clients/resolvespec-go/README.md create mode 100644 clients/resolvespec-go/client.go create mode 100644 clients/resolvespec-go/funcspec.go create mode 100644 clients/resolvespec-go/funcspec_test.go create mode 100644 clients/resolvespec-go/go.mod create mode 100644 clients/resolvespec-go/resolvespec.go create mode 100644 clients/resolvespec-go/resolvespec_test.go create mode 100644 clients/resolvespec-go/types.go create mode 100644 clients/resolvespec-rs/.gitignore create mode 100644 clients/resolvespec-rs/Cargo.toml create mode 100644 clients/resolvespec-rs/README.md create mode 100644 clients/resolvespec-rs/src/client.rs create mode 100644 clients/resolvespec-rs/src/error.rs create mode 100644 clients/resolvespec-rs/src/funcspec.rs create mode 100644 clients/resolvespec-rs/src/lib.rs create mode 100644 clients/resolvespec-rs/src/resolvespec.rs create mode 100644 clients/resolvespec-rs/src/types.rs create mode 100644 clients/resolvespec-rs/tests/client.rs diff --git a/clients/README.md b/clients/README.md new file mode 100644 index 0000000..1039b0d --- /dev/null +++ b/clients/README.md @@ -0,0 +1,12 @@ +# Clients + +| Dir | Language | Specs | Verified | +|---|---|---|---| +| `resolvespec-js` | TypeScript | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | yes | +| `resolvespec-python` | Python >= 3.11 | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | yes (61 tests) | +| `resolvespec-go` | Go | ResolveSpec, FunctionSpec | yes (`go test`) | +| `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | yes (`cargo test`) | +| `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | **not compiled** | +| `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | **not compiled** | + +Wire behaviour is identical across clients; FunctionSpec server quirks are listed in each README. diff --git a/clients/resolvespec-cs/.gitignore b/clients/resolvespec-cs/.gitignore new file mode 100644 index 0000000..cd42ee3 --- /dev/null +++ b/clients/resolvespec-cs/.gitignore @@ -0,0 +1,2 @@ +bin/ +obj/ diff --git a/clients/resolvespec-cs/README.md b/clients/resolvespec-cs/README.md new file mode 100644 index 0000000..41488e5 --- /dev/null +++ b/clients/resolvespec-cs/README.md @@ -0,0 +1,43 @@ +# ResolveSpec.Client (C#) + +.NET 8 client for ResolveSpec (JSON body) and FunctionSpec. `System.Text.Json`, no other dependencies. + +> Not compiled or tested yet (no .NET SDK was available). Run `dotnet test tests/` first. + +## Clients + +| Type | Constructor | Methods | +|---|---|---| +| `ResolveSpecClient` | `(baseUrl, ClientOptions?)` | `GetMetadataAsync` `ReadAsync` `CreateAsync` `UpdateAsync` `DeleteAsync` | +| `FuncSpecClient` | `(baseUrl, ClientOptions?)` | `QueryAsync` `QueryListAsync` | + +`ClientOptions`: `Token`, `Headers`, `Timeout`, `HttpClient`. Precedence: Content-Type < custom headers < bearer token. + +## ResolveSpec + +- `id`: int/long/string → URL, `IEnumerable` → body. +- `Options` with nullable properties; wire names via `JsonPropertyName`. +- Result: `Response{Success, Data (JsonElement), Metadata}`; `resp.Decode()`. + +## FunctionSpec + +- Routes are server-defined: pass the `path`. +- Params (`IDictionary`) → query string (enumerable → repeated keys, bool → `true`/`false`, null skipped). +- `FuncSpecOptions` → `X-*` headers: `Filters`, `SearchFilters`, `CustomSqlWhere`, `CustomSqlOr`, `Sort`, `Limit`, `Offset`, `Distinct`, `SkipCount`, `SkipCache`, `ResponseFormat`. +- `QueryListAsync` fills `Metadata` from `Content-Range`; 206 is success. +- Static helpers: `BuildHeaders`, `BuildQuery`, `EncodeHeaderValue`, `DecodeHeaderValue`. + +## Server quirks + +- `Sort` is raw SQL in ORDER BY (client sends `col ASC|DESC`). +- One search operator per column. +- Values starting `ZIP_` / `__` are base64-decoded by the server. +- Non-ASCII, control chars and edge spaces are auto-encoded (`ZIP_`). + +## Errors + +`ResolveSpecException{StatusCode, Message, Error{Code, Detail, Sql}}`. + +## Test + +`dotnet test tests/` diff --git a/clients/resolvespec-cs/src/FuncSpecClient.cs b/clients/resolvespec-cs/src/FuncSpecClient.cs new file mode 100644 index 0000000..acdc588 --- /dev/null +++ b/clients/resolvespec-cs/src/FuncSpecClient.cs @@ -0,0 +1,185 @@ +using System.Globalization; +using System.Text; +using System.Text.Json; +using System.Text.RegularExpressions; + +namespace ResolveSpec; + +/// +/// Options sent to funcspec endpoints as X-* headers. +/// Server behaviour (pkg/funcspec): Sort is inserted raw into ORDER BY (so it is sent as SQL terms); +/// only one search operator per column is kept; values starting with "ZIP_" or "__" are +/// base64-decoded by the server, so such plaintext values cannot be sent faithfully. +/// +public sealed class FuncSpecOptions +{ + /// eq+AND -> X-FieldFilter; others X-SearchOp / X-SearchOr. + public List? Filters { get; set; } + /// X-SearchFilter-{col}: text ILIKE. + public Dictionary? SearchFilters { get; set; } + public string? CustomSqlWhere { get; set; } + public string? CustomSqlOr { get; set; } + public List? Sort { get; set; } + public int? Limit { get; set; } + public int? Offset { get; set; } + public bool? Distinct { get; set; } + public bool? SkipCount { get; set; } + public bool? SkipCache { get; set; } + /// simple | detail | syncfusion + public string? ResponseFormat { get; set; } +} + +/// Client for user-defined SQL endpoints. Routes are defined by the server application. +public sealed class FuncSpecClient +{ + readonly Transport _t; + + public FuncSpecClient(string baseUrl, ClientOptions? options = null) => _t = new Transport(baseUrl, options); + + static readonly Dictionary OperatorMap = new() + { + ["eq"] = "equals", ["neq"] = "notequals", ["gt"] = "greaterthan", ["gte"] = "greaterthanorequal", + ["lt"] = "lessthan", ["lte"] = "lessthanorequal", ["like"] = "contains", ["ilike"] = "contains", + ["contains"] = "contains", ["startswith"] = "beginswith", ["endswith"] = "endswith", ["in"] = "in", + ["between"] = "between", ["between_inclusive"] = "betweeninclusive", + ["is_null"] = "empty", ["is_not_null"] = "notempty", + }; + + static string Scalar(object? v) => v switch + { + null => "", + string s => s, + bool b => b ? "true" : "false", + JsonElement { ValueKind: JsonValueKind.Null } => "", + JsonElement e => e.ValueKind == JsonValueKind.String ? e.GetString() ?? "" : e.ToString(), + IFormattable f => f.ToString(null, CultureInfo.InvariantCulture), + _ => v.ToString() ?? "", + }; + + static string FilterValue(object? v) => + v is System.Collections.IEnumerable list and not string + ? string.Join(",", list.Cast().Select(Scalar)) + : Scalar(v); + + /// Base64 (UTF-8) with the ZIP_ prefix. + public static string EncodeHeaderValue(string v) => "ZIP_" + Convert.ToBase64String(Encoding.UTF8.GetBytes(v)); + + /// Decode a value that may carry a ZIP_ or __ prefix (nested allowed). + public static string DecodeHeaderValue(string v) + { + foreach (var p in new[] { "ZIP_", "__" }) + { + if (!v.StartsWith(p, StringComparison.Ordinal)) continue; + var b64 = Regex.Replace(v[p.Length..], "[\n\r ]", ""); + b64 = b64.PadRight(b64.Length + (4 - b64.Length % 4) % 4, '='); + try { return DecodeHeaderValue(Encoding.UTF8.GetString(Convert.FromBase64String(b64))); } + catch (FormatException) { return v; } + } + return v; + } + + /// Encode values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces). + static string Safe(string v) => + v != v.Trim() || v.Any(c => c > 127 || char.IsControl(c)) ? EncodeHeaderValue(v) : v; + + /// Build the X-* headers understood by funcspec.ParseParameters. + public static Dictionary BuildHeaders(FuncSpecOptions? o) + { + var h = new Dictionary(); + if (o == null) return h; + + foreach (var f in o.Filters ?? new()) + { + var logic = string.IsNullOrEmpty(f.LogicOperator) ? "AND" : f.LogicOperator; + var v = Safe(FilterValue(f.Value)); + if (f.Operator == "eq" && logic == "AND") { h[$"X-FieldFilter-{f.Column}"] = v; continue; } + var op = OperatorMap.TryGetValue(f.Operator, out var m) ? m : f.Operator; + h[$"{(logic == "OR" ? "X-SearchOr" : "X-SearchOp")}-{op}-{f.Column}"] = v; + } + foreach (var (col, text) in o.SearchFilters ?? new()) h[$"X-SearchFilter-{col}"] = Safe(text); + if (!string.IsNullOrEmpty(o.CustomSqlWhere)) h["X-Custom-SQL-W"] = Safe(o.CustomSqlWhere); + if (!string.IsNullOrEmpty(o.CustomSqlOr)) h["X-Custom-SQL-Or"] = Safe(o.CustomSqlOr); + if (o.Sort is { Count: > 0 }) + { + // funcspec puts this verbatim into ORDER BY + h["X-Sort"] = Safe(string.Join(",", o.Sort.Select(s => + $"{s.Column} {(string.Equals(s.Direction, "desc", StringComparison.OrdinalIgnoreCase) ? "DESC" : "ASC")}"))); + } + if (o.Limit != null) h["X-Limit"] = o.Limit.Value.ToString(CultureInfo.InvariantCulture); + if (o.Offset != null) h["X-Offset"] = o.Offset.Value.ToString(CultureInfo.InvariantCulture); + if (o.Distinct != null) h["X-Distinct"] = Bool(o.Distinct.Value); + if (o.SkipCount != null) h["X-SkipCount"] = Bool(o.SkipCount.Value); + if (o.SkipCache != null) h["X-SkipCache"] = Bool(o.SkipCache.Value); + switch (o.ResponseFormat) + { + case "simple": h["X-SimpleApi"] = "true"; break; + case "detail": h["X-DetailApi"] = "true"; break; + case "syncfusion": h["X-Syncfusion"] = "true"; break; + } + return h; + } + + static string Bool(bool b) => b ? "true" : "false"; + + /// Build query-string pairs: bools -> true/false, lists -> repeated keys, null skipped. + public static List> BuildQuery(IDictionary? p) + { + var o = new List>(); + foreach (var (k, v) in p ?? new Dictionary()) + { + if (v == null) continue; + if (v is System.Collections.IEnumerable list and not string) + foreach (var e in list) o.Add(new(k, Safe(Scalar(e)))); + else o.Add(new(k, Safe(Scalar(v)))); + } + return o; + } + + static readonly Regex ContentRange = new(@"(\d+)-(\d+)/(\d+)"); + + static Metadata MetadataFrom(string? contentRange, FuncSpecOptions? o) + { + var m = new Metadata { Limit = o?.Limit ?? 0 }; + var g = ContentRange.Match(contentRange ?? ""); + if (g.Success) + { + var start = long.Parse(g.Groups[1].Value, CultureInfo.InvariantCulture); + var end = long.Parse(g.Groups[2].Value, CultureInfo.InvariantCulture); + var total = long.Parse(g.Groups[3].Value, CultureInfo.InvariantCulture); + m.Total = total; m.Filtered = total; m.Count = end - start; m.Offset = start; + } + return m; + } + + async Task CallAsync(HttpMethod method, string path, IDictionary? p, FuncSpecOptions? o, bool list, CancellationToken ct) + { + var url = $"{_t.BaseUrl}/{path.TrimStart('/')}"; + var q = BuildQuery(p); + if (q.Count > 0) + url += "?" + string.Join("&", q.Select(kv => $"{Uri.EscapeDataString(kv.Key)}={Uri.EscapeDataString(kv.Value)}")); + + var (resp, text) = await _t.SendAsync(method, url, null, BuildHeaders(o), ct).ConfigureAwait(false); + var status = (int)resp.StatusCode; + if (!resp.IsSuccessStatusCode) throw Transport.ErrorFrom(status, text, resp.ReasonPhrase); // 206 is success + + var r = new Response + { + Success = true, + Data = string.IsNullOrWhiteSpace(text) ? JsonDocument.Parse("null").RootElement.Clone() : JsonDocument.Parse(text).RootElement.Clone(), + }; + if (list) + { + resp.Headers.TryGetValues("Content-Range", out var cr); + r.Metadata = MetadataFrom(cr?.FirstOrDefault(), o); + } + return r; + } + + /// Single-record endpoint (SqlQuery). Data is the row object. + public Task QueryAsync(string path, IDictionary? p = null, FuncSpecOptions? o = null, HttpMethod? method = null, CancellationToken ct = default) => + CallAsync(method ?? HttpMethod.Get, path, p, o, false, ct); + + /// List endpoint (SqlQueryList). Metadata comes from Content-Range. + public Task QueryListAsync(string path, IDictionary? p = null, FuncSpecOptions? o = null, HttpMethod? method = null, CancellationToken ct = default) => + CallAsync(method ?? HttpMethod.Get, path, p, o, true, ct); +} diff --git a/clients/resolvespec-cs/src/Http.cs b/clients/resolvespec-cs/src/Http.cs new file mode 100644 index 0000000..21e4b90 --- /dev/null +++ b/clients/resolvespec-cs/src/Http.cs @@ -0,0 +1,78 @@ +using System.Net.Http.Headers; +using System.Text; +using System.Text.Json; + +namespace ResolveSpec; + +/// Shared HTTP configuration for both clients. +public sealed class ClientOptions +{ + public string? Token { get; set; } + public Dictionary Headers { get; } = new(StringComparer.OrdinalIgnoreCase); + public TimeSpan Timeout { get; set; } = TimeSpan.FromSeconds(30); + /// Supply your own HttpClient (tests, pooling). Its BaseAddress is ignored. + public HttpClient? HttpClient { get; set; } +} + +internal sealed class Transport +{ + public readonly string BaseUrl; + readonly ClientOptions _o; + readonly HttpClient _http; + + public Transport(string baseUrl, ClientOptions? o) + { + BaseUrl = baseUrl.TrimEnd('/'); + _o = o ?? new ClientOptions(); + _http = _o.HttpClient ?? new HttpClient { Timeout = _o.Timeout }; + } + + /// Content-Type < custom headers < per-call headers < bearer token. + public async Task<(HttpResponseMessage resp, string body)> SendAsync( + HttpMethod method, string url, string? json, IDictionary? extra, CancellationToken ct) + { + using var req = new HttpRequestMessage(method, url); + if (json != null) req.Content = new StringContent(json, Encoding.UTF8, "application/json"); + foreach (var (k, v) in _o.Headers) Set(req, k, v); + if (extra != null) foreach (var (k, v) in extra) Set(req, k, v); + if (!string.IsNullOrEmpty(_o.Token)) req.Headers.Authorization = new AuthenticationHeaderValue("Bearer", _o.Token); + var resp = await _http.SendAsync(req, ct).ConfigureAwait(false); + var body = await resp.Content.ReadAsStringAsync(ct).ConfigureAwait(false); + return (resp, body); + } + + static void Set(HttpRequestMessage req, string name, string value) + { + req.Headers.Remove(name); + if (!req.Headers.TryAddWithoutValidation(name, value) && req.Content != null) + { + req.Content.Headers.Remove(name); + req.Content.Headers.TryAddWithoutValidation(name, value); + } + } + + public static ResolveSpecException ErrorFrom(int status, string body, string? reason) + { + ApiError? err = null; + var isJson = false; + try + { + using var doc = JsonDocument.Parse(body); + isJson = true; + if (doc.RootElement.ValueKind == JsonValueKind.Object && doc.RootElement.TryGetProperty("error", out var e) && e.ValueKind == JsonValueKind.Object) + err = e.Deserialize(); + } + catch (JsonException) { } + + var message = err?.Message; + if (string.IsNullOrEmpty(message)) + { + var text = isJson ? "" : body.Trim(); + if (text.Length > 200) text = text[..200]; + message = text.Length > 0 ? text : $"{reason ?? "Error"} ({status})"; + } + return new ResolveSpecException(message, status, err); + } + + public static string Segment(string s) => Uri.EscapeDataString(s); +} diff --git a/clients/resolvespec-cs/src/ResolveSpec.csproj b/clients/resolvespec-cs/src/ResolveSpec.csproj new file mode 100644 index 0000000..bcbcd46 --- /dev/null +++ b/clients/resolvespec-cs/src/ResolveSpec.csproj @@ -0,0 +1,11 @@ + + + net8.0 + enable + enable + ResolveSpec + ResolveSpec.Client + 0.1.0 + Client for ResolveSpec (JSON body) and FunctionSpec endpoints + + diff --git a/clients/resolvespec-cs/src/ResolveSpecClient.cs b/clients/resolvespec-cs/src/ResolveSpecClient.cs new file mode 100644 index 0000000..c57e891 --- /dev/null +++ b/clients/resolvespec-cs/src/ResolveSpecClient.cs @@ -0,0 +1,65 @@ +using System.Text.Json; +using System.Text.Json.Serialization; + +namespace ResolveSpec; + +/// Client for the ResolveSpec JSON body protocol: POST {operation, data, options}. +public sealed class ResolveSpecClient +{ + static readonly JsonSerializerOptions Json = new() { DefaultIgnoreCondition = JsonIgnoreCondition.WhenWritingNull }; + readonly Transport _t; + + public ResolveSpecClient(string baseUrl, ClientOptions? options = null) => _t = new Transport(baseUrl, options); + + sealed class Request + { + [JsonPropertyName("operation")] public string Operation { get; set; } = ""; + [JsonPropertyName("id")] public string[]? Id { get; set; } + [JsonPropertyName("data")] public object? Data { get; set; } + [JsonPropertyName("options")] public Options? Options { get; set; } + } + + // A single id (int/long/string) goes in the URL; string[] / IEnumerable goes in the body. + static string? UrlId(object? id) => id switch + { + null => null, + string s => s, + IEnumerable => null, + _ => Convert.ToString(id, System.Globalization.CultureInfo.InvariantCulture), + }; + + static string[]? BodyId(object? id) => id is IEnumerable e and not string ? e.ToArray() : null; + + string Url(string schema, string entity, string? id) + { + var u = $"{_t.BaseUrl}/{Transport.Segment(schema)}/{Transport.Segment(entity)}"; + return string.IsNullOrEmpty(id) ? u : $"{u}/{Transport.Segment(id)}"; + } + + async Task SendAsync(HttpMethod method, string url, Request? body, CancellationToken ct) + { + var json = body == null ? null : JsonSerializer.Serialize(body, Json); + var (resp, text) = await _t.SendAsync(method, url, json, null, ct).ConfigureAwait(false); + var status = (int)resp.StatusCode; + if (!resp.IsSuccessStatusCode) throw Transport.ErrorFrom(status, text, resp.ReasonPhrase); + var r = JsonSerializer.Deserialize(text, Json) ?? new Response(); + if (!r.Success && r.Error != null) throw new ResolveSpecException(r.Error.Message, status, r.Error); + return r; + } + + /// GET /{schema}/{entity} + public Task GetMetadataAsync(string schema, string entity, CancellationToken ct = default) => + SendAsync(HttpMethod.Get, Url(schema, entity, null), null, ct); + + public Task ReadAsync(string schema, string entity, object? id = null, Options? options = null, CancellationToken ct = default) => + SendAsync(HttpMethod.Post, Url(schema, entity, UrlId(id)), new Request { Operation = "read", Id = BodyId(id), Options = options }, ct); + + public Task CreateAsync(string schema, string entity, object data, Options? options = null, CancellationToken ct = default) => + SendAsync(HttpMethod.Post, Url(schema, entity, null), new Request { Operation = "create", Data = data, Options = options }, ct); + + public Task UpdateAsync(string schema, string entity, object data, object? id = null, Options? options = null, CancellationToken ct = default) => + SendAsync(HttpMethod.Post, Url(schema, entity, UrlId(id)), new Request { Operation = "update", Id = BodyId(id), Data = data, Options = options }, ct); + + public Task DeleteAsync(string schema, string entity, object id, CancellationToken ct = default) => + SendAsync(HttpMethod.Post, Url(schema, entity, UrlId(id)), new Request { Operation = "delete" }, ct); +} diff --git a/clients/resolvespec-cs/src/Types.cs b/clients/resolvespec-cs/src/Types.cs new file mode 100644 index 0000000..5ab7fd1 --- /dev/null +++ b/clients/resolvespec-cs/src/Types.cs @@ -0,0 +1,137 @@ +using System.Text.Json; +using System.Text.Json.Serialization; + +namespace ResolveSpec; + +// Types aligned with Go pkg/common/types.go. JsonPropertyName values are the wire names. + +public sealed class FilterOption +{ + [JsonPropertyName("column")] public string Column { get; set; } = ""; + /// eq neq gt gte lt lte like ilike in contains startswith endswith between between_inclusive is_null is_not_null + [JsonPropertyName("operator")] public string Operator { get; set; } = "eq"; + [JsonPropertyName("value")] public object? Value { get; set; } + /// AND | OR + [JsonPropertyName("logic_operator")] public string? LogicOperator { get; set; } +} + +public sealed class SortOption +{ + [JsonPropertyName("column")] public string Column { get; set; } = ""; + /// asc | desc + [JsonPropertyName("direction")] public string Direction { get; set; } = "asc"; +} + +public sealed class Parameter +{ + [JsonPropertyName("name")] public string Name { get; set; } = ""; + [JsonPropertyName("value")] public string Value { get; set; } = ""; + [JsonPropertyName("sequence")] public int? Sequence { get; set; } +} + +public sealed class CustomOperator +{ + [JsonPropertyName("name")] public string Name { get; set; } = ""; + [JsonPropertyName("sql")] public string Sql { get; set; } = ""; +} + +public sealed class ComputedColumn +{ + [JsonPropertyName("name")] public string Name { get; set; } = ""; + [JsonPropertyName("expression")] public string Expression { get; set; } = ""; +} + +public sealed class PreloadOption +{ + [JsonPropertyName("relation")] public string? Relation { get; set; } + [JsonPropertyName("table_name")] public string? TableName { get; set; } + [JsonPropertyName("columns")] public List? Columns { get; set; } + [JsonPropertyName("omit_columns")] public List? OmitColumns { get; set; } + [JsonPropertyName("sort")] public List? Sort { get; set; } + [JsonPropertyName("filters")] public List? Filters { get; set; } + [JsonPropertyName("where")] public string? Where { get; set; } + [JsonPropertyName("limit")] public int? Limit { get; set; } + [JsonPropertyName("offset")] public int? Offset { get; set; } + [JsonPropertyName("updateable")] public bool? Updateable { get; set; } + [JsonPropertyName("computed_ql")] public Dictionary? ComputedQl { get; set; } + [JsonPropertyName("recursive")] public bool? Recursive { get; set; } + [JsonPropertyName("primary_key")] public string? PrimaryKey { get; set; } + [JsonPropertyName("related_key")] public string? RelatedKey { get; set; } + [JsonPropertyName("foreign_key")] public string? ForeignKey { get; set; } + [JsonPropertyName("recursive_child_key")] public string? RecursiveChildKey { get; set; } + [JsonPropertyName("sql_joins")] public List? SqlJoins { get; set; } + [JsonPropertyName("join_aliases")] public List? JoinAliases { get; set; } +} + +public sealed class VectorSearchOption +{ + [JsonPropertyName("column")] public string Column { get; set; } = ""; + [JsonPropertyName("vector")] public List Vector { get; set; } = new(); + /// l2 (default) | cosine | ip + [JsonPropertyName("metric")] public string? Metric { get; set; } + /// Distance column alias, default _distance. + [JsonPropertyName("as")] public string? As { get; set; } + [JsonPropertyName("direction")] public string? Direction { get; set; } +} + +/// ResolveSpec request options object. +public sealed class Options +{ + [JsonPropertyName("preload")] public List? Preload { get; set; } + [JsonPropertyName("columns")] public List? Columns { get; set; } + [JsonPropertyName("omit_columns")] public List? OmitColumns { get; set; } + [JsonPropertyName("filters")] public List? Filters { get; set; } + [JsonPropertyName("sort")] public List? Sort { get; set; } + [JsonPropertyName("limit")] public int? Limit { get; set; } + [JsonPropertyName("offset")] public int? Offset { get; set; } + [JsonPropertyName("customOperators")] public List? CustomOperators { get; set; } + [JsonPropertyName("computedColumns")] public List? ComputedColumns { get; set; } + [JsonPropertyName("parameters")] public List? Parameters { get; set; } + [JsonPropertyName("cursor_forward")] public string? CursorForward { get; set; } + [JsonPropertyName("cursor_backward")] public string? CursorBackward { get; set; } + [JsonPropertyName("fetch_row_number")] public string? FetchRowNumber { get; set; } + [JsonPropertyName("vector_search")] public VectorSearchOption? VectorSearch { get; set; } +} + +public sealed class Metadata +{ + [JsonPropertyName("total")] public long Total { get; set; } + [JsonPropertyName("count")] public long Count { get; set; } + [JsonPropertyName("filtered")] public long Filtered { get; set; } + [JsonPropertyName("limit")] public long Limit { get; set; } + [JsonPropertyName("offset")] public long Offset { get; set; } +} + +public sealed class ApiError +{ + [JsonPropertyName("code")] public string Code { get; set; } = ""; + [JsonPropertyName("message")] public string Message { get; set; } = ""; + [JsonPropertyName("details")] public JsonElement? Details { get; set; } + /// Server-side reason (funcspec / restheadspec). + [JsonPropertyName("detail")] public string? Detail { get; set; } + [JsonPropertyName("sql")] public string? Sql { get; set; } +} + +/// ResolveSpec envelope. is raw JSON; use . +public sealed class Response +{ + [JsonPropertyName("success")] public bool Success { get; set; } + [JsonPropertyName("data")] public JsonElement Data { get; set; } + [JsonPropertyName("metadata")] public Metadata? Metadata { get; set; } + [JsonPropertyName("error")] public ApiError? Error { get; set; } + + public T? Decode() => Data.ValueKind == JsonValueKind.Undefined ? default : Data.Deserialize(); +} + +/// Thrown on a non-2xx response or an unsuccessful API result. +public sealed class ResolveSpecException : Exception +{ + public int StatusCode { get; } + public ApiError Error { get; } + + public ResolveSpecException(string message, int statusCode, ApiError? error = null) : base(message) + { + StatusCode = statusCode; + Error = error ?? new ApiError { Message = message }; + } +} diff --git a/clients/resolvespec-cs/tests/ClientTests.cs b/clients/resolvespec-cs/tests/ClientTests.cs new file mode 100644 index 0000000..4a88311 --- /dev/null +++ b/clients/resolvespec-cs/tests/ClientTests.cs @@ -0,0 +1,178 @@ +using System.Net; +using System.Text; +using System.Text.Json; +using ResolveSpec; +using Xunit; + +public class Stub : HttpMessageHandler +{ + public HttpRequestMessage? Request; + public string Body = ""; + readonly HttpStatusCode _status; + readonly string _json; + readonly Dictionary _headers; + + public Stub(HttpStatusCode status, string json, Dictionary? headers = null) + { + _status = status; _json = json; _headers = headers ?? new(); + } + + protected override async Task SendAsync(HttpRequestMessage request, CancellationToken ct) + { + Request = request; + Body = request.Content == null ? "" : await request.Content.ReadAsStringAsync(ct); + var r = new HttpResponseMessage(_status) { Content = new StringContent(_json, Encoding.UTF8, "application/json") }; + foreach (var (k, v) in _headers) r.Headers.TryAddWithoutValidation(k, v); + return r; + } +} + +public class ResolveSpecTests +{ + static (ResolveSpecClient, Stub) Make(HttpStatusCode s, string json) + { + var stub = new Stub(s, json); + var o = new ClientOptions { Token = "tok", HttpClient = new HttpClient(stub) }; + o.Headers["X-Tenant"] = "a"; + return (new ResolveSpecClient("http://localhost:3000/", o), stub); + } + + [Fact] + public async Task ReadPostsBody() + { + var (c, s) = Make(HttpStatusCode.OK, """{"success":true,"data":[{"id":1}]}"""); + var r = await c.ReadAsync("public", "users", null, new Options { Limit = 5, Filters = new() { new FilterOption { Column = "a", Operator = "eq", Value = 1 } } }); + Assert.Equal(HttpMethod.Post, s.Request!.Method); + Assert.Equal("/public/users", s.Request.RequestUri!.AbsolutePath); + Assert.Equal("Bearer tok", s.Request.Headers.Authorization!.ToString()); + Assert.Equal("a", s.Request.Headers.GetValues("X-Tenant").Single()); + using var body = JsonDocument.Parse(s.Body); + Assert.Equal("read", body.RootElement.GetProperty("operation").GetString()); + Assert.Equal(5, body.RootElement.GetProperty("options").GetProperty("limit").GetInt32()); + Assert.False(body.RootElement.TryGetProperty("id", out _)); + Assert.Single(r.Decode>>()!); + } + + [Fact] + public async Task IdPlacement() + { + var (c, s) = Make(HttpStatusCode.OK, """{"success":true,"data":{}}"""); + await c.ReadAsync("s", "e", 7); + Assert.Equal("/s/e/7", s.Request!.RequestUri!.AbsolutePath); + await c.UpdateAsync("s", "e", new { a = 1 }, new[] { "1", "2" }); + Assert.Equal("/s/e", s.Request!.RequestUri!.AbsolutePath); + using (var b = JsonDocument.Parse(s.Body)) + { + Assert.Equal(2, b.RootElement.GetProperty("id").GetArrayLength()); + Assert.Equal("update", b.RootElement.GetProperty("operation").GetString()); + } + await c.DeleteAsync("s", "e", "a/b"); + Assert.Equal("/s/e/a%2Fb", s.Request!.RequestUri!.AbsoluteUri[(s.Request.RequestUri.AbsoluteUri.IndexOf("/s/e", StringComparison.Ordinal))..]); + Assert.Contains("\"delete\"", s.Body); + } + + [Fact] + public async Task Errors() + { + var (c, _) = Make(HttpStatusCode.BadRequest, """{"success":false,"error":{"code":"x","message":"bad","detail":"why"}}"""); + var e = await Assert.ThrowsAsync(() => c.ReadAsync("s", "e")); + Assert.Equal((400, "x", "bad", "why"), (e.StatusCode, e.Error.Code, e.Message, e.Error.Detail)); + + var (c2, _) = Make(HttpStatusCode.BadGateway, "bad gateway"); + var e2 = await Assert.ThrowsAsync(() => c2.ReadAsync("s", "e")); + Assert.Equal((502, "bad gateway"), (e2.StatusCode, e2.Message)); + + var (c3, _) = Make(HttpStatusCode.OK, """{"success":false,"error":{"code":"c","message":"nope"}}"""); + var e3 = await Assert.ThrowsAsync(() => c3.ReadAsync("s", "e")); + Assert.Equal("nope", e3.Message); + } +} + +public class FuncSpecTests +{ + [Fact] + public void HeaderFilters() + { + var h = FuncSpecClient.BuildHeaders(new FuncSpecOptions + { + Filters = new() + { + new() { Column = "status", Operator = "eq", Value = "active" }, + new() { Column = "age", Operator = "gte", Value = 18 }, + new() { Column = "name", Operator = "contains", Value = "x", LogicOperator = "OR" }, + new() { Column = "deleted", Operator = "is_null" }, + new() { Column = "id", Operator = "in", Value = new[] { 1, 2 } }, + new() { Column = "p", Operator = "between_inclusive", Value = new[] { 1, 5 } }, + }, + }); + Assert.Equal(new Dictionary + { + ["X-FieldFilter-status"] = "active", + ["X-SearchOp-greaterthanorequal-age"] = "18", + ["X-SearchOr-contains-name"] = "x", + ["X-SearchOp-empty-deleted"] = "", + ["X-SearchOp-in-id"] = "1,2", + ["X-SearchOp-betweeninclusive-p"] = "1,5", + }, h); + } + + [Fact] + public void HeaderMiscAndEncoding() + { + var h = FuncSpecClient.BuildHeaders(new FuncSpecOptions + { + SearchFilters = new() { ["name"] = "bob" }, CustomSqlWhere = "a = 1", CustomSqlOr = "b = 2", + Sort = new() { new() { Column = "name", Direction = "asc" }, new() { Column = "created_at", Direction = "DESC" } }, + Limit = 5, Offset = 10, Distinct = true, SkipCount = true, SkipCache = false, ResponseFormat = "syncfusion", + }); + Assert.Equal("name ASC,created_at DESC", h["X-Sort"]); + Assert.Equal("bob", h["X-SearchFilter-name"]); + Assert.Equal("a = 1", h["X-Custom-SQL-W"]); + Assert.Equal("false", h["X-SkipCache"]); + Assert.Equal("true", h["X-Syncfusion"]); + + h = FuncSpecClient.BuildHeaders(new FuncSpecOptions { Filters = new() + { + new() { Column = "n", Operator = "eq", Value = "héllo" }, + new() { Column = "m", Operator = "eq", Value = " pad" }, + } }); + Assert.StartsWith("ZIP_", h["X-FieldFilter-n"]); + Assert.Equal("héllo", FuncSpecClient.DecodeHeaderValue(h["X-FieldFilter-n"])); + Assert.Equal(" pad", FuncSpecClient.DecodeHeaderValue(h["X-FieldFilter-m"])); + } + + [Fact] + public void QueryBuilding() + { + var q = FuncSpecClient.BuildQuery(new Dictionary { ["a"] = true, ["b"] = new[] { "x", "y" }, ["c"] = null, ["d"] = 3 }); + Assert.Equal(new[] { "a=true", "b=x", "b=y", "d=3" }, q.Select(kv => $"{kv.Key}={kv.Value}")); + } + + [Fact] + public async Task QueryListMetadata() + { + var stub = new Stub((HttpStatusCode)206, """[{"id":1},{"id":2}]""", new() { ["Content-Range"] = "items 10-12/50" }); + var c = new FuncSpecClient("http://x", new ClientOptions { Token = "tok", HttpClient = new HttpClient(stub) }); + var r = await c.QueryListAsync("/api/users", new Dictionary { ["org"] = 1 }, new FuncSpecOptions { Limit = 2 }); + Assert.Equal("GET", stub.Request!.Method.Method); + Assert.Equal("/api/users", stub.Request.RequestUri!.AbsolutePath); + Assert.Equal("?org=1", stub.Request.RequestUri.Query); + Assert.Equal("2", stub.Request.Headers.GetValues("X-Limit").Single()); + Assert.Equal((50L, 2L, 50L, 2L, 10L), (r.Metadata!.Total, r.Metadata.Count, r.Metadata.Filtered, r.Metadata.Limit, r.Metadata.Offset)); + Assert.Equal(2, r.Data.GetArrayLength()); + } + + [Fact] + public async Task QuerySingleAndError() + { + var ok = new FuncSpecClient("http://x", new ClientOptions { HttpClient = new HttpClient(new Stub(HttpStatusCode.OK, """{"id":1}""")) }); + var r = await ok.QueryAsync("api/u"); + Assert.Null(r.Metadata); + Assert.Equal(1, r.Data.GetProperty("id").GetInt32()); + + var bad = new FuncSpecClient("http://x", new ClientOptions { HttpClient = new HttpClient(new Stub(HttpStatusCode.BadRequest, + """{"success":false,"error":{"code":"hook_error","message":"Hook execution failed","detail":"authentication required"}}""")) }); + var e = await Assert.ThrowsAsync(() => bad.QueryAsync("api/u")); + Assert.Equal(("hook_error", "authentication required"), (e.Error.Code, e.Error.Detail)); + } +} diff --git a/clients/resolvespec-cs/tests/ResolveSpec.Tests.csproj b/clients/resolvespec-cs/tests/ResolveSpec.Tests.csproj new file mode 100644 index 0000000..c8a7373 --- /dev/null +++ b/clients/resolvespec-cs/tests/ResolveSpec.Tests.csproj @@ -0,0 +1,16 @@ + + + net8.0 + enable + enable + false + + + + + + + + + + diff --git a/clients/resolvespec-dart/.gitignore b/clients/resolvespec-dart/.gitignore new file mode 100644 index 0000000..315dbaf --- /dev/null +++ b/clients/resolvespec-dart/.gitignore @@ -0,0 +1,3 @@ +.dart_tool/ +pubspec.lock +build/ diff --git a/clients/resolvespec-dart/README.md b/clients/resolvespec-dart/README.md new file mode 100644 index 0000000..825f1af --- /dev/null +++ b/clients/resolvespec-dart/README.md @@ -0,0 +1,43 @@ +# resolvespec (Dart) + +Dart / Flutter client for ResolveSpec (JSON body) and FunctionSpec. Depends on `package:http`. Dart >= 3.3. + +> Not compiled or tested yet (no Dart SDK was available). Run `dart pub get && dart test` first. + +## Clients + +| Type | Constructor | Methods | +|---|---|---| +| `ResolveSpecClient` | `(baseUrl, [ClientOptions])` | `getMetadata` `read` `create` `update` `delete` `close` | +| `FuncSpecClient` | `(baseUrl, [ClientOptions])` | `query` `queryList` `close` | + +`ClientOptions(token:, headers:, timeout:, httpClient:)`. Precedence: Content-Type < custom headers < bearer token. + +## ResolveSpec + +- `id`: `int`/`String` → URL, `List` → body. Named args: `id:`, `options:`. +- `Options`, `FilterOption(column, operator, [value, logic])`, `SortOption(column, [direction])`. +- Result: `Response{success, data (decoded JSON), metadata}`. + +## FunctionSpec + +- Routes are server-defined: pass the `path`. +- `params:` map → query string (list → repeated keys, null skipped). +- `FuncSpecOptions` → `X-*` headers: `filters`, `searchFilters`, `customSqlWhere`, `customSqlOr`, `sort`, `limit`, `offset`, `distinct`, `skipCount`, `skipCache`, `responseFormat`. +- `queryList` fills `metadata` from `Content-Range`; 206 is success. +- Helpers: `buildHeaders`, `buildQuery`, `encodeHeaderValue`, `decodeHeaderValue`. + +## Server quirks + +- `sort` is raw SQL in ORDER BY (client sends `col ASC|DESC`). +- One search operator per column. +- Values starting `ZIP_` / `__` are base64-decoded by the server. +- Non-ASCII, control chars and edge spaces are auto-encoded (`ZIP_`). + +## Errors + +`ResolveSpecException{statusCode, message, error: ApiError{code, detail, sql}}`. + +## Test + +`dart test` diff --git a/clients/resolvespec-dart/analysis_options.yaml b/clients/resolvespec-dart/analysis_options.yaml new file mode 100644 index 0000000..572dd23 --- /dev/null +++ b/clients/resolvespec-dart/analysis_options.yaml @@ -0,0 +1 @@ +include: package:lints/recommended.yaml diff --git a/clients/resolvespec-dart/lib/resolvespec.dart b/clients/resolvespec-dart/lib/resolvespec.dart new file mode 100644 index 0000000..8381929 --- /dev/null +++ b/clients/resolvespec-dart/lib/resolvespec.dart @@ -0,0 +1,7 @@ +/// Client for ResolveSpec (JSON body) and FunctionSpec endpoints. +library; + +export 'src/client.dart' show ClientOptions, ResolveSpecException; +export 'src/funcspec.dart'; +export 'src/resolvespec.dart'; +export 'src/types.dart'; diff --git a/clients/resolvespec-dart/lib/src/client.dart b/clients/resolvespec-dart/lib/src/client.dart new file mode 100644 index 0000000..3c2baa8 --- /dev/null +++ b/clients/resolvespec-dart/lib/src/client.dart @@ -0,0 +1,84 @@ +import 'dart:convert'; + +import 'package:http/http.dart' as http; + +import 'types.dart'; + +/// Thrown on a non-2xx response or an unsuccessful API result. +class ResolveSpecException implements Exception { + final int statusCode; + final String message; + final ApiError error; + + ResolveSpecException(this.message, this.statusCode, [ApiError? error]) : error = error ?? ApiError(message: message); + + @override + String toString() => 'ResolveSpecException($statusCode): $message'; +} + +/// Shared HTTP configuration for both clients. +class ClientOptions { + final String? token; + final Map headers; + final Duration timeout; + + /// Supply your own client (tests, pooling). + final http.Client? httpClient; + + const ClientOptions({this.token, this.headers = const {}, this.timeout = const Duration(seconds: 30), this.httpClient}); +} + +class Transport { + final String baseUrl; + final ClientOptions options; + final http.Client _http; + + Transport(String baseUrl, ClientOptions? options) + : baseUrl = baseUrl.replaceAll(RegExp(r'/+$'), ''), + options = options ?? const ClientOptions(), + _http = options?.httpClient ?? http.Client(); + + /// Content-Type < custom headers < per-call headers < bearer token. + Future send(String method, Uri uri, {String? body, Map? extra}) { + final headers = {'Content-Type': 'application/json'}; + void merge(Map src) { + for (final e in src.entries) { + headers.removeWhere((k, _) => k.toLowerCase() == e.key.toLowerCase()); + headers[e.key] = e.value; + } + } + + merge(options.headers); + if (extra != null) merge(extra); + final token = options.token; + if (token != null && token.isNotEmpty) merge({'Authorization': 'Bearer $token'}); + + final req = http.Request(method, uri)..headers.addAll(headers); + if (body != null) req.body = body; + return _http.send(req).timeout(options.timeout).then(http.Response.fromStream); + } + + void close() => _http.close(); + + static ResolveSpecException errorFrom(http.Response resp) { + final body = utf8.decode(resp.bodyBytes, allowMalformed: true); + ApiError? err; + var isJson = false; + try { + final parsed = jsonDecode(body); + isJson = true; + if (parsed is Map && parsed['error'] is Map) { + err = ApiError.fromJson(parsed['error'] as Map); + } + } on FormatException { + // not JSON + } + var message = err?.message ?? ''; + if (message.isEmpty) { + var text = isJson ? '' : body.trim(); + if (text.length > 200) text = text.substring(0, 200); + message = text.isNotEmpty ? text : '${resp.reasonPhrase ?? 'Error'} (${resp.statusCode})'; + } + return ResolveSpecException(message, resp.statusCode, err); + } +} diff --git a/clients/resolvespec-dart/lib/src/funcspec.dart b/clients/resolvespec-dart/lib/src/funcspec.dart new file mode 100644 index 0000000..49e9226 --- /dev/null +++ b/clients/resolvespec-dart/lib/src/funcspec.dart @@ -0,0 +1,186 @@ +import 'dart:convert'; + +import 'client.dart'; +import 'types.dart'; + +/// Options sent to funcspec endpoints as X-* headers. +/// +/// Server behaviour (pkg/funcspec): [sort] is inserted raw into ORDER BY (so it is sent as SQL +/// terms); only one search operator per column is kept; values starting with `ZIP_` or `__` +/// are base64-decoded by the server, so such plaintext values cannot be sent faithfully. +class FuncSpecOptions { + /// eq+AND -> X-FieldFilter; others X-SearchOp / X-SearchOr. + final List? filters; + + /// X-SearchFilter-{col}: text ILIKE. + final Map? searchFilters; + final String? customSqlWhere; + final String? customSqlOr; + final List? sort; + final int? limit; + final int? offset; + final bool? distinct; + final bool? skipCount; + final bool? skipCache; + + /// simple | detail | syncfusion + final String? responseFormat; + + const FuncSpecOptions({ + this.filters, + this.searchFilters, + this.customSqlWhere, + this.customSqlOr, + this.sort, + this.limit, + this.offset, + this.distinct, + this.skipCount, + this.skipCache, + this.responseFormat, + }); +} + +const _operatorMap = { + 'eq': 'equals', + 'neq': 'notequals', + 'gt': 'greaterthan', + 'gte': 'greaterthanorequal', + 'lt': 'lessthan', + 'lte': 'lessthanorequal', + 'like': 'contains', + 'ilike': 'contains', + 'contains': 'contains', + 'startswith': 'beginswith', + 'endswith': 'endswith', + 'in': 'in', + 'between': 'between', + 'between_inclusive': 'betweeninclusive', + 'is_null': 'empty', + 'is_not_null': 'notempty', +}; + +String _scalar(Object? v) => v == null ? '' : v.toString(); + +String _filterValue(Object? v) => v is Iterable ? v.map(_scalar).join(',') : _scalar(v); + +/// Base64 (UTF-8) with the `ZIP_` prefix. +String encodeHeaderValue(String v) => 'ZIP_${base64.encode(utf8.encode(v))}'; + +/// Decode a value that may carry a `ZIP_` or `__` prefix (nested allowed). +String decodeHeaderValue(String v) { + for (final p in const ['ZIP_', '__']) { + if (v.startsWith(p)) { + var b64 = v.substring(p.length).replaceAll(RegExp(r'[\n\r ]'), ''); + b64 = b64.padRight(b64.length + (4 - b64.length % 4) % 4, '='); + try { + return decodeHeaderValue(utf8.decode(base64.decode(b64))); + } on FormatException { + return v; + } + } + } + return v; +} + +/// Encode values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces). +String _safe(String v) { + final unsafe = v != v.trim() || v.runes.any((c) => c > 127 || c < 32 || c == 127); + return unsafe ? encodeHeaderValue(v) : v; +} + +/// Build the X-* headers understood by funcspec.ParseParameters. +Map buildHeaders(FuncSpecOptions? o) { + final h = {}; + if (o == null) return h; + + for (final f in o.filters ?? const []) { + final logic = f.logicOperator ?? 'AND'; + final v = _safe(_filterValue(f.value)); + if (f.operator == 'eq' && logic == 'AND') { + h['X-FieldFilter-${f.column}'] = v; + } else { + final kind = logic == 'OR' ? 'X-SearchOr' : 'X-SearchOp'; + h['$kind-${_operatorMap[f.operator] ?? f.operator}-${f.column}'] = v; + } + } + o.searchFilters?.forEach((col, text) => h['X-SearchFilter-$col'] = _safe(text)); + if (o.customSqlWhere != null && o.customSqlWhere!.isNotEmpty) h['X-Custom-SQL-W'] = _safe(o.customSqlWhere!); + if (o.customSqlOr != null && o.customSqlOr!.isNotEmpty) h['X-Custom-SQL-Or'] = _safe(o.customSqlOr!); + if (o.sort != null && o.sort!.isNotEmpty) { + // funcspec puts this verbatim into ORDER BY + h['X-Sort'] = _safe(o.sort!.map((s) => '${s.column} ${s.direction.toLowerCase() == 'desc' ? 'DESC' : 'ASC'}').join(',')); + } + if (o.limit != null) h['X-Limit'] = '${o.limit}'; + if (o.offset != null) h['X-Offset'] = '${o.offset}'; + if (o.distinct != null) h['X-Distinct'] = '${o.distinct}'; + if (o.skipCount != null) h['X-SkipCount'] = '${o.skipCount}'; + if (o.skipCache != null) h['X-SkipCache'] = '${o.skipCache}'; + switch (o.responseFormat) { + case 'simple': + h['X-SimpleApi'] = 'true'; + case 'detail': + h['X-DetailApi'] = 'true'; + case 'syncfusion': + h['X-Syncfusion'] = 'true'; + } + return h; +} + +/// Build query-string pairs: lists -> repeated keys, null skipped, bools -> true/false. +Map> buildQuery(Map? params) { + final out = >{}; + params?.forEach((k, v) { + if (v == null) return; + out[k] = v is Iterable ? v.map((e) => _safe(_scalar(e))).toList() : [_safe(_scalar(v))]; + }); + return out; +} + +final _contentRange = RegExp(r'(\d+)-(\d+)/(\d+)'); + +Metadata _metadata(String? contentRange, FuncSpecOptions? o) { + final m = _contentRange.firstMatch(contentRange ?? ''); + if (m == null) return Metadata(limit: o?.limit ?? 0); + final start = int.parse(m.group(1)!); + final end = int.parse(m.group(2)!); + final total = int.parse(m.group(3)!); + return Metadata(total: total, count: end - start, filtered: total, limit: o?.limit ?? 0, offset: start); +} + +/// Client for user-defined SQL endpoints. Routes are defined by the server application. +class FuncSpecClient { + final Transport _t; + + FuncSpecClient(String baseUrl, [ClientOptions? options]) : _t = Transport(baseUrl, options); + + void close() => _t.close(); + + Future _call(String method, String path, Map? params, FuncSpecOptions? o, bool list) async { + final base = Uri.parse('${_t.baseUrl}/${path.replaceAll(RegExp(r'^/+'), '')}'); + final pairs = []; + buildQuery(params).forEach((k, vs) { + for (final v in vs) { + pairs.add('${Uri.encodeQueryComponent(k)}=${Uri.encodeQueryComponent(v)}'); + } + }); + final uri = pairs.isEmpty ? base : base.replace(query: pairs.join('&')); + + final resp = await _t.send(method, uri, extra: buildHeaders(o)); + if (resp.statusCode < 200 || resp.statusCode > 299) throw Transport.errorFrom(resp); // 206 is success + final text = utf8.decode(resp.bodyBytes); + return Response( + success: true, + data: text.trim().isEmpty ? null : jsonDecode(text), + metadata: list ? _metadata(resp.headers['content-range'], o) : null, + ); + } + + /// Single-record endpoint (SqlQuery). `data` is the row object. + Future query(String path, {Map? params, FuncSpecOptions? options, String method = 'GET'}) => + _call(method.toUpperCase(), path, params, options, false); + + /// List endpoint (SqlQueryList). Metadata comes from Content-Range. + Future queryList(String path, {Map? params, FuncSpecOptions? options, String method = 'GET'}) => + _call(method.toUpperCase(), path, params, options, true); +} diff --git a/clients/resolvespec-dart/lib/src/resolvespec.dart b/clients/resolvespec-dart/lib/src/resolvespec.dart new file mode 100644 index 0000000..1ec5cf8 --- /dev/null +++ b/clients/resolvespec-dart/lib/src/resolvespec.dart @@ -0,0 +1,63 @@ +import 'dart:convert'; + +import 'client.dart'; +import 'types.dart'; + +/// Client for the ResolveSpec JSON body protocol: POST {operation, data, options}. +/// +/// A record `id` of type `int` or `String` goes in the URL; a `List` goes in the body. +class ResolveSpecClient { + final Transport _t; + + ResolveSpecClient(String baseUrl, [ClientOptions? options]) : _t = Transport(baseUrl, options); + + void close() => _t.close(); + + static String? _urlId(Object? id) => id == null || id is List ? null : id.toString(); + + static List? _bodyId(Object? id) => id is List ? id.map((e) => e.toString()).toList() : null; + + Uri _url(String schema, String entity, String? id) { + var u = '${_t.baseUrl}/${Uri.encodeComponent(schema)}/${Uri.encodeComponent(entity)}'; + if (id != null && id.isNotEmpty) u += '/${Uri.encodeComponent(id)}'; + return Uri.parse(u); + } + + Future _send(String method, Uri url, Map? body) async { + final resp = await _t.send(method, url, body: body == null ? null : jsonEncode(body)); + if (resp.statusCode < 200 || resp.statusCode > 299) throw Transport.errorFrom(resp); + final decoded = jsonDecode(utf8.decode(resp.bodyBytes)); + final r = Response.fromJson(decoded as Map); + if (!r.success && r.error != null) throw ResolveSpecException(r.error!.message, resp.statusCode, r.error); + return r; + } + + /// GET /{schema}/{entity} + Future getMetadata(String schema, String entity) => _send('GET', _url(schema, entity, null), null); + + Future read(String schema, String entity, {Object? id, Options? options}) => _send( + 'POST', + _url(schema, entity, _urlId(id)), + {'operation': 'read', if (_bodyId(id) != null) 'id': _bodyId(id), if (options != null) 'options': options.toJson()}, + ); + + Future create(String schema, String entity, Object data, {Options? options}) => _send( + 'POST', + _url(schema, entity, null), + {'operation': 'create', 'data': data, if (options != null) 'options': options.toJson()}, + ); + + Future update(String schema, String entity, Object data, {Object? id, Options? options}) => _send( + 'POST', + _url(schema, entity, _urlId(id)), + { + 'operation': 'update', + if (_bodyId(id) != null) 'id': _bodyId(id), + 'data': data, + if (options != null) 'options': options.toJson(), + }, + ); + + Future delete(String schema, String entity, Object id) => + _send('POST', _url(schema, entity, _urlId(id)), {'operation': 'delete'}); +} diff --git a/clients/resolvespec-dart/lib/src/types.dart b/clients/resolvespec-dart/lib/src/types.dart new file mode 100644 index 0000000..f56d613 --- /dev/null +++ b/clients/resolvespec-dart/lib/src/types.dart @@ -0,0 +1,268 @@ +// Types aligned with Go pkg/common/types.go. toJson() emits the wire names. + +Map _compact(Map m) { + m.removeWhere((_, v) => v == null); + return m; +} + +class FilterOption { + final String column; + + /// eq neq gt gte lt lte like ilike in contains startswith endswith between + /// between_inclusive is_null is_not_null + final String operator; + final Object? value; + + /// AND | OR + final String? logicOperator; + + const FilterOption(this.column, this.operator, [this.value, this.logicOperator]); + + Map toJson() => _compact({ + 'column': column, + 'operator': operator, + 'value': value, + 'logic_operator': logicOperator, + }); +} + +class SortOption { + final String column; + + /// asc | desc + final String direction; + + const SortOption(this.column, [this.direction = 'asc']); + + Map toJson() => {'column': column, 'direction': direction}; +} + +class Parameter { + final String name; + final String value; + final int? sequence; + + const Parameter(this.name, this.value, [this.sequence]); + + Map toJson() => _compact({'name': name, 'value': value, 'sequence': sequence}); +} + +class CustomOperator { + final String name; + final String sql; + + const CustomOperator(this.name, this.sql); + + Map toJson() => {'name': name, 'sql': sql}; +} + +class ComputedColumn { + final String name; + final String expression; + + const ComputedColumn(this.name, this.expression); + + Map toJson() => {'name': name, 'expression': expression}; +} + +class PreloadOption { + final String? relation; + final String? tableName; + final List? columns; + final List? omitColumns; + final List? sort; + final List? filters; + final String? where; + final int? limit; + final int? offset; + final bool? updateable; + final Map? computedQl; + final bool? recursive; + final String? primaryKey; + final String? relatedKey; + final String? foreignKey; + final String? recursiveChildKey; + final List? sqlJoins; + final List? joinAliases; + + const PreloadOption({ + this.relation, + this.tableName, + this.columns, + this.omitColumns, + this.sort, + this.filters, + this.where, + this.limit, + this.offset, + this.updateable, + this.computedQl, + this.recursive, + this.primaryKey, + this.relatedKey, + this.foreignKey, + this.recursiveChildKey, + this.sqlJoins, + this.joinAliases, + }); + + Map toJson() => _compact({ + 'relation': relation, + 'table_name': tableName, + 'columns': columns, + 'omit_columns': omitColumns, + 'sort': sort?.map((e) => e.toJson()).toList(), + 'filters': filters?.map((e) => e.toJson()).toList(), + 'where': where, + 'limit': limit, + 'offset': offset, + 'updateable': updateable, + 'computed_ql': computedQl, + 'recursive': recursive, + 'primary_key': primaryKey, + 'related_key': relatedKey, + 'foreign_key': foreignKey, + 'recursive_child_key': recursiveChildKey, + 'sql_joins': sqlJoins, + 'join_aliases': joinAliases, + }); +} + +class VectorSearchOption { + final String column; + final List vector; + + /// l2 (default) | cosine | ip + final String? metric; + + /// Distance column alias, default _distance. + final String? as; + final String? direction; + + const VectorSearchOption(this.column, this.vector, {this.metric, this.as, this.direction}); + + Map toJson() => + _compact({'column': column, 'vector': vector, 'metric': metric, 'as': as, 'direction': direction}); +} + +/// ResolveSpec request options object. +class Options { + final List? preload; + final List? columns; + final List? omitColumns; + final List? filters; + final List? sort; + final int? limit; + final int? offset; + final List? customOperators; + final List? computedColumns; + final List? parameters; + final String? cursorForward; + final String? cursorBackward; + final String? fetchRowNumber; + final VectorSearchOption? vectorSearch; + + const Options({ + this.preload, + this.columns, + this.omitColumns, + this.filters, + this.sort, + this.limit, + this.offset, + this.customOperators, + this.computedColumns, + this.parameters, + this.cursorForward, + this.cursorBackward, + this.fetchRowNumber, + this.vectorSearch, + }); + + Map toJson() => _compact({ + 'preload': preload?.map((e) => e.toJson()).toList(), + 'columns': columns, + 'omit_columns': omitColumns, + 'filters': filters?.map((e) => e.toJson()).toList(), + 'sort': sort?.map((e) => e.toJson()).toList(), + 'limit': limit, + 'offset': offset, + 'customOperators': customOperators?.map((e) => e.toJson()).toList(), + 'computedColumns': computedColumns?.map((e) => e.toJson()).toList(), + 'parameters': parameters?.map((e) => e.toJson()).toList(), + 'cursor_forward': cursorForward, + 'cursor_backward': cursorBackward, + 'fetch_row_number': fetchRowNumber, + 'vector_search': vectorSearch?.toJson(), + }); +} + +class Metadata { + final int total; + final int count; + final int filtered; + final int limit; + final int offset; + + const Metadata({this.total = 0, this.count = 0, this.filtered = 0, this.limit = 0, this.offset = 0}); + + factory Metadata.fromJson(Map j) => Metadata( + total: (j['total'] as num?)?.toInt() ?? 0, + count: (j['count'] as num?)?.toInt() ?? 0, + filtered: (j['filtered'] as num?)?.toInt() ?? 0, + limit: (j['limit'] as num?)?.toInt() ?? 0, + offset: (j['offset'] as num?)?.toInt() ?? 0, + ); + + @override + bool operator ==(Object other) => + other is Metadata && + other.total == total && + other.count == count && + other.filtered == filtered && + other.limit == limit && + other.offset == offset; + + @override + int get hashCode => Object.hash(total, count, filtered, limit, offset); + + @override + String toString() => 'Metadata(total: $total, count: $count, filtered: $filtered, limit: $limit, offset: $offset)'; +} + +class ApiError { + final String code; + final String message; + final Object? details; + + /// Server-side reason (funcspec / restheadspec). + final String? detail; + final String? sql; + + const ApiError({this.code = '', this.message = '', this.details, this.detail, this.sql}); + + factory ApiError.fromJson(Map j) => ApiError( + code: (j['code'] as String?) ?? '', + message: (j['message'] as String?) ?? '', + details: j['details'], + detail: j['detail'] as String?, + sql: j['sql'] as String?, + ); +} + +/// ResolveSpec envelope. [data] is the decoded JSON value (Map, List or scalar). +class Response { + final bool success; + final Object? data; + final Metadata? metadata; + final ApiError? error; + + const Response({required this.success, this.data, this.metadata, this.error}); + + factory Response.fromJson(Map j) => Response( + success: j['success'] == true, + data: j['data'], + metadata: j['metadata'] is Map ? Metadata.fromJson(j['metadata'] as Map) : null, + error: j['error'] is Map ? ApiError.fromJson(j['error'] as Map) : null, + ); +} diff --git a/clients/resolvespec-dart/pubspec.yaml b/clients/resolvespec-dart/pubspec.yaml new file mode 100644 index 0000000..41989ba --- /dev/null +++ b/clients/resolvespec-dart/pubspec.yaml @@ -0,0 +1,14 @@ +name: resolvespec +description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints. +version: 0.1.0 +publish_to: none + +environment: + sdk: ">=3.3.0 <4.0.0" + +dependencies: + http: ^1.2.0 + +dev_dependencies: + lints: ^4.0.0 + test: ^1.25.0 diff --git a/clients/resolvespec-dart/test/client_test.dart b/clients/resolvespec-dart/test/client_test.dart new file mode 100644 index 0000000..7e213a1 --- /dev/null +++ b/clients/resolvespec-dart/test/client_test.dart @@ -0,0 +1,179 @@ +import 'dart:convert'; + +import 'package:http/http.dart' as http; +import 'package:http/testing.dart'; +import 'package:resolvespec/resolvespec.dart'; +import 'package:test/test.dart'; + +(http.Client, List) stub(int status, Object body, {Map headers = const {}}) { + final seen = []; + final client = MockClient((req) async { + seen.add(req); + final text = body is String ? body : jsonEncode(body); + return http.Response(text, status, headers: {'content-type': 'application/json', ...headers}); + }); + return (client, seen); +} + +void main() { + group('resolvespec', () { + test('read posts body with headers', () async { + final (c, seen) = stub(200, {'success': true, 'data': [{'id': 1}]}); + final client = ResolveSpecClient( + 'http://localhost:3000/', + ClientOptions(token: 'tok', headers: {'X-Tenant': 'a'}, httpClient: c), + ); + final r = await client.read('public', 'users', + options: const Options(limit: 5, filters: [FilterOption('a', 'eq', 1)])); + final req = seen.single; + expect(req.method, 'POST'); + expect(req.url.path, '/public/users'); + expect(req.headers['authorization'], 'Bearer tok'); + expect(req.headers['x-tenant'], 'a'); + final body = jsonDecode(req.body) as Map; + expect(body['operation'], 'read'); + expect(body['options']['limit'], 5); + expect(body.containsKey('id'), isFalse); + expect((r.data as List).length, 1); + }); + + test('id placement', () async { + final (c, seen) = stub(200, {'success': true, 'data': {}}); + final client = ResolveSpecClient('http://x', ClientOptions(httpClient: c)); + await client.read('s', 'e', id: 7); + expect(seen.last.url.path, '/s/e/7'); + await client.update('s', 'e', {'a': 1}, id: ['1', '2']); + expect(seen.last.url.path, '/s/e'); + final b = jsonDecode(seen.last.body) as Map; + expect(b['id'], ['1', '2']); + expect(b['operation'], 'update'); + await client.delete('s', 'e', 'a/b'); + expect(seen.last.url.toString(), 'http://x/s/e/a%2Fb'); + expect(jsonDecode(seen.last.body)['operation'], 'delete'); + }); + + test('errors', () async { + final client = ResolveSpecClient( + 'http://x', + ClientOptions( + httpClient: stub(400, { + 'success': false, + 'error': {'code': 'x', 'message': 'bad', 'detail': 'why'} + }).$1), + ); + await expectLater( + client.read('s', 'e'), + throwsA(isA() + .having((e) => e.statusCode, 'status', 400) + .having((e) => e.error.code, 'code', 'x') + .having((e) => e.message, 'message', 'bad') + .having((e) => e.error.detail, 'detail', 'why')), + ); + final plain = ResolveSpecClient('http://x', ClientOptions(httpClient: stub(502, 'bad gateway').$1)); + await expectLater( + plain.read('s', 'e'), + throwsA(isA().having((e) => e.message, 'message', 'bad gateway')), + ); + final soft = ResolveSpecClient( + 'http://x', + ClientOptions( + httpClient: stub(200, { + 'success': false, + 'error': {'code': 'c', 'message': 'nope'} + }).$1), + ); + await expectLater(soft.read('s', 'e'), throwsA(isA().having((e) => e.message, 'message', 'nope'))); + }); + }); + + group('funcspec', () { + test('header filters', () { + final h = buildHeaders(const FuncSpecOptions(filters: [ + FilterOption('status', 'eq', 'active'), + FilterOption('age', 'gte', 18), + FilterOption('name', 'contains', 'x', 'OR'), + FilterOption('deleted', 'is_null'), + FilterOption('id', 'in', [1, 2]), + FilterOption('p', 'between_inclusive', [1, 5]), + ])); + expect(h, { + 'X-FieldFilter-status': 'active', + 'X-SearchOp-greaterthanorequal-age': '18', + 'X-SearchOr-contains-name': 'x', + 'X-SearchOp-empty-deleted': '', + 'X-SearchOp-in-id': '1,2', + 'X-SearchOp-betweeninclusive-p': '1,5', + }); + }); + + test('misc headers and encoding', () { + var h = buildHeaders(const FuncSpecOptions( + searchFilters: {'name': 'bob'}, + customSqlWhere: 'a = 1', + customSqlOr: 'b = 2', + sort: [SortOption('name', 'asc'), SortOption('created_at', 'DESC')], + limit: 5, + offset: 10, + distinct: true, + skipCount: true, + skipCache: false, + responseFormat: 'syncfusion', + )); + expect(h['X-Sort'], 'name ASC,created_at DESC'); + expect(h['X-SearchFilter-name'], 'bob'); + expect(h['X-Custom-SQL-W'], 'a = 1'); + expect(h['X-SkipCache'], 'false'); + expect(h['X-Syncfusion'], 'true'); + + h = buildHeaders(const FuncSpecOptions(filters: [ + FilterOption('n', 'eq', 'héllo'), + FilterOption('m', 'eq', ' pad'), + ])); + expect(h['X-FieldFilter-n'], startsWith('ZIP_')); + expect(decodeHeaderValue(h['X-FieldFilter-n']!), 'héllo'); + expect(decodeHeaderValue(h['X-FieldFilter-m']!), ' pad'); + }); + + test('query building', () { + expect(buildQuery({'a': true, 'b': ['x', 'y'], 'c': null, 'd': 3}), { + 'a': ['true'], + 'b': ['x', 'y'], + 'd': ['3'], + }); + }); + + test('queryList metadata', () async { + final (c, seen) = stub(206, [{'id': 1}, {'id': 2}], headers: {'Content-Range': 'items 10-12/50'}); + final client = FuncSpecClient('http://x', ClientOptions(token: 'tok', httpClient: c)); + final r = await client.queryList('/api/users', params: {'org': 1}, options: const FuncSpecOptions(limit: 2)); + expect(seen.single.method, 'GET'); + expect(seen.single.url.path, '/api/users'); + expect(seen.single.url.query, 'org=1'); + expect(seen.single.headers['x-limit'], '2'); + expect(r.metadata, const Metadata(total: 50, count: 2, filtered: 50, limit: 2, offset: 10)); + expect((r.data as List).length, 2); + }); + + test('query single and error', () async { + final ok = FuncSpecClient('http://x', ClientOptions(httpClient: stub(200, {'id': 1}).$1)); + final r = await ok.query('api/u'); + expect(r.metadata, isNull); + expect((r.data as Map)['id'], 1); + + final bad = FuncSpecClient( + 'http://x', + ClientOptions( + httpClient: stub(400, { + 'success': false, + 'error': {'code': 'hook_error', 'message': 'Hook execution failed', 'detail': 'authentication required'} + }).$1), + ); + await expectLater( + bad.query('api/u'), + throwsA(isA() + .having((e) => e.error.code, 'code', 'hook_error') + .having((e) => e.error.detail, 'detail', 'authentication required')), + ); + }); + }); +} diff --git a/clients/resolvespec-go/README.md b/clients/resolvespec-go/README.md new file mode 100644 index 0000000..f05effe --- /dev/null +++ b/clients/resolvespec-go/README.md @@ -0,0 +1,40 @@ +# resolvespec-go + +Go client for ResolveSpec (JSON body) and FunctionSpec. Module: `github.com/bitechdev/ResolveSpec/clients/resolvespec-go`. Stdlib only. + +## Clients + +| Type | Constructor | Methods | +|---|---|---| +| `Client` | `NewClient(baseURL, opts...)` | `GetMetadata` `Read` `Create` `Update` `Delete` | +| `FuncSpecClient` | `NewFuncSpecClient(baseURL, opts...)` | `Query` `QueryList` `Do` | + +Client options: `WithToken`, `WithHeader`, `WithHTTPClient`. Precedence: Content-Type < custom headers < bearer token. + +## ResolveSpec + +- All methods take `ctx`; `Read`/`Update`/`Delete` take `RecordID` (`nil`, int/string → URL, `[]string` → body). +- `Options` fields use pointers for optional ints/bools (`Int(n)`, `Bool(b)`). +- Result: `*Response{Success, Data (raw JSON), Metadata}`; `resp.Decode(&v)`. + +## FunctionSpec + +- Routes are server-defined: pass the `path`. +- `Params` → query string (slice → repeated keys, bool → `true`/`false`). +- `FuncSpecOptions` → `X-*` headers: `Filters`, `SearchFilters`, `CustomSQLWhere`, `CustomSQLOr`, `Sort`, `Limit`, `Offset`, `Distinct`, `SkipCount`, `SkipCache`, `ResponseFormat`. +- `QueryList` fills `Metadata` from `Content-Range` (`items a-b/total`); 206 is success. + +## Server quirks + +- `Sort` is raw SQL in ORDER BY (client sends `col ASC|DESC`). +- One search operator per column. +- Values starting `ZIP_` / `__` are base64-decoded by the server. +- Non-ASCII, control chars and edge spaces are auto-encoded (`ZIP_`). + +## Errors + +`*Error{StatusCode, APIError{Code, Message, Detail, SQL}}`. + +## Test + +`go test ./...` diff --git a/clients/resolvespec-go/client.go b/clients/resolvespec-go/client.go new file mode 100644 index 0000000..104c182 --- /dev/null +++ b/clients/resolvespec-go/client.go @@ -0,0 +1,121 @@ +package resolvespec + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" +) + +// APIError is the server error object. +type APIError struct { + Code string `json:"code"` + Message string `json:"message"` + Details any `json:"details,omitempty"` + Detail string `json:"detail,omitempty"` // server-side reason (funcspec / restheadspec) + SQL string `json:"sql,omitempty"` +} + +// Error is returned on a non-2xx response or an unsuccessful result. +type Error struct { + StatusCode int + APIError +} + +func (e *Error) Error() string { + if e.Message != "" { + return e.Message + } + return fmt.Sprintf("http %d", e.StatusCode) +} + +type config struct { + baseURL string + token string + headers http.Header + http *http.Client +} + +// Option configures a client. +type Option func(*config) + +func WithToken(token string) Option { return func(c *config) { c.token = token } } +func WithHTTPClient(h *http.Client) Option { return func(c *config) { c.http = h } } +func WithHeader(name, value string) Option { + return func(c *config) { c.headers.Set(name, value) } +} + +func newConfig(baseURL string, opts []Option) config { + c := config{baseURL: strings.TrimRight(baseURL, "/"), headers: http.Header{}, http: &http.Client{Timeout: 30 * time.Second}} + for _, o := range opts { + o(&c) + } + return c +} + +// headers: Content-Type < custom headers < bearer token. +func (c *config) newRequest(ctx context.Context, method, u string, body []byte) (*http.Request, error) { + var r io.Reader + if body != nil { + r = bytes.NewReader(body) + } + req, err := http.NewRequestWithContext(ctx, method, u, r) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + for k, vs := range c.headers { + req.Header[k] = append([]string(nil), vs...) + } + if c.token != "" { + req.Header.Set("Authorization", "Bearer "+c.token) + } + return req, nil +} + +func (c *config) do(req *http.Request) (*http.Response, []byte, error) { + resp, err := c.http.Do(req) + if err != nil { + return nil, nil, err + } + defer resp.Body.Close() + b, err := io.ReadAll(resp.Body) + return resp, b, err +} + +func errorFrom(status int, body []byte) *Error { + e := &Error{StatusCode: status} + var env struct { + Error *APIError `json:"error"` + } + if json.Unmarshal(body, &env) == nil && env.Error != nil { + e.APIError = *env.Error + } + if e.Message == "" { + text := "" + if !json.Valid(body) { + text = strings.TrimSpace(string(body)) + if len(text) > 200 { + text = text[:200] + } + } + if text == "" { + text = fmt.Sprintf("%s (%d)", http.StatusText(status), status) + } + e.Message = text + } + return e +} + +func buildURL(base, schema, entity string, id string) string { + u := base + "/" + url.PathEscape(schema) + "/" + url.PathEscape(entity) + if id != "" { + u += "/" + url.PathEscape(id) + } + return u +} diff --git a/clients/resolvespec-go/funcspec.go b/clients/resolvespec-go/funcspec.go new file mode 100644 index 0000000..f4a63ca --- /dev/null +++ b/clients/resolvespec-go/funcspec.go @@ -0,0 +1,274 @@ +package resolvespec + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/url" + "regexp" + "strconv" + "strings" + "unicode" +) + +// FuncSpecOptions are sent to funcspec endpoints as X-* headers. +// +// Server behaviour (pkg/funcspec): Sort is inserted raw into ORDER BY (so it is sent as SQL +// terms); only one search operator per column is kept; values starting with "ZIP_" or "__" +// are base64-decoded by the server, so such plaintext values cannot be sent faithfully. +type FuncSpecOptions struct { + Filters []FilterOption // eq+AND -> X-FieldFilter; others X-SearchOp / X-SearchOr + SearchFilters map[string]string // X-SearchFilter-{col}: text ILIKE + CustomSQLWhere string // X-Custom-SQL-W + CustomSQLOr string // X-Custom-SQL-Or + Sort []SortOption + Limit *int + Offset *int + Distinct *bool + SkipCount *bool + SkipCache *bool + ResponseFormat string // simple | detail | syncfusion +} + +// Params are query-string values. Slice values are sent as repeated keys (server: IN filter). +type Params map[string]any + +// FuncSpecClient calls user-defined SQL endpoints. Routes are defined by the server app. +type FuncSpecClient struct{ cfg config } + +func NewFuncSpecClient(baseURL string, opts ...Option) *FuncSpecClient { + return &FuncSpecClient{cfg: newConfig(baseURL, opts)} +} + +var operatorMap = map[string]string{ + "eq": "equals", "neq": "notequals", "gt": "greaterthan", "gte": "greaterthanorequal", + "lt": "lessthan", "lte": "lessthanorequal", "like": "contains", "ilike": "contains", + "contains": "contains", "startswith": "beginswith", "endswith": "endswith", "in": "in", + "between": "between", "between_inclusive": "betweeninclusive", + "is_null": "empty", "is_not_null": "notempty", +} + +func scalar(v any) string { + switch x := v.(type) { + case nil: + return "" + case bool: + return strconv.FormatBool(x) + case string: + return x + case fmt.Stringer: + return x.String() + } + return fmt.Sprint(v) +} + +func filterValue(v any) string { + switch x := v.(type) { + case nil: + return "" + case []string: + return strings.Join(x, ",") + case []int: + parts := make([]string, len(x)) + for i, n := range x { + parts[i] = strconv.Itoa(n) + } + return strings.Join(parts, ",") + case []any: + parts := make([]string, len(x)) + for i, n := range x { + parts[i] = scalar(n) + } + return strings.Join(parts, ",") + } + return scalar(v) +} + +// EncodeHeaderValue base64-encodes (UTF-8) with the ZIP_ prefix. +func EncodeHeaderValue(v string) string { return "ZIP_" + base64.StdEncoding.EncodeToString([]byte(v)) } + +// DecodeHeaderValue decodes a value that may carry a ZIP_ or __ prefix (nested allowed). +func DecodeHeaderValue(v string) string { + for _, p := range []string{"ZIP_", "__"} { + if strings.HasPrefix(v, p) { + b64 := strings.NewReplacer("\n", "", "\r", "", " ", "").Replace(v[len(p):]) + for len(b64)%4 != 0 { + b64 += "=" + } + raw, err := base64.StdEncoding.DecodeString(b64) + if err != nil { + return v + } + return DecodeHeaderValue(string(raw)) + } + } + return v +} + +// safe encodes values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces). +func safe(v string) string { + if v != strings.TrimSpace(v) { + return EncodeHeaderValue(v) + } + for _, r := range v { + if r > unicode.MaxASCII || !unicode.IsPrint(r) { + return EncodeHeaderValue(v) + } + } + return v +} + +// BuildHeaders builds the X-* headers understood by funcspec.ParseParameters. +func BuildHeaders(o *FuncSpecOptions) map[string]string { + h := map[string]string{} + if o == nil { + return h + } + for _, f := range o.Filters { + logic := f.LogicOperator + if logic == "" { + logic = "AND" + } + v := safe(filterValue(f.Value)) + if f.Operator == "eq" && logic == "AND" { + h["X-FieldFilter-"+f.Column] = v + continue + } + op := operatorMap[f.Operator] + if op == "" { + op = f.Operator + } + kind := "X-SearchOp" + if logic == "OR" { + kind = "X-SearchOr" + } + h[kind+"-"+op+"-"+f.Column] = v + } + for col, text := range o.SearchFilters { + h["X-SearchFilter-"+col] = safe(text) + } + if o.CustomSQLWhere != "" { + h["X-Custom-SQL-W"] = safe(o.CustomSQLWhere) + } + if o.CustomSQLOr != "" { + h["X-Custom-SQL-Or"] = safe(o.CustomSQLOr) + } + if len(o.Sort) > 0 { + terms := make([]string, len(o.Sort)) + for i, s := range o.Sort { + dir := "ASC" + if strings.EqualFold(s.Direction, "desc") { + dir = "DESC" + } + terms[i] = s.Column + " " + dir // funcspec puts this verbatim into ORDER BY + } + h["X-Sort"] = safe(strings.Join(terms, ",")) + } + if o.Limit != nil { + h["X-Limit"] = strconv.Itoa(*o.Limit) + } + if o.Offset != nil { + h["X-Offset"] = strconv.Itoa(*o.Offset) + } + for name, v := range map[string]*bool{"X-Distinct": o.Distinct, "X-SkipCount": o.SkipCount, "X-SkipCache": o.SkipCache} { + if v != nil { + h[name] = strconv.FormatBool(*v) + } + } + switch o.ResponseFormat { + case "simple": + h["X-SimpleApi"] = "true" + case "detail": + h["X-DetailApi"] = "true" + case "syncfusion": + h["X-Syncfusion"] = "true" + } + return h +} + +// BuildQuery builds query-string values: bools -> true/false, slices -> repeated keys, nil skipped. +func BuildQuery(p Params) url.Values { + q := url.Values{} + for k, v := range p { + switch x := v.(type) { + case nil: + case []string: + for _, e := range x { + q.Add(k, safe(e)) + } + case []int: + for _, e := range x { + q.Add(k, strconv.Itoa(e)) + } + case []any: + for _, e := range x { + q.Add(k, safe(scalar(e))) + } + default: + q.Add(k, safe(scalar(v))) + } + } + return q +} + +var contentRange = regexp.MustCompile(`(\d+)-(\d+)/(\d+)`) + +func metadata(h http.Header, o *FuncSpecOptions) *Metadata { + m := &Metadata{} + if g := contentRange.FindStringSubmatch(h.Get("Content-Range")); g != nil { + start, _ := strconv.ParseInt(g[1], 10, 64) + end, _ := strconv.ParseInt(g[2], 10, 64) + total, _ := strconv.ParseInt(g[3], 10, 64) + m.Total, m.Count, m.Filtered, m.Offset = total, end-start, total, int(start) + } + if o != nil && o.Limit != nil { + m.Limit = *o.Limit + } + return m +} + +func (c *FuncSpecClient) call(ctx context.Context, method, path string, p Params, o *FuncSpecOptions, withMeta bool) (*Response, error) { + u := c.cfg.baseURL + "/" + strings.TrimLeft(path, "/") + if q := BuildQuery(p); len(q) > 0 { + u += "?" + q.Encode() + } + req, err := c.cfg.newRequest(ctx, strings.ToUpper(method), u, nil) + if err != nil { + return nil, err + } + for k, v := range BuildHeaders(o) { + req.Header.Set(k, v) + } + resp, b, err := c.cfg.do(req) + if err != nil { + return nil, err + } + if resp.StatusCode < 200 || resp.StatusCode > 299 { // 206 is success + return nil, errorFrom(resp.StatusCode, b) + } + out := &Response{Success: true, Data: json.RawMessage(b)} + if len(b) == 0 { + out.Data = json.RawMessage("null") + } + if withMeta { + out.Metadata = metadata(resp.Header, o) + } + return out, nil +} + +// Query calls a single-record endpoint (SqlQuery). Data is the row object. +func (c *FuncSpecClient) Query(ctx context.Context, path string, p Params, o *FuncSpecOptions) (*Response, error) { + return c.call(ctx, http.MethodGet, path, p, o, false) +} + +// QueryList calls a list endpoint (SqlQueryList). Metadata comes from Content-Range. +func (c *FuncSpecClient) QueryList(ctx context.Context, path string, p Params, o *FuncSpecOptions) (*Response, error) { + return c.call(ctx, http.MethodGet, path, p, o, true) +} + +// Do is like Query/QueryList with an explicit HTTP method (routes are app-defined). +func (c *FuncSpecClient) Do(ctx context.Context, method, path string, p Params, o *FuncSpecOptions, list bool) (*Response, error) { + return c.call(ctx, method, path, p, o, list) +} diff --git a/clients/resolvespec-go/funcspec_test.go b/clients/resolvespec-go/funcspec_test.go new file mode 100644 index 0000000..5ffff1e --- /dev/null +++ b/clients/resolvespec-go/funcspec_test.go @@ -0,0 +1,98 @@ +package resolvespec + +import ( + "context" + "reflect" + "testing" +) + +func TestBuildHeadersFilters(t *testing.T) { + got := BuildHeaders(&FuncSpecOptions{Filters: []FilterOption{ + {Column: "status", Operator: "eq", Value: "active"}, + {Column: "age", Operator: "gte", Value: 18}, + {Column: "name", Operator: "contains", Value: "x", LogicOperator: "OR"}, + {Column: "deleted", Operator: "is_null"}, + {Column: "id", Operator: "in", Value: []int{1, 2}}, + {Column: "p", Operator: "between_inclusive", Value: []any{1, 5}}, + }}) + want := map[string]string{ + "X-FieldFilter-status": "active", + "X-SearchOp-greaterthanorequal-age": "18", + "X-SearchOr-contains-name": "x", + "X-SearchOp-empty-deleted": "", + "X-SearchOp-in-id": "1,2", + "X-SearchOp-betweeninclusive-p": "1,5", + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("%v", got) + } +} + +func TestBuildHeadersMisc(t *testing.T) { + got := BuildHeaders(&FuncSpecOptions{ + SearchFilters: map[string]string{"name": "bob"}, CustomSQLWhere: "a = 1", CustomSQLOr: "b = 2", + Sort: []SortOption{{"name", "asc"}, {"created_at", "DESC"}}, + Limit: Int(5), Offset: Int(10), Distinct: Bool(true), SkipCount: Bool(true), SkipCache: Bool(false), + ResponseFormat: "syncfusion", + }) + want := map[string]string{ + "X-SearchFilter-name": "bob", "X-Custom-SQL-W": "a = 1", "X-Custom-SQL-Or": "b = 2", + "X-Sort": "name ASC,created_at DESC", "X-Limit": "5", "X-Offset": "10", "X-Distinct": "true", + "X-SkipCount": "true", "X-SkipCache": "false", "X-Syncfusion": "true", + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("%v", got) + } +} + +func TestEncodeUnsafe(t *testing.T) { + h := BuildHeaders(&FuncSpecOptions{Filters: []FilterOption{{Column: "n", Operator: "eq", Value: "héllo"}, {Column: "m", Operator: "eq", Value: " pad"}}}) + for _, k := range []string{"X-FieldFilter-n", "X-FieldFilter-m"} { + if len(h[k]) < 4 || h[k][:4] != "ZIP_" { + t.Fatalf("%s=%q", k, h[k]) + } + } + if DecodeHeaderValue(h["X-FieldFilter-n"]) != "héllo" || DecodeHeaderValue(h["X-FieldFilter-m"]) != " pad" { + t.Fatal("roundtrip") + } +} + +func TestBuildQuery(t *testing.T) { + q := BuildQuery(Params{"a": true, "b": []string{"x", "y"}, "c": nil, "d": 3}) + if q.Get("a") != "true" || !reflect.DeepEqual(q["b"], []string{"x", "y"}) || q.Has("c") || q.Get("d") != "3" { + t.Fatalf("%v", q) + } +} + +func TestQueryListMetadata(t *testing.T) { + srv, s := server(t, 206, `[{"id":1},{"id":2}]`, map[string]string{"Content-Range": "items 10-12/50"}) + c := NewFuncSpecClient(srv.URL, WithToken("tok")) + resp, err := c.QueryList(context.Background(), "/api/users", Params{"org": 1}, &FuncSpecOptions{Limit: Int(2)}) + if err != nil { + t.Fatal(err) + } + if s.method != "GET" || s.path != "/api/users?org=1" || s.header.Get("X-Limit") != "2" { + t.Fatalf("%s %v", s.path, s.header) + } + m := resp.Metadata + if m.Total != 50 || m.Count != 2 || m.Offset != 10 || m.Limit != 2 || m.Filtered != 50 { + t.Fatalf("%+v", m) + } + var rows []map[string]any + if err := resp.Decode(&rows); err != nil || len(rows) != 2 { + t.Fatal(err) + } +} + +func TestQuerySingleNoMetadataAndError(t *testing.T) { + srv, _ := server(t, 200, `{"id":1}`, nil) + resp, err := NewFuncSpecClient(srv.URL).Query(context.Background(), "api/u", nil, nil) + if err != nil || resp.Metadata != nil { + t.Fatalf("%v %v", resp, err) + } + srv2, _ := server(t, 400, `{"success":false,"error":{"code":"hook_error","message":"Hook execution failed","detail":"authentication required"}}`, nil) + _, err = NewFuncSpecClient(srv2.URL).Query(context.Background(), "api/u", nil, nil) + if e := err.(*Error); e.Code != "hook_error" || e.Detail != "authentication required" { + t.Fatalf("%#v", e) + } +} diff --git a/clients/resolvespec-go/go.mod b/clients/resolvespec-go/go.mod new file mode 100644 index 0000000..952f301 --- /dev/null +++ b/clients/resolvespec-go/go.mod @@ -0,0 +1,3 @@ +module github.com/bitechdev/ResolveSpec/clients/resolvespec-go + +go 1.22 diff --git a/clients/resolvespec-go/resolvespec.go b/clients/resolvespec-go/resolvespec.go new file mode 100644 index 0000000..2f57b3f --- /dev/null +++ b/clients/resolvespec-go/resolvespec.go @@ -0,0 +1,94 @@ +package resolvespec + +import ( + "context" + "encoding/json" + "fmt" + "net/http" +) + +// Client speaks the ResolveSpec JSON body protocol: POST {operation, data, options}. +type Client struct{ cfg config } + +func NewClient(baseURL string, opts ...Option) *Client { + return &Client{cfg: newConfig(baseURL, opts)} +} + +// RecordID is a single id (int or string, sent in the URL) or a []string (sent in the body). +type RecordID any + +func urlID(id RecordID) string { + switch v := id.(type) { + case nil: + return "" + case []string: + return "" + case string: + return v + default: + return fmt.Sprint(v) + } +} + +func bodyID(id RecordID) []string { + ids, _ := id.([]string) + return ids +} + +type request struct { + Operation string `json:"operation"` + ID []string `json:"id,omitempty"` + Data any `json:"data,omitempty"` + Options *Options `json:"options,omitempty"` +} + +func (c *Client) send(ctx context.Context, method, schema, entity, id string, body any) (*Response, error) { + var payload []byte + if body != nil { + var err error + if payload, err = json.Marshal(body); err != nil { + return nil, err + } + } + req, err := c.cfg.newRequest(ctx, method, buildURL(c.cfg.baseURL, schema, entity, id), payload) + if err != nil { + return nil, err + } + resp, b, err := c.cfg.do(req) + if err != nil { + return nil, err + } + if resp.StatusCode < 200 || resp.StatusCode > 299 { + return nil, errorFrom(resp.StatusCode, b) + } + var out Response + if err := json.Unmarshal(b, &out); err != nil { + return nil, err + } + if !out.Success && out.Error != nil { + return nil, &Error{StatusCode: resp.StatusCode, APIError: *out.Error} + } + return &out, nil +} + +// GetMetadata returns table metadata (GET /{schema}/{entity}). +func (c *Client) GetMetadata(ctx context.Context, schema, entity string) (*Response, error) { + return c.send(ctx, http.MethodGet, schema, entity, "", nil) +} + +// Read reads records; id may be nil, an int/string (URL) or []string (body). +func (c *Client) Read(ctx context.Context, schema, entity string, id RecordID, opts *Options) (*Response, error) { + return c.send(ctx, http.MethodPost, schema, entity, urlID(id), request{Operation: "read", ID: bodyID(id), Options: opts}) +} + +func (c *Client) Create(ctx context.Context, schema, entity string, data any, opts *Options) (*Response, error) { + return c.send(ctx, http.MethodPost, schema, entity, "", request{Operation: "create", Data: data, Options: opts}) +} + +func (c *Client) Update(ctx context.Context, schema, entity string, data any, id RecordID, opts *Options) (*Response, error) { + return c.send(ctx, http.MethodPost, schema, entity, urlID(id), request{Operation: "update", ID: bodyID(id), Data: data, Options: opts}) +} + +func (c *Client) Delete(ctx context.Context, schema, entity string, id RecordID) (*Response, error) { + return c.send(ctx, http.MethodPost, schema, entity, urlID(id), request{Operation: "delete"}) +} diff --git a/clients/resolvespec-go/resolvespec_test.go b/clients/resolvespec-go/resolvespec_test.go new file mode 100644 index 0000000..91598bd --- /dev/null +++ b/clients/resolvespec-go/resolvespec_test.go @@ -0,0 +1,99 @@ +package resolvespec + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "reflect" + "testing" +) + +type seen struct { + method, path string + header http.Header + body map[string]any +} + +func server(t *testing.T, status int, body string, hdr map[string]string) (*httptest.Server, *seen) { + t.Helper() + s := &seen{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + s.method, s.path, s.header = r.Method, r.URL.EscapedPath()+"?"+r.URL.RawQuery, r.Header + b, _ := io.ReadAll(r.Body) + if len(b) > 0 { + _ = json.Unmarshal(b, &s.body) + } + for k, v := range hdr { + w.Header().Set(k, v) + } + w.WriteHeader(status) + _, _ = w.Write([]byte(body)) + })) + t.Cleanup(srv.Close) + return srv, s +} + +func TestReadBody(t *testing.T) { + srv, s := server(t, 200, `{"success":true,"data":[{"id":1}]}`, nil) + c := NewClient(srv.URL+"/", WithToken("tok"), WithHeader("X-Tenant", "a")) + resp, err := c.Read(context.Background(), "public", "users", nil, &Options{Limit: Int(5), Filters: []FilterOption{{Column: "a", Operator: "eq", Value: 1}}}) + if err != nil { + t.Fatal(err) + } + if s.method != "POST" || s.path != "/public/users?" { + t.Fatalf("got %s %s", s.method, s.path) + } + if s.header.Get("Authorization") != "Bearer tok" || s.header.Get("X-Tenant") != "a" { + t.Fatalf("headers %v", s.header) + } + if s.body["operation"] != "read" || s.body["options"].(map[string]any)["limit"] != float64(5) { + t.Fatalf("body %v", s.body) + } + var rows []map[string]any + if err := resp.Decode(&rows); err != nil || len(rows) != 1 { + t.Fatalf("decode %v %v", rows, err) + } +} + +func TestIDPlacement(t *testing.T) { + srv, s := server(t, 200, `{"success":true,"data":{}}`, nil) + c := NewClient(srv.URL) + ctx := context.Background() + _, _ = c.Read(ctx, "s", "e", 7, nil) + if s.path != "/s/e/7?" || s.body["id"] != nil { + t.Fatalf("%s %v", s.path, s.body) + } + _, _ = c.Update(ctx, "s", "e", map[string]any{"a": 1}, []string{"1", "2"}, nil) + if s.path != "/s/e?" || !reflect.DeepEqual(s.body["id"], []any{"1", "2"}) || s.body["operation"] != "update" { + t.Fatalf("%s %v", s.path, s.body) + } + _, _ = c.Delete(ctx, "s", "e", "a/b") + if s.path != "/s/e/a%2Fb?" || s.body["operation"] != "delete" { + t.Fatalf("%s %v", s.path, s.body) + } + _, _ = c.GetMetadata(ctx, "s", "e") + if s.method != "GET" { + t.Fatal(s.method) + } +} + +func TestErrors(t *testing.T) { + srv, _ := server(t, 400, `{"success":false,"error":{"code":"x","message":"bad","detail":"why"}}`, nil) + _, err := NewClient(srv.URL).Read(context.Background(), "s", "e", nil, nil) + e, ok := err.(*Error) + if !ok || e.StatusCode != 400 || e.Code != "x" || e.Message != "bad" || e.Detail != "why" { + t.Fatalf("%#v", err) + } + srv2, _ := server(t, 502, "bad gateway", nil) + _, err = NewClient(srv2.URL).Read(context.Background(), "s", "e", nil, nil) + if e := err.(*Error); e.StatusCode != 502 || e.Message != "bad gateway" { + t.Fatalf("%#v", e) + } + srv3, _ := server(t, 200, `{"success":false,"error":{"code":"c","message":"nope"}}`, nil) + _, err = NewClient(srv3.URL).Read(context.Background(), "s", "e", nil, nil) + if e := err.(*Error); e.Message != "nope" { + t.Fatalf("%#v", e) + } +} diff --git a/clients/resolvespec-go/types.go b/clients/resolvespec-go/types.go new file mode 100644 index 0000000..e00c2a4 --- /dev/null +++ b/clients/resolvespec-go/types.go @@ -0,0 +1,107 @@ +// Package resolvespec is a client for ResolveSpec (JSON body) and FunctionSpec endpoints. +package resolvespec + +import "encoding/json" + +// FilterOption mirrors common.FilterOption. Operator: eq neq gt gte lt lte like ilike in +// contains startswith endswith between between_inclusive is_null is_not_null. +type FilterOption struct { + Column string `json:"column"` + Operator string `json:"operator"` + Value any `json:"value"` + LogicOperator string `json:"logic_operator,omitempty"` // AND | OR +} + +type SortOption struct { + Column string `json:"column"` + Direction string `json:"direction"` // asc | desc +} + +type Parameter struct { + Name string `json:"name"` + Value string `json:"value"` + Sequence int `json:"sequence,omitempty"` +} + +type CustomOperator struct { + Name string `json:"name"` + SQL string `json:"sql"` +} + +type ComputedColumn struct { + Name string `json:"name"` + Expression string `json:"expression"` +} + +type PreloadOption struct { + Relation string `json:"relation,omitempty"` + TableName string `json:"table_name,omitempty"` + Columns []string `json:"columns,omitempty"` + OmitColumns []string `json:"omit_columns,omitempty"` + Sort []SortOption `json:"sort,omitempty"` + Filters []FilterOption `json:"filters,omitempty"` + Where string `json:"where,omitempty"` + Limit *int `json:"limit,omitempty"` + Offset *int `json:"offset,omitempty"` + Updateable *bool `json:"updateable,omitempty"` + ComputedQL map[string]string `json:"computed_ql,omitempty"` + Recursive bool `json:"recursive,omitempty"` + PrimaryKey string `json:"primary_key,omitempty"` + RelatedKey string `json:"related_key,omitempty"` + ForeignKey string `json:"foreign_key,omitempty"` + RecursiveChildKey string `json:"recursive_child_key,omitempty"` + SQLJoins []string `json:"sql_joins,omitempty"` + JoinAliases []string `json:"join_aliases,omitempty"` +} + +type VectorSearchOption struct { + Column string `json:"column"` + Vector []float64 `json:"vector"` + Metric string `json:"metric,omitempty"` // l2 (default) | cosine | ip + As string `json:"as,omitempty"` // distance alias, default _distance + Direction string `json:"direction,omitempty"` +} + +// Options is the ResolveSpec request options object. +type Options struct { + Preload []PreloadOption `json:"preload,omitempty"` + Columns []string `json:"columns,omitempty"` + OmitColumns []string `json:"omit_columns,omitempty"` + Filters []FilterOption `json:"filters,omitempty"` + Sort []SortOption `json:"sort,omitempty"` + Limit *int `json:"limit,omitempty"` + Offset *int `json:"offset,omitempty"` + CustomOperators []CustomOperator `json:"customOperators,omitempty"` + ComputedColumns []ComputedColumn `json:"computedColumns,omitempty"` + Parameters []Parameter `json:"parameters,omitempty"` + CursorForward string `json:"cursor_forward,omitempty"` + CursorBackward string `json:"cursor_backward,omitempty"` + FetchRowNumber string `json:"fetch_row_number,omitempty"` + VectorSearch *VectorSearchOption `json:"vector_search,omitempty"` +} + +// Metadata of a list response. +type Metadata struct { + Total int64 `json:"total"` + Count int64 `json:"count"` + Filtered int64 `json:"filtered"` + Limit int `json:"limit"` + Offset int `json:"offset"` +} + +// Response is the ResolveSpec envelope. Data is left raw for the caller to decode. +type Response struct { + Success bool `json:"success"` + Data json.RawMessage `json:"data"` + Metadata *Metadata `json:"metadata,omitempty"` + Error *APIError `json:"error,omitempty"` +} + +// Decode unmarshals Data into v. +func (r *Response) Decode(v any) error { return json.Unmarshal(r.Data, v) } + +// Int returns a pointer to n, for optional Options fields. +func Int(n int) *int { return &n } + +// Bool returns a pointer to b. +func Bool(b bool) *bool { return &b } diff --git a/clients/resolvespec-rs/.gitignore b/clients/resolvespec-rs/.gitignore new file mode 100644 index 0000000..2c96eb1 --- /dev/null +++ b/clients/resolvespec-rs/.gitignore @@ -0,0 +1,2 @@ +target/ +Cargo.lock diff --git a/clients/resolvespec-rs/Cargo.toml b/clients/resolvespec-rs/Cargo.toml new file mode 100644 index 0000000..20aa5e3 --- /dev/null +++ b/clients/resolvespec-rs/Cargo.toml @@ -0,0 +1,18 @@ +[package] +name = "resolvespec" +version = "0.1.0" +edition = "2021" +rust-version = "1.80" +description = "Client for ResolveSpec (JSON body) and FunctionSpec endpoints" +license = "MIT" + +[dependencies] +reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] } +serde = { version = "1", features = ["derive"] } +serde_json = "1" +base64 = "0.22" +thiserror = "1" + +[dev-dependencies] +tokio = { version = "1", features = ["macros", "rt-multi-thread"] } +wiremock = "0.6" diff --git a/clients/resolvespec-rs/README.md b/clients/resolvespec-rs/README.md new file mode 100644 index 0000000..cf8a83c --- /dev/null +++ b/clients/resolvespec-rs/README.md @@ -0,0 +1,40 @@ +# resolvespec (Rust) + +Rust client for ResolveSpec (JSON body) and FunctionSpec. Async (`reqwest` + `tokio`). MSRV 1.80. + +## Clients + +| Type | Constructor | Methods | +|---|---|---| +| `ResolveSpecClient` | `new(base_url)` / `from_builder(ClientBuilder)` | `get_metadata` `read` `create` `update` `delete` | +| `FuncSpecClient` | `new(base_url)` / `from_builder(ClientBuilder)` | `query` `query_list` `request` | + +`ClientBuilder::new(url).token().header().timeout().http_client()`. Precedence: Content-Type < custom headers < bearer token. + +## ResolveSpec + +- `RecordId`: `Int`/`Str` → URL, `Many(Vec)` → body (`From` impls provided). +- `Options` (`Default` + struct update), optional fields are `Option`/empty `Vec`. +- Result: `Response{success, data: serde_json::Value, metadata}`; `resp.decode::()`. + +## FunctionSpec + +- Routes are server-defined: pass the `path`. +- `Params = BTreeMap` → query string (`Param::List` → repeated keys). +- `FuncSpecOptions` → `X-*` headers: `filters`, `search_filters`, `custom_sql_where`, `custom_sql_or`, `sort`, `limit`, `offset`, `distinct`, `skip_count`, `skip_cache`, `response_format`. +- `query_list` fills `metadata` from `Content-Range`; 206 is success. + +## Server quirks + +- `sort` is raw SQL in ORDER BY (client sends `col ASC|DESC`). +- One search operator per column. +- Values starting `ZIP_` / `__` are base64-decoded by the server. +- Non-ASCII, control chars and edge spaces are auto-encoded (`ZIP_`). + +## Errors + +`Error::Api { status, message, error: ApiError{code, message, detail, sql} }`, `Error::Http`, `Error::Json`. + +## Test + +`cargo test` diff --git a/clients/resolvespec-rs/src/client.rs b/clients/resolvespec-rs/src/client.rs new file mode 100644 index 0000000..8d9c876 --- /dev/null +++ b/clients/resolvespec-rs/src/client.rs @@ -0,0 +1,90 @@ +use std::collections::HashMap; +use std::time::Duration; + +use reqwest::header::{HeaderMap, HeaderName, HeaderValue, AUTHORIZATION, CONTENT_TYPE}; + +use crate::error::Result; + +/// Shared HTTP configuration. +#[derive(Clone)] +pub(crate) struct Config { + pub base_url: String, + pub token: Option, + pub headers: HashMap, + pub http: reqwest::Client, +} + +/// Builder options shared by both clients. +#[derive(Default, Clone)] +pub struct ClientBuilder { + base_url: String, + token: Option, + headers: HashMap, + timeout: Option, + http: Option, +} + +impl ClientBuilder { + pub fn new(base_url: &str) -> Self { + Self { base_url: base_url.trim_end_matches('/').into(), timeout: Some(Duration::from_secs(30)), ..Default::default() } + } + pub fn token(mut self, token: &str) -> Self { + self.token = Some(token.into()); + self + } + pub fn header(mut self, name: &str, value: &str) -> Self { + self.headers.insert(name.into(), value.into()); + self + } + pub fn timeout(mut self, t: Duration) -> Self { + self.timeout = Some(t); + self + } + pub fn http_client(mut self, c: reqwest::Client) -> Self { + self.http = Some(c); + self + } + pub(crate) fn config(self) -> Result { + let http = match self.http { + Some(c) => c, + None => { + let mut b = reqwest::Client::builder(); + if let Some(t) = self.timeout { + b = b.timeout(t); + } + b.build()? + } + }; + Ok(Config { base_url: self.base_url, token: self.token, headers: self.headers, http }) + } +} + +impl Config { + /// Content-Type < custom headers < extra (per-call) < bearer token. + pub fn headers(&self, extra: &HashMap) -> HeaderMap { + let mut m = HeaderMap::new(); + m.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + for (k, v) in self.headers.iter().chain(extra.iter()) { + if let (Ok(n), Ok(v)) = (HeaderName::try_from(k.as_str()), HeaderValue::from_str(v)) { + m.insert(n, v); + } + } + if let Some(t) = &self.token { + if let Ok(v) = HeaderValue::from_str(&format!("Bearer {t}")) { + m.insert(AUTHORIZATION, v); + } + } + m + } +} + +pub(crate) fn path_segment(s: &str) -> String { + let mut out = String::new(); + for b in s.bytes() { + match b { + b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => out.push(b as char), + _ => out.push_str(&format!("%{b:02X}")), + } + } + out +} diff --git a/clients/resolvespec-rs/src/error.rs b/clients/resolvespec-rs/src/error.rs new file mode 100644 index 0000000..cbf00e2 --- /dev/null +++ b/clients/resolvespec-rs/src/error.rs @@ -0,0 +1,34 @@ +use crate::types::ApiError; + +/// Returned on transport failure, a non-2xx response or an unsuccessful API result. +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("{message}")] + Api { status: u16, message: String, error: ApiError }, + #[error(transparent)] + Http(#[from] reqwest::Error), + #[error(transparent)] + Json(#[from] serde_json::Error), +} + +pub type Result = std::result::Result; + +pub(crate) fn error_from(status: u16, body: &str) -> Error { + let parsed: Option = serde_json::from_str(body).ok(); + let err: ApiError = parsed + .as_ref() + .and_then(|v| v.get("error")) + .and_then(|e| serde_json::from_value(e.clone()).ok()) + .unwrap_or_default(); + let message = if !err.message.is_empty() { + err.message.clone() + } else { + let text = if parsed.is_none() { body.trim().chars().take(200).collect::() } else { String::new() }; + if text.is_empty() { + format!("{} ({})", reqwest::StatusCode::from_u16(status).ok().and_then(|s| s.canonical_reason()).unwrap_or("Error"), status) + } else { + text + } + }; + Error::Api { status, message, error: err } +} diff --git a/clients/resolvespec-rs/src/funcspec.rs b/clients/resolvespec-rs/src/funcspec.rs new file mode 100644 index 0000000..4d3c12d --- /dev/null +++ b/clients/resolvespec-rs/src/funcspec.rs @@ -0,0 +1,263 @@ +use std::collections::{BTreeMap, HashMap}; + +use base64::{engine::general_purpose::STANDARD, Engine}; +use reqwest::Method; +use serde_json::Value; + +use crate::client::{ClientBuilder, Config}; +use crate::error::{error_from, Result}; +use crate::types::{FilterOption, Metadata, Response, SortOption}; + +/// Options sent to funcspec endpoints as `X-*` headers. +/// +/// Server behaviour (`pkg/funcspec`): `sort` is inserted raw into ORDER BY (so it is sent as SQL +/// terms); only one search operator per column is kept; values starting with `ZIP_` or `__` +/// are base64-decoded by the server, so such plaintext values cannot be sent faithfully. +#[derive(Debug, Clone, Default)] +pub struct FuncSpecOptions { + /// eq+AND -> X-FieldFilter; others X-SearchOp / X-SearchOr. + pub filters: Vec, + /// X-SearchFilter-{col}: text ILIKE. + pub search_filters: BTreeMap, + pub custom_sql_where: Option, + pub custom_sql_or: Option, + pub sort: Vec, + pub limit: Option, + pub offset: Option, + pub distinct: Option, + pub skip_count: Option, + pub skip_cache: Option, + /// simple | detail | syncfusion + pub response_format: Option, +} + +/// Query-string parameter value. `List` is sent as repeated keys (server: IN filter). +#[derive(Debug, Clone)] +pub enum Param { + Str(String), + Int(i64), + Bool(bool), + List(Vec), +} + +impl From<&str> for Param { + fn from(v: &str) -> Self { + Self::Str(v.into()) + } +} +impl From for Param { + fn from(v: String) -> Self { + Self::Str(v) + } +} +impl From for Param { + fn from(v: i64) -> Self { + Self::Int(v) + } +} +impl From for Param { + fn from(v: bool) -> Self { + Self::Bool(v) + } +} +impl From> for Param { + fn from(v: Vec) -> Self { + Self::List(v) + } +} + +pub type Params = BTreeMap; + +fn operator(op: &str) -> &str { + match op { + "eq" => "equals", + "neq" => "notequals", + "gt" => "greaterthan", + "gte" => "greaterthanorequal", + "lt" => "lessthan", + "lte" => "lessthanorequal", + "like" | "ilike" | "contains" => "contains", + "startswith" => "beginswith", + "endswith" => "endswith", + "in" => "in", + "between" => "between", + "between_inclusive" => "betweeninclusive", + "is_null" => "empty", + "is_not_null" => "notempty", + other => other, + } +} + +fn scalar(v: &Value) -> String { + match v { + Value::Null => String::new(), + Value::String(s) => s.clone(), + Value::Array(a) => a.iter().map(scalar).collect::>().join(","), + other => other.to_string(), + } +} + +/// Base64 (UTF-8) with the `ZIP_` prefix. +pub fn encode_header_value(v: &str) -> String { + format!("ZIP_{}", STANDARD.encode(v.as_bytes())) +} + +/// Decode a value that may carry a `ZIP_` or `__` prefix (nested allowed). +pub fn decode_header_value(v: &str) -> String { + for p in ["ZIP_", "__"] { + if let Some(rest) = v.strip_prefix(p) { + let mut b64: String = rest.chars().filter(|c| !matches!(c, '\n' | '\r' | ' ')).collect(); + while b64.len() % 4 != 0 { + b64.push('='); + } + return match STANDARD.decode(b64).ok().and_then(|b| String::from_utf8(b).ok()) { + Some(s) => decode_header_value(&s), + None => v.to_string(), + }; + } + } + v.to_string() +} + +/// Encode values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces). +fn safe(v: &str) -> String { + if v != v.trim() || v.chars().any(|c| !c.is_ascii() || c.is_ascii_control()) { + encode_header_value(v) + } else { + v.to_string() + } +} + +/// Build the `X-*` headers understood by `funcspec.ParseParameters`. +pub fn build_headers(o: &FuncSpecOptions) -> BTreeMap { + let mut h = BTreeMap::new(); + for f in &o.filters { + let logic = f.logic_operator.as_deref().unwrap_or("AND"); + let v = safe(&scalar(&f.value)); + if f.operator == "eq" && logic == "AND" { + h.insert(format!("X-FieldFilter-{}", f.column), v); + } else { + let kind = if logic == "OR" { "X-SearchOr" } else { "X-SearchOp" }; + h.insert(format!("{kind}-{}-{}", operator(&f.operator), f.column), v); + } + } + for (col, text) in &o.search_filters { + h.insert(format!("X-SearchFilter-{col}"), safe(text)); + } + if let Some(v) = o.custom_sql_where.as_deref().filter(|s| !s.is_empty()) { + h.insert("X-Custom-SQL-W".into(), safe(v)); + } + if let Some(v) = o.custom_sql_or.as_deref().filter(|s| !s.is_empty()) { + h.insert("X-Custom-SQL-Or".into(), safe(v)); + } + if !o.sort.is_empty() { + let terms: Vec = o + .sort + .iter() + .map(|s| format!("{} {}", s.column, if s.direction.eq_ignore_ascii_case("desc") { "DESC" } else { "ASC" })) + .collect(); + h.insert("X-Sort".into(), safe(&terms.join(","))); // funcspec puts this verbatim into ORDER BY + } + if let Some(n) = o.limit { + h.insert("X-Limit".into(), n.to_string()); + } + if let Some(n) = o.offset { + h.insert("X-Offset".into(), n.to_string()); + } + for (name, v) in [("X-Distinct", o.distinct), ("X-SkipCount", o.skip_count), ("X-SkipCache", o.skip_cache)] { + if let Some(b) = v { + h.insert(name.into(), b.to_string()); + } + } + match o.response_format.as_deref() { + Some("simple") => h.insert("X-SimpleApi".into(), "true".into()), + Some("detail") => h.insert("X-DetailApi".into(), "true".into()), + Some("syncfusion") => h.insert("X-Syncfusion".into(), "true".into()), + _ => None, + }; + h +} + +/// Build query-string pairs: bools -> true/false, lists -> repeated keys. +pub fn build_query(p: &Params) -> Vec<(String, String)> { + let mut out = Vec::new(); + for (k, v) in p { + match v { + Param::Str(s) => out.push((k.clone(), safe(s))), + Param::Int(n) => out.push((k.clone(), n.to_string())), + Param::Bool(b) => out.push((k.clone(), b.to_string())), + Param::List(l) => out.extend(l.iter().map(|s| (k.clone(), safe(s)))), + } + } + out +} + +fn metadata(content_range: Option<&str>, limit: Option) -> Metadata { + let mut m = Metadata { limit: limit.unwrap_or(0), ..Default::default() }; + if let Some(cr) = content_range { + // "items {start}-{end}/{total}" + let rest = cr.rsplit(' ').next().unwrap_or(""); + if let Some((range, total)) = rest.split_once('/') { + if let (Some((s, e)), Ok(t)) = (range.split_once('-'), total.parse::()) { + if let (Ok(s), Ok(e)) = (s.parse::(), e.parse::()) { + m.total = t; + m.filtered = t; + m.count = e - s; + m.offset = s; + } + } + } + } + m +} + +/// Client for user-defined SQL endpoints. Routes are defined by the server application. +#[derive(Clone)] +pub struct FuncSpecClient { + cfg: Config, +} + +impl FuncSpecClient { + pub fn new(base_url: &str) -> Result { + Self::from_builder(ClientBuilder::new(base_url)) + } + + pub fn from_builder(b: ClientBuilder) -> Result { + Ok(Self { cfg: b.config()? }) + } + + async fn call(&self, method: Method, path: &str, params: &Params, options: Option<&FuncSpecOptions>, list: bool) -> Result { + let url = format!("{}/{}", self.cfg.base_url, path.trim_start_matches('/')); + let extra: HashMap = options.map(|o| build_headers(o).into_iter().collect()).unwrap_or_default(); + let resp = self.cfg.http.request(method, url).headers(self.cfg.headers(&extra)).query(&build_query(params)).send().await?; + let status = resp.status(); + let cr = resp.headers().get("content-range").and_then(|v| v.to_str().ok()).map(str::to_owned); + let text = resp.text().await?; + if !status.is_success() { + // 206 Partial Content is success + return Err(error_from(status.as_u16(), &text)); + } + let data = if text.trim().is_empty() { Value::Null } else { serde_json::from_str(&text)? }; + Ok(Response { + success: true, + data, + metadata: list.then(|| metadata(cr.as_deref(), options.and_then(|o| o.limit))), + error: None, + }) + } + + /// Single-record endpoint (`SqlQuery`). `data` is the row object. + pub async fn query(&self, path: &str, params: &Params, options: Option<&FuncSpecOptions>) -> Result { + self.call(Method::GET, path, params, options, false).await + } + + /// List endpoint (`SqlQueryList`). Metadata comes from Content-Range. + pub async fn query_list(&self, path: &str, params: &Params, options: Option<&FuncSpecOptions>) -> Result { + self.call(Method::GET, path, params, options, true).await + } + + /// Like `query` / `query_list` with an explicit HTTP method (routes are app-defined). + pub async fn request(&self, method: Method, path: &str, params: &Params, options: Option<&FuncSpecOptions>, list: bool) -> Result { + self.call(method, path, params, options, list).await + } +} diff --git a/clients/resolvespec-rs/src/lib.rs b/clients/resolvespec-rs/src/lib.rs new file mode 100644 index 0000000..3ff20c5 --- /dev/null +++ b/clients/resolvespec-rs/src/lib.rs @@ -0,0 +1,12 @@ +//! Client for ResolveSpec (JSON body) and FunctionSpec endpoints. +mod client; +mod error; +mod funcspec; +mod resolvespec; +pub mod types; + +pub use client::ClientBuilder; +pub use error::{Error, Result}; +pub use funcspec::{build_headers, build_query, decode_header_value, encode_header_value, FuncSpecClient, FuncSpecOptions, Param, Params}; +pub use resolvespec::{RecordId, ResolveSpecClient}; +pub use types::*; diff --git a/clients/resolvespec-rs/src/resolvespec.rs b/clients/resolvespec-rs/src/resolvespec.rs new file mode 100644 index 0000000..5757cf0 --- /dev/null +++ b/clients/resolvespec-rs/src/resolvespec.rs @@ -0,0 +1,130 @@ +use std::collections::HashMap; + +use reqwest::Method; +use serde::Serialize; +use serde_json::Value; + +use crate::client::{path_segment, ClientBuilder, Config}; +use crate::error::{error_from, Error, Result}; +use crate::types::{Options, Response}; + +/// A record id: a single value goes in the URL, a list goes in the body. +#[derive(Debug, Clone)] +pub enum RecordId { + Int(i64), + Str(String), + Many(Vec), +} + +impl From for RecordId { + fn from(v: i64) -> Self { + Self::Int(v) + } +} +impl From<&str> for RecordId { + fn from(v: &str) -> Self { + Self::Str(v.into()) + } +} +impl From> for RecordId { + fn from(v: Vec) -> Self { + Self::Many(v) + } +} + +fn url_id(id: &Option) -> Option { + match id { + Some(RecordId::Int(n)) => Some(n.to_string()), + Some(RecordId::Str(s)) => Some(s.clone()), + _ => None, + } +} + +#[derive(Serialize)] +struct Request<'a> { + operation: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + id: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + data: Option, + #[serde(skip_serializing_if = "Option::is_none")] + options: Option<&'a Options>, +} + +/// Client for the ResolveSpec JSON body protocol. +#[derive(Clone)] +pub struct ResolveSpecClient { + cfg: Config, +} + +impl ResolveSpecClient { + pub fn new(base_url: &str) -> Result { + Self::from_builder(ClientBuilder::new(base_url)) + } + + pub fn from_builder(b: ClientBuilder) -> Result { + Ok(Self { cfg: b.config()? }) + } + + fn url(&self, schema: &str, entity: &str, id: Option) -> String { + let mut u = format!("{}/{}/{}", self.cfg.base_url, path_segment(schema), path_segment(entity)); + if let Some(id) = id.filter(|i| !i.is_empty()) { + u.push('/'); + u.push_str(&path_segment(&id)); + } + u + } + + async fn send(&self, method: Method, url: String, body: Option>) -> Result { + let mut req = self.cfg.http.request(method, url).headers(self.cfg.headers(&HashMap::new())); + if let Some(b) = body { + req = req.body(serde_json::to_vec(&b)?); + } + let resp = req.send().await?; + let status = resp.status(); + let text = resp.text().await?; + if !status.is_success() { + return Err(error_from(status.as_u16(), &text)); + } + let out: Response = serde_json::from_str(&text)?; + if !out.success { + if let Some(e) = out.error.clone() { + return Err(Error::Api { status: status.as_u16(), message: e.message.clone(), error: e }); + } + } + Ok(out) + } + + /// GET /{schema}/{entity} + pub async fn get_metadata(&self, schema: &str, entity: &str) -> Result { + self.send(Method::GET, self.url(schema, entity, None), None).await + } + + pub async fn read(&self, schema: &str, entity: &str, id: Option, options: Option<&Options>) -> Result { + let body = Request { operation: "read", id: many(&id), data: None, options }; + self.send(Method::POST, self.url(schema, entity, url_id(&id)), Some(body)).await + } + + pub async fn create(&self, schema: &str, entity: &str, data: Value, options: Option<&Options>) -> Result { + let body = Request { operation: "create", id: None, data: Some(data), options }; + self.send(Method::POST, self.url(schema, entity, None), Some(body)).await + } + + pub async fn update(&self, schema: &str, entity: &str, data: Value, id: Option, options: Option<&Options>) -> Result { + let body = Request { operation: "update", id: many(&id), data: Some(data), options }; + self.send(Method::POST, self.url(schema, entity, url_id(&id)), Some(body)).await + } + + pub async fn delete(&self, schema: &str, entity: &str, id: impl Into) -> Result { + let id = Some(id.into()); + let body = Request { operation: "delete", id: None, data: None, options: None }; + self.send(Method::POST, self.url(schema, entity, url_id(&id)), Some(body)).await + } +} + +fn many(id: &Option) -> Option> { + match id { + Some(RecordId::Many(v)) => Some(v.clone()), + _ => None, + } +} diff --git a/clients/resolvespec-rs/src/types.rs b/clients/resolvespec-rs/src/types.rs new file mode 100644 index 0000000..6f09e5a --- /dev/null +++ b/clients/resolvespec-rs/src/types.rs @@ -0,0 +1,192 @@ +//! Types aligned with Go `pkg/common/types.go`. Field names are the wire names. +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::HashMap; + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct FilterOption { + pub column: String, + /// eq neq gt gte lt lte like ilike in contains startswith endswith between + /// between_inclusive is_null is_not_null + pub operator: String, + #[serde(default)] + pub value: Value, + #[serde(skip_serializing_if = "Option::is_none")] + pub logic_operator: Option, // AND | OR +} + +impl FilterOption { + pub fn new(column: &str, operator: &str, value: impl Into) -> Self { + Self { column: column.into(), operator: operator.into(), value: value.into(), logic_operator: None } + } + pub fn or(mut self) -> Self { + self.logic_operator = Some("OR".into()); + self + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SortOption { + pub column: String, + pub direction: String, // asc | desc +} + +impl SortOption { + pub fn new(column: &str, direction: &str) -> Self { + Self { column: column.into(), direction: direction.into() } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Parameter { + pub name: String, + pub value: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub sequence: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CustomOperator { + pub name: String, + pub sql: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ComputedColumn { + pub name: String, + pub expression: String, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct PreloadOption { + #[serde(skip_serializing_if = "Option::is_none")] + pub relation: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub table_name: Option, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub columns: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub omit_columns: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub sort: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub filters: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub r#where: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub offset: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub updateable: Option, + #[serde(skip_serializing_if = "HashMap::is_empty", default)] + pub computed_ql: HashMap, + #[serde(skip_serializing_if = "Option::is_none")] + pub recursive: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub primary_key: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub related_key: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub foreign_key: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub recursive_child_key: Option, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub sql_joins: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub join_aliases: Vec, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct VectorSearchOption { + pub column: String, + pub vector: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub metric: Option, // l2 (default) | cosine | ip + #[serde(rename = "as", skip_serializing_if = "Option::is_none")] + pub alias: Option, // distance alias, default _distance + #[serde(skip_serializing_if = "Option::is_none")] + pub direction: Option, +} + +/// ResolveSpec request options object. +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct Options { + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub preload: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub columns: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub omit_columns: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub filters: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub sort: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub offset: Option, + #[serde(rename = "customOperators", skip_serializing_if = "Vec::is_empty", default)] + pub custom_operators: Vec, + #[serde(rename = "computedColumns", skip_serializing_if = "Vec::is_empty", default)] + pub computed_columns: Vec, + #[serde(skip_serializing_if = "Vec::is_empty", default)] + pub parameters: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub cursor_forward: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cursor_backward: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub fetch_row_number: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub vector_search: Option, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)] +pub struct Metadata { + #[serde(default)] + pub total: i64, + #[serde(default)] + pub count: i64, + #[serde(default)] + pub filtered: i64, + #[serde(default)] + pub limit: i64, + #[serde(default)] + pub offset: i64, +} + +/// ResolveSpec envelope. `data` is left as JSON for the caller to decode. +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct Response { + #[serde(default)] + pub success: bool, + #[serde(default)] + pub data: Value, + #[serde(skip_serializing_if = "Option::is_none")] + pub metadata: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub error: Option, +} + +impl Response { + /// Decode `data` into `T`. + pub fn decode(&self) -> Result { + serde_json::from_value(self.data.clone()) + } +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct ApiError { + #[serde(default)] + pub code: String, + #[serde(default)] + pub message: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub details: Option, + /// Server-side reason (funcspec / restheadspec). + #[serde(skip_serializing_if = "Option::is_none")] + pub detail: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub sql: Option, +} diff --git a/clients/resolvespec-rs/tests/client.rs b/clients/resolvespec-rs/tests/client.rs new file mode 100644 index 0000000..a0f2522 --- /dev/null +++ b/clients/resolvespec-rs/tests/client.rs @@ -0,0 +1,164 @@ +use resolvespec::*; +use serde_json::json; +use wiremock::matchers::{header, method, path, query_param}; +use wiremock::{Mock, MockServer, ResponseTemplate}; + +#[tokio::test] +async fn read_posts_body_with_headers() { + let srv = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/public/users")) + .and(header("authorization", "Bearer tok")) + .and(header("x-tenant", "a")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"success": true, "data": [{"id": 1}]}))) + .expect(1) + .mount(&srv) + .await; + let c = ResolveSpecClient::from_builder(ClientBuilder::new(&format!("{}/", srv.uri())).token("tok").header("X-Tenant", "a")).unwrap(); + let opts = Options { limit: Some(5), filters: vec![FilterOption::new("a", "eq", 1)], ..Default::default() }; + let r = c.read("public", "users", None, Some(&opts)).await.unwrap(); + let rows: Vec = r.decode().unwrap(); + assert_eq!(rows.len(), 1); + let body: serde_json::Value = serde_json::from_slice(&srv.received_requests().await.unwrap()[0].body).unwrap(); + assert_eq!(body["operation"], "read"); + assert_eq!(body["options"]["limit"], 5); + assert!(body.get("id").is_none()); +} + +#[tokio::test] +async fn id_placement() { + let srv = MockServer::start().await; + Mock::given(method("POST")).respond_with(ResponseTemplate::new(200).set_body_json(json!({"success": true, "data": {}}))).mount(&srv).await; + let c = ResolveSpecClient::new(&srv.uri()).unwrap(); + c.read("s", "e", Some(7.into()), None).await.unwrap(); + c.update("s", "e", json!({"a": 1}), Some(vec!["1".to_string(), "2".to_string()].into()), None).await.unwrap(); + c.delete("s", "e", "a/b").await.unwrap(); + let reqs = srv.received_requests().await.unwrap(); + assert_eq!(reqs[0].url.path(), "/s/e/7"); + assert_eq!(reqs[1].url.path(), "/s/e"); + let b: serde_json::Value = serde_json::from_slice(&reqs[1].body).unwrap(); + assert_eq!(b["id"], json!(["1", "2"])); + assert_eq!(b["operation"], "update"); + assert_eq!(reqs[2].url.path(), "/s/e/a%2Fb"); +} + +#[tokio::test] +async fn errors() { + let srv = MockServer::start().await; + Mock::given(path("/s/a")).respond_with(ResponseTemplate::new(400).set_body_json(json!({"success": false, "error": {"code": "x", "message": "bad", "detail": "why"}}))).mount(&srv).await; + Mock::given(path("/s/b")).respond_with(ResponseTemplate::new(502).set_body_string("bad gateway")).mount(&srv).await; + Mock::given(path("/s/c")).respond_with(ResponseTemplate::new(200).set_body_json(json!({"success": false, "error": {"code": "c", "message": "nope"}}))).mount(&srv).await; + let c = ResolveSpecClient::new(&srv.uri()).unwrap(); + match c.read("s", "a", None, None).await.unwrap_err() { + Error::Api { status, message, error } => { + assert_eq!((status, message.as_str(), error.code.as_str(), error.detail.as_deref()), (400, "bad", "x", Some("why"))) + } + e => panic!("{e:?}"), + } + match c.read("s", "b", None, None).await.unwrap_err() { + Error::Api { status, message, .. } => assert_eq!((status, message.as_str()), (502, "bad gateway")), + e => panic!("{e:?}"), + } + assert_eq!(c.read("s", "c", None, None).await.unwrap_err().to_string(), "nope"); +} + +#[test] +fn headers_filters() { + let o = FuncSpecOptions { + filters: vec![ + FilterOption::new("status", "eq", "active"), + FilterOption::new("age", "gte", 18), + FilterOption::new("name", "contains", "x").or(), + FilterOption::new("deleted", "is_null", serde_json::Value::Null), + FilterOption::new("id", "in", json!([1, 2])), + FilterOption::new("p", "between_inclusive", json!([1, 5])), + ], + ..Default::default() + }; + let h = build_headers(&o); + let want: std::collections::BTreeMap = [ + ("X-FieldFilter-status", "active"), + ("X-SearchOp-greaterthanorequal-age", "18"), + ("X-SearchOr-contains-name", "x"), + ("X-SearchOp-empty-deleted", ""), + ("X-SearchOp-in-id", "1,2"), + ("X-SearchOp-betweeninclusive-p", "1,5"), + ] + .into_iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + assert_eq!(h, want); +} + +#[test] +fn headers_misc_and_encoding() { + let o = FuncSpecOptions { + search_filters: [("name".to_string(), "bob".to_string())].into(), + custom_sql_where: Some("a = 1".into()), + custom_sql_or: Some("b = 2".into()), + sort: vec![SortOption::new("name", "asc"), SortOption::new("created_at", "DESC")], + limit: Some(5), + offset: Some(10), + distinct: Some(true), + skip_count: Some(true), + skip_cache: Some(false), + response_format: Some("syncfusion".into()), + ..Default::default() + }; + let h = build_headers(&o); + assert_eq!(h["X-Sort"], "name ASC,created_at DESC"); + assert_eq!(h["X-SearchFilter-name"], "bob"); + assert_eq!(h["X-Custom-SQL-W"], "a = 1"); + assert_eq!(h["X-Limit"], "5"); + assert_eq!(h["X-SkipCache"], "false"); + assert_eq!(h["X-Syncfusion"], "true"); + + let o = FuncSpecOptions { filters: vec![FilterOption::new("n", "eq", "héllo"), FilterOption::new("m", "eq", " pad")], ..Default::default() }; + let h = build_headers(&o); + assert!(h["X-FieldFilter-n"].starts_with("ZIP_")); + assert_eq!(decode_header_value(&h["X-FieldFilter-n"]), "héllo"); + assert_eq!(decode_header_value(&h["X-FieldFilter-m"]), " pad"); +} + +#[test] +fn query_building() { + let mut p = Params::new(); + p.insert("a".into(), true.into()); + p.insert("b".into(), vec!["x".to_string(), "y".to_string()].into()); + p.insert("d".into(), 3i64.into()); + assert_eq!(build_query(&p), vec![("a".into(), "true".into()), ("b".into(), "x".into()), ("b".into(), "y".into()), ("d".into(), "3".into())]); +} + +#[tokio::test] +async fn query_list_metadata() { + let srv = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/users")) + .and(query_param("org", "1")) + .and(header("x-limit", "2")) + .respond_with(ResponseTemplate::new(206).insert_header("Content-Range", "items 10-12/50").set_body_json(json!([{"id": 1}, {"id": 2}]))) + .expect(1) + .mount(&srv) + .await; + let c = FuncSpecClient::from_builder(ClientBuilder::new(&srv.uri()).token("tok")).unwrap(); + let mut p = Params::new(); + p.insert("org".into(), 1i64.into()); + let r = c.query_list("/api/users", &p, Some(&FuncSpecOptions { limit: Some(2), ..Default::default() })).await.unwrap(); + assert_eq!(r.metadata.unwrap(), Metadata { total: 50, count: 2, filtered: 50, limit: 2, offset: 10 }); + assert_eq!(r.data.as_array().unwrap().len(), 2); +} + +#[tokio::test] +async fn query_single_and_error() { + let srv = MockServer::start().await; + Mock::given(path("/api/ok")).respond_with(ResponseTemplate::new(200).set_body_json(json!({"id": 1}))).mount(&srv).await; + Mock::given(path("/api/bad")).respond_with(ResponseTemplate::new(400).set_body_json(json!({"success": false, "error": {"code": "hook_error", "message": "Hook execution failed", "detail": "authentication required"}}))).mount(&srv).await; + let c = FuncSpecClient::new(&srv.uri()).unwrap(); + let r = c.query("api/ok", &Params::new(), None).await.unwrap(); + assert!(r.metadata.is_none()); + assert_eq!(r.data["id"], 1); + match c.query("api/bad", &Params::new(), None).await.unwrap_err() { + Error::Api { error, .. } => assert_eq!((error.code.as_str(), error.detail.as_deref()), ("hook_error", Some("authentication required"))), + e => panic!("{e:?}"), + } +} From ce706bacda3c1dfde25971d8cccd7b5b58732ec0 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:42:33 +0200 Subject: [PATCH 08/23] feat(tx): run update re-fetch and post-commit hooks in a second transaction Re-fetch, BeforeScan and AfterUpdate/AfterCreate now run on a short transaction that fires OnTxBegin. Existence selects inside the first transaction use tx instead of the pool. --- audit/single_tran.md | 8 +- pkg/resolvespec/handler.go | 89 ++++++++++++--------- pkg/resolvespec/update_tx_test.go | 48 +++++++++++ pkg/restheadspec/handler.go | 89 ++++++++++++--------- pkg/restheadspec/update_tx_test.go | 123 +++++++++++++++++++++++++++++ 5 files changed, 278 insertions(+), 79 deletions(-) create mode 100644 pkg/resolvespec/update_tx_test.go create mode 100644 pkg/restheadspec/update_tx_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index aaecd87..2bbec0c 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -74,7 +74,7 @@ | 0 | DONE | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above | | 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | | 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 | -| 3 | TODO | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | +| 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | | 4 | TODO | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | @@ -86,8 +86,10 @@ - DONE infra: `sqlmock` delete tx tests (both specs); compose test server + `scripts/testserver-smoke.sh` (podman first); testmodels ids now serial. - NOTE: restheadspec single delete still does the lookup before `BeforeDelete`; safe once `OnTxBegin` (P2) exists. An `AfterDelete` failure now rolls the delete back. - 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). -- FOUND (P3 scope): resolvespec batch update (`handler.go` ~`:1377`, `:1529`) reads existing record via `h.db.NewSelect()` inside the tx = pool connection; should be `tx`. -- NEXT: P3. +- 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. +- 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. +- NEXT: P4. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index 9c7e9f7..8827a15 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -1210,7 +1210,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Now read the existing record from the database existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...) + selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...) // Apply conditions to select, based on the resolved target ID // (URL ID, request ID, or the "id" field embedded in the data payload). @@ -1292,22 +1292,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url return } - // Fetch the updated record after the transaction commits to capture any trigger changes + // Fetch the updated record in a second short transaction after the first + // commit to capture any trigger changes (OnTxBegin re-applies RLS state). updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - fetchQuery := h.db.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...) - if urlID != "" { - fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID) - } else if reqID != nil { - switch id := reqID.(type) { - case string: - fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) - case []string: - if len(id) > 0 { - fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id) + if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { + fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...) + if urlID != "" { + fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID) + } else if reqID != nil { + switch id := reqID.(type) { + case string: + fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) + case []string: + if len(id) > 0 { + fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id) + } } } - } - if err := fetchQuery.ScanModel(ctx); err != nil { + return fetchQuery.ScanModel(ctx) + }); err != nil { logger.Error("Failed to fetch updated record: %v", err) h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) return @@ -1375,7 +1378,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // First, read the existing record existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { continue // Skip if record not found @@ -1441,19 +1444,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url return } - // Fetch updated records after the transaction commits to capture any trigger changes + // Fetch updated records in a second short transaction after the first commit + // to capture any trigger changes (OnTxBegin re-applies RLS state). fetchedUpdates := make([]interface{}, 0, len(updates)) - for _, item := range updates { - if itemID, ok := item["id"]; ok && itemID != nil { - fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) - if err := fetchQuery.ScanModel(ctx); err != nil { - logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err) - h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) - return + if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { + for _, item := range updates { + if itemID, ok := item["id"]; ok && itemID != nil { + fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() + fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + if err := fetchQuery.ScanModel(ctx); err != nil { + return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err) + } + fetchedUpdates = append(fetchedUpdates, fetchedRecord) } - fetchedUpdates = append(fetchedUpdates, fetchedRecord) } + return nil + }); err != nil { + logger.Error("Failed to fetch updated records: %v", err) + h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) + return } logger.Info("Successfully updated %d records", len(fetchedUpdates)) @@ -1524,7 +1533,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // First, read the existing record existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { continue // Skip if record not found @@ -1593,21 +1602,27 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url return } - // Fetch updated records after the transaction commits to capture any trigger changes + // Fetch updated records in a second short transaction after the first commit + // to capture any trigger changes (OnTxBegin re-applies RLS state). fetchedList := make([]interface{}, 0, len(list)) - for _, item := range list { - if itemMap, ok := item.(map[string]interface{}); ok { - if itemID, ok := itemMap["id"]; ok && itemID != nil { - fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) - if err := fetchQuery.ScanModel(ctx); err != nil { - logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err) - h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) - return + if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { + for _, item := range list { + if itemMap, ok := item.(map[string]interface{}); ok { + if itemID, ok := itemMap["id"]; ok && itemID != nil { + fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() + fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) + if err := fetchQuery.ScanModel(ctx); err != nil { + return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err) + } + fetchedList = append(fetchedList, fetchedRecord) } - fetchedList = append(fetchedList, fetchedRecord) } } + return nil + }); err != nil { + logger.Error("Failed to fetch updated records: %v", err) + h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) + return } logger.Info("Successfully updated %d records", len(fetchedList)) diff --git a/pkg/resolvespec/update_tx_test.go b/pkg/resolvespec/update_tx_test.go new file mode 100644 index 0000000..2922c84 --- /dev/null +++ b/pkg/resolvespec/update_tx_test.go @@ -0,0 +1,48 @@ +package resolvespec + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + var begins []common.Database + h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error { + begins = append(begins, ctx.Tx) + return nil + }) + + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(begins) != 2 || begins[0] == begins[1] { + t.Fatalf("OnTxBegin must fire once per transaction (2), got %d", len(begins)) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 3daf210..89e9d93 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -1459,10 +1459,8 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat } } - // Execute AfterCreate hooks (runs after the transaction commits, against the - // pooled db — hookCtx.Tx was pointed at the now-closed transaction inside the - // RunInTransaction closure above and must not be reused here). - hookCtx.Tx = h.db + // Execute AfterCreate hooks in a second short transaction (the first has + // committed); OnTxBegin re-applies transaction-local state to it. var responseData interface{} if len(mergedResults) == 1 { responseData = mergedResults[0] @@ -1473,7 +1471,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat } hookCtx.Error = nil - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + if err := h.runInTx(ctx, hookCtx, func(common.Database) error { + return h.hooks.Execute(AfterCreate, hookCtx) + }); err != nil { logger.Error("AfterCreate hook failed: %v", err) h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) return @@ -1571,7 +1571,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // Now read the existing record from the database existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() - selectQuery := h.db.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) + selectQuery := tx.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) if err := selectQuery.ScanModel(ctx); err != nil { if err == sql.ErrNoRows { return fmt.Errorf("record not found with ID: %v", targetID) @@ -1657,43 +1657,54 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id return } - // Fetch the updated record after the transaction commits to capture any trigger changes - fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface() - selectQuery := h.db.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) + // Second short transaction: fetch the updated record after the first commit to + // capture any trigger changes, then run AfterUpdate. OnTxBegin re-applies + // transaction-local state (e.g. RLS settings) to this transaction. + var mergedData interface{} + var errCode, errMsg string + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { + fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface() + selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) - // Execute BeforeScan hooks so row security is re-applied to the post-update - // re-fetch, same as it is for the initial read and the update query itself. - // Without this, the re-fetch can return a row the caller isn't authorized to see. - // The transaction has already committed by this point, so hooks must use the - // pooled connection rather than the now-dead tx. - hookCtx.Tx = h.db - hookCtx.Query = selectQuery - if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil { - logger.Error("BeforeScan hook failed: %v", err) - h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) - return - } - if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok { - selectQuery = modifiedQuery - } + // Execute BeforeScan hooks so row security is re-applied to the post-update + // re-fetch, same as it is for the initial read and the update query itself. + // Without this, the re-fetch can return a row the caller isn't authorized to see. + hookCtx.Query = selectQuery + if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil { + logger.Error("BeforeScan hook failed: %v", err) + errCode, errMsg = "hook_error", "Hook execution failed" + return err + } + if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok { + selectQuery = modifiedQuery + } - if err := selectQuery.ScanModel(ctx); err != nil { - logger.Error("Failed to fetch updated record: %v", err) - h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) - return - } - updatedRecord = fetchedRecord + if err := selectQuery.ScanModel(ctx); err != nil { + logger.Error("Failed to fetch updated record: %v", err) + errCode, errMsg = "fetch_error", "Failed to fetch updated record" + return err + } + updatedRecord = fetchedRecord - // Merge the updated record with the original request data - // This preserves extra keys from the request and updates values from the database - mergedData := h.mergeRecordWithRequest(updatedRecord, dataMap) + // Merge the updated record with the original request data + // This preserves extra keys from the request and updates values from the database + mergedData = h.mergeRecordWithRequest(updatedRecord, dataMap) - // Execute AfterUpdate hooks - hookCtx.Result = mergedData - hookCtx.Error = nil - if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { - logger.Error("AfterUpdate hook failed: %v", err) - h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) + // Execute AfterUpdate hooks + hookCtx.Result = mergedData + hookCtx.Error = nil + if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { + logger.Error("AfterUpdate hook failed: %v", err) + errCode, errMsg = "hook_error", "Hook execution failed" + return err + } + return nil + }) + if err != nil { + if errCode == "" { + errCode, errMsg = "fetch_error", "Failed to fetch updated record" + } + h.sendError(w, http.StatusInternalServerError, errCode, errMsg, err) return } diff --git a/pkg/restheadspec/update_tx_test.go b/pkg/restheadspec/update_tx_test.go new file mode 100644 index 0000000..7069a21 --- /dev/null +++ b/pkg/restheadspec/update_tx_test.go @@ -0,0 +1,123 @@ +package restheadspec + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/pgdialect" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" + "github.com/bitechdev/ResolveSpec/pkg/modelregistry" +) + +func TestUpdateRefetchAndAfterUpdateRunInSecondTransaction(t *testing.T) { + // The bun adapter builds model-based updates; the pgsql adapter does not. + sqlDB, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry()) + var begins []common.Database + var order []string + var afterTx common.Database + h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error { + begins = append(begins, ctx.Tx) + order = append(order, "begin") + return nil + }) + h.Hooks().Register(AfterUpdate, func(ctx *HookContext) error { + afterTx = ctx.Tx + order = append(order, "after_update") + return nil + }) + + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + ctx = WithModel(ctx, delItem{}) + h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, ExtendedRequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(begins) != 2 || begins[0] == begins[1] { + t.Fatalf("OnTxBegin must fire once per transaction (2), got %d", len(begins)) + } + if afterTx == nil || afterTx == h.db || afterTx != begins[1] { + t.Fatalf("AfterUpdate must run on the second transaction") + } + if len(order) != 3 || order[2] != "after_update" { + t.Fatalf("unexpected hook order %v", order) + } +} + +func TestAfterCreateRunsInSecondTransaction(t *testing.T) { + sqlDB, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry()) + + var begins []common.Database + var afterTx common.Database + h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error { + begins = append(begins, ctx.Tx) + return nil + }) + h.Hooks().Register(AfterCreate, func(ctx *HookContext) error { + afterTx = ctx.Tx + return nil + }) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + ctx = WithModel(ctx, delItem{}) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, ExtendedRequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(begins) != 2 || afterTx == nil || afterTx == h.db || afterTx != begins[1] { + t.Fatalf("AfterCreate must run on the second transaction, begins=%d", len(begins)) + } +} From 17bb6ea76deaee4ed79433e6eb1d2d3a79bc91ca Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:42:55 +0200 Subject: [PATCH 09/23] fix(clients): verify Dart client, case-insensitive Content-Range lookup --- clients/README.md | 2 +- clients/resolvespec-dart/README.md | 2 - clients/resolvespec-dart/lib/src/client.dart | 28 +++++-- .../resolvespec-dart/lib/src/funcspec.dart | 69 +++++++++++++---- .../resolvespec-dart/lib/src/resolvespec.dart | 53 +++++++++---- clients/resolvespec-dart/lib/src/types.dart | 39 +++++++--- .../resolvespec-dart/test/client_test.dart | 77 ++++++++++++++----- 7 files changed, 200 insertions(+), 70 deletions(-) diff --git a/clients/README.md b/clients/README.md index 1039b0d..b6964d2 100644 --- a/clients/README.md +++ b/clients/README.md @@ -7,6 +7,6 @@ | `resolvespec-go` | Go | ResolveSpec, FunctionSpec | yes (`go test`) | | `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | yes (`cargo test`) | | `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | **not compiled** | -| `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | **not compiled** | +| `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | yes (`dart test`) | Wire behaviour is identical across clients; FunctionSpec server quirks are listed in each README. diff --git a/clients/resolvespec-dart/README.md b/clients/resolvespec-dart/README.md index 825f1af..3a00124 100644 --- a/clients/resolvespec-dart/README.md +++ b/clients/resolvespec-dart/README.md @@ -2,8 +2,6 @@ Dart / Flutter client for ResolveSpec (JSON body) and FunctionSpec. Depends on `package:http`. Dart >= 3.3. -> Not compiled or tested yet (no Dart SDK was available). Run `dart pub get && dart test` first. - ## Clients | Type | Constructor | Methods | diff --git a/clients/resolvespec-dart/lib/src/client.dart b/clients/resolvespec-dart/lib/src/client.dart index 3c2baa8..6bf0367 100644 --- a/clients/resolvespec-dart/lib/src/client.dart +++ b/clients/resolvespec-dart/lib/src/client.dart @@ -10,7 +10,8 @@ class ResolveSpecException implements Exception { final String message; final ApiError error; - ResolveSpecException(this.message, this.statusCode, [ApiError? error]) : error = error ?? ApiError(message: message); + ResolveSpecException(this.message, this.statusCode, [ApiError? error]) + : error = error ?? ApiError(message: message); @override String toString() => 'ResolveSpecException($statusCode): $message'; @@ -25,7 +26,11 @@ class ClientOptions { /// Supply your own client (tests, pooling). final http.Client? httpClient; - const ClientOptions({this.token, this.headers = const {}, this.timeout = const Duration(seconds: 30), this.httpClient}); + const ClientOptions( + {this.token, + this.headers = const {}, + this.timeout = const Duration(seconds: 30), + this.httpClient}); } class Transport { @@ -39,7 +44,8 @@ class Transport { _http = options?.httpClient ?? http.Client(); /// Content-Type < custom headers < per-call headers < bearer token. - Future send(String method, Uri uri, {String? body, Map? extra}) { + Future send(String method, Uri uri, + {String? body, Map? extra}) { final headers = {'Content-Type': 'application/json'}; void merge(Map src) { for (final e in src.entries) { @@ -51,11 +57,16 @@ class Transport { merge(options.headers); if (extra != null) merge(extra); final token = options.token; - if (token != null && token.isNotEmpty) merge({'Authorization': 'Bearer $token'}); + if (token != null && token.isNotEmpty) { + merge({'Authorization': 'Bearer $token'}); + } final req = http.Request(method, uri)..headers.addAll(headers); if (body != null) req.body = body; - return _http.send(req).timeout(options.timeout).then(http.Response.fromStream); + return _http + .send(req) + .timeout(options.timeout) + .then(http.Response.fromStream); } void close() => _http.close(); @@ -67,7 +78,8 @@ class Transport { try { final parsed = jsonDecode(body); isJson = true; - if (parsed is Map && parsed['error'] is Map) { + if (parsed is Map && + parsed['error'] is Map) { err = ApiError.fromJson(parsed['error'] as Map); } } on FormatException { @@ -77,7 +89,9 @@ class Transport { if (message.isEmpty) { var text = isJson ? '' : body.trim(); if (text.length > 200) text = text.substring(0, 200); - message = text.isNotEmpty ? text : '${resp.reasonPhrase ?? 'Error'} (${resp.statusCode})'; + message = text.isNotEmpty + ? text + : '${resp.reasonPhrase ?? 'Error'} (${resp.statusCode})'; } return ResolveSpecException(message, resp.statusCode, err); } diff --git a/clients/resolvespec-dart/lib/src/funcspec.dart b/clients/resolvespec-dart/lib/src/funcspec.dart index 49e9226..3868c5f 100644 --- a/clients/resolvespec-dart/lib/src/funcspec.dart +++ b/clients/resolvespec-dart/lib/src/funcspec.dart @@ -62,7 +62,8 @@ const _operatorMap = { String _scalar(Object? v) => v == null ? '' : v.toString(); -String _filterValue(Object? v) => v is Iterable ? v.map(_scalar).join(',') : _scalar(v); +String _filterValue(Object? v) => + v is Iterable ? v.map(_scalar).join(',') : _scalar(v); /// Base64 (UTF-8) with the `ZIP_` prefix. String encodeHeaderValue(String v) => 'ZIP_${base64.encode(utf8.encode(v))}'; @@ -85,7 +86,8 @@ String decodeHeaderValue(String v) { /// Encode values that are unsafe as raw header/query text (non-ASCII, control chars, edge spaces). String _safe(String v) { - final unsafe = v != v.trim() || v.runes.any((c) => c > 127 || c < 32 || c == 127); + final unsafe = + v != v.trim() || v.runes.any((c) => c > 127 || c < 32 || c == 127); return unsafe ? encodeHeaderValue(v) : v; } @@ -104,12 +106,20 @@ Map buildHeaders(FuncSpecOptions? o) { h['$kind-${_operatorMap[f.operator] ?? f.operator}-${f.column}'] = v; } } - o.searchFilters?.forEach((col, text) => h['X-SearchFilter-$col'] = _safe(text)); - if (o.customSqlWhere != null && o.customSqlWhere!.isNotEmpty) h['X-Custom-SQL-W'] = _safe(o.customSqlWhere!); - if (o.customSqlOr != null && o.customSqlOr!.isNotEmpty) h['X-Custom-SQL-Or'] = _safe(o.customSqlOr!); + o.searchFilters + ?.forEach((col, text) => h['X-SearchFilter-$col'] = _safe(text)); + if (o.customSqlWhere != null && o.customSqlWhere!.isNotEmpty) { + h['X-Custom-SQL-W'] = _safe(o.customSqlWhere!); + } + if (o.customSqlOr != null && o.customSqlOr!.isNotEmpty) { + h['X-Custom-SQL-Or'] = _safe(o.customSqlOr!); + } if (o.sort != null && o.sort!.isNotEmpty) { // funcspec puts this verbatim into ORDER BY - h['X-Sort'] = _safe(o.sort!.map((s) => '${s.column} ${s.direction.toLowerCase() == 'desc' ? 'DESC' : 'ASC'}').join(',')); + h['X-Sort'] = _safe(o.sort! + .map((s) => + '${s.column} ${s.direction.toLowerCase() == 'desc' ? 'DESC' : 'ASC'}') + .join(',')); } if (o.limit != null) h['X-Limit'] = '${o.limit}'; if (o.offset != null) h['X-Offset'] = '${o.offset}'; @@ -132,11 +142,20 @@ Map> buildQuery(Map? params) { final out = >{}; params?.forEach((k, v) { if (v == null) return; - out[k] = v is Iterable ? v.map((e) => _safe(_scalar(e))).toList() : [_safe(_scalar(v))]; + out[k] = v is Iterable + ? v.map((e) => _safe(_scalar(e))).toList() + : [_safe(_scalar(v))]; }); return out; } +String? _header(Map headers, String name) { + for (final e in headers.entries) { + if (e.key.toLowerCase() == name) return e.value; + } + return null; +} + final _contentRange = RegExp(r'(\d+)-(\d+)/(\d+)'); Metadata _metadata(String? contentRange, FuncSpecOptions? o) { @@ -145,42 +164,60 @@ Metadata _metadata(String? contentRange, FuncSpecOptions? o) { final start = int.parse(m.group(1)!); final end = int.parse(m.group(2)!); final total = int.parse(m.group(3)!); - return Metadata(total: total, count: end - start, filtered: total, limit: o?.limit ?? 0, offset: start); + return Metadata( + total: total, + count: end - start, + filtered: total, + limit: o?.limit ?? 0, + offset: start); } /// Client for user-defined SQL endpoints. Routes are defined by the server application. class FuncSpecClient { final Transport _t; - FuncSpecClient(String baseUrl, [ClientOptions? options]) : _t = Transport(baseUrl, options); + FuncSpecClient(String baseUrl, [ClientOptions? options]) + : _t = Transport(baseUrl, options); void close() => _t.close(); - Future _call(String method, String path, Map? params, FuncSpecOptions? o, bool list) async { - final base = Uri.parse('${_t.baseUrl}/${path.replaceAll(RegExp(r'^/+'), '')}'); + Future _call(String method, String path, + Map? params, FuncSpecOptions? o, bool list) async { + final base = + Uri.parse('${_t.baseUrl}/${path.replaceAll(RegExp(r'^/+'), '')}'); final pairs = []; buildQuery(params).forEach((k, vs) { for (final v in vs) { - pairs.add('${Uri.encodeQueryComponent(k)}=${Uri.encodeQueryComponent(v)}'); + pairs.add( + '${Uri.encodeQueryComponent(k)}=${Uri.encodeQueryComponent(v)}'); } }); final uri = pairs.isEmpty ? base : base.replace(query: pairs.join('&')); final resp = await _t.send(method, uri, extra: buildHeaders(o)); - if (resp.statusCode < 200 || resp.statusCode > 299) throw Transport.errorFrom(resp); // 206 is success + if (resp.statusCode < 200 || resp.statusCode > 299) { + throw Transport.errorFrom(resp); // 206 is success + } final text = utf8.decode(resp.bodyBytes); return Response( success: true, data: text.trim().isEmpty ? null : jsonDecode(text), - metadata: list ? _metadata(resp.headers['content-range'], o) : null, + metadata: + list ? _metadata(_header(resp.headers, 'content-range'), o) : null, ); } /// Single-record endpoint (SqlQuery). `data` is the row object. - Future query(String path, {Map? params, FuncSpecOptions? options, String method = 'GET'}) => + Future query(String path, + {Map? params, + FuncSpecOptions? options, + String method = 'GET'}) => _call(method.toUpperCase(), path, params, options, false); /// List endpoint (SqlQueryList). Metadata comes from Content-Range. - Future queryList(String path, {Map? params, FuncSpecOptions? options, String method = 'GET'}) => + Future queryList(String path, + {Map? params, + FuncSpecOptions? options, + String method = 'GET'}) => _call(method.toUpperCase(), path, params, options, true); } diff --git a/clients/resolvespec-dart/lib/src/resolvespec.dart b/clients/resolvespec-dart/lib/src/resolvespec.dart index 1ec5cf8..b3cb73b 100644 --- a/clients/resolvespec-dart/lib/src/resolvespec.dart +++ b/clients/resolvespec-dart/lib/src/resolvespec.dart @@ -9,45 +9,70 @@ import 'types.dart'; class ResolveSpecClient { final Transport _t; - ResolveSpecClient(String baseUrl, [ClientOptions? options]) : _t = Transport(baseUrl, options); + ResolveSpecClient(String baseUrl, [ClientOptions? options]) + : _t = Transport(baseUrl, options); void close() => _t.close(); - static String? _urlId(Object? id) => id == null || id is List ? null : id.toString(); + static String? _urlId(Object? id) => + id == null || id is List ? null : id.toString(); - static List? _bodyId(Object? id) => id is List ? id.map((e) => e.toString()).toList() : null; + static List? _bodyId(Object? id) => + id is List ? id.map((e) => e.toString()).toList() : null; Uri _url(String schema, String entity, String? id) { - var u = '${_t.baseUrl}/${Uri.encodeComponent(schema)}/${Uri.encodeComponent(entity)}'; + var u = + '${_t.baseUrl}/${Uri.encodeComponent(schema)}/${Uri.encodeComponent(entity)}'; if (id != null && id.isNotEmpty) u += '/${Uri.encodeComponent(id)}'; return Uri.parse(u); } - Future _send(String method, Uri url, Map? body) async { - final resp = await _t.send(method, url, body: body == null ? null : jsonEncode(body)); - if (resp.statusCode < 200 || resp.statusCode > 299) throw Transport.errorFrom(resp); + Future _send( + String method, Uri url, Map? body) async { + final resp = await _t.send(method, url, + body: body == null ? null : jsonEncode(body)); + if (resp.statusCode < 200 || resp.statusCode > 299) { + throw Transport.errorFrom(resp); + } final decoded = jsonDecode(utf8.decode(resp.bodyBytes)); final r = Response.fromJson(decoded as Map); - if (!r.success && r.error != null) throw ResolveSpecException(r.error!.message, resp.statusCode, r.error); + if (!r.success && r.error != null) { + throw ResolveSpecException(r.error!.message, resp.statusCode, r.error); + } return r; } /// GET /{schema}/{entity} - Future getMetadata(String schema, String entity) => _send('GET', _url(schema, entity, null), null); + Future getMetadata(String schema, String entity) => + _send('GET', _url(schema, entity, null), null); - Future read(String schema, String entity, {Object? id, Options? options}) => _send( + Future read(String schema, String entity, + {Object? id, Options? options}) => + _send( 'POST', _url(schema, entity, _urlId(id)), - {'operation': 'read', if (_bodyId(id) != null) 'id': _bodyId(id), if (options != null) 'options': options.toJson()}, + { + 'operation': 'read', + if (_bodyId(id) != null) 'id': _bodyId(id), + if (options != null) 'options': options.toJson() + }, ); - Future create(String schema, String entity, Object data, {Options? options}) => _send( + Future create(String schema, String entity, Object data, + {Options? options}) => + _send( 'POST', _url(schema, entity, null), - {'operation': 'create', 'data': data, if (options != null) 'options': options.toJson()}, + { + 'operation': 'create', + 'data': data, + if (options != null) 'options': options.toJson() + }, ); - Future update(String schema, String entity, Object data, {Object? id, Options? options}) => _send( + Future update(String schema, String entity, Object data, + {Object? id, Options? options}) => + _send( 'POST', _url(schema, entity, _urlId(id)), { diff --git a/clients/resolvespec-dart/lib/src/types.dart b/clients/resolvespec-dart/lib/src/types.dart index f56d613..1aaac6f 100644 --- a/clients/resolvespec-dart/lib/src/types.dart +++ b/clients/resolvespec-dart/lib/src/types.dart @@ -16,7 +16,8 @@ class FilterOption { /// AND | OR final String? logicOperator; - const FilterOption(this.column, this.operator, [this.value, this.logicOperator]); + const FilterOption(this.column, this.operator, + [this.value, this.logicOperator]); Map toJson() => _compact({ 'column': column, @@ -44,7 +45,8 @@ class Parameter { const Parameter(this.name, this.value, [this.sequence]); - Map toJson() => _compact({'name': name, 'value': value, 'sequence': sequence}); + Map toJson() => + _compact({'name': name, 'value': value, 'sequence': sequence}); } class CustomOperator { @@ -139,10 +141,16 @@ class VectorSearchOption { final String? as; final String? direction; - const VectorSearchOption(this.column, this.vector, {this.metric, this.as, this.direction}); + const VectorSearchOption(this.column, this.vector, + {this.metric, this.as, this.direction}); - Map toJson() => - _compact({'column': column, 'vector': vector, 'metric': metric, 'as': as, 'direction': direction}); + Map toJson() => _compact({ + 'column': column, + 'vector': vector, + 'metric': metric, + 'as': as, + 'direction': direction + }); } /// ResolveSpec request options object. @@ -204,7 +212,12 @@ class Metadata { final int limit; final int offset; - const Metadata({this.total = 0, this.count = 0, this.filtered = 0, this.limit = 0, this.offset = 0}); + const Metadata( + {this.total = 0, + this.count = 0, + this.filtered = 0, + this.limit = 0, + this.offset = 0}); factory Metadata.fromJson(Map j) => Metadata( total: (j['total'] as num?)?.toInt() ?? 0, @@ -227,7 +240,8 @@ class Metadata { int get hashCode => Object.hash(total, count, filtered, limit, offset); @override - String toString() => 'Metadata(total: $total, count: $count, filtered: $filtered, limit: $limit, offset: $offset)'; + String toString() => + 'Metadata(total: $total, count: $count, filtered: $filtered, limit: $limit, offset: $offset)'; } class ApiError { @@ -239,7 +253,8 @@ class ApiError { final String? detail; final String? sql; - const ApiError({this.code = '', this.message = '', this.details, this.detail, this.sql}); + const ApiError( + {this.code = '', this.message = '', this.details, this.detail, this.sql}); factory ApiError.fromJson(Map j) => ApiError( code: (j['code'] as String?) ?? '', @@ -262,7 +277,11 @@ class Response { factory Response.fromJson(Map j) => Response( success: j['success'] == true, data: j['data'], - metadata: j['metadata'] is Map ? Metadata.fromJson(j['metadata'] as Map) : null, - error: j['error'] is Map ? ApiError.fromJson(j['error'] as Map) : null, + metadata: j['metadata'] is Map + ? Metadata.fromJson(j['metadata'] as Map) + : null, + error: j['error'] is Map + ? ApiError.fromJson(j['error'] as Map) + : null, ); } diff --git a/clients/resolvespec-dart/test/client_test.dart b/clients/resolvespec-dart/test/client_test.dart index 7e213a1..afe1310 100644 --- a/clients/resolvespec-dart/test/client_test.dart +++ b/clients/resolvespec-dart/test/client_test.dart @@ -5,12 +5,14 @@ import 'package:http/testing.dart'; import 'package:resolvespec/resolvespec.dart'; import 'package:test/test.dart'; -(http.Client, List) stub(int status, Object body, {Map headers = const {}}) { +(http.Client, List) stub(int status, Object body, + {Map headers = const {}}) { final seen = []; final client = MockClient((req) async { seen.add(req); final text = body is String ? body : jsonEncode(body); - return http.Response(text, status, headers: {'content-type': 'application/json', ...headers}); + return http.Response(text, status, + headers: {'content-type': 'application/json', ...headers}); }); return (client, seen); } @@ -18,13 +20,19 @@ import 'package:test/test.dart'; void main() { group('resolvespec', () { test('read posts body with headers', () async { - final (c, seen) = stub(200, {'success': true, 'data': [{'id': 1}]}); + final (c, seen) = stub(200, { + 'success': true, + 'data': [ + {'id': 1} + ] + }); final client = ResolveSpecClient( 'http://localhost:3000/', ClientOptions(token: 'tok', headers: {'X-Tenant': 'a'}, httpClient: c), ); final r = await client.read('public', 'users', - options: const Options(limit: 5, filters: [FilterOption('a', 'eq', 1)])); + options: + const Options(limit: 5, filters: [FilterOption('a', 'eq', 1)])); final req = seen.single; expect(req.method, 'POST'); expect(req.url.path, '/public/users'); @@ -39,7 +47,8 @@ void main() { test('id placement', () async { final (c, seen) = stub(200, {'success': true, 'data': {}}); - final client = ResolveSpecClient('http://x', ClientOptions(httpClient: c)); + final client = + ResolveSpecClient('http://x', ClientOptions(httpClient: c)); await client.read('s', 'e', id: 7); expect(seen.last.url.path, '/s/e/7'); await client.update('s', 'e', {'a': 1}, id: ['1', '2']); @@ -69,10 +78,12 @@ void main() { .having((e) => e.message, 'message', 'bad') .having((e) => e.error.detail, 'detail', 'why')), ); - final plain = ResolveSpecClient('http://x', ClientOptions(httpClient: stub(502, 'bad gateway').$1)); + final plain = ResolveSpecClient( + 'http://x', ClientOptions(httpClient: stub(502, 'bad gateway').$1)); await expectLater( plain.read('s', 'e'), - throwsA(isA().having((e) => e.message, 'message', 'bad gateway')), + throwsA(isA() + .having((e) => e.message, 'message', 'bad gateway')), ); final soft = ResolveSpecClient( 'http://x', @@ -82,7 +93,10 @@ void main() { 'error': {'code': 'c', 'message': 'nope'} }).$1), ); - await expectLater(soft.read('s', 'e'), throwsA(isA().having((e) => e.message, 'message', 'nope'))); + await expectLater( + soft.read('s', 'e'), + throwsA(isA() + .having((e) => e.message, 'message', 'nope'))); }); }); @@ -135,27 +149,45 @@ void main() { }); test('query building', () { - expect(buildQuery({'a': true, 'b': ['x', 'y'], 'c': null, 'd': 3}), { - 'a': ['true'], - 'b': ['x', 'y'], - 'd': ['3'], - }); + expect( + buildQuery({ + 'a': true, + 'b': ['x', 'y'], + 'c': null, + 'd': 3 + }), + { + 'a': ['true'], + 'b': ['x', 'y'], + 'd': ['3'], + }); }); test('queryList metadata', () async { - final (c, seen) = stub(206, [{'id': 1}, {'id': 2}], headers: {'Content-Range': 'items 10-12/50'}); - final client = FuncSpecClient('http://x', ClientOptions(token: 'tok', httpClient: c)); - final r = await client.queryList('/api/users', params: {'org': 1}, options: const FuncSpecOptions(limit: 2)); + final (c, seen) = stub(206, [ + {'id': 1}, + {'id': 2} + ], headers: { + 'Content-Range': 'items 10-12/50' + }); + final client = FuncSpecClient( + 'http://x', ClientOptions(token: 'tok', httpClient: c)); + final r = await client.queryList('/api/users', + params: {'org': 1}, options: const FuncSpecOptions(limit: 2)); expect(seen.single.method, 'GET'); expect(seen.single.url.path, '/api/users'); expect(seen.single.url.query, 'org=1'); expect(seen.single.headers['x-limit'], '2'); - expect(r.metadata, const Metadata(total: 50, count: 2, filtered: 50, limit: 2, offset: 10)); + expect( + r.metadata, + const Metadata( + total: 50, count: 2, filtered: 50, limit: 2, offset: 10)); expect((r.data as List).length, 2); }); test('query single and error', () async { - final ok = FuncSpecClient('http://x', ClientOptions(httpClient: stub(200, {'id': 1}).$1)); + final ok = FuncSpecClient( + 'http://x', ClientOptions(httpClient: stub(200, {'id': 1}).$1)); final r = await ok.query('api/u'); expect(r.metadata, isNull); expect((r.data as Map)['id'], 1); @@ -165,14 +197,19 @@ void main() { ClientOptions( httpClient: stub(400, { 'success': false, - 'error': {'code': 'hook_error', 'message': 'Hook execution failed', 'detail': 'authentication required'} + 'error': { + 'code': 'hook_error', + 'message': 'Hook execution failed', + 'detail': 'authentication required' + } }).$1), ); await expectLater( bad.query('api/u'), throwsA(isA() .having((e) => e.error.code, 'code', 'hook_error') - .having((e) => e.error.detail, 'detail', 'authentication required')), + .having( + (e) => e.error.detail, 'detail', 'authentication required')), ); }); }); From 97bcb44fdc5479518585980e9565c7b11c0f2261 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:44:15 +0200 Subject: [PATCH 10/23] fix(clients): verify C# client, read Content-Range from content headers --- clients/README.md | 2 +- clients/resolvespec-cs/README.md | 2 +- clients/resolvespec-cs/src/FuncSpecClient.cs | 4 +++- clients/resolvespec-cs/src/ResolveSpecClient.cs | 2 +- clients/resolvespec-cs/tests/ClientTests.cs | 3 ++- 5 files changed, 8 insertions(+), 5 deletions(-) diff --git a/clients/README.md b/clients/README.md index b6964d2..f456603 100644 --- a/clients/README.md +++ b/clients/README.md @@ -6,7 +6,7 @@ | `resolvespec-python` | Python >= 3.11 | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | yes (61 tests) | | `resolvespec-go` | Go | ResolveSpec, FunctionSpec | yes (`go test`) | | `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | yes (`cargo test`) | -| `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | **not compiled** | +| `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | yes (`dotnet test`) | | `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | yes (`dart test`) | Wire behaviour is identical across clients; FunctionSpec server quirks are listed in each README. diff --git a/clients/resolvespec-cs/README.md b/clients/resolvespec-cs/README.md index 41488e5..4853af8 100644 --- a/clients/resolvespec-cs/README.md +++ b/clients/resolvespec-cs/README.md @@ -2,7 +2,7 @@ .NET 8 client for ResolveSpec (JSON body) and FunctionSpec. `System.Text.Json`, no other dependencies. -> Not compiled or tested yet (no .NET SDK was available). Run `dotnet test tests/` first. +> Tests run with `DOTNET_ROLL_FORWARD=Major` when only a newer runtime than 8.0 is installed. ## Clients diff --git a/clients/resolvespec-cs/src/FuncSpecClient.cs b/clients/resolvespec-cs/src/FuncSpecClient.cs index acdc588..84df0de 100644 --- a/clients/resolvespec-cs/src/FuncSpecClient.cs +++ b/clients/resolvespec-cs/src/FuncSpecClient.cs @@ -169,7 +169,9 @@ public sealed class FuncSpecClient }; if (list) { - resp.Headers.TryGetValues("Content-Range", out var cr); + // Content-Range is a content header in HttpClient; fall back to response headers. + IEnumerable? cr = null; + if (!resp.Content.Headers.TryGetValues("Content-Range", out cr)) resp.Headers.TryGetValues("Content-Range", out cr); r.Metadata = MetadataFrom(cr?.FirstOrDefault(), o); } return r; diff --git a/clients/resolvespec-cs/src/ResolveSpecClient.cs b/clients/resolvespec-cs/src/ResolveSpecClient.cs index c57e891..6b7bfb0 100644 --- a/clients/resolvespec-cs/src/ResolveSpecClient.cs +++ b/clients/resolvespec-cs/src/ResolveSpecClient.cs @@ -28,7 +28,7 @@ public sealed class ResolveSpecClient _ => Convert.ToString(id, System.Globalization.CultureInfo.InvariantCulture), }; - static string[]? BodyId(object? id) => id is IEnumerable e and not string ? e.ToArray() : null; + static string[]? BodyId(object? id) => id is IEnumerable e ? e.ToArray() : null; string Url(string schema, string entity, string? id) { diff --git a/clients/resolvespec-cs/tests/ClientTests.cs b/clients/resolvespec-cs/tests/ClientTests.cs index 4a88311..075f7fc 100644 --- a/clients/resolvespec-cs/tests/ClientTests.cs +++ b/clients/resolvespec-cs/tests/ClientTests.cs @@ -22,7 +22,8 @@ public class Stub : HttpMessageHandler Request = request; Body = request.Content == null ? "" : await request.Content.ReadAsStringAsync(ct); var r = new HttpResponseMessage(_status) { Content = new StringContent(_json, Encoding.UTF8, "application/json") }; - foreach (var (k, v) in _headers) r.Headers.TryAddWithoutValidation(k, v); + foreach (var (k, v) in _headers) + if (!r.Headers.TryAddWithoutValidation(k, v)) r.Content.Headers.TryAddWithoutValidation(k, v); return r; } } From ed457eb14afedce01a22f8155b7d4831adab13d5 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:46:46 +0200 Subject: [PATCH 11/23] feat(tx): run websocketspec and mqttspec operations in per-message transactions Reads and deletes run in one transaction; create and update write in one and re-fetch plus After hooks in a second. OnTxBegin fires first in each. --- audit/single_tran.md | 6 +- pkg/mqttspec/handler.go | 256 ++++++++++++++++++------------- pkg/mqttspec/handler_test.go | 9 +- pkg/mqttspec/hooks.go | 5 + pkg/mqttspec/tx_test.go | 91 +++++++++++ pkg/websocketspec/handler.go | 288 ++++++++++++++++++++--------------- pkg/websocketspec/hooks.go | 10 ++ pkg/websocketspec/tx_test.go | 173 +++++++++++++++++++++ 8 files changed, 611 insertions(+), 227 deletions(-) create mode 100644 pkg/mqttspec/tx_test.go create mode 100644 pkg/websocketspec/tx_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 2bbec0c..83dbe22 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -75,7 +75,7 @@ | 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | | 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 | | 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | -| 4 | TODO | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | +| 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | | 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | @@ -89,7 +89,9 @@ - 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. - 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. -- NEXT: P4. +- 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"). +- NEXT: P5. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/mqttspec/handler.go b/pkg/mqttspec/handler.go index 841a0c9..98185cf 100644 --- a/pkg/mqttspec/handler.go +++ b/pkg/mqttspec/handler.go @@ -3,6 +3,7 @@ package mqttspec import ( "context" "encoding/json" + "errors" "fmt" "reflect" "strings" @@ -313,42 +314,73 @@ func (h *Handler) handleRequest(client *Client, msg *Message) { } } -// handleRead processes a read operation +// stageError marks which stage of an operation failed inside a transaction so the +// right error response is sent once the transaction has rolled back. +type stageError struct { + code string + err error +} + +func (e *stageError) Error() string { return e.err.Error() } +func (e *stageError) Unwrap() error { return e.err } + +func hookStage(err error) error { return &stageError{code: "hook_error", err: err} } +func opStage(code string, err error) error { return &stageError{code: code, err: err} } + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Transactions are per message, never per client connection. +func (h *Handler) runInTx(hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(hookCtx.Context, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} + +// sendTxError sends the response for an error returned by runInTx. Errors that are +// not stage errors (begin, OnTxBegin, commit) carry no detail to the client. +func (h *Handler) sendTxError(client *Client, msgID string, err error) { + var stage *stageError + if errors.As(err, &stage) { + logger.Error("[MQTTSpec] %s: %v", stage.code, stage.err) + h.sendError(client.ID, msgID, stage.code, stage.err.Error()) + return + } + logger.Error("[MQTTSpec] Transaction failed: %v", err) + h.sendError(client.ID, msgID, "transaction_error", "Transaction failed") +} + +// handleRead processes a read operation; hooks and queries share one transaction. func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeRead hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } - - // Perform read operation - var data interface{} var metadata map[string]interface{} - var err error - if hookCtx.ID != "" { - // Read single record by ID - data, err = h.readByID(hookCtx) - metadata = map[string]interface{}{"total": 1} - } else { - // Read multiple records - data, metadata, err = h.readMultiple(hookCtx) - } + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { + return hookStage(err) + } + var data interface{} + var err error + if hookCtx.ID != "" { + // Read single record by ID + data, err = h.readByID(hookCtx) + metadata = map[string]interface{}{"total": 1} + } else { + // Read multiple records + data, metadata, err = h.readMultiple(hookCtx) + } + if err != nil { + return opStage("read_error", err) + } + + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Read operation failed: %v", err) - h.sendError(client.ID, msg.ID, "read_error", err.Error()) - return - } - - // Update hook context - hookCtx.Result = data - - // Execute after hook - if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterRead hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } @@ -356,30 +388,49 @@ func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) h.sendResponse(client.ID, msg.ID, hookCtx.Result, metadata) } -// handleCreate processes a create operation +// handleCreate processes a create operation. The insert runs in the first +// transaction; the re-fetch (to capture DB defaults/triggers) and AfterCreate run +// in a second short transaction. func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeCreate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } + var data interface{} - // Perform create operation - data, err := h.create(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { + return hookStage(err) + } + + var err error + data, err = h.create(hookCtx) + if err != nil { + return opStage("create_error", err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Create operation failed: %v", err) - h.sendError(client.ID, msg.ID, "create_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { + hookCtx.ID = fmt.Sprintf("%v", pkVal) + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("create_error", err) + } + } - // Execute after hook - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterCreate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -390,30 +441,42 @@ func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContex h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationCreate, data) } -// handleUpdate processes an update operation +// handleUpdate processes an update operation. The update runs in the first +// transaction; the re-fetch and AfterUpdate run in a second short transaction. func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeUpdate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } + var data interface{} - // Perform update operation - data, err := h.update(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { + return hookStage(err) + } + if err := h.update(hookCtx); err != nil { + return opStage("update_error", err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Update operation failed: %v", err) - h.sendError(client.ID, msg.ID, "update_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("update_error", err) + } - // Execute after hook - if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterUpdate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -424,26 +487,22 @@ func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContex h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationUpdate, data) } -// handleDelete processes a delete operation +// handleDelete processes a delete operation; hooks and delete share one transaction. func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeDelete hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } - - // Perform delete operation - if err := h.delete(hookCtx); err != nil { - logger.Error("[MQTTSpec] Delete operation failed: %v", err) - h.sendError(client.ID, msg.ID, "delete_error", err.Error()) - return - } - - // Execute after hook - if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterDelete hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { + return hookStage(err) + } + if err := h.delete(hookCtx); err != nil { + return opStage("delete_error", err) + } + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -671,7 +730,7 @@ func (h *Handler) getTableName(schema, entity string, model interface{}) string // readByID reads a single record by ID func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) @@ -711,7 +770,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { // readMultiple reads multiple records func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata map[string]interface{}, err error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Apply options if hookCtx.Options != nil { @@ -786,7 +845,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata // Get count metadata = make(map[string]interface{}) - countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + countQuery := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if hookCtx.Options != nil { for _, filter := range hookCtx.Options.Filters { if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok { @@ -835,22 +894,16 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { } // Insert record - query := h.db.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } - // Re-fetch the created record to capture DB-generated defaults/triggers. - if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { - hookCtx.ID = fmt.Sprintf("%v", pkVal) - return h.readByID(hookCtx) - } - return hookCtx.ModelPtr, nil } // update updates an existing record -func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { +func (h *Handler) update(hookCtx *HookContext) error { // Convert request data to a map var updates map[string]interface{} if m, ok := hookCtx.Data.(map[string]interface{}); ok { @@ -858,10 +911,10 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { } else { dataBytes, err := json.Marshal(hookCtx.Data) if err != nil { - return nil, fmt.Errorf("failed to marshal data: %w", err) + return fmt.Errorf("failed to marshal data: %w", err) } if err := json.Unmarshal(dataBytes, &updates); err != nil { - return nil, fmt.Errorf("failed to unmarshal data into map: %w", err) + return fmt.Errorf("failed to unmarshal data into map: %w", err) } } @@ -872,21 +925,20 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) if len(values) > 0 { - query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values). + query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) if _, err := query.Exec(hookCtx.Context); err != nil { - return nil, fmt.Errorf("failed to update record: %w", err) + return fmt.Errorf("failed to update record: %w", err) } } - // Fetch updated record - return h.readByID(hookCtx) + return nil } // delete deletes a record func (h *Handler) delete(hookCtx *HookContext) error { - query := h.db.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) diff --git a/pkg/mqttspec/handler_test.go b/pkg/mqttspec/handler_test.go index db30ae2..6cc4bf0 100644 --- a/pkg/mqttspec/handler_test.go +++ b/pkg/mqttspec/handler_test.go @@ -773,7 +773,7 @@ func TestHandler_HandleIncomingMessage_ValidMessage(t *testing.T) { } func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { - newHook := func(id string, data map[string]interface{}) *HookContext { + newHook := func(handler *Handler, id string, data map[string]interface{}) *HookContext { return &HookContext{ Context: context.Background(), TableName: "users", @@ -784,6 +784,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { ID: id, Data: data, Options: &common.RequestOptions{}, + Tx: handler.db, } } seed := func(t *testing.T, db *gorm.DB) { @@ -794,7 +795,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler, db := setupTestHandler(t) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"name": ""})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"name": ""})) require.NoError(t, err) var got TestUser @@ -808,7 +809,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler, db := setupTestHandler(t) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"status": "inactive"})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"status": "inactive"})) require.NoError(t, err) var got TestUser @@ -823,7 +824,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler.SetDisallowNulls(true) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"name": nil, "status": ""})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"name": nil, "status": ""})) require.NoError(t, err) var got TestUser diff --git a/pkg/mqttspec/hooks.go b/pkg/mqttspec/hooks.go index 205be3c..3277441 100644 --- a/pkg/mqttspec/hooks.go +++ b/pkg/mqttspec/hooks.go @@ -50,6 +50,11 @@ const ( // BeforeOp fires immediately before every SQL operation (read, create, update, delete) BeforeOp = websocketspec.BeforeOp + + // OnTxBegin fires once, first, inside every transaction the handler opens for a + // read/create/update/delete message (including the second short transaction for + // post-commit work). hookCtx.Tx is the transaction. + OnTxBegin = websocketspec.OnTxBegin ) // NewHookRegistry creates a new hook registry diff --git a/pkg/mqttspec/tx_test.go b/pkg/mqttspec/tx_test.go new file mode 100644 index 0000000..4293e92 --- /dev/null +++ b/pkg/mqttspec/tx_test.go @@ -0,0 +1,91 @@ +package mqttspec + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func newTxHook(handler *Handler, data interface{}) *HookContext { + return &HookContext{ + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Schema: "public", + Entity: "users", + ID: "1", + Data: data, + Options: &common.RequestOptions{}, + Metadata: map[string]interface{}{}, + Tx: handler.db, + } +} + +func TestHandler_OnTxBeginFiresFirstOnTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var order []string + var txs []common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { + order = append(order, "begin") + txs = append(txs, c.Tx) + return nil + }) + handler.hooks.Register(BeforeDelete, func(c *HookContext) error { + order = append(order, "before") + return nil + }) + + handler.handleDelete(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + var count int64 + require.NoError(t, db.Model(&TestUser{}).Where("id = 1").Count(&count).Error) + assert.Zero(t, count) + require.Len(t, txs, 1) + assert.NotEqual(t, handler.db, txs[0]) + assert.Equal(t, []string{"begin", "before"}, order) +} + +func TestHandler_OnTxBeginErrorAbortsWithoutWriting(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { return errors.New("no user") }) + + handler.handleDelete(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + var count int64 + require.NoError(t, db.Model(&TestUser{}).Where("id = 1").Count(&count).Error) + assert.Equal(t, int64(1), count) +} + +func TestHandler_UpdateRunsAfterHookOnSecondTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var begins []common.Database + var afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { + begins = append(begins, c.Tx) + return nil + }) + handler.hooks.Register(AfterUpdate, func(c *HookContext) error { + afterTx = c.Tx + return nil + }) + + handler.handleUpdate(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, map[string]interface{}{"name": "b"})) + + var got TestUser + require.NoError(t, db.First(&got, 1).Error) + assert.Equal(t, "b", got.Name) + require.Len(t, begins, 2) + assert.NotEqual(t, begins[0], begins[1]) + assert.Equal(t, begins[1], afterTx) +} diff --git a/pkg/websocketspec/handler.go b/pkg/websocketspec/handler.go index 0dbe2f4..352d1d8 100644 --- a/pkg/websocketspec/handler.go +++ b/pkg/websocketspec/handler.go @@ -221,49 +221,84 @@ func (h *Handler) handleRequest(conn *Connection, msg *Message) { } } -// handleRead processes a read operation +// stageError marks which stage of an operation failed inside a transaction so the +// right error response is sent once the transaction has rolled back. +type stageError struct { + code string + hook bool + err error +} + +func (e *stageError) Error() string { return e.err.Error() } +func (e *stageError) Unwrap() error { return e.err } + +func hookStage(err error) error { return &stageError{code: "hook_error", hook: true, err: err} } +func opStage(code string, err error) error { + return &stageError{code: code, err: err} +} + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Transactions are per message, never per connection. +func (h *Handler) runInTx(hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(hookCtx.Context, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} + +// sendTxError sends the response for an error returned by runInTx. Errors that are +// not stage errors (begin, OnTxBegin, commit) carry no detail to the client. +func (h *Handler) sendTxError(conn *Connection, msgID string, err error) { + var stage *stageError + switch { + case errors.As(err, &stage) && stage.hook: + logger.Error("[WebSocketSpec] %s: %v", stage.code, stage.err) + _ = conn.SendJSON(NewErrorResponse(msgID, stage.code, stage.err.Error())) + case errors.As(err, &stage): + logger.Error("[WebSocketSpec] %s: %v", stage.code, stage.err) + _ = conn.SendJSON(newErrorResponseFromErr(msgID, stage.code, stage.err)) + default: + logger.Error("[WebSocketSpec] Transaction failed: %v", err) + _ = conn.SendJSON(NewErrorResponse(msgID, "transaction_error", "Transaction failed")) + } +} + +// handleRead processes a read operation; hooks and queries share one transaction. func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeRead hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } - - // Perform read operation - var data interface{} var metadata map[string]interface{} - var err error - // Check if FetchRowNumber is specified (treat as single record read) - isFetchRowNumber := hookCtx.Options != nil && hookCtx.Options.FetchRowNumber != nil && *hookCtx.Options.FetchRowNumber != "" + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { + return hookStage(err) + } - if hookCtx.ID != "" || isFetchRowNumber { - // Read single record by ID or FetchRowNumber - data, err = h.readByID(hookCtx) - metadata = map[string]interface{}{"total": 1} - // The row number is already set on the record itself via setRowNumbersOnRecords - } else { - // Read multiple records - data, metadata, err = h.readMultiple(hookCtx) - } + var data interface{} + var err error + // Check if FetchRowNumber is specified (treat as single record read) + isFetchRowNumber := hookCtx.Options != nil && hookCtx.Options.FetchRowNumber != nil && *hookCtx.Options.FetchRowNumber != "" + if hookCtx.ID != "" || isFetchRowNumber { + // Read single record by ID or FetchRowNumber + data, err = h.readByID(hookCtx) + metadata = map[string]interface{}{"total": 1} + // The row number is already set on the record itself via setRowNumbersOnRecords + } else { + // Read multiple records + data, metadata, err = h.readMultiple(hookCtx) + } + if err != nil { + return opStage("read_error", err) + } + + // Update hook context with result + hookCtx.Result = data + + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Read operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "read_error", err) - _ = conn.SendJSON(errResp) - return - } - - // Update hook context with result - hookCtx.Result = data - - // Execute after hook - if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterRead hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } @@ -273,33 +308,49 @@ func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContex _ = conn.SendJSON(resp) } -// handleCreate processes a create operation +// handleCreate processes a create operation. The insert runs in the first +// transaction; the re-fetch (to capture DB defaults/triggers) and AfterCreate run +// in a second short transaction. func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeCreate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } + var data interface{} - // Perform create operation - data, err := h.create(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { + return hookStage(err) + } + + var err error + data, err = h.create(hookCtx) + if err != nil { + return opStage("create_error", err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Create operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "create_error", err) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + if reflection.GetPrimaryKeyValue(hookCtx.ModelPtr) != nil { + hookCtx.ID = fmt.Sprintf("%v", reflection.GetPrimaryKeyValue(hookCtx.ModelPtr)) + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("create_error", err) + } + } - // Execute after hook - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterCreate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(conn, msg.ID, err) return } @@ -311,33 +362,42 @@ func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookCont h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationCreate, data) } -// handleUpdate processes an update operation +// handleUpdate processes an update operation. The update runs in the first +// transaction; the re-fetch and AfterUpdate run in a second short transaction. func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeUpdate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } + var data interface{} - // Perform update operation - data, err := h.update(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { + return hookStage(err) + } + if err := h.update(hookCtx); err != nil { + return opStage("update_error", err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Update operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "update_error", err) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("update_error", err) + } - // Execute after hook - if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterUpdate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(conn, msg.ID, err) return } @@ -349,30 +409,22 @@ func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookCont h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationUpdate, data) } -// handleDelete processes a delete operation +// handleDelete processes a delete operation; hooks and delete share one transaction. func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeDelete hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } - - // Perform delete operation - err := h.delete(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { + return hookStage(err) + } + if err := h.delete(hookCtx); err != nil { + return opStage("delete_error", err) + } + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Delete operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "delete_error", err) - _ = conn.SendJSON(errResp) - return - } - - // Execute after hook - if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterDelete hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } @@ -548,7 +600,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { fetchRowNumberPKValue := *hookCtx.Options.FetchRowNumber logger.Debug("[WebSocketSpec] FetchRowNumber: Fetching row number for PK %s = %s", pkName, fetchRowNumberPKValue) - rowNum, err := h.FetchRowNumber(hookCtx.Context, hookCtx.TableName, pkName, fetchRowNumberPKValue, hookCtx.Options, hookCtx.Model) + rowNum, err := h.fetchRowNumber(hookCtx.Context, hookCtx.Tx, hookCtx.TableName, pkName, fetchRowNumberPKValue, hookCtx.Options, hookCtx.Model) if err != nil { return nil, fmt.Errorf("failed to fetch row number: %w", err) } @@ -560,7 +612,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { hookCtx.ID = fetchRowNumberPKValue } - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID) @@ -604,7 +656,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { } func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata map[string]interface{}, err error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Apply options (simplified implementation) if hookCtx.Options != nil { @@ -669,7 +721,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata // Get count metadata = make(map[string]interface{}) - countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + countQuery := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if hookCtx.Options != nil { for _, filter := range hookCtx.Options.Filters { cond, args := h.buildFilterCondition(filter, hookCtx.Model) @@ -705,21 +757,15 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { } // Insert record - query := h.db.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } - // Re-fetch the created record to capture DB-generated defaults/triggers. - if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { - hookCtx.ID = fmt.Sprintf("%v", pkVal) - return h.readByID(hookCtx) - } - return hookCtx.ModelPtr, nil } -func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { +func (h *Handler) update(hookCtx *HookContext) error { // Convert request data to a map var updates map[string]interface{} if m, ok := hookCtx.Data.(map[string]interface{}); ok { @@ -727,10 +773,10 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { } else { dataBytes, err := json.Marshal(hookCtx.Data) if err != nil { - return nil, fmt.Errorf("failed to marshal data: %w", err) + return fmt.Errorf("failed to marshal data: %w", err) } if err := json.Unmarshal(dataBytes, &updates); err != nil { - return nil, fmt.Errorf("failed to unmarshal data into map: %w", err) + return fmt.Errorf("failed to unmarshal data into map: %w", err) } } @@ -741,20 +787,19 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) if len(values) > 0 { - query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values). + query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) if _, err := query.Exec(hookCtx.Context); err != nil { - return nil, fmt.Errorf("failed to update record: %w", err) + return fmt.Errorf("failed to update record: %w", err) } } - // Fetch updated record - return h.readByID(hookCtx) + return nil } func (h *Handler) delete(hookCtx *HookContext) error { - query := h.db.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) @@ -966,6 +1011,11 @@ func (h *Handler) getOperatorSQL(operator string) string { // FetchRowNumber calculates the row number of a specific record based on sorting and filtering // Returns the 1-based row number of the record with the given primary key value func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName string, pkValue string, options *common.RequestOptions, model interface{}) (int64, error) { + return h.fetchRowNumber(ctx, h.db, tableName, pkName, pkValue, options, model) +} + +// fetchRowNumber is FetchRowNumber on the given database or transaction. +func (h *Handler) fetchRowNumber(ctx context.Context, db common.Database, tableName string, pkName string, pkValue string, options *common.RequestOptions, model interface{}) (int64, error) { defer func() { if r := recover(); r != nil { logger.Error("[WebSocketSpec] Panic during FetchRowNumber: %v", r) @@ -1033,7 +1083,7 @@ func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName s var result []struct { RN int64 `bun:"rn"` } - err := h.db.Query(ctx, &result, queryStr, whereArgs...) + err := db.Query(ctx, &result, queryStr, whereArgs...) if err != nil { return 0, fmt.Errorf("failed to fetch row number: %w", err) } diff --git a/pkg/websocketspec/hooks.go b/pkg/websocketspec/hooks.go index 171470d..5d4e55c 100644 --- a/pkg/websocketspec/hooks.go +++ b/pkg/websocketspec/hooks.go @@ -64,6 +64,13 @@ const ( // Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each // individual SQL-operation hook point, so it runs once per statement executed. BeforeOp HookType = "before_op" + + // OnTxBegin fires once, first, inside every transaction the handler opens for a + // read/create/update/delete message (including the second short transaction for + // post-commit work). hookCtx.Tx is the transaction; use it to stamp + // transaction-local state such as RLS settings. An error or abort rolls the + // transaction back. + OnTxBegin HookType = common.TxHookName ) // HookContext contains context information for hook execution @@ -128,6 +135,9 @@ type HookContext struct { Metadata map[string]interface{} } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is a function that processes a hook type HookFunc func(*HookContext) error diff --git a/pkg/websocketspec/tx_test.go b/pkg/websocketspec/tx_test.go new file mode 100644 index 0000000..12e410d --- /dev/null +++ b/pkg/websocketspec/tx_test.go @@ -0,0 +1,173 @@ +package websocketspec + +import ( + "context" + "encoding/json" + "errors" + "testing" + "time" + + "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/modelregistry" +) + +type txItem struct { + ID int `json:"id" bun:"id,pk"` + Name string `json:"name" bun:"name"` +} + +func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, *Connection, *HookContext) { + t.Helper() + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + // One connection: any statement that bypasses the transaction while it is + // open cannot get a connection and fails on the context timeout. + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry()) + conn := NewConnection("c1", nil, h) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + hookCtx := &HookContext{ + Context: ctx, + Handler: h, + Schema: "public", + Entity: "items", + TableName: "public.items", + Model: txItem{}, + ModelPtr: &txItem{}, + ID: "7", + Tx: h.db, + Metadata: map[string]interface{}{}, + } + return h, mock, conn, hookCtx +} + +func response(t *testing.T, conn *Connection) ResponseMessage { + t.Helper() + select { + case raw := <-conn.send: + var resp ResponseMessage + if err := json.Unmarshal(raw, &resp); err != nil { + t.Fatal(err) + } + return resp + default: + t.Fatal("no response sent") + return ResponseMessage{} + } +} + +func TestDeleteRunsHooksAndDeleteInOneTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + var order []string + var txs []common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { + order = append(order, "begin") + txs = append(txs, c.Tx) + return nil + }) + h.Hooks().Register(BeforeDelete, func(c *HookContext) error { + order = append(order, "before") + return nil + }) + h.Hooks().Register(AfterDelete, func(c *HookContext) error { + order = append(order, "after") + return nil + }) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(txs) != 1 || txs[0] == h.db || len(order) != 3 || order[0] != "begin" { + t.Fatalf("OnTxBegin must fire once, first, on the transaction: %v", order) + } +} + +func TestUpdateRefetchAndAfterHookRunInSecondTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + hookCtx.Data = map[string]interface{}{"name": "b"} + var begins []common.Database + var afterTx common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { + begins = append(begins, c.Tx) + return nil + }) + h.Hooks().Register(AfterUpdate, func(c *HookContext) error { + afterTx = c.Tx + return nil + }) + + mock.ExpectBegin() + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "b")) + mock.ExpectCommit() + + h.handleUpdate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(begins) != 2 || begins[0] == begins[1] || afterTx != begins[1] { + t.Fatalf("OnTxBegin must fire per transaction and AfterUpdate run on the second, got %d", len(begins)) + } +} + +func TestReadRunsInOneTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + var beforeTx common.Database + h.Hooks().Register(BeforeRead, func(c *HookContext) error { + beforeTx = c.Tx + return nil + }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + h.handleRead(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success || beforeTx == nil || beforeTx == h.db { + t.Fatal("BeforeRead must run on the transaction") + } +} + +func TestOnTxBeginErrorRollsBackWithoutDetail(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { return errors.New("secret detail") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + resp := response(t, conn) + if resp.Success || resp.Error == nil || resp.Error.Code != "transaction_error" || resp.Error.Message != "Transaction failed" { + t.Fatalf("unexpected response %+v", resp) + } +} From 4cbe4f597dc486a5b9374905c540c18e0cc8c15c Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:49:47 +0200 Subject: [PATCH 12/23] feat(tx): run resolvemcp operations in per-operation transactions with OnTxBegin --- audit/single_tran.md | 5 +- pkg/resolvemcp/handler.go | 236 +++++++++++++++++++++----------------- pkg/resolvemcp/hooks.go | 9 ++ pkg/resolvemcp/tx_test.go | 207 +++++++++++++++++++++++++++++++++ 4 files changed, 350 insertions(+), 107 deletions(-) create mode 100644 pkg/resolvemcp/tx_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 83dbe22..0ec5e90 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -76,7 +76,7 @@ | 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 | | 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | | 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | -| 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | +| 5 | DONE | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | | 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | @@ -91,7 +91,8 @@ - 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. - 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"). -- NEXT: P5. +- 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). +- NEXT: P6. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 94741ad..5723466 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -247,10 +247,27 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op return nil, nil, err } + // Hooks and queries share one transaction so transaction-local state set by + // hooks (e.g. RLS settings) applies to every statement. + var data interface{} + var metadata *common.Metadata + err = h.runInTx(ctx, hookCtx, func(common.Database) error { + var err error + data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options) + return err + }) + if err != nil { + return nil, nil, err + } + return data, metadata, nil +} + +// readInTx runs the read hooks and queries on hookCtx.Tx. +func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) { sliceType := reflect.SliceOf(reflect.PointerTo(modelType)) modelPtr := reflect.New(sliceType).Interface() - query := h.db.NewSelect().Model(modelPtr) + query := hookCtx.Tx.NewSelect().Model(modelPtr) tempInstance := reflect.New(modelType).Interface() if provider, ok := tempInstance.(common.TableNameProvider); !ok || provider.TableName() == "" { @@ -431,96 +448,83 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil { return nil, err } - if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil { - return nil, err - } - - // Use potentially modified data - data = hookCtx.Data pkName := reflection.GetPrimaryKeyName(model) + modelType := reflect.TypeOf(model) + if modelType.Kind() == reflect.Pointer { + modelType = modelType.Elem() + } - switch v := data.(type) { - case map[string]interface{}: - query := h.db.NewInsert().Table(tableName) - for key, value := range v { - query = query.Value(key, value) + // Transaction 1: BeforeCreate + inserts. + var ( + single bool + originals []map[string]interface{} + insertedIDs []interface{} + ) + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { + if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil { + return err } - if pkName != "" { - var insertedID interface{} - if err := query.Returning(pkName).Scan(ctx, &insertedID); err != nil { - return nil, fmt.Errorf("create error: %w", err) - } - // Re-fetch after insert to capture DB-generated defaults/triggers. - modelType := reflect.TypeOf(model) - if modelType.Kind() == reflect.Pointer { - modelType = modelType.Elem() - } - fetchedRecord := reflect.New(modelType).Interface() - if err := h.db.NewSelect().Model(fetchedRecord). - Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), insertedID). - ScanModel(ctx); err == nil { - v = mergeWithInput(fetchedRecord, v) - } else { - logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, insertedID, err) - } - } else { - if _, err := query.Exec(ctx); err != nil { - return nil, fmt.Errorf("create error: %w", err) - } - } - hookCtx.Result = v - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - return nil, fmt.Errorf("AfterCreate hook failed: %w", err) - } - return v, nil - - case []interface{}: - modelType := reflect.TypeOf(model) - if modelType.Kind() == reflect.Pointer { - modelType = modelType.Elem() - } - originals := make([]map[string]interface{}, 0, len(v)) - insertedIDs := make([]interface{}, 0, len(v)) - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { + // Use potentially modified data + switch v := hookCtx.Data.(type) { + case map[string]interface{}: + single = true + originals = []map[string]interface{}{v} + case []interface{}: + originals = make([]map[string]interface{}, 0, len(v)) for _, item := range v { itemMap, ok := item.(map[string]interface{}) if !ok { return fmt.Errorf("each item must be an object") } - q := tx.NewInsert().Table(tableName) - for key, value := range itemMap { - q = q.Value(key, value) - } - if pkName == "" { - if _, err := q.Exec(ctx); err != nil { - return err - } - originals = append(originals, itemMap) - insertedIDs = append(insertedIDs, nil) - continue - } - var returnedID interface{} - if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil { + originals = append(originals, itemMap) + } + default: + return fmt.Errorf("data must be an object or array of objects") + } + + insertedIDs = make([]interface{}, 0, len(originals)) + for _, itemMap := range originals { + q := tx.NewInsert().Table(tableName) + for key, value := range itemMap { + q = q.Value(key, value) + } + if pkName == "" { + if _, err := q.Exec(ctx); err != nil { return err } - originals = append(originals, itemMap) - insertedIDs = append(insertedIDs, returnedID) + insertedIDs = append(insertedIDs, nil) + continue } - return nil - }) - if err != nil { + var returnedID interface{} + if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil { + return err + } + insertedIDs = append(insertedIDs, returnedID) + } + return nil + }) + if err != nil { + if single { + return nil, fmt.Errorf("create error: %w", err) + } + if _, ok := hookCtx.Data.([]interface{}); ok { return nil, fmt.Errorf("batch create error: %w", err) } - // Re-fetch each record after transaction commits; fall back to input on failure. - results := make([]interface{}, 0, len(insertedIDs)) + return nil, err + } + + // Transaction 2: re-fetch to capture DB-generated defaults/triggers, then AfterCreate. + results := make([]interface{}, 0, len(insertedIDs)) + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { + results = results[:0] for i, pkVal := range insertedIDs { if pkVal == nil { results = append(results, originals[i]) continue } fetchedRecord := reflect.New(modelType).Interface() - if err := h.db.NewSelect().Model(fetchedRecord). + if err := tx.NewSelect().Model(fetchedRecord). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), pkVal). ScanModel(ctx); err == nil { results = append(results, mergeWithInput(fetchedRecord, originals[i])) @@ -529,15 +533,23 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data results = append(results, originals[i]) } } - hookCtx.Result = results - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - return nil, fmt.Errorf("AfterCreate hook failed: %w", err) + if single { + hookCtx.Result = results[0] + } else { + hookCtx.Result = results } - return results, nil - - default: - return nil, fmt.Errorf("data must be an object or array of objects") + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return fmt.Errorf("AfterCreate hook failed: %w", err) + } + return nil + }) + if err != nil { + return nil, err } + if single { + return results[0], nil + } + return results, nil } // executeUpdate updates a record by ID. @@ -573,8 +585,20 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, pkName := reflection.GetPrimaryKeyName(model) + hookCtx := &HookContext{ + Context: ctx, + Handler: h, + Schema: schema, + Entity: entity, + Model: model, + Operation: "update", + ID: id, + Data: updates, + Tx: h.db, + } + var updateResult interface{} - err = h.db.RunInTransaction(ctx, func(tx common.Database) error { + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { // Read existing record modelType := reflect.TypeOf(model) if modelType.Kind() == reflect.Pointer { @@ -601,17 +625,6 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, return fmt.Errorf("error unmarshaling existing record: %w", err) } - hookCtx := &HookContext{ - Context: ctx, - Handler: h, - Schema: schema, - Entity: entity, - Model: model, - Operation: "update", - ID: id, - Data: updates, - Tx: tx, - } if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil { return err } @@ -649,22 +662,28 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, return nil, err } - // Re-fetch the record after transaction commits to capture DB-generated changes. + // Transaction 2: re-fetch to capture DB-generated changes. modelType := reflect.TypeOf(model) if modelType.Kind() == reflect.Pointer { modelType = modelType.Elem() } - fetchedRecord := reflect.New(modelType).Interface() - if err := h.db.NewSelect().Model(fetchedRecord). - Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id). - ScanModel(ctx); err == nil { - jsonData, marshalErr := json.Marshal(fetchedRecord) - if marshalErr == nil { - var fetchedMap map[string]interface{} - if json.Unmarshal(jsonData, &fetchedMap) == nil { - updateResult = fetchedMap + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { + fetchedRecord := reflect.New(modelType).Interface() + if err := tx.NewSelect().Model(fetchedRecord). + Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id). + ScanModel(ctx); err == nil { + jsonData, marshalErr := json.Marshal(fetchedRecord) + if marshalErr == nil { + var fetchedMap map[string]interface{} + if json.Unmarshal(jsonData, &fetchedMap) == nil { + updateResult = fetchedMap + } } } + return nil + }) + if err != nil { + return nil, err } return updateResult, nil @@ -706,9 +725,6 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil { return nil, err } - if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil { - return nil, err - } modelType := reflect.TypeOf(model) if modelType.Kind() == reflect.Pointer { @@ -717,7 +733,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) var recordToDelete interface{} - err = h.db.RunInTransaction(ctx, func(tx common.Database) error { + err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { + if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil { + return err + } record := reflect.New(modelType).Interface() selectQuery := tx.NewSelect().Model(record). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) @@ -739,7 +758,6 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) } recordToDelete = record - hookCtx.Tx = tx hookCtx.Result = record return h.hooks.Execute(AfterDelete, hookCtx) }) @@ -873,3 +891,11 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre } return query, nil } + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Every transaction the handler opens goes through here. +func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(ctx, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} diff --git a/pkg/resolvemcp/hooks.go b/pkg/resolvemcp/hooks.go index 327242b..11aa4fb 100644 --- a/pkg/resolvemcp/hooks.go +++ b/pkg/resolvemcp/hooks.go @@ -26,6 +26,12 @@ const ( BeforeDelete HookType = "before_delete" AfterDelete HookType = "after_delete" + + // OnTxBegin fires once, first, inside every transaction the handler opens + // (including the second short transaction for post-commit work). hookCtx.Tx is + // the transaction; use it to stamp transaction-local state such as RLS + // settings. An error or abort rolls the transaction back. + OnTxBegin HookType = common.TxHookName ) // HookContext contains all the data available to a hook @@ -48,6 +54,9 @@ type HookContext struct { Tx common.Database } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is the signature for hook functions type HookFunc func(*HookContext) error diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go new file mode 100644 index 0000000..80f4945 --- /dev/null +++ b/pkg/resolvemcp/tx_test.go @@ -0,0 +1,207 @@ +package resolvemcp + +import ( + "context" + "database/sql" + "testing" + "time" + + "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/modelregistry" +) + +type txItem struct { + ID int `json:"id" bun:"id,pk"` + Name string `json:"name" bun:"name"` +} + +func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context) { + t.Helper() + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + // One connection: any statement bypassing the open tx cannot get a + // connection and fails on the context timeout. + db.SetMaxOpenConns(1) + t.Cleanup(func() { _ = db.Close() }) + h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{}) + if err := h.RegisterModel("public", "items", &txItem{}); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + return h, mock, ctx +} + +// trace records hook firing order and the Tx each hook saw. +type trace struct { + order []string + txs map[string][]common.Database +} + +func traceHooks(h *Handler, types ...HookType) *trace { + tr := &trace{txs: map[string][]common.Database{}} + for _, ht := range types { + ht := ht + h.Hooks().Register(ht, func(c *HookContext) error { + tr.order = append(tr.order, string(ht)) + tr.txs[string(ht)] = append(tr.txs[string(ht)], c.Tx) + return nil + }) + } + return tr +} + +func (tr *trace) assertOrder(t *testing.T, want ...string) { + t.Helper() + if len(tr.order) != len(want) { + t.Fatalf("hook order %v, want %v", tr.order, want) + } + for i := range want { + if tr.order[i] != want[i] { + t.Fatalf("hook order %v, want %v", tr.order, want) + } + } +} + +func TestDeleteRunsHooksInOneTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeDelete, AfterDelete) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + if _, err := h.executeDelete(ctx, "public", "items", "7"); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + tr.assertOrder(t, "on_tx_begin", "before_delete", "after_delete") + if tr.txs["before_delete"][0] != tr.txs["on_tx_begin"][0] || tr.txs["after_delete"][0] != tr.txs["on_tx_begin"][0] { + t.Fatal("OnTxBegin, BeforeDelete and AfterDelete must share one transaction") + } +} + +func TestDeleteBeforeHookErrorRollsBack(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(BeforeDelete, func(*HookContext) error { return sql.ErrConnDone }) + + mock.ExpectBegin() + mock.ExpectRollback() + + if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestReadRunsInOneTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeRead, AfterRead) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + if _, _, err := h.executeRead(ctx, "public", "items", "7", common.RequestOptions{}); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + tr.assertOrder(t, "on_tx_begin", "before_read", "after_read") + for _, ht := range []string{"before_read", "after_read"} { + if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] { + t.Fatalf("%s must run on the OnTxBegin transaction", ht) + } + } +} + +func TestCreateSingleUsesTwoTransactions(t *testing.T) { + h, mock, ctx := newTxHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + tr.assertOrder(t, "on_tx_begin", "before_create", "on_tx_begin", "after_create") + if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] { + t.Fatal("BeforeCreate must run on the first transaction") + } + if tr.txs["after_create"][0] != tr.txs["on_tx_begin"][1] || tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] { + t.Fatal("AfterCreate must run on a second, distinct transaction") + } +} + +func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + tr := traceHooks(h, OnTxBegin, AfterCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a")) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "b")) + mock.ExpectCommit() + + items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}} + if _, err := h.executeCreate(ctx, "public", "items", items); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + tr.assertOrder(t, "on_tx_begin", "on_tx_begin", "after_create") +} + +func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate) + + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"}); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + tr.assertOrder(t, "on_tx_begin", "before_update", "after_update", "on_tx_begin") + if tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] { + t.Fatal("re-fetch must run on a second transaction") + } + for _, ht := range []string{"before_update", "after_update"} { + if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] { + t.Fatalf("%s must run on the first transaction", ht) + } + } +} From 6b6f540ab0b617f2fdce9a4055d2d1ffeb64452d Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:50:41 +0200 Subject: [PATCH 13/23] feat(tx): run funcspec OnTxBegin and BeforeResponse in transactions --- audit/single_tran.md | 5 +- pkg/funcspec/function_api.go | 388 +++++++++++++++++++---------------- pkg/funcspec/hooks.go | 9 + pkg/funcspec/tx_test.go | 102 +++++++++ 4 files changed, 321 insertions(+), 183 deletions(-) create mode 100644 pkg/funcspec/tx_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 0ec5e90..1039c5e 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -77,7 +77,7 @@ | 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | | 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 5 | DONE | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | -| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | +| 6 | DONE | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | | 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | ## Progress @@ -92,7 +92,8 @@ - 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"). - 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). -- NEXT: P6. +- DONE P6: funcspec. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx` for `SqlQuery` and `SqlQueryList`. `BeforeResponse` now runs in a second short tx (`Tx` is no longer the pool). `BeforeOp` is unchanged (still per statement). A begin/`OnTxBegin`/commit failure answers 500 `transaction_error` / "Transaction failed" (before, it returned with no response); body failures still answer via `sendError`. Tests: `pkg/funcspec/tx_test.go`. +- NEXT: P7. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/funcspec/function_api.go b/pkg/funcspec/function_api.go index 1c5f6ec..9afabe5 100644 --- a/pkg/funcspec/function_api.go +++ b/pkg/funcspec/function_api.go @@ -192,131 +192,140 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun hookCtx.InputVars = inputvars // Execute query within transaction - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { - // Set transaction in hook context for hooks to use - hookCtx.Tx = tx + // bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent + // yet) from a body failure (sendError already answered). + var bodyRan, bodyFailed bool + err := h.runInTx(ctx, hookCtx, func(tx common.Database) error { + bodyRan = true + berr := func() error { - // Execute BeforeQueryList hook (inside transaction) - if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil { - logger.Error("BeforeQueryList hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - - // Check if hook aborted the operation - if hookCtx.Abort { - if hookCtx.AbortCode == 0 { - hookCtx.AbortCode = http.StatusBadRequest - } - sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil) - return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage) - } - - // Use potentially modified SQL query from hook - sqlquery = hookCtx.SQLQuery - sqlqueryCnt := sqlquery - - // Parse sorting and pagination parameters - sortcols, limit, offset := h.parsePaginationParams(r) - - // Override with parsed parameters if available - if reqParams.SortColumns != "" { - sortcols = reqParams.SortColumns - } - if reqParams.Limit > 0 { - limit = reqParams.Limit - } - if reqParams.Offset > 0 { - offset = reqParams.Offset - } - - hookCtx.SortColumns = sortcols - hookCtx.Limit = limit - hookCtx.Offset = offset - fromPos := strings.Index(strings.ToLower(sqlquery), "from ") - orderbyPos := strings.Index(strings.ToLower(sqlquery), "order by") - - if len(sortcols) > 0 && (orderbyPos < 0 || (orderbyPos > 0 && orderbyPos < fromPos)) { - sqlquery = fmt.Sprintf("%s \nORDER BY %s", sqlquery, ValidSQL(sortcols, "select")) - } - - if !options.NoCount { - if limit > 0 && offset > 0 { - sqlquery = fmt.Sprintf("%s \nLIMIT %d OFFSET %d", sqlquery, limit, offset) - } else if limit > 0 { - sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, limit) - } else { - sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, 20000) - } - - // Get total count - countQuery := fmt.Sprintf("SELECT COUNT(1) FROM (%s) cnts", sqlqueryCnt) - var countResult struct{ Count int64 } - if err := tx.Query(ctx, &countResult, countQuery); err != nil { - sendError(w, http.StatusBadRequest, "count_failed", "Failed to retrieve record count", err) + // Execute BeforeQueryList hook (inside transaction) + if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil { + logger.Error("BeforeQueryList hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) return err } - total = countResult.Count - } - // Execute BeforeSQLExec hook - hookCtx.SQLQuery = sqlquery - if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil { - logger.Error("BeforeSQLExec hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified SQL query from hook - sqlquery = hookCtx.SQLQuery + // Check if hook aborted the operation + if hookCtx.Abort { + if hookCtx.AbortCode == 0 { + hookCtx.AbortCode = http.StatusBadRequest + } + sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil) + return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage) + } - // Execute main query - rows := make([]map[string]interface{}, 0) - if err := tx.Query(ctx, &rows, sqlquery); err != nil { - sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err) - return err - } + // Use potentially modified SQL query from hook + sqlquery = hookCtx.SQLQuery + sqlqueryCnt := sqlquery - // Normalize PostgreSQL types for proper JSON marshaling - dbobjlist = normalizePostgresTypesList(rows) + // Parse sorting and pagination parameters + sortcols, limit, offset := h.parsePaginationParams(r) - if options.NoCount { - total = int64(len(dbobjlist)) - } + // Override with parsed parameters if available + if reqParams.SortColumns != "" { + sortcols = reqParams.SortColumns + } + if reqParams.Limit > 0 { + limit = reqParams.Limit + } + if reqParams.Offset > 0 { + offset = reqParams.Offset + } - // Execute AfterSQLExec hook - hookCtx.Result = dbobjlist - hookCtx.Total = total - if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil { - logger.Error("AfterSQLExec hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified result from hook - if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok { - dbobjlist = modifiedResult - } - total = hookCtx.Total + hookCtx.SortColumns = sortcols + hookCtx.Limit = limit + hookCtx.Offset = offset + fromPos := strings.Index(strings.ToLower(sqlquery), "from ") + orderbyPos := strings.Index(strings.ToLower(sqlquery), "order by") - // Execute AfterQueryList hook (inside transaction) - hookCtx.Result = dbobjlist - hookCtx.Total = total - hookCtx.Error = nil - if err := h.hooks.Execute(AfterQueryList, hookCtx); err != nil { - logger.Error("AfterQueryList hook failed: %v", err) - sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified result from hook - if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok { - dbobjlist = modifiedResult - } - total = hookCtx.Total + if len(sortcols) > 0 && (orderbyPos < 0 || (orderbyPos > 0 && orderbyPos < fromPos)) { + sqlquery = fmt.Sprintf("%s \nORDER BY %s", sqlquery, ValidSQL(sortcols, "select")) + } - return nil + if !options.NoCount { + if limit > 0 && offset > 0 { + sqlquery = fmt.Sprintf("%s \nLIMIT %d OFFSET %d", sqlquery, limit, offset) + } else if limit > 0 { + sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, limit) + } else { + sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, 20000) + } + + // Get total count + countQuery := fmt.Sprintf("SELECT COUNT(1) FROM (%s) cnts", sqlqueryCnt) + var countResult struct{ Count int64 } + if err := tx.Query(ctx, &countResult, countQuery); err != nil { + sendError(w, http.StatusBadRequest, "count_failed", "Failed to retrieve record count", err) + return err + } + total = countResult.Count + } + + // Execute BeforeSQLExec hook + hookCtx.SQLQuery = sqlquery + if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil { + logger.Error("BeforeSQLExec hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified SQL query from hook + sqlquery = hookCtx.SQLQuery + + // Execute main query + rows := make([]map[string]interface{}, 0) + if err := tx.Query(ctx, &rows, sqlquery); err != nil { + sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err) + return err + } + + // Normalize PostgreSQL types for proper JSON marshaling + dbobjlist = normalizePostgresTypesList(rows) + + if options.NoCount { + total = int64(len(dbobjlist)) + } + + // Execute AfterSQLExec hook + hookCtx.Result = dbobjlist + hookCtx.Total = total + if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil { + logger.Error("AfterSQLExec hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified result from hook + if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok { + dbobjlist = modifiedResult + } + total = hookCtx.Total + + // Execute AfterQueryList hook (inside transaction) + hookCtx.Result = dbobjlist + hookCtx.Total = total + hookCtx.Error = nil + if err := h.hooks.Execute(AfterQueryList, hookCtx); err != nil { + logger.Error("AfterQueryList hook failed: %v", err) + sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified result from hook + if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok { + dbobjlist = modifiedResult + } + total = hookCtx.Total + + return nil + }() + bodyFailed = berr != nil + return berr }) if err != nil { logger.Error("Transaction failed: %v", err) + if !bodyRan || !bodyFailed { + sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil) + } return } @@ -331,13 +340,13 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total)) logger.Info("Serving: Records %d of %d", len(dbobjlist), total) - // Execute BeforeResponse hook. The transaction has already committed by - // this point, so hooks must use the pooled connection rather than the - // now-dead tx. - hookCtx.Tx = h.db + // Execute BeforeResponse hook in a second short transaction: the main one + // has already committed, and hooks must never get the pooled connection. hookCtx.Result = dbobjlist hookCtx.Total = total - if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil { + if err := h.runInTx(ctx, hookCtx, func(common.Database) error { + return h.hooks.Execute(BeforeResponse, hookCtx) + }); err != nil { logger.Error("BeforeResponse hook failed: %v", err) sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) return @@ -558,88 +567,97 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp hookCtx.InputVars = inputvars // Execute query within transaction - err := h.db.RunInTransaction(ctx, func(tx common.Database) error { - // Set transaction in hook context for hooks to use - hookCtx.Tx = tx + // bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent + // yet) from a body failure (sendError already answered). + var bodyRan, bodyFailed bool + err := h.runInTx(ctx, hookCtx, func(tx common.Database) error { + bodyRan = true + berr := func() error { - // Execute BeforeQuery hook (inside transaction) - if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil { - logger.Error("BeforeQuery hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - - // Check if hook aborted the operation - if hookCtx.Abort { - if hookCtx.AbortCode == 0 { - hookCtx.AbortCode = http.StatusBadRequest + // Execute BeforeQuery hook (inside transaction) + if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil { + logger.Error("BeforeQuery hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) + return err } - sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil) - return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage) - } - // Use potentially modified SQL query from hook - sqlquery = hookCtx.SQLQuery + // Check if hook aborted the operation + if hookCtx.Abort { + if hookCtx.AbortCode == 0 { + hookCtx.AbortCode = http.StatusBadRequest + } + sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil) + return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage) + } - // Execute BeforeSQLExec hook - if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil { - logger.Error("BeforeSQLExec hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified SQL query from hook - sqlquery = hookCtx.SQLQuery + // Use potentially modified SQL query from hook + sqlquery = hookCtx.SQLQuery - // Execute main query - rows := make([]map[string]interface{}, 0) - if err := tx.Query(ctx, &rows, sqlquery); err != nil { - sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err) - return err - } + // Execute BeforeSQLExec hook + if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil { + logger.Error("BeforeSQLExec hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified SQL query from hook + sqlquery = hookCtx.SQLQuery - if len(rows) > 0 { - dbobj = normalizePostgresTypes(rows[0]) - } + // Execute main query + rows := make([]map[string]interface{}, 0) + if err := tx.Query(ctx, &rows, sqlquery); err != nil { + sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err) + return err + } - // Execute AfterSQLExec hook - hookCtx.Result = dbobj - if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil { - logger.Error("AfterSQLExec hook failed: %v", err) - sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified result from hook - if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok { - dbobj = modifiedResult - } + if len(rows) > 0 { + dbobj = normalizePostgresTypes(rows[0]) + } - // Execute AfterQuery hook (inside transaction) - hookCtx.Result = dbobj - hookCtx.Error = nil - if err := h.hooks.Execute(AfterQuery, hookCtx); err != nil { - logger.Error("AfterQuery hook failed: %v", err) - sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) - return err - } - // Use potentially modified result from hook - if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok { - dbobj = modifiedResult - } + // Execute AfterSQLExec hook + hookCtx.Result = dbobj + if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil { + logger.Error("AfterSQLExec hook failed: %v", err) + sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified result from hook + if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok { + dbobj = modifiedResult + } - return nil + // Execute AfterQuery hook (inside transaction) + hookCtx.Result = dbobj + hookCtx.Error = nil + if err := h.hooks.Execute(AfterQuery, hookCtx); err != nil { + logger.Error("AfterQuery hook failed: %v", err) + sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) + return err + } + // Use potentially modified result from hook + if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok { + dbobj = modifiedResult + } + + return nil + }() + bodyFailed = berr != nil + return berr }) if err != nil { logger.Error("Transaction failed: %v", err) + if !bodyRan || !bodyFailed { + sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil) + } return } - // Execute BeforeResponse hook. The transaction has already committed by - // this point, so hooks must use the pooled connection rather than the - // now-dead tx. - hookCtx.Tx = h.db + // Execute BeforeResponse hook in a second short transaction: the main one + // has already committed, and hooks must never get the pooled connection. hookCtx.Result = dbobj - if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil { + if err := h.runInTx(ctx, hookCtx, func(common.Database) error { + return h.hooks.Execute(BeforeResponse, hookCtx) + }); err != nil { logger.Error("BeforeResponse hook failed: %v", err) sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) return @@ -1249,3 +1267,11 @@ func normalizePostgresValue(value interface{}) interface{} { return v } } + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Every transaction the handler opens goes through here. +func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(ctx, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} diff --git a/pkg/funcspec/hooks.go b/pkg/funcspec/hooks.go index e417b6c..de7e7d9 100644 --- a/pkg/funcspec/hooks.go +++ b/pkg/funcspec/hooks.go @@ -32,6 +32,12 @@ const ( // BeforeOp fires immediately before every SQL operation (query, query list, SQL exec). // It fires at each individual SQL-operation hook point, so it runs once per statement executed. BeforeOp HookType = "before_op" + + // OnTxBegin fires once, first, inside every transaction the handler opens + // (the main transaction and the short one that runs BeforeResponse). + // hookCtx.Tx is the transaction; use it to stamp transaction-local state such + // as RLS settings. An error or abort rolls the transaction back. + OnTxBegin HookType = common.TxHookName ) // HookContext contains all the data available to a hook @@ -75,6 +81,9 @@ type HookContext struct { AbortCode int // HTTP status code if aborted } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is the signature for hook functions // It receives a HookContext and can modify it or return an error // If an error is returned, the operation will be aborted diff --git a/pkg/funcspec/tx_test.go b/pkg/funcspec/tx_test.go new file mode 100644 index 0000000..3ad946f --- /dev/null +++ b/pkg/funcspec/tx_test.go @@ -0,0 +1,102 @@ +package funcspec + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// txFactory returns a pool whose every transaction is a distinct MockDatabase. +func txFactory(queries *int) *MockDatabase { + return &MockDatabase{ + RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error { + return fn(&MockDatabase{ + QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error { + *queries++ + if rows, ok := dest.(*[]map[string]interface{}); ok { + *rows = []map[string]interface{}{{"id": float64(1)}} + } + return nil + }, + }) + }, + } +} + +func TestOnTxBeginFirstAndBeforeResponseOnSecondTx(t *testing.T) { + var queries int + h := NewHandler(txFactory(&queries)) + var order []string + txs := map[HookType][]common.Database{} + for _, ht := range []HookType{OnTxBegin, BeforeQuery, AfterQuery, BeforeResponse} { + ht := ht + h.Hooks().Register(ht, func(c *HookContext) error { + order = append(order, string(ht)) + txs[ht] = append(txs[ht], c.Tx) + return nil + }) + } + + w := httptest.NewRecorder() + h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil)) + + if w.Code != http.StatusOK { + t.Fatalf("status %d body %s", w.Code, w.Body) + } + if got := strings.Join(order, ","); got != "on_tx_begin,before_query,after_query,on_tx_begin,before_response" { + t.Fatalf("hook order %s", got) + } + if txs[BeforeQuery][0] != txs[OnTxBegin][0] || txs[AfterQuery][0] != txs[OnTxBegin][0] { + t.Fatal("query hooks must run on the OnTxBegin transaction") + } + if txs[OnTxBegin][0] == txs[OnTxBegin][1] || txs[BeforeResponse][0] != txs[OnTxBegin][1] { + t.Fatal("BeforeResponse must run on a second, distinct transaction") + } +} + +func TestOnTxBeginListBeforeResponseOnSecondTx(t *testing.T) { + var queries int + h := NewHandler(txFactory(&queries)) + var txs []common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { txs = append(txs, c.Tx); return nil }) + var respTx common.Database + h.Hooks().Register(BeforeResponse, func(c *HookContext) error { respTx = c.Tx; return nil }) + + w := httptest.NewRecorder() + h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true})(w, createTestRequest("GET", "/t", nil, nil, nil)) + + if w.Code != http.StatusOK { + t.Fatalf("status %d body %s", w.Code, w.Body) + } + if len(txs) != 2 || txs[0] == txs[1] || respTx != txs[1] { + t.Fatalf("expected 2 distinct transactions with BeforeResponse on the second, got %d", len(txs)) + } +} + +func TestOnTxBeginErrorAnswersTransactionError(t *testing.T) { + var queries int + h := NewHandler(txFactory(&queries)) + h.Hooks().Register(OnTxBegin, func(*HookContext) error { return errors.New("secret detail") }) + + for name, run := range map[string]HTTPFuncType{ + "single": h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{}), + "list": h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true}), + } { + w := httptest.NewRecorder() + run(w, createTestRequest("GET", "/t", nil, nil, nil)) + if w.Code != http.StatusInternalServerError { + t.Fatalf("%s: status %d body %s", name, w.Code, w.Body) + } + if strings.Contains(w.Body.String(), "secret detail") { + t.Fatalf("%s: hook error leaked to client: %s", name, w.Body) + } + } + if queries != 0 { + t.Fatalf("no query may run after a failed OnTxBegin, ran %d", queries) + } +} From 3b93802a250375ac0339173c2f8e7f042a8b4be3 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:53:07 +0200 Subject: [PATCH 14/23] fix(tx): run restheadspec AfterRead in a second short transaction --- pkg/restheadspec/handler.go | 8 +++-- pkg/restheadspec/read_tx_test.go | 62 ++++++++++++++++++++++++++++++++ 2 files changed, 67 insertions(+), 3 deletions(-) create mode 100644 pkg/restheadspec/read_tx_test.go diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 89e9d93..649901f 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -1000,12 +1000,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st logger.Debug("FetchRowNumber: Row number %d set in metadata", *fetchedRowNumber) } - // Execute AfterRead hooks (runs after the transaction commits, against the pooled db) - hookCtx.Tx = h.db + // Execute AfterRead hooks in a second short transaction: the read tx has + // already committed, and hooks must never get the pooled connection. hookCtx.Result = modelPtr hookCtx.Error = nil - if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + if err := h.runInTx(ctx, hookCtx, func(common.Database) error { + return h.hooks.Execute(AfterRead, hookCtx) + }); err != nil { logger.Error("AfterRead hook failed: %v", err) h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) return diff --git a/pkg/restheadspec/read_tx_test.go b/pkg/restheadspec/read_tx_test.go new file mode 100644 index 0000000..22ed88b --- /dev/null +++ b/pkg/restheadspec/read_tx_test.go @@ -0,0 +1,62 @@ +package restheadspec + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/pgdialect" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" + "github.com/bitechdev/ResolveSpec/pkg/modelregistry" +) + +func TestAfterReadRunsInSecondTransaction(t *testing.T) { + sqlDB, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry()) + var begins []common.Database + var readTx, afterTx common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + h.Hooks().Register(BeforeRead, func(c *HookContext) error { readTx = c.Tx; return nil }) + h.Hooks().Register(AfterRead, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + ctx = WithModel(ctx, delItem{}) + h.handleRead(ctx, w, "7", ExtendedRequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(begins) != 2 || begins[0] == begins[1] { + t.Fatalf("OnTxBegin must fire once per transaction (2), got %d", len(begins)) + } + if readTx != begins[0] || afterTx != begins[1] || afterTx == h.db { + t.Fatal("BeforeRead must run on the first tx and AfterRead on the second") + } +} From ff76eb8e1f17a4521d71f23c1b121667db3abfc6 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:55:26 +0200 Subject: [PATCH 15/23] feat(security): stamp transaction-local settings on OnTxBegin in all specs --- audit/single_tran.md | 6 +- pkg/common/TRANSACTIONS.md | 45 ++++++++++ pkg/funcspec/security_adapter.go | 6 ++ pkg/mqttspec/security_hooks.go | 6 ++ pkg/resolvemcp/security_hooks.go | 6 ++ pkg/resolvespec/security_hooks.go | 6 ++ pkg/resolvespec/tx_settings_test.go | 78 +++++++++++++++++ pkg/restheadspec/security_hooks.go | 6 ++ pkg/security/README.md | 2 + pkg/security/provider.go | 4 + pkg/security/txsettings.go | 86 ++++++++++++++++++ pkg/security/txsettings_test.go | 131 ++++++++++++++++++++++++++++ pkg/websocketspec/security_hooks.go | 6 ++ 13 files changed, 386 insertions(+), 2 deletions(-) create mode 100644 pkg/common/TRANSACTIONS.md create mode 100644 pkg/resolvespec/tx_settings_test.go create mode 100644 pkg/security/txsettings.go create mode 100644 pkg/security/txsettings_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 1039c5e..fa0b2b4 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -78,7 +78,7 @@ | 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 5 | DONE | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 6 | DONE | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | -| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | +| 7 | DONE | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | ## Progress - DONE P0: baseline via `dbtrace` on real Postgres (commit `cd96404`): create/read/delete `pooled=0`; update `pooled=1` (re-fetch) = P3 target. websocketspec/mqttspec/resolvemcp not measured. @@ -93,7 +93,9 @@ - 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 P6: funcspec. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx` for `SqlQuery` and `SqlQueryList`. `BeforeResponse` now runs in a second short tx (`Tx` is no longer the pool). `BeforeOp` is unchanged (still per statement). A begin/`OnTxBegin`/commit failure answers 500 `transaction_error` / "Transaction failed" (before, it returned with no response); body failures still answer via `sendError`. Tests: `pkg/funcspec/tx_test.go`. -- NEXT: P7. +- DONE (AfterRead, decided by user): restheadspec `AfterRead` now runs in a second short tx. Test: `pkg/restheadspec/read_tx_test.go`. +- DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. +- NEXT: extra tests (create/update for other specs, `dbtrace` `pooled == 0` on real Postgres). ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/common/TRANSACTIONS.md b/pkg/common/TRANSACTIONS.md new file mode 100644 index 0000000..df2f8d2 --- /dev/null +++ b/pkg/common/TRANSACTIONS.md @@ -0,0 +1,45 @@ +# Request Transactions (cheatsheet) + +Every DB statement and every DB-touching hook of one request runs on one transaction. Hooks never get the pool. + +## Rules +- `hookCtx.Tx` is always the open transaction (never `h.db`), except `BeforeHandle`, which runs before any tx and must not touch the DB. +- `OnTxBegin` fires once, first, in every tx the handler opens (incl. the second short tx). +- `OnTxBegin` error or abort: rollback, client gets a generic error, nothing leaked. +- Begin or commit failure: generic error (`transaction_error` / "Transaction failed" in websocketspec, mqttspec, funcspec). +- Transaction-local state (`set_config(..., true)`, RLS GUCs) is only visible on that tx. Set it in `OnTxBegin`. + +## Transactions per operation +| Operation | Tx 1 | Tx 2 (short, after commit) | +|---|---|---| +| read | `BeforeRead`, count, scan, `AfterRead`* | restheadspec: `AfterRead` | +| create | `Before*`, insert | re-fetch, `BeforeScan`, `AfterCreate` | +| update | `Before*`, select, update, `AfterUpdate`† | re-fetch (+ `AfterUpdate` where noted) | +| delete (single/batch) | `BeforeDelete`, select, delete, `AfterDelete` | none | +| funcspec query | `BeforeQuery*`, `BeforeSQLExec`, SQL, `After*` | `BeforeResponse` | + +\* resolvespec, websocketspec, mqttspec, resolvemcp. restheadspec runs `AfterRead` in tx 2. +† restheadspec, websocketspec, mqttspec run `AfterUpdate` in tx 2. resolvespec, resolvemcp run it in tx 1. + +Tx 2 exists so the re-fetch sees trigger changes from the committed write. + +## Per spec +| Spec | `OnTxBegin` | Helper | +|---|---|---| +| resolvespec, restheadspec, websocketspec, resolvemcp, funcspec | own `HookType` = `common.TxHookName` | `Handler.runInTx` | +| mqttspec | re-exports `websocketspec.OnTxBegin` | `Handler.runInTx` | + +- `common.RunRequestTx(ctx, db, TxContext, onBegin, body)`: open tx, `SetTx`, `onBegin`, `body`. +- `common.TxContext`: `SetTx(tx)`; implemented by each spec's `HookContext`. + +## RLS / transaction settings (pkg/security) +- `SecurityList.SetTxSettings(fn)`: `fn(SecurityContext) (map[string]string, error)`; nil disables. +- Every spec's `RegisterSecurityHooks` registers `OnTxBegin` → `security.StampTxSettings`. `fn` is read per call, so set order does not matter. +- Stamps via `set_config(name, value, true)` in name order, before any other SQL. +- Fail closed: `fn` error, invalid name, or non-Postgres driver with a non-empty map aborts the tx. +- Name: dotted identifier (`ns.name`). Value is hex-encoded in SQL, never inlined. +- Low level: `security.ApplyTxSettings(secCtx, tx, map)`. + +## Test notes +- sqlmock + `SetMaxOpenConns(1)`: any pool use inside an open tx blocks and fails. +- restheadspec model-based updates/reads need the bun adapter. diff --git a/pkg/funcspec/security_adapter.go b/pkg/funcspec/security_adapter.go index f0201cc..5648b2b 100644 --- a/pkg/funcspec/security_adapter.go +++ b/pkg/funcspec/security_adapter.go @@ -12,6 +12,12 @@ import ( // Note: funcspec operates on SQL queries directly, so row-level security is not directly applicable // We provide auth enforcement and audit logging for data access tracking func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newFuncSpecSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeQueryList - Auth check before list query execution handler.Hooks().Register(BeforeQueryList, func(hookCtx *HookContext) error { if hookCtx.UserContext == nil || hookCtx.UserContext.UserID == 0 { diff --git a/pkg/mqttspec/security_hooks.go b/pkg/mqttspec/security_hooks.go index a92462c..bcdec15 100644 --- a/pkg/mqttspec/security_hooks.go +++ b/pkg/mqttspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the MQTT handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvemcp/security_hooks.go b/pkg/resolvemcp/security_hooks.go index cb8f2ab..f7d8bd4 100644 --- a/pkg/resolvemcp/security_hooks.go +++ b/pkg/resolvemcp/security_hooks.go @@ -19,6 +19,12 @@ import ( // - Column-level security: sensitive columns masked/hidden in read results. // - Audit logging after each read. func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // BeforeHandle: enforce model-level operation rules (auth check). handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvespec/security_hooks.go b/pkg/resolvespec/security_hooks.go index 450e33e..5db28a5 100644 --- a/pkg/resolvespec/security_hooks.go +++ b/pkg/resolvespec/security_hooks.go @@ -11,6 +11,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvespec/tx_settings_test.go b/pkg/resolvespec/tx_settings_test.go new file mode 100644 index 0000000..42d1918 --- /dev/null +++ b/pkg/resolvespec/tx_settings_test.go @@ -0,0 +1,78 @@ +package resolvespec + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" +) + +// stubProvider satisfies security.SecurityProvider; its methods are never called +// because the model rules are public and no rules are loaded for this test. +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsFirstOnEveryTx(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + base = context.WithValue(base, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleDelete(ctx, w, "7", nil) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestSecurityHooksTxSettingsErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + list, _ := security.NewSecurityList(stubProvider{}) + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"bad name": "1"}, nil + }) + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + base = context.WithValue(base, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleDelete(ctx, w, "7", nil) + + if rec.Code == http.StatusOK { + t.Fatal("a failed stamp must not let the delete proceed") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/restheadspec/security_hooks.go b/pkg/restheadspec/security_hooks.go index b9365b6..e1c18f8 100644 --- a/pkg/restheadspec/security_hooks.go +++ b/pkg/restheadspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/security/README.md b/pkg/security/README.md index 0fd0a69..7f11ace 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -1213,6 +1213,8 @@ The main changes: ## Documentation +- [Request transactions and RLS stamping](../common/TRANSACTIONS.md) + | File | Description | |------|-------------| | **QUICK_REFERENCE.md** | Quick reference guide with examples | diff --git a/pkg/security/provider.go b/pkg/security/provider.go index f7a2ee4..547e70b 100644 --- a/pkg/security/provider.go +++ b/pkg/security/provider.go @@ -132,6 +132,10 @@ type SecurityList struct { lastColPrune time.Time lastRowPrune time.Time + // txSettings stamps transaction-local settings at OnTxBegin (see txsettings.go). + txSettingsMu sync.RWMutex + txSettings TxSettingsFunc + // loads collapses concurrent provider calls for the same key (cold-cache stampede). loads singleflight.Group } diff --git a/pkg/security/txsettings.go b/pkg/security/txsettings.go new file mode 100644 index 0000000..ebec0cc --- /dev/null +++ b/pkg/security/txsettings.go @@ -0,0 +1,86 @@ +package security + +import ( + "encoding/hex" + "fmt" + "regexp" + "sort" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// TxSettingsFunc returns the transaction-local settings (e.g. RLS GUCs such as +// "app.user_id") to stamp on a transaction. It runs once per transaction, at +// OnTxBegin, before any other SQL. Returning an error rolls the transaction back. +type TxSettingsFunc func(secCtx SecurityContext) (map[string]string, error) + +// settingNameRE matches a custom GUC name: two or more dot-separated identifiers. +var settingNameRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)+$`) + +// SetTxSettings sets the function that provides transaction-local settings for +// every transaction opened by a spec that registered its security hooks with this +// list. Pass nil to disable. May be called before or after RegisterSecurityHooks. +func (m *SecurityList) SetTxSettings(fn TxSettingsFunc) { + m.txSettingsMu.Lock() + defer m.txSettingsMu.Unlock() + m.txSettings = fn +} + +// TxSettings returns the configured TxSettingsFunc, or nil. +func (m *SecurityList) TxSettings() TxSettingsFunc { + m.txSettingsMu.RLock() + defer m.txSettingsMu.RUnlock() + return m.txSettings +} + +// StampTxSettings runs the list's TxSettingsFunc and applies the result to tx as +// transaction-local settings (set_config(name, value, true)). No-op when no +// function is configured or it returns no settings. tx must be the transaction +// itself, never the pool: the settings are lost on any other connection. +func StampTxSettings(secCtx SecurityContext, list *SecurityList, tx common.Database) error { + if list == nil { + return nil + } + fn := list.TxSettings() + if fn == nil { + return nil + } + settings, err := fn(secCtx) + if err != nil { + return err + } + return ApplyTxSettings(secCtx, tx, settings) +} + +// 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(secCtx SecurityContext, tx common.Database, settings map[string]string) error { + if len(settings) == 0 { + return nil + } + if tx == nil { + return fmt.Errorf("tx settings: no transaction") + } + if drv := tx.DriverName(); drv != "postgres" && drv != "pgsql" { + return fmt.Errorf("tx settings: unsupported driver %q", drv) + } + names := make([]string, 0, len(settings)) + for name := range settings { + if !settingNameRE.MatchString(name) { + return fmt.Errorf("tx settings: invalid setting name %q", name) + } + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + // The value is hex-encoded so it needs no quoting and cannot be read as a + // bind placeholder by any adapter. + query := fmt.Sprintf("SELECT set_config('%s', convert_from(decode('%s', 'hex'), 'UTF8'), true)", + name, hex.EncodeToString([]byte(settings[name]))) + if _, err := tx.Exec(secCtx.GetContext(), query); err != nil { + return fmt.Errorf("tx settings: set %s: %w", name, err) + } + } + return nil +} diff --git a/pkg/security/txsettings_test.go b/pkg/security/txsettings_test.go new file mode 100644 index 0000000..1b1bd37 --- /dev/null +++ b/pkg/security/txsettings_test.go @@ -0,0 +1,131 @@ +package security + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" +) + +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) + sc := &mockSecurityContext{ctx: context.Background()} + + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.tenant', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec(`set_config\('app\.user_id', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + + err := pool.RunInTransaction(context.Background(), func(tx common.Database) error { + return ApplyTxSettings(sc, tx, map[string]string{"app.user_id": "7", "app.tenant": "o'x?"}) + }) + if err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestApplyTxSettingsValueIsNeverInlined(t *testing.T) { + pool, mock := txSettingsDB(t) + sc := &mockSecurityContext{ctx: context.Background()} + var seen string + mock.ExpectBegin() + mock.ExpectExec(`set_config`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + _ = pool.RunInTransaction(context.Background(), func(tx common.Database) error { + // Capture via a wrapper so the raw statement can be inspected. + return ApplyTxSettings(sc, &queryRecorder{Database: tx, got: &seen}, map[string]string{"app.v": "'; DROP TABLE x; --?"}) + }) + if strings.Contains(seen, "DROP") || strings.Contains(seen, "?") { + t.Fatalf("value leaked into SQL text: %s", seen) + } +} + +type queryRecorder struct { + common.Database + got *string +} + +func (q *queryRecorder) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) { + *q.got = query + return q.Database.Exec(ctx, query, args...) +} + +func TestApplyTxSettingsRejectsBadNameAndDriver(t *testing.T) { + pool, _ := txSettingsDB(t) + sc := &mockSecurityContext{ctx: context.Background()} + + for _, name := range []string{"user_id", "app.x'); DROP", "app..x", "app.x y"} { + if err := ApplyTxSettings(sc, pool, map[string]string{name: "1"}); err == nil { + t.Fatalf("name %q must be rejected", name) + } + } + if err := ApplyTxSettings(sc, pool, nil); err != nil { + t.Fatalf("empty settings must be a no-op: %v", err) + } + if err := ApplyTxSettings(sc, &driverStub{Database: pool, name: "sqlite"}, map[string]string{"app.x": "1"}); err == nil { + t.Fatal("non-postgres driver must fail closed") + } +} + +type driverStub struct { + common.Database + name string +} + +func (d *driverStub) DriverName() string { return d.name } + +func TestStampTxSettingsUsesConfiguredFunc(t *testing.T) { + pool, mock := txSettingsDB(t) + sc := &mockSecurityContext{ctx: context.Background()} + list, err := NewSecurityList(&mockSecurityProvider{}) + if err != nil { + t.Fatal(err) + } + + // Nil list and unset func are no-ops (no SQL expected). + if err := StampTxSettings(sc, nil, pool); err != nil { + t.Fatal(err) + } + if err := StampTxSettings(sc, list, pool); err != nil { + t.Fatal(err) + } + + list.SetTxSettings(func(SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + if err := pool.RunInTransaction(context.Background(), func(tx common.Database) error { + return StampTxSettings(sc, list, tx) + }); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + + boom := errors.New("no tenant") + list.SetTxSettings(func(SecurityContext) (map[string]string, error) { return nil, boom }) + if err := StampTxSettings(sc, list, pool); !errors.Is(err, boom) { + t.Fatalf("func error must propagate, got %v", err) + } +} diff --git a/pkg/websocketspec/security_hooks.go b/pkg/websocketspec/security_hooks.go index f5596a0..41d5c21 100644 --- a/pkg/websocketspec/security_hooks.go +++ b/pkg/websocketspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { From da0b1f51233fd26e5a31d124949ab01781c805a8 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:01:52 +0200 Subject: [PATCH 16/23] chore(testserver): host networking, ports 8123/8124, smoke read+update, dbtrace pooled=0 verified --- audit/single_tran.md | 3 ++- docker-compose.yml | 20 +++++++------------- docker/Dockerfile.testserver | 2 +- docker/testserver.config.yaml | 6 +++--- pkg/resolvespec/integration_test.go | 2 +- pkg/restheadspec/integration_test.go | 2 +- scripts/testserver-smoke.sh | 4 +++- 7 files changed, 18 insertions(+), 21 deletions(-) diff --git a/audit/single_tran.md b/audit/single_tran.md index fa0b2b4..d212f67 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -95,7 +95,8 @@ - DONE P6: funcspec. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx` for `SqlQuery` and `SqlQueryList`. `BeforeResponse` now runs in a second short tx (`Tx` is no longer the pool). `BeforeOp` is unchanged (still per statement). A begin/`OnTxBegin`/commit failure answers 500 `transaction_error` / "Transaction failed" (before, it returned with no response); body failures still answer via `sendError`. Tests: `pkg/funcspec/tx_test.go`. - DONE (AfterRead, decided by user): restheadspec `AfterRead` now runs in a second short tx. Test: `pkg/restheadspec/read_tx_test.go`. - DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. -- NEXT: extra tests (create/update for other specs, `dbtrace` `pooled == 0` on real Postgres). +- DONE real-Postgres check (resolvespec, testserver via compose): create `tx=1 pooled=0`, read `tx=1 pooled=0`, update `tx=2 pooled=0` (was `pooled=1`), single delete `tx=1 pooled=0`, batch create/delete `tx=1 pooled=0`. Compose now uses host networking (bridge fails here): testserver on 8123, Postgres on 8124 (was 8080/5434); integration test DSNs updated. Smoke script covers read and update. websocketspec/mqttspec/resolvemcp/restheadspec/funcspec not measured on real Postgres. +- NEXT: extra per-spec create/update tests (optional). ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/docker-compose.yml b/docker-compose.yml index 4f5c819..4304ead 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -6,17 +6,17 @@ services: POSTGRES_USER: postgres POSTGRES_PASSWORD: postgres POSTGRES_DB: postgres - ports: - - "5434:5432" + # Host networking (bridge networks are unavailable in some environments): + # postgres listens directly on host port 8124. + network_mode: host + command: ["postgres", "-p", "8124"] volumes: - postgres-test-data:/var/lib/postgresql/data healthcheck: - test: ["CMD-SHELL", "pg_isready -U postgres"] + test: ["CMD-SHELL", "pg_isready -U postgres -p 8124"] interval: 5s timeout: 5s retries: 5 - networks: - - resolvespec-test testserver: build: @@ -27,18 +27,12 @@ services: RESOLVESPEC_DB_TRACE_ENABLED: "true" RESOLVESPEC_DB_TRACE_MIN_CALLS: "1" RESOLVESPEC_DB_TRACE_POOL_LOG: "true" - ports: - - "8080:8080" + # Serves on host port 8123 (docker/testserver.config.yaml). + network_mode: host depends_on: postgres-test: condition: service_healthy - networks: - - resolvespec-test volumes: postgres-test-data: driver: local - -networks: - resolvespec-test: - driver: bridge diff --git a/docker/Dockerfile.testserver b/docker/Dockerfile.testserver index f937a3b..50c7a50 100644 --- a/docker/Dockerfile.testserver +++ b/docker/Dockerfile.testserver @@ -9,5 +9,5 @@ FROM alpine:3.20 RUN apk add --no-cache ca-certificates COPY --from=build /out/testserver /usr/local/bin/testserver COPY docker/testserver.config.yaml /etc/resolvespec/config.yaml -EXPOSE 8080 +EXPOSE 8123 ENTRYPOINT ["testserver"] diff --git a/docker/testserver.config.yaml b/docker/testserver.config.yaml index 5a8fa98..9613e53 100644 --- a/docker/testserver.config.yaml +++ b/docker/testserver.config.yaml @@ -12,7 +12,7 @@ servers: main: name: "main" host: "0.0.0.0" - port: 8080 + port: 8123 description: "Main server instance" gzip: true tags: @@ -79,8 +79,8 @@ dbmanager: default: name: "default" type: "postgres" - host: "postgres-test" - port: 5432 + host: "localhost" + port: 8124 user: "postgres" password: "postgres" database: "postgres" diff --git a/pkg/resolvespec/integration_test.go b/pkg/resolvespec/integration_test.go index b7b8c6c..01042e0 100644 --- a/pkg/resolvespec/integration_test.go +++ b/pkg/resolvespec/integration_test.go @@ -67,7 +67,7 @@ func setupTestDB(t *testing.T) *gorm.DB { // Get connection string from environment or use default dsn := os.Getenv("TEST_DATABASE_URL") if dsn == "" { - dsn = "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5434 sslmode=disable" + dsn = "host=localhost user=postgres password=postgres dbname=resolvespec_test port=8124 sslmode=disable" } db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ diff --git a/pkg/restheadspec/integration_test.go b/pkg/restheadspec/integration_test.go index 4cdbf96..1a1d2d6 100644 --- a/pkg/restheadspec/integration_test.go +++ b/pkg/restheadspec/integration_test.go @@ -67,7 +67,7 @@ func setupTestDB(t *testing.T) *gorm.DB { // Get connection string from environment or use default dsn := os.Getenv("TEST_DATABASE_URL") if dsn == "" { - dsn = "host=localhost user=postgres password=postgres dbname=restheadspec_test port=5434 sslmode=disable" + dsn = "host=localhost user=postgres password=postgres dbname=restheadspec_test port=8124 sslmode=disable" } db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ diff --git a/scripts/testserver-smoke.sh b/scripts/testserver-smoke.sh index de72346..30967f1 100755 --- a/scripts/testserver-smoke.sh +++ b/scripts/testserver-smoke.sh @@ -3,7 +3,7 @@ # Usage: scripts/testserver-smoke.sh [base_url] (COMPOSE overrides the compose command) set -euo pipefail -BASE="${1:-http://localhost:8080}" +BASE="${1:-http://localhost:8123}" if [ -z "${COMPOSE:-}" ]; then if command -v podman >/dev/null 2>&1; then COMPOSE="podman compose"; else COMPOSE="docker compose"; fi fi @@ -24,6 +24,8 @@ ids() { grep -o '"id":[0-9]*' "$BODY" | cut -d: -f2; } expect create 200 "$(call "{\"operation\":\"create\",\"data\":{\"name\":\"Smoke\",\"code\":\"S$TS\"}}")" ID="$(ids | head -1)" +expect read 200 "$(call '{"operation":"read"}' "/$ID")" +expect update 200 "$(call '{"operation":"update","data":{"name":"Smoke2"}}' "/$ID")" expect delete 200 "$(call '{"operation":"delete"}' "/$ID")" expect delete-again 404 "$(call '{"operation":"delete"}' "/$ID")" From 129c1a043db9c3750bf4b600914c139125d77062 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:03:50 +0200 Subject: [PATCH 17/23] docs(readme): document single transaction per request, OnTxBegin, test server --- README.md | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/README.md b/README.md index 4256c31..5d638a5 100644 --- a/README.md +++ b/README.md @@ -43,6 +43,7 @@ All share the same core architecture and provide dynamic data querying, relation * **Pagination**: Built-in limit/offset and cursor-based pagination (both ResolveSpec and RestHeadSpec) * **Computed Columns**: Define virtual columns for complex calculations * **Custom Operators**: Add custom SQL conditions when needed +* **🆕 One Transaction Per Request**: Every statement and DB-touching hook of a request runs on one transaction; `OnTxBegin` hook stamps transaction-local settings (RLS) first. See [pkg/common/TRANSACTIONS.md](pkg/common/TRANSACTIONS.md) * **🆕 Recursive CRUD Handler**: Automatically handle nested object graphs with foreign key resolution and per-record operation control via `_request` field ### Architecture (v2.0+) @@ -370,6 +371,14 @@ ResolveSpec is designed for testability with mockable interfaces. For testing ex - [RestHeadSpec Testing](pkg/restheadspec/README.md#testing) - [WebSocketSpec Testing](pkg/websocketspec/README.md) +### Test Server (dbtrace, real PostgreSQL) + +* `make testserver-up` / `make testserver-down`: testserver + PostgreSQL via compose (host networking) +* `make testserver-smoke`: create, read, update, delete, batch create/delete against the testserver +* Ports: testserver `8123`, PostgreSQL `8124` +* `dbtrace` logs per request `tx`, `tx_queries`, `pooled`, `raw`; `pooled=0` is the target +* Integration tests default to PostgreSQL on `localhost:8124` + ## Continuous Integration ResolveSpec uses GitHub Actions for automated testing and quality checks. The CI pipeline runs on every push and pull request. @@ -683,6 +692,24 @@ This project is licensed under the MIT License - see the [LICENSE](LICENSE) file ## What's New +### Unreleased + +**Single transaction per request**: + +* **One tx per request**: hooks get the transaction in `hookCtx.Tx`, never the pool (`BeforeHandle` runs before any tx and must not touch the DB) +* **`OnTxBegin` hook**: all specs (mqttspec re-exports websocketspec's); fires once, first, in every tx; error or abort rolls back with no detail to the client +* **Second short tx**: create/update re-fetch, `BeforeScan` and post-commit hooks (`AfterCreate`, `AfterUpdate`, restheadspec `AfterRead`, funcspec `BeforeResponse`) run on a new tx after the first commits +* **Delete**: single and batch delete, hooks included, in one tx +* **websocketspec / mqttspec**: one tx per message; begin/commit failures answer `transaction_error` +* **resolvemcp**: read, create, update, delete transactional +* **RLS stamping**: `SecurityList.SetTxSettings(fn)`; `RegisterSecurityHooks` of every spec stamps `set_config(name, value, true)` on `OnTxBegin`; fails closed +* **New**: `common.RunRequestTx`, `common.TxContext`, `common.TxHookName` +* **Behavior changes**: `AfterDelete` failure now rolls the delete back; funcspec begin/commit failure answers 500 `transaction_error` + +**Clients**: Go, Rust, C# and Dart clients for ResolveSpec and FunctionSpec under `clients/`. + +**Test server**: compose uses host networking; ports `8123` (testserver) and `8124` (PostgreSQL), previously `8080` and `5434`. + ### v3.2 (Latest - March 2026) **ResolveMCP - Model Context Protocol Server (🆕)**: From 6335bfe87e015fb14db9153c86db47b7999f4f4b Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:04:56 +0200 Subject: [PATCH 18/23] docs(readme): reference all pkg packages and clients --- README.md | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/README.md b/README.md index 5d638a5..89f7139 100644 --- a/README.md +++ b/README.md @@ -491,6 +491,19 @@ Execute SQL functions and queries through a simple HTTP API with header-based pa For complete documentation, see [pkg/funcspec/](pkg/funcspec/). +#### Clients + +All clients are under [clients/](clients/README.md); wire behaviour is identical across them. + +| Client | Language | Specs | Docs | +|---|---|---|---| +| `resolvespec-js` | TypeScript | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-js/README.md) | +| `resolvespec-python` | Python >= 3.11 | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-python/README.md) | +| `resolvespec-go` | Go | ResolveSpec, FunctionSpec | [README](clients/resolvespec-go/README.md) | +| `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | [README](clients/resolvespec-rs/README.md) | +| `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | [README](clients/resolvespec-cs/README.md) | +| `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | [README](clients/resolvespec-dart/README.md) | + #### ResolveSpec JS - TypeScript Client Library TypeScript/JavaScript client library supporting all three REST and WebSocket protocols. @@ -669,6 +682,23 @@ Configuration management with support for multiple formats and environments. For documentation, see [pkg/config/README.md](pkg/config/README.md). +#### DB Trace + +Per-request DB call counting (`tx`, `tx_queries`, `pooled`, `raw`) and pool logging. Off by default. + +For documentation, see [pkg/dbtrace/README.md](pkg/dbtrace/README.md). + +### Core Libraries + +| Package | Purpose | +|---|---| +| [`pkg/common`](pkg/common/) | Shared interfaces (database, request/response adapters), validation, recursive CRUD, request transactions ([TRANSACTIONS.md](pkg/common/TRANSACTIONS.md)) | +| [`pkg/modelregistry`](pkg/modelregistry/) | Model registration by schema/entity and per-model access rules | +| [`pkg/reflection`](pkg/reflection/) | Model/struct reflection helpers (primary keys, columns, relations) | +| [`pkg/spectypes`](pkg/spectypes/) | SQL-aware types (nullable, JSONB, PostGIS, vector) | +| [`pkg/logger`](pkg/logger/) | Logging used by all packages | +| [`pkg/testmodels`](pkg/testmodels/) | Shared test models and data for tests and the testserver | + ## Security Considerations * Implement proper authentication and authorization From 204220581794a6987b5db6242b048b416113c33f Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:11:30 +0200 Subject: [PATCH 19/23] fix(pgsql): return subquery preload errors instead of logging and continuing --- pkg/common/adapters/database/pgsql.go | 6 +- .../adapters/database/pgsql_preload_test.go | 56 +++++++++++++++++++ 2 files changed, 60 insertions(+), 2 deletions(-) create mode 100644 pkg/common/adapters/database/pgsql_preload_test.go diff --git a/pkg/common/adapters/database/pgsql.go b/pkg/common/adapters/database/pgsql.go index 96b3b5c..8d7ff7a 100644 --- a/pkg/common/adapters/database/pgsql.go +++ b/pkg/common/adapters/database/pgsql.go @@ -1208,7 +1208,7 @@ func (p *PgSQLSelectQuery) applySubqueryPreloads(ctx context.Context, dest inter for i := 0; i < destValue.Len(); i++ { elem := destValue.Index(i) if err := p.loadPreloadsForRecord(ctx, elem, subqueryPreloads); err != nil { - logger.Warn("Failed to load preloads for record %d: %v", i, err) + return fmt.Errorf("record %d: %w", i, err) } } return nil @@ -1256,7 +1256,9 @@ func (p *PgSQLSelectQuery) loadPreloadsForRecord(ctx context.Context, record ref // Build and execute the preload query err := p.executePreloadQuery(ctx, field, meta, fkValue, preload) if err != nil { - logger.Warn("Failed to execute preload query for '%s': %v", preload.relation, err) + // Inside a transaction a failed statement aborts it, so carrying on would + // only turn into a misleading "transaction is aborted" on the next query. + return fmt.Errorf("preload %s: %w", preload.relation, err) } } diff --git a/pkg/common/adapters/database/pgsql_preload_test.go b/pkg/common/adapters/database/pgsql_preload_test.go new file mode 100644 index 0000000..ee3f6e5 --- /dev/null +++ b/pkg/common/adapters/database/pgsql_preload_test.go @@ -0,0 +1,56 @@ +package database + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +type preloadChild struct { + ID int `db:"id"` + UserID int `db:"user_id"` +} + +func (preloadChild) TableName() string { return "children" } + +type preloadParent struct { + ID int `db:"id"` + Children []preloadChild `bun:"rel:has-many,join:ID=user_id"` +} + +func (preloadParent) TableName() string { return "parents" } + +func TestSubqueryPreloadErrorIsReturned(t *testing.T) { + for name, dest := range map[string]interface{}{ + "slice": &[]preloadParent{}, + "single": &preloadParent{}, + } { + t.Run(name, func(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + defer db.Close() + + mock.ExpectBegin() + mock.ExpectQuery(`FROM parents`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`FROM children`).WillReturnError(errors.New("boom")) + mock.ExpectRollback() + + err = NewPgSQLAdapter(db).RunInTransaction(context.Background(), func(tx common.Database) error { + return tx.NewSelect().Model(&preloadParent{}).PreloadRelation("Children").Scan(context.Background(), dest) + }) + if err == nil || !strings.Contains(err.Error(), "preload Children") { + t.Fatalf("expected the preload error, got %v", err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} From 7f84debdc53095a4d1f1139ec0cc3bfa027b3849 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:19:58 +0200 Subject: [PATCH 20/23] test(tx): regression tests for per-request transactions across all specs --- audit/single_tran.md | 5 +- .../adapters/database/pgsql_preload_test.go | 33 ++++ pkg/common/tx_guard_test.go | 99 ++++++++++ pkg/funcspec/tx_test.go | 43 ++++ pkg/mqttspec/tx_test.go | 51 +++++ pkg/resolvemcp/tx_test.go | 119 +++++++++++ pkg/resolvespec/ops_tx_test.go | 186 ++++++++++++++++++ pkg/restheadspec/read_tx_test.go | 99 ++++++++++ pkg/websocketspec/tx_test.go | 78 ++++++++ 9 files changed, 712 insertions(+), 1 deletion(-) create mode 100644 pkg/common/tx_guard_test.go create mode 100644 pkg/resolvespec/ops_tx_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index d212f67..380f30e 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -96,7 +96,10 @@ - DONE (AfterRead, decided by user): restheadspec `AfterRead` now runs in a second short tx. Test: `pkg/restheadspec/read_tx_test.go`. - DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. - DONE real-Postgres check (resolvespec, testserver via compose): create `tx=1 pooled=0`, read `tx=1 pooled=0`, update `tx=2 pooled=0` (was `pooled=1`), single delete `tx=1 pooled=0`, batch create/delete `tx=1 pooled=0`. Compose now uses host networking (bridge fails here): testserver on 8123, Postgres on 8124 (was 8080/5434); integration test DSNs updated. Smoke script covers read and update. websocketspec/mqttspec/resolvemcp/restheadspec/funcspec not measured on real Postgres. -- NEXT: extra per-spec create/update tests (optional). +- DONE regression tests: per-spec read/create/update/delete hook-on-tx, failure-rollback (Before*/After*/`OnTxBegin`) and second-tx tests in all six specs (`ops_tx_test.go`, `tx_test.go`, `read_tx_test.go`); stamping tests for resolvespec, resolvemcp, funcspec; pgsql adapter preload tests (same connection, error returned); source guard `pkg/common/tx_guard_test.go` (no direct `RunInTransaction`/`BeginTx`, no `Tx = h.db` beyond the allowlisted BeforeHandle placeholders, no pool statements in spec handlers). +- NOTE: resolvespec never fires `AfterRead`/`AfterCreate` (hook types exist, no call site); not tested, pre-existing. +- NOTE: restheadspec total-count cache is process-wide and ignores the record id; read tests call `resetTotalCache`. +- NOTE: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (also on `cd96404`, before this work); use `-count=1`. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/common/adapters/database/pgsql_preload_test.go b/pkg/common/adapters/database/pgsql_preload_test.go index ee3f6e5..6d93103 100644 --- a/pkg/common/adapters/database/pgsql_preload_test.go +++ b/pkg/common/adapters/database/pgsql_preload_test.go @@ -5,6 +5,7 @@ import ( "errors" "strings" "testing" + "time" "github.com/DATA-DOG/go-sqlmock" @@ -54,3 +55,35 @@ func TestSubqueryPreloadErrorIsReturned(t *testing.T) { }) } } + +// With a single pooled connection, a preload that escaped the transaction would +// block on the pool and fail on the context timeout. +func TestSubqueryPreloadRunsOnTheTransaction(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + + mock.ExpectBegin() + mock.ExpectQuery(`FROM parents`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`FROM children`).WillReturnRows(sqlmock.NewRows([]string{"id", "user_id"}).AddRow(10, 1)) + mock.ExpectCommit() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + var parents []preloadParent + err = NewPgSQLAdapter(db).RunInTransaction(ctx, func(tx common.Database) error { + return tx.NewSelect().Model(&preloadParent{}).PreloadRelation("Children").Scan(ctx, &parents) + }) + if err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(parents) != 1 || len(parents[0].Children) != 1 { + t.Fatalf("preloaded children missing: %+v", parents) + } +} diff --git a/pkg/common/tx_guard_test.go b/pkg/common/tx_guard_test.go new file mode 100644 index 0000000..2c943f5 --- /dev/null +++ b/pkg/common/tx_guard_test.go @@ -0,0 +1,99 @@ +package common + +import ( + "os" + "path/filepath" + "regexp" + "strings" + "testing" +) + +// Source-level regression guard for the single-transaction-per-request rule +// (audit/single_tran.md). Runtime tests prove the current paths; this catches a +// new code path that quietly reaches for the pool. + +var guardedSpecs = []string{"resolvespec", "restheadspec", "websocketspec", "mqttspec", "resolvemcp", "funcspec"} + +var ( + directTxRE = regexp.MustCompile(`\.(RunInTransaction|BeginTx)\(`) + poolHookTxRE = regexp.MustCompile(`\bTx(:\s+|\s*=\s*)(h|handler|h\.handler)\.db\b`) + poolQueryRE = regexp.MustCompile(`\b(h|handler)\.db\.(NewSelect|NewInsert|NewUpdate|NewDelete|Exec|Query)\(`) +) + +// allowedPoolHookTx: hook contexts that start life on the pool before the handler +// opens its transaction (BeforeHandle runs before any tx and must be DB-free). +// runInTx replaces Tx with the transaction before any other hook runs. +var allowedPoolHookTx = map[string]int{ + "resolvespec/handler.go": 1, + "websocketspec/handler.go": 1, + "resolvemcp/handler.go": 4, +} + +// allowedPoolQuery: statements outside the request path. +var allowedPoolQuery = map[string]int{ + "resolvemcp/annotation.go": 2, // tool annotations, not a data request +} + +func guardedFiles(t *testing.T) map[string][]string { + t.Helper() + out := map[string][]string{} + for _, spec := range guardedSpecs { + files, err := filepath.Glob(filepath.Join("..", spec, "*.go")) + if err != nil || len(files) == 0 { + t.Fatalf("no sources found for %s: %v", spec, err) + } + for _, f := range files { + if strings.HasSuffix(f, "_test.go") { + continue + } + raw, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + var lines []string + for _, l := range strings.Split(string(raw), "\n") { + if s := strings.TrimSpace(l); strings.HasPrefix(s, "//") { + continue + } + lines = append(lines, l) + } + out[spec+"/"+filepath.Base(f)] = lines + } + } + return out +} + +func countMatches(lines []string, re *regexp.Regexp) int { + n := 0 + for _, l := range lines { + if re.MatchString(l) { + n++ + } + } + return n +} + +func TestNoDirectTransactionsInSpecHandlers(t *testing.T) { + for file, lines := range guardedFiles(t) { + if n := countMatches(lines, directTxRE); n > 0 { + t.Errorf("%s opens a transaction directly (%d): use the handler's runInTx so OnTxBegin fires", file, n) + } + } +} + +func TestHookContextsDoNotRetainThePool(t *testing.T) { + files := guardedFiles(t) + for file, lines := range files { + if got, want := countMatches(lines, poolHookTxRE), allowedPoolHookTx[file]; got != want { + t.Errorf("%s has %d hook contexts set to the pool, allowed %d: hooks must get the transaction", file, got, want) + } + } +} + +func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) { + for file, lines := range guardedFiles(t) { + if got, want := countMatches(lines, poolQueryRE), allowedPoolQuery[file]; got != want { + t.Errorf("%s runs %d statements on the pool, allowed %d: use the transaction", file, got, want) + } + } +} diff --git a/pkg/funcspec/tx_test.go b/pkg/funcspec/tx_test.go index 3ad946f..df95ba2 100644 --- a/pkg/funcspec/tx_test.go +++ b/pkg/funcspec/tx_test.go @@ -9,6 +9,7 @@ import ( "testing" "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" ) // txFactory returns a pool whose every transaction is a distinct MockDatabase. @@ -100,3 +101,45 @@ func TestOnTxBeginErrorAnswersTransactionError(t *testing.T) { t.Fatalf("no query may run after a failed OnTxBegin, ran %d", queries) } } + +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { + var execs []string + h := NewHandler(&MockDatabase{ + RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error { + return fn(&MockDatabase{ + ExecFunc: func(ctx context.Context, query string, args ...interface{}) (common.Result, error) { + execs = append(execs, query) + return &MockResult{}, nil + }, + QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error { + execs = append(execs, "QUERY") + if rows, ok := dest.(*[]map[string]interface{}); ok { + *rows = []map[string]interface{}{{"id": float64(1)}} + } + return nil + }, + }) + }, + }) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "1"}, nil + }) + RegisterSecurityHooks(h, list) + + w := httptest.NewRecorder() + h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil)) + + if w.Code != http.StatusOK { + t.Fatalf("status %d body %s", w.Code, w.Body) + } + // tx 1: stamp, then the query; tx 2 (BeforeResponse): stamp again. + if len(execs) != 3 || !strings.Contains(execs[0], "set_config('app.user_id'") || execs[1] != "QUERY" || !strings.Contains(execs[2], "set_config('app.user_id'") { + t.Fatalf("each transaction must be stamped before any other SQL, got %v", execs) + } +} diff --git a/pkg/mqttspec/tx_test.go b/pkg/mqttspec/tx_test.go index 4293e92..8ae5acc 100644 --- a/pkg/mqttspec/tx_test.go +++ b/pkg/mqttspec/tx_test.go @@ -89,3 +89,54 @@ func TestHandler_UpdateRunsAfterHookOnSecondTransaction(t *testing.T) { assert.NotEqual(t, begins[0], begins[1]) assert.Equal(t, begins[1], afterTx) } + +func TestHandler_ReadRunsHooksOnOneTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var begins []common.Database + var beforeTx, afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + handler.hooks.Register(BeforeRead, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + handler.hooks.Register(AfterRead, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + handler.handleRead(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + require.Len(t, begins, 1) + assert.NotEqual(t, handler.db, begins[0]) + assert.Equal(t, begins[0], beforeTx) + assert.Equal(t, begins[0], afterTx) +} + +func TestHandler_CreateRunsHooksOnTwoTransactions(t *testing.T) { + handler, db := setupTestHandler(t) + + var begins []common.Database + var beforeTx, afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + handler.hooks.Register(BeforeCreate, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + handler.hooks.Register(AfterCreate, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + hook := newTxHook(handler, map[string]interface{}{"id": 5, "name": "n", "email": "n@example.com", "status": "active"}) + hook.ID = "" + handler.handleCreate(&Client{ID: "c1"}, &Message{ID: "m1"}, hook) + + var got TestUser + require.NoError(t, db.First(&got, 5).Error) + require.Len(t, begins, 2) + assert.NotEqual(t, begins[0], begins[1]) + assert.Equal(t, begins[0], beforeTx) + assert.Equal(t, begins[1], afterTx) +} + +func TestHandler_BeforeHookErrorAbortsUpdateWithoutWriting(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + handler.hooks.Register(BeforeUpdate, func(c *HookContext) error { return errors.New("denied") }) + + handler.handleUpdate(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, map[string]interface{}{"name": "b"})) + + var got TestUser + require.NoError(t, db.First(&got, 1).Error) + assert.Equal(t, "a", got.Name) +} diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go index 80f4945..3bff822 100644 --- a/pkg/resolvemcp/tx_test.go +++ b/pkg/resolvemcp/tx_test.go @@ -11,6 +11,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" + "github.com/bitechdev/ResolveSpec/pkg/security" ) type txItem struct { @@ -205,3 +206,121 @@ func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) { } } } + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return sql.ErrConnDone }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil { + t.Fatal("a failing AfterDelete must fail the request") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateBeforeHookErrorRollsBackWithoutInsert(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(BeforeCreate, func(*HookContext) error { return sql.ErrConnDone }) + + mock.ExpectBegin() + mock.ExpectRollback() + + if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateInsertErrorRollsBackBatch(t *testing.T) { + h, mock, ctx := newTxHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`INSERT`).WillReturnError(sql.ErrConnDone) + mock.ExpectRollback() + + items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}} + if _, err := h.executeCreate(ctx, "public", "items", items); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestOnTxBeginErrorRollsBackEveryOperation(t *testing.T) { + ops := map[string]func(h *Handler, ctx context.Context) error{ + "read": func(h *Handler, ctx context.Context) error { + _, _, err := h.executeRead(ctx, "public", "items", "7", common.RequestOptions{}) + return err + }, + "create": func(h *Handler, ctx context.Context) error { + _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}) + return err + }, + "update": func(h *Handler, ctx context.Context) error { + _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "a"}) + return err + }, + "delete": func(h *Handler, ctx context.Context) error { + _, err := h.executeDelete(ctx, "public", "items", "7") + return err + }, + } + for name, op := range ops { + t.Run(name, func(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(OnTxBegin, func(*HookContext) error { return sql.ErrConnDone }) + mock.ExpectBegin() + mock.ExpectRollback() + if err := op(h, ctx); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + RegisterSecurityHooks(h, list) + ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + + // Update opens two transactions; each must be stamped before any other SQL. + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"}); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/resolvespec/ops_tx_test.go b/pkg/resolvespec/ops_tx_test.go new file mode 100644 index 0000000..aa75cba --- /dev/null +++ b/pkg/resolvespec/ops_tx_test.go @@ -0,0 +1,186 @@ +package resolvespec + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func opCtx(t *testing.T) (context.Context, common.ResponseWriter, *httptest.ResponseRecorder) { + t.Helper() + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + return WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}), w, rec +} + +// hookTrace records hook order and the Tx each hook saw. +type hookTrace struct { + order []string + tx map[HookType][]common.Database +} + +func traceHooks(h *Handler, types ...HookType) *hookTrace { + tr := &hookTrace{tx: map[HookType][]common.Database{}} + for _, ht := range types { + ht := ht + h.Hooks().Register(ht, func(c *HookContext) error { + tr.order = append(tr.order, string(ht)) + tr.tx[ht] = append(tr.tx[ht], c.Tx) + return nil + }) + } + return tr +} + +func (tr *hookTrace) mustBeOn(t *testing.T, tx common.Database, types ...HookType) { + t.Helper() + for _, ht := range types { + if len(tr.tx[ht]) == 0 || tr.tx[ht][0] != tx { + t.Fatalf("%s must run on the OnTxBegin transaction", ht) + } + } +} + +func TestReadRunsHooksOnOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeRead) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[OnTxBegin]) != 1 || tr.order[0] != "on_tx_begin" { + t.Fatalf("OnTxBegin must fire once and first, got %v", tr.order) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeRead) +} + +func TestReadBeforeHookErrorRollsBackWithoutQueries(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(BeforeRead, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeRead must not return data: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateRunsHooksOnTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if tr.order[0] != "on_tx_begin" { + t.Fatalf("OnTxBegin must fire first, got %v", tr.order) + } + for _, tx := range tr.tx[BeforeCreate] { + if tx == nil || tx == h.db { + t.Fatal("BeforeCreate must not get the pool") + } + } +} + +func TestCreateBeforeHookErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(BeforeCreate, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeCreate must not create: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestUpdateAfterHookErrorRollsBackAndSkipsRefetch(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterUpdate, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing AfterUpdate must fail the request: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestUpdateRunsBeforeAndAfterOnFirstTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate) + + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[OnTxBegin]) != 2 { + t.Fatalf("expected OnTxBegin twice, got %d", len(tr.tx[OnTxBegin])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeUpdate, AfterUpdate) +} diff --git a/pkg/restheadspec/read_tx_test.go b/pkg/restheadspec/read_tx_test.go index 22ed88b..ac05387 100644 --- a/pkg/restheadspec/read_tx_test.go +++ b/pkg/restheadspec/read_tx_test.go @@ -11,12 +11,14 @@ import ( "github.com/uptrace/bun" "github.com/uptrace/bun/dialect/pgdialect" + "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" ) func TestAfterReadRunsInSecondTransaction(t *testing.T) { + resetTotalCache(t) sqlDB, mock, err := sqlmock.New() if err != nil { t.Fatal(err) @@ -60,3 +62,100 @@ func TestAfterReadRunsInSecondTransaction(t *testing.T) { t.Fatal("BeforeRead must run on the first tx and AfterRead on the second") } } + +// resetTotalCache empties the process-wide query-total cache. Its key ignores the +// record id, so a cached total would skip the count query and desync the mock. +func resetTotalCache(t *testing.T) { + t.Helper() + _ = cache.GetDefaultCache().Clear(context.Background()) + t.Cleanup(func() { _ = cache.GetDefaultCache().Clear(context.Background()) }) +} + +func newBunHarness(t *testing.T) (*Handler, sqlmock.Sqlmock) { + t.Helper() + sqlDB, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + return NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry()), mock +} + +func itemCtx(t *testing.T) context.Context { + t.Helper() + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + return WithModel(ctx, delItem{}) +} + +func TestAfterReadErrorFailsRequestOnSecondTransaction(t *testing.T) { + resetTotalCache(t) + h, mock := newBunHarness(t) + h.Hooks().Register(AfterRead, func(*HookContext) error { return http.ErrAbortHandler }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, "7", ExtendedRequestOptions{}) + + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestBeforeReadErrorRollsBackWithoutQueries(t *testing.T) { + h, mock := newBunHarness(t) + h.Hooks().Register(BeforeRead, func(*HookContext) error { return http.ErrAbortHandler }) + + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, "7", ExtendedRequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeRead must not return data: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestAfterUpdateErrorFailsRequest(t *testing.T) { + h, mock := newBunHarness(t) + h.Hooks().Register(AfterUpdate, func(*HookContext) error { return http.ErrAbortHandler }) + + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil)) + h.handleUpdate(itemCtx(t), w, "7", nil, map[string]interface{}{"name": "b"}, ExtendedRequestOptions{}) + + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/websocketspec/tx_test.go b/pkg/websocketspec/tx_test.go index 12e410d..979ec7e 100644 --- a/pkg/websocketspec/tx_test.go +++ b/pkg/websocketspec/tx_test.go @@ -8,6 +8,8 @@ import ( "time" "github.com/DATA-DOG/go-sqlmock" + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/pgdialect" "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" @@ -171,3 +173,79 @@ func TestOnTxBeginErrorRollsBackWithoutDetail(t *testing.T) { t.Fatalf("unexpected response %+v", resp) } } + +func TestCreateRunsHooksOnTwoTransactions(t *testing.T) { + // The bun adapter builds model-based inserts; the pgsql adapter does not. + h, mock, conn, hookCtx := newTxHarness(t) + sqlDB, bunMock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + h.db = database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())) + conn = NewConnection("c2", nil, h) + mock = bunMock + hookCtx.ID = "" + hookCtx.Data = map[string]interface{}{"id": 7, "name": "a"} + var begins []common.Database + var beforeTx, afterTx common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + h.Hooks().Register(BeforeCreate, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + h.Hooks().Register(AfterCreate, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + mock.ExpectBegin() + mock.ExpectExec(`INSERT`).WillReturnResult(sqlmock.NewResult(7, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + h.handleCreate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(begins) != 2 || begins[0] == begins[1] || beforeTx != begins[0] || afterTx != begins[1] { + t.Fatalf("BeforeCreate must run on tx 1 and AfterCreate on tx 2, begins=%d", len(begins)) + } +} + +func TestBeforeHookErrorRollsBackWithoutWriting(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + hookCtx.Data = map[string]interface{}{"name": "b"} + h.Hooks().Register(BeforeUpdate, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + h.handleUpdate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if resp := response(t, conn); resp.Success || resp.Error == nil || resp.Error.Code != "hook_error" { + t.Fatalf("unexpected response %+v", resp) + } +} + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if response(t, conn).Success { + t.Fatal("a failing AfterDelete must fail the request") + } +} From 5933637a8811e92cfe699eaf73066dadada8c5f1 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:35:08 +0200 Subject: [PATCH 21/23] docs(readme): update slogan placement in README --- README.md | 5 +- audit/mcp_plan.md | 145 ++++++++++++++++++++++++++++++++++ audit/pkg/resolvemcp.audit.md | 121 ++++++++++++++++++++++++++++ 3 files changed, 270 insertions(+), 1 deletion(-) create mode 100644 audit/mcp_plan.md create mode 100644 audit/pkg/resolvemcp.audit.md diff --git a/README.md b/README.md index 89f7139..788f301 100644 --- a/README.md +++ b/README.md @@ -13,7 +13,7 @@ ResolveSpec is a flexible and powerful REST API specification and implementation All share the same core architecture and provide dynamic data querying, relationship preloading, and complex filtering. -![1.00](./generated_slogan.webp) + ## Table of Contents @@ -859,3 +859,6 @@ This project is licensed under the MIT License - see the [LICENSE](LICENSE) file * Slogan generated using DALL-E * AI used for documentation checking and correction * Community feedback and contributions that made v2.0 and v2.1 possible + + +![1.00](./generated_slogan.webp) \ No newline at end of file diff --git a/audit/mcp_plan.md b/audit/mcp_plan.md new file mode 100644 index 0000000..61d6676 --- /dev/null +++ b/audit/mcp_plan.md @@ -0,0 +1,145 @@ +# resolvemcp rewrite plan + +Source: `audit/pkg/resolvemcp.audit.md`. Status: plan only, no code changed. + +## Goal + +Replace 4 tools + 1 resource per model with a fixed set of meta tools. +Endpoint guarded by OAuth / session token / API key; tools run as the authenticated caller. +Same rules as resolvespec CRUD, plus guardrails. + +## Decisions + +| Topic | Decision | +|---|---| +| Tools | Fixed meta tools; per-model tools/resources removed (breaking) | +| Functions | Explicit registry `Handler.RegisterFunction`; two kinds: Go callback (`func(ctx, tx, args)` + JSON-schema params) and SQL procedure by name (declared params); both behind `call_function`, run in tx with hooks | +| Create | `insert_into_table` included | +| Writes | update/delete by id **or** filters | +| Guardrails | require id or filters, max rows, `dry_run`, confirm token | +| Token scope | filter writes only; id writes = single row, no token | +| Confirm token store | in-memory, TTL, bound to user/table/filter hash; lost on restart, single instance | +| Read limits | server caps: limit, offset, batch, preload depth, timeout | +| Model exposure | all registered models visible; rules only restrict operations | +| Visibility | list tools show only what the caller may do (rules) | +| Identity | the authenticated caller's `UserContext`; **no fixed MCP user, no `SetUsername`, no service session, no background refresh** | +| Guard | endpoint always requires one of: OAuth bearer, session token, API key; no guest/optional mode | +| API key login | new `DatabaseAuthenticator.LoginWithAPIKey(ctx, rawKey)` + procedure `resolvespec_login_api_key`; validates key via keystore, creates session, returns `LoginResponse` | +| Session SQL | procedure mode + direct-SQL fallback (`ShouldUseProcedure`), same as `Login` | +| OAuth routes | `oauth2.go`/`oauth2_server.go` kept, part of the guard | +| Annotations | opt-in `Config.EnableAnnotations`, via `BeforeHandle` | + +## Open + +- None. + +## Tools + +| Tool | Purpose | +|---|---| +| `list_tables` | visible `schema.entity` + allowed ops | +| `describe_table` | columns, PK, relations, writable columns, rules, limits | +| `select_table` | filters, sort, columns, preloads, cursor; capped | +| `insert_into_table` | one or batch (capped); column allowlist | +| `update_table` | validated keys; id or filters; guardrails | +| `delete_from_table` | id or filters; guardrails | +| `list_functions` | registered functions + parameter schemas | +| `call_function` | validated args; tx + hooks + rules | + +## Config additions + +| Field | Purpose | +|---|---| +| `DefaultLimit`, `MaxLimit`, `MaxOffset` | read paging caps | +| `MaxBatch` | insert batch cap | +| `MaxPreloadDepth` | preload cap | +| `MaxWriteRows` | filter-write row cap | +| `QueryTimeout` | per-call context timeout | +| `ConfirmTTL` | confirm token lifetime | +| `EnableAnnotations` | opt-in annotate tool | + +## Guardrail rules + +| Rule | Behaviour | +|---|---| +| Target required | update/delete with neither id nor filters rejected | +| Max rows | count matches inside tx; abort above `MaxWriteRows` | +| `dry_run` | returns match count + preview, no write | +| Confirm token | filter write: first call returns token + preview; second call with token executes; bound to user, table, filter hash; expires at `ConfirmTTL` | +| Id write | single row, no token | + +## Work items + +### 1. API key login (`pkg/security`) +- Existing: keystore has `ValidateKey` and `KeyStoreAuthenticator`; `Login` needs a password; no key-to-session path. +- Add `resolvespec_login_api_key` to `SQLNames` (default + override) and a SQL script beside the existing procedures. Contract: `p_success, p_error, p_data`, input raw key; hashes, validates active/non-expired key, creates session for the key's user. +- Add `DatabaseAuthenticator.LoginWithAPIKey(ctx, rawKey)`; procedure first, direct-SQL fallback via `ShouldUseProcedure`. +- Hashed lookup; same generic error for unknown, expired or inactive key; no key material in logs. +- Expose through the chain/composite authenticators so the middleware can accept it. + +### 2. Endpoint guard +- Wire `security.NewAuthMiddleware` with a chain of OAuth bearer, session token (header/cookie) and API key. +- `SetupMux*`/`SetupBunRouter*` helpers require the guard; unauthenticated serving only when explicitly constructed without it, logged loudly. Remove `OptionalAuth*` from the MCP path. +- Caller `UserContext` flows to every tool call context; rules, RLS and `OnTxBegin` apply to that user. + +### 3. Security fixes (audit #1-6) +- Put model rules in request context (`withRequestData`) and/or `AddRegistry` on construction. +- Add `BeforeCreate` -> `CheckModelCreateAllowed` (new in `pkg/security`). +- Call `BeforeHandle` first in `executeUpdate`. +- Validate create/update keys against `ColumnValidator`; reject unknown; resolve column names from model, not json tags. +- Apply row security to update/delete pre-read; fail if row not visible. +- Annotate tool: opt-in + `BeforeHandle`. + +### 4. Limits (audit #7, #13) +- Apply default/max limit, max offset, batch cap, preload depth cap, timeout. +- Validate preload names against model relations. +- Count only when requested. + +### 5. Error and panic surface (audit #9, #17) +- Stable error codes + short message to client. +- Details and stack logged server-side. +- Recover hook panics. + +### 6. Update/create semantics (audit #10-12) +- `SET` from validated incoming keys only; explicit null supported. +- Lock row (`FOR UPDATE`) on update. +- Refetch and `After*` hooks inside the same tx, or report committed write with warning if not possible (see `audit/single_tran.md`). + +### 7. Smaller fixes (audit #8, #14, #15) +- SSE pool: require `BaseURL` or cap/evict; allowlist Host. +- Uniform not-found vs hook error text. +- Add mutex to `HookRegistry`. + +### 8. Meta tools +- New file for meta tools; reuse parse helpers and `buildModelInfo` for `describe_table`. +- Remove per-model register functions and resources. +- `RegisterModel` only registers to registry. +- Function registry (Go callback kind + SQL procedure kind) + validation of args against declared schema. + +### 9. Tests +- Update `tx_test.go` (calls `executeRead/Create/Update/Delete`) and `tools_test.go`. +- New: rule enforcement, unknown keys, limits, guardrails (cap, dry_run, token expiry/binding), guard rejects unauthenticated, API key login (valid, expired, inactive, unknown), visibility filtering, `-race`. +- Check for existing test data first; ask before generating any. + +### 10. Docs +- Rewrite `pkg/resolvemcp/README.md` cheatsheet style. +- Document `resolvespec_login_api_key` in `pkg/security` docs. +- Update root README references. +- Update audit file when findings are closed. + +## Order + +1. `LoginWithAPIKey` + procedure in `pkg/security` (1) +2. Guard + security fixes (2-3) +3. Limits, errors, update/create semantics (4-6) +4. Meta tools + function registry (8) +5. Smaller fixes (7) +6. Tests (9), docs (10) + +## Breaking changes + +- Per-model tools and resources gone. +- MCP endpoint requires authentication. +- Annotate tool off by default. +- Update/create reject unknown keys. +- Reads capped by default. diff --git a/audit/pkg/resolvemcp.audit.md b/audit/pkg/resolvemcp.audit.md new file mode 100644 index 0000000..6f1d8df --- /dev/null +++ b/audit/pkg/resolvemcp.audit.md @@ -0,0 +1,121 @@ +# Audit: `pkg/resolvemcp` + +| | | +|---|---| +| **Package** | `github.com/bitechdev/ResolveSpec/pkg/resolvemcp` | +| **Files** | `handler.go` (901), `tools.go` (720), `cursor.go`, `oauth2.go`, `oauth2_server.go`, `annotation.go`, `hooks.go`, `security_hooks.go`, `context.go`, `resolvemcp.go` | +| **Tests** | `tools_test.go` (34), `tx_test.go` (207); `go test` passes. No hostile-input tests, no `-race` | +| **Audit date** | 2026-09-30 | +| **Axes** | thread locking/waiting, slowness, security, panic handling & logging, agent usability | +| **Threat model** | hostile or confused MCP client (LLM agent, possibly prompt-injected); tool arguments are attacker-controlled | +| **Depth** | targeted (request path, security wiring; verified against source) | + +## Summary + +Every model registers 4 tools + 1 resource (`read_/create_/update_/delete__`), each with an +inlined column list, relation list and schema doc. Tool list grows 4N; context cost is +paid on every session whether or not the table is used. Replace with fixed meta tools +(see Rewrite). + +Security wiring fails open in several places: model rules never reach the hooks, +`create` has no rule check, `update` skips `BeforeHandle`, update/delete skip row-level +security, and create/update write client-chosen column names. Reads have no size cap. +`resolvespec_annotate` is an unauthenticated write channel into agent-visible text. + +## Findings + +| # | Severity | Axis | Finding | +|---|---|---|---| +| 1 | **High** | security | Model rules set via `RegisterModelWithRules` never reach `security.Check*`: handler uses a private registry that is not `modelregistry.AddRegistry`'d; hooks look up the global list | +| 2 | **High** | security | `create` has no rule check: `CheckModelAuthAllowed` only tests `CanPublicCreate`/auth; no `BeforeCreate` hook registered, `CanCreate` never read | +| 3 | **High** | security | `executeUpdate` never fires `BeforeHandle` (create/read/delete do); auth + public-rule check skipped, only `BeforeUpdate` (`CanUpdate`) runs | +| 4 | **High** | security | Create/update data keys are not validated against model columns (`q.Value(key,…)`, `SetMap(existingMap)`): mass assignment of any column, arbitrary identifiers | +| 5 | **High** | security | Row-level security (`ApplyRowSecurity`) is wired to `BeforeRead` only; update/delete by id bypass app-level RLS (DB-level RLS via `OnTxBegin` still applies) | +| 6 | **High** | security | `resolvespec_annotate` has no auth/rule check, writes through `h.db` (outside tx, no `OnTxBegin`), any `tool_name` key; annotations are agent-facing text, so it is a prompt-injection store | +| 7 | **High** | slowness | No default/max `limit`, no max `offset`, `COUNT(*)` on every read, `[]` batch create unbounded, no statement timeout | +| 8 | **Medium** | security | `dynamicSSEHandler.pool` keyed by `Host` + `X-Forwarded-Proto` (attacker-controlled): unbounded map growth and poisoned `message` endpoint URL sent to the client | +| 9 | **Medium** | security / logging | Raw `err.Error()` (DB errors, hook errors, panic value `"internal error: %s"`) returned as tool text; `logger.Error` of the same forwards to Sentry (X8) | +| 10 | **Medium** | correctness | Update reads row, merges **json-tag keys** into `SetMap` as column names, writes every column back; breaks when json tag ≠ db column, clobbers concurrent edits (no `FOR UPDATE`) | +| 11 | **Medium** | correctness | Update ignores `nil` and `""` values: a column cannot be set to NULL or empty | +| 12 | **Medium** | correctness | Create/update commit tx 1, then run tx 2 (refetch + `AfterCreate`). Tx 2 failure returns an error for a committed write; an agent retry duplicates the insert | +| 13 | **Medium** | security | Preload relation names are passed straight to `PreloadRelation` without checking the model's relations; no depth/breadth cap | +| 14 | **Low** | security | Update/delete distinguish `record not found` from hook errors, so ids can be enumerated by error text | +| 15 | **Medium** | locking | `HookRegistry.hooks` map unsynchronized; `Register`/`Clear*` race with `Execute` (same as funcspec #9) | +| 16 | **Low** | security | Filter columns are validated by `ColumnValidator` for reads only; sort/column values are interpolated unquoted after validation (relies on validator being exact); `CustomOperators`/`ComputedColumns` unreachable from tools today, keep it that way | +| 17 | **Low** | panic | `recoverPanic` returns the panic value to the client and loses the stack; hook panics in `Execute` are not recovered before the handler-level recover | +| 18 | **Low** | agent usability | Tool names embed schema+entity (`read_public_users`); no discovery tool, so clients cannot list tables without loading every tool schema | +| 19 | **Info** | testing | No tests for auth/rule enforcement, hostile filters, key validation, limits, or `-race` | + +## Details + +### 1. Rules invisible to hooks (High) +`NewHandlerWithGORM/Bun/DB` call `modelregistry.NewModelRegistry()`. `security` resolves rules +via `GetModelRulesFromContext` then `modelregistry.GetModelRulesByName`, which walks the +**global** list (`registries`). The handler registry is never added, so +`ErrModelNotFound` → `CheckModelUpdate/DeleteAllowed` return `nil` (allow) and +`CheckModelAuthAllowed` falls back to "auth required, public flags ignored". +`CanUpdate=false`, `CanDelete=false` are not enforced. Fix: put rules into the +request context in `withRequestData` (`security.ModelRulesKey`) and/or `AddRegistry` on +construction. + +### 2-3. Create/update gating (High) +`CheckModelAuthAllowed(op)` handles public flags only. Add `BeforeCreate` → +`CheckModelCreateAllowed` (new, mirrors update/delete), and call `BeforeHandle` at the +top of `executeUpdate`. + +### 4. Column allowlist (High) +Validate every key in create/update `data` against `common.NewColumnValidator(model)`; +reject unknown keys with an error (do not silently drop on writes). Also consider a +per-model writable-column set (excluding PK, `CanPublic*`-guarded columns) for agents. + +### 5. RLS on writes (High) +Run `LoadSecurityRules` + a row predicate on the update/delete pre-read query; fail the +write when the row is not visible to the user. + +### 6. Annotation tool (High) +Remove from default registration or gate behind `BeforeHandle` + explicit rule. Values +returned to the agent must be treated as data, not instructions. + +### 7. Limits (High) +Server config: `DefaultLimit` (e.g. 50), `MaxLimit`, `MaxOffset`, `MaxBatch`, `MaxPreloadDepth`, +per-call `context.WithTimeout`. Skip `COUNT(*)` unless requested (`with_count`). + +### 8. SSE pool (Medium) +Require `Config.BaseURL` for SSE, or cap/evict `pool`, and validate `Host` against an +allowlist. + +### 9/17. Error surface (Medium/Low) +Map errors to stable codes + short message; log details server-side with stack +(`logger.HandlePanic`). + +### 10-12. Update/create semantics (Medium) +Build the `SET` only from validated incoming keys (column names resolved from model, +not json tags); use `NULL` for explicit null; single tx including refetch and `After*` +hooks (see `audit/single_tran.md`), or return success + warning when tx 2 fails. + +## Rewrite (agreed design) + +Replace per-model tools with fixed meta tools. Decisions recorded 2026-09-30: + +| Decision | Choice | +|---|---| +| Functions source | Explicit registry: `Handler.RegisterFunction(name, meta, fn)`; only registered functions visible/callable | +| Old tools/resources | Removed (breaking) | +| Discovery | `list_tables`, `describe_table`, `list_functions` | +| Create | `insert_into_table` added | +| Write scope | update/delete by id **or** filters; max-rows cap, `dry_run` and confirm token apply to filter writes; id writes are single-row, no token | +| Guardrails | require id/filter, max rows affected, `dry_run`, confirm token | +| Read limits / ACL | server caps (limit, offset, preload depth); list tools filtered per caller rules | +| Identity | Authenticated caller's `UserContext`; no fixed MCP user. Endpoint guarded by OAuth / session token / API key (new `resolvespec_login_api_key`); no guest mode | +| Annotations | `resolvespec_annotate` becomes opt-in (`Config.EnableAnnotations`) and goes through `BeforeHandle` | + +| Tool | Purpose | +|---|---| +| `list_tables` | registered `schema.entity` visible to caller, with allowed ops | +| `describe_table` | columns, PK, relations, writable columns, rules, limits for one table | +| `select_table` | filters/sort/columns/preloads/cursor; capped | +| `insert_into_table` | one or batch (capped); column allowlist | +| `update_table` | validated keys; guardrails | +| `delete_from_table` | guardrails | +| `list_functions` | registered functions + parameter schemas | +| `call_function` | validated args, tx + hooks | From a65ca5f5ce877514c5faeb6855a58a294da610bf Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:40:32 +0200 Subject: [PATCH 22/23] fix: fire resolvespec AfterRead/AfterCreate/AfterDelete, key restheadspec total cache by record id --- audit/single_tran.md | 5 +- pkg/common/TRANSACTIONS.md | 2 + pkg/common/tx_guard_test.go | 45 ++++++ pkg/resolvespec/handler.go | 98 +++++++++++- pkg/resolvespec/ops_tx_test.go | 237 +++++++++++++++++++++++++++++ pkg/restheadspec/cache_helpers.go | 12 +- pkg/restheadspec/cache_key_test.go | 61 ++++++++ pkg/restheadspec/handler.go | 1 + 8 files changed, 448 insertions(+), 13 deletions(-) create mode 100644 pkg/restheadspec/cache_key_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 380f30e..13bb6a6 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -97,8 +97,8 @@ - DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. - DONE real-Postgres check (resolvespec, testserver via compose): create `tx=1 pooled=0`, read `tx=1 pooled=0`, update `tx=2 pooled=0` (was `pooled=1`), single delete `tx=1 pooled=0`, batch create/delete `tx=1 pooled=0`. Compose now uses host networking (bridge fails here): testserver on 8123, Postgres on 8124 (was 8080/5434); integration test DSNs updated. Smoke script covers read and update. websocketspec/mqttspec/resolvemcp/restheadspec/funcspec not measured on real Postgres. - DONE regression tests: per-spec read/create/update/delete hook-on-tx, failure-rollback (Before*/After*/`OnTxBegin`) and second-tx tests in all six specs (`ops_tx_test.go`, `tx_test.go`, `read_tx_test.go`); stamping tests for resolvespec, resolvemcp, funcspec; pgsql adapter preload tests (same connection, error returned); source guard `pkg/common/tx_guard_test.go` (no direct `RunInTransaction`/`BeginTx`, no `Tx = h.db` beyond the allowlisted BeforeHandle placeholders, no pool statements in spec handlers). -- NOTE: resolvespec never fires `AfterRead`/`AfterCreate` (hook types exist, no call site); not tested, pre-existing. -- NOTE: restheadspec total-count cache is process-wide and ignores the record id; read tests call `resetTotalCache`. +- FIXED: resolvespec now fires `AfterRead` (in the read tx; `Result` = scanned slice for single and list) `AfterCreate` (in the create tx, per record, all four create paths) and `AfterDelete` (in the delete tx, once per request; a failure rolls the delete back). Before, neither fired, so `AfterRead` column-level security masking (registered by `RegisterSecurityHooks`) was silently skipped on resolvespec reads. Failing `AfterRead`/`AfterCreate` fails the request and rolls back. Tests: `pkg/resolvespec/ops_tx_test.go` (incl. end-to-end column hiding). +- FIXED: restheadspec total-count cache key now includes the record id; a read by id (total 1) no longer poisons the list total for the 2-minute TTL. Tests: `pkg/restheadspec/cache_key_test.go`. resolvespec is unaffected (its count runs before the id filter, so the total is the list total by design). The cache is still process-wide, so read tests call `resetTotalCache`. - NOTE: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (also on `cd96404`, before this work); use `-count=1`. ## Tests @@ -118,3 +118,4 @@ - `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. - 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. diff --git a/pkg/common/TRANSACTIONS.md b/pkg/common/TRANSACTIONS.md index df2f8d2..5b72d59 100644 --- a/pkg/common/TRANSACTIONS.md +++ b/pkg/common/TRANSACTIONS.md @@ -21,6 +21,8 @@ Every DB statement and every DB-touching hook of one request runs on one transac \* resolvespec, websocketspec, mqttspec, resolvemcp. restheadspec runs `AfterRead` in tx 2. † restheadspec, websocketspec, mqttspec run `AfterUpdate` in tx 2. resolvespec, resolvemcp run it in tx 1. +resolvespec has no tx 2 for create: `AfterCreate` runs in tx 1, once per record (per item in a batch), after the insert and re-fetch. Its `AfterRead` gets the scanned slice (single and list reads) and may mask in place; a failing `AfterRead` fails the read (fail closed). + Tx 2 exists so the re-fetch sees trigger changes from the committed write. ## Per spec diff --git a/pkg/common/tx_guard_test.go b/pkg/common/tx_guard_test.go index 2c943f5..40a3820 100644 --- a/pkg/common/tx_guard_test.go +++ b/pkg/common/tx_guard_test.go @@ -97,3 +97,48 @@ func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) { } } } + +// Hook types that are defined but deliberately or knowingly never executed. +// 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 +// AfterRead skipped column-level security masking until it was wired). +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*=`) + +func TestEveryDefinedHookHasACallSite(t *testing.T) { + for _, spec := range []string{"resolvespec", "restheadspec", "websocketspec", "resolvemcp", "funcspec"} { + raw, err := os.ReadFile(filepath.Join("..", spec, "hooks.go")) + if err != nil { + t.Fatal(err) + } + var src strings.Builder + files, _ := filepath.Glob(filepath.Join("..", spec, "*.go")) + for _, f := range files { + if strings.HasSuffix(f, "_test.go") || strings.HasSuffix(f, "hooks.go") || strings.HasSuffix(f, "hooks_example.go") { + continue + } + b, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + src.Write(b) + } + for _, m := range hookConstRE.FindAllStringSubmatch(string(raw), -1) { + name := m[1] + if name == "BeforeOp" || name == "OnTxBegin" { // fired by the registry / runInTx + continue + } + if _, ok := unwiredHooks[spec+"/"+name]; ok { + continue + } + call := regexp.MustCompile(`Execute(BeforeOp)?\(` + name + `\b`) + if !call.MatchString(src.String()) { + t.Errorf("%s: hook %s is defined but never executed", spec, name) + } + } + } +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index 8827a15..da30452 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -666,6 +666,17 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st result = reflect.ValueOf(modelPtr).Elem().Interface() } + // AfterRead runs inside the read transaction (e.g. column-level security + // masking). Result is the scanned slice for single and multi-record reads + // alike; hooks mutate the records in place, which `result` shares. + hookCtx.Result = modelPtr + hookCtx.Error = nil + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + logger.Error("AfterRead hook failed: %v", err) + statusCode, errCode, errMsg = http.StatusInternalServerError, "hook_error", "Hook execution failed" + return err + } + logger.Info("Successfully retrieved records") return nil }) @@ -762,7 +773,17 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat var procErr error nestedResult, procErr = h.nestedProcessor.ProcessNestedCUD(ctx, "insert", v, model, make(map[string]interface{}), tableName) - return procErr + if procErr != nil { + return procErr + } + res, err := h.afterCreate(hookCtx, nestedResult.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + nestedResult.Data = m + } + return nil }) if err != nil { logger.Error("Error in nested create: %v", err) @@ -814,6 +835,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } logger.Info("Successfully created record, rows affected: %d", result.RowsAffected()) + res, err := h.afterCreate(hookCtx, responseData) + if err != nil { + return err + } + responseData = res return nil } @@ -830,6 +856,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, insertedID, fetchErr) } + res, err := h.afterCreate(hookCtx, responseData) + if err != nil { + return err + } + responseData = res return nil }) if err != nil { @@ -889,6 +920,13 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if err != nil { return fmt.Errorf("failed to process item: %w", err) } + res, err := h.afterCreate(hookCtx, result.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + result.Data = m + } results = append(results, result.Data) } return nil @@ -941,7 +979,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if _, err := txQuery.Exec(ctx); err != nil { return err } - responseItems = append(responseItems, item) + res, err := h.afterCreate(hookCtx, item) + if err != nil { + return err + } + responseItems = append(responseItems, res) continue } var returnedID interface{} @@ -949,14 +991,20 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } fetchedRecord := reflect.New(modelElemType).Interface() + var created interface{} if fetchErr := tx.NewSelect().Model(fetchedRecord). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), returnedID). ScanModel(ctx); fetchErr == nil { - responseItems = append(responseItems, mergeWithInput(fetchedRecord, item)) + created = mergeWithInput(fetchedRecord, item) } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, returnedID, fetchErr) - responseItems = append(responseItems, item) + created = item } + res, err := h.afterCreate(hookCtx, created) + if err != nil { + return err + } + responseItems = append(responseItems, res) } return nil }) @@ -1022,6 +1070,13 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if err != nil { return fmt.Errorf("failed to process item: %w", err) } + res, err := h.afterCreate(hookCtx, result.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + result.Data = m + } results = append(results, result.Data) } } @@ -1080,7 +1135,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if _, err := txQuery.Exec(ctx); err != nil { return err } - responseItems = append(responseItems, itemMap) + res, err := h.afterCreate(hookCtx, itemMap) + if err != nil { + return err + } + responseItems = append(responseItems, res) continue } var returnedID interface{} @@ -1088,14 +1147,20 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } fetchedRecord := reflect.New(modelElemType).Interface() + var created interface{} if fetchErr := tx.NewSelect().Model(fetchedRecord). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), returnedID). ScanModel(ctx); fetchErr == nil { - responseItems = append(responseItems, mergeWithInput(fetchedRecord, itemMap)) + created = mergeWithInput(fetchedRecord, itemMap) } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, returnedID, fetchErr) - responseItems = append(responseItems, itemMap) + created = itemMap } + res, err := h.afterCreate(hookCtx, created) + if err != nil { + return err + } + responseItems = append(responseItems, res) } return nil }) @@ -1677,6 +1742,15 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id if failure != nil { return failure } + // AfterDelete runs inside the delete transaction: a failing hook rolls the + // delete back. Result is the deleted record, or the batch summary. + hookCtx.Result = payload + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + logger.Error("AfterDelete hook failed: %v", err) + failure = &deleteFailure{http.StatusInternalServerError, "hook_error", "Hook execution failed", err} + return failure + } + payload = hookCtx.Result return nil }) if failure != nil { @@ -2598,3 +2672,13 @@ func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(t return h.hooks.Execute(OnTxBegin, hookCtx) }, body) } + +// afterCreate runs the AfterCreate hooks inside the create transaction with the +// created record as Result and returns the (possibly replaced) result. +func (h *Handler) afterCreate(hookCtx *HookContext, result interface{}) (interface{}, error) { + hookCtx.Result = result + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return nil, fmt.Errorf("AfterCreate hook failed: %w", err) + } + return hookCtx.Result, nil +} diff --git a/pkg/resolvespec/ops_tx_test.go b/pkg/resolvespec/ops_tx_test.go index aa75cba..b5da745 100644 --- a/pkg/resolvespec/ops_tx_test.go +++ b/pkg/resolvespec/ops_tx_test.go @@ -3,16 +3,28 @@ package resolvespec import ( "context" "errors" + "fmt" "net/http" "net/http/httptest" + "strings" "testing" "time" "github.com/DATA-DOG/go-sqlmock" + "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" ) +// resetTotalCache empties the process-wide query-total cache so a total cached by +// another test cannot skip the count query and desync the mock. +func resetTotalCache(t *testing.T) { + t.Helper() + _ = cache.GetDefaultCache().Clear(context.Background()) + t.Cleanup(func() { _ = cache.GetDefaultCache().Clear(context.Background()) }) +} + func opCtx(t *testing.T) (context.Context, common.ResponseWriter, *httptest.ResponseRecorder) { t.Helper() rec := httptest.NewRecorder() @@ -51,6 +63,7 @@ func (tr *hookTrace) mustBeOn(t *testing.T, tx common.Database, types ...HookTyp } func TestReadRunsHooksOnOneTransaction(t *testing.T) { + resetTotalCache(t) h, mock, _ := newDeleteHarness(t) tr := traceHooks(h, OnTxBegin, BeforeRead) @@ -184,3 +197,227 @@ func TestUpdateRunsBeforeAndAfterOnFirstTransaction(t *testing.T) { } tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeUpdate, AfterUpdate) } + +func TestAfterReadRunsOnReadTransactionForSingleAndList(t *testing.T) { + for name, id := range map[string]string{"single": "7", "list": ""} { + t.Run(name, func(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterRead) + var resultType string + h.Hooks().Register(AfterRead, func(c *HookContext) error { + resultType = fmt.Sprintf("%T", c.Result) + return nil + }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, id, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[AfterRead]) != 1 { + t.Fatalf("AfterRead must fire once, got %d", len(tr.tx[AfterRead])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterRead) + if !strings.HasPrefix(resultType, "*[]") { + t.Fatalf("AfterRead Result must be the scanned slice, got %s", resultType) + } + }) + } +} + +func TestAfterReadErrorFailsReadAndRollsBack(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterRead, func(*HookContext) error { return errors.New("masking failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing AfterRead must not return data (fail closed): %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +type columnSecProvider struct{ security.SecurityProvider } + +func (columnSecProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) { + return []security.ColumnSecurity{{ + Schema: schema, Tablename: table, Path: []string{"name"}, UserID: userID, Accesstype: "hide", + }}, nil +} + +func (columnSecProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (security.RowSecurity, error) { + return security.RowSecurity{}, nil +} + +// Column-level security is applied by an AfterRead hook; before AfterRead was wired +// into resolvespec reads, the configured masking was silently skipped. +func TestColumnSecurityHidesColumnOnRead(t *testing.T) { + for name, id := range map[string]string{"single": "7", "list": ""} { + t.Run(name, func(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + list, err := security.NewSecurityList(columnSecProvider{}) + if err != nil { + t.Fatal(err) + } + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "secret")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx = context.WithValue(ctx, security.UserIDKey, 7) + h.handleRead(ctx, w, id, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if strings.Contains(rec.Body.String(), "secret") { + t.Fatalf("hidden column leaked: %s", rec.Body) + } + }) + } +} + +func TestAfterCreateRunsInsideCreateTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[AfterCreate]) != 1 { + t.Fatalf("AfterCreate must fire once, got %d", len(tr.tx[AfterCreate])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterCreate) +} + +func TestAfterCreateFiresPerItemInBatch(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterCreate) + + mock.ExpectBegin() + for i := 1; i <= 2; i++ { + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(i)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(i, "a")) + } + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}} + h.handleCreate(ctx, w, items, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[AfterCreate]) != 2 || len(tr.tx[OnTxBegin]) != 1 { + t.Fatalf("AfterCreate must fire per item in one transaction, got %d in %d tx", len(tr.tx[AfterCreate]), len(tr.tx[OnTxBegin])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterCreate) +} + +func TestAfterCreateErrorRollsBackCreate(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterCreate, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing AfterCreate must fail the request: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestAfterDeleteRunsInsideDeleteTransaction(t *testing.T) { + for name, tc := range map[string]struct { + id string + data interface{} + exec int + }{"single": {id: "7", exec: 1}, "batch": {data: []interface{}{"1", "2"}, exec: 2}} { + t.Run(name, func(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterDelete) + + mock.ExpectBegin() + if tc.id != "" { + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + } + for i := 0; i < tc.exec; i++ { + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + } + mock.ExpectCommit() + + if rec := runDelete(h, tc.id, tc.data); rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(tr.tx[AfterDelete]) != 1 { + t.Fatalf("AfterDelete must fire once per request, got %d", len(tr.tx[AfterDelete])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterDelete) + }) + } +} + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusInternalServerError { + t.Fatalf("a failing AfterDelete must fail and roll back the delete: status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/restheadspec/cache_helpers.go b/pkg/restheadspec/cache_helpers.go index 1e81187..c0950b9 100644 --- a/pkg/restheadspec/cache_helpers.go +++ b/pkg/restheadspec/cache_helpers.go @@ -22,6 +22,7 @@ type expandOptionKey struct { // queryCacheKey represents the components used to build a cache key for query total count type queryCacheKey struct { TableName string `json:"table_name"` + ID string `json:"id,omitempty"` Filters []common.FilterOption `json:"filters"` Sort []common.SortOption `json:"sort"` CustomSQLWhere string `json:"custom_sql_where,omitempty"` @@ -39,12 +40,15 @@ type cachedTotal struct { } // buildExtendedQueryCacheKey builds a cache key for extended query options (restheadspec) -// Includes expand, distinct, and cursor pagination options -func buildExtendedQueryCacheKey(tableName string, filters []common.FilterOption, sort []common.SortOption, +// Includes expand, distinct, and cursor pagination options. id is the record id of a +// single-record read: it constrains the counted query, so a count cached for one id +// (or for the whole list) must never be served for another. +func buildExtendedQueryCacheKey(tableName, id string, filters []common.FilterOption, sort []common.SortOption, customWhere, customOr string, customJoin []string, expandOpts []interface{}, distinct bool, cursorFwd, cursorBwd string) string { key := queryCacheKey{ TableName: tableName, + ID: id, Filters: filters, Sort: sort, CustomSQLWhere: customWhere, @@ -77,8 +81,8 @@ func buildExtendedQueryCacheKey(tableName string, filters []common.FilterOption, jsonData, err := json.Marshal(key) if err != nil { // Fallback to simple string concatenation if JSON fails - return hashString(fmt.Sprintf("%s_%v_%v_%s_%s_%v_%v_%v_%s_%s", - tableName, filters, sort, customWhere, customOr, customJoin, expandOpts, distinct, cursorFwd, cursorBwd)) + return hashString(fmt.Sprintf("%s_%s_%v_%v_%s_%s_%v_%v_%v_%s_%s", + tableName, id, filters, sort, customWhere, customOr, customJoin, expandOpts, distinct, cursorFwd, cursorBwd)) } return hashString(string(jsonData)) diff --git a/pkg/restheadspec/cache_key_test.go b/pkg/restheadspec/cache_key_test.go new file mode 100644 index 0000000..cd396ff --- /dev/null +++ b/pkg/restheadspec/cache_key_test.go @@ -0,0 +1,61 @@ +package restheadspec + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func TestQueryTotalCacheKeyIncludesRecordID(t *testing.T) { + list := buildExtendedQueryCacheKey("items", "", nil, nil, "", "", nil, nil, false, "", "") + one := buildExtendedQueryCacheKey("items", "7", nil, nil, "", "", nil, nil, false, "", "") + other := buildExtendedQueryCacheKey("items", "8", nil, nil, "", "", nil, nil, false, "", "") + if list == one || one == other { + t.Fatal("list, id 7 and id 8 must not share a cache key") + } + if one != buildExtendedQueryCacheKey("items", "7", nil, nil, "", "", nil, nil, false, "", "") { + t.Fatal("the key must be stable for the same query") + } +} + +// A read by id counts 1 row; before the id was part of the key, that total was +// cached under the list query's key and served as the list total for 2 minutes. +func TestReadByIDDoesNotPoisonListTotal(t *testing.T) { + resetTotalCache(t) + h, mock := newBunHarness(t) + cols := []string{"id", "name"} + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + // The list must run its own count query, not reuse the id read's total. + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(5)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a").AddRow(8, "b")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + + read := func(id string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, id, ExtendedRequestOptions{}) + return rec + } + if rec := read("7"); rec.Code != http.StatusOK { + t.Fatalf("read by id: status %d body %s", rec.Code, rec.Body) + } + if rec := read(""); rec.Code != http.StatusOK { + t.Fatalf("list: status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatalf("the list must run its own count query: %v", err) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 649901f..29a59ac 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -824,6 +824,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st cacheKeyHash := buildExtendedQueryCacheKey( tableName, + id, options.Filters, options.Sort, options.CustomSQLWhere, From b35399fdfa6412ae9982b92f120261dce169f70c Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:48:09 +0200 Subject: [PATCH 23/23] feat(security): exclude hidden/masked columns from create and update payloads --- audit/single_tran.md | 1 + pkg/mqttspec/security_hooks.go | 23 ++++ pkg/resolvemcp/handler.go | 3 + pkg/resolvemcp/security_hooks.go | 23 ++++ pkg/resolvemcp/tx_test.go | 5 + pkg/resolvespec/ops_tx_test.go | 62 +++++++++++ pkg/resolvespec/security_hooks.go | 17 +++ pkg/restheadspec/security_hooks.go | 17 +++ pkg/security/README.md | 6 ++ pkg/security/hooks.go | 21 +++- pkg/security/writesecurity.go | 159 ++++++++++++++++++++++++++++ pkg/security/writesecurity_test.go | 100 +++++++++++++++++ pkg/websocketspec/security_hooks.go | 23 ++++ 13 files changed, 456 insertions(+), 4 deletions(-) create mode 100644 pkg/security/writesecurity.go create mode 100644 pkg/security/writesecurity_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 13bb6a6..300b84a 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -119,3 +119,4 @@ - 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. - OPEN: websocketspec `BeforeDisconnect`/`AfterDisconnect` are defined but never executed (connection lifecycle, not DB). Allowlisted in `TestEveryDefinedHookHasACallSite`; wire them to remove the entry. +- 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`. diff --git a/pkg/mqttspec/security_hooks.go b/pkg/mqttspec/security_hooks.go index bcdec15..587a5c7 100644 --- a/pkg/mqttspec/security_hooks.go +++ b/pkg/mqttspec/security_hooks.go @@ -27,6 +27,21 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList return nil }) + // BeforeHandle: preload column rules for writes before the handler opens its + // transaction; the write hooks below only read the cache. + handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { + return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation) + }) + + // BeforeCreate/BeforeUpdate: drop columns hidden or masked for the user from the + // write payload, so they cannot be inserted or updated. + handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + // Hook 1: BeforeRead - Load security rules handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error { secCtx := newSecurityContext(hookCtx) @@ -122,6 +137,14 @@ func (s *securityContext) SetQuery(query interface{}) { s.ctx.Metadata["query"] = query } +func (s *securityContext) GetData() interface{} { + return s.ctx.Data +} + +func (s *securityContext) SetData(data interface{}) { + s.ctx.Data = data +} + func (s *securityContext) GetResult() interface{} { return s.ctx.Result } diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 5723466..4feb95a 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -596,6 +596,9 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, Data: updates, Tx: h.db, } + if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil { + return nil, err + } var updateResult interface{} err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { diff --git a/pkg/resolvemcp/security_hooks.go b/pkg/resolvemcp/security_hooks.go index f7d8bd4..a26b109 100644 --- a/pkg/resolvemcp/security_hooks.go +++ b/pkg/resolvemcp/security_hooks.go @@ -36,6 +36,21 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList return nil }) + // BeforeHandle: preload column rules for writes before the handler opens its + // transaction; the write hooks below only read the cache. + handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { + return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation) + }) + + // BeforeCreate/BeforeUpdate: drop columns hidden or masked for the user from the + // write payload, so they cannot be inserted or updated. + handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + // BeforeRead (1st): load RLS + CLS rules from the provider into SecurityList. handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error { return security.LoadSecurityRules(newSecurityContext(hookCtx), securityList) @@ -123,6 +138,14 @@ func (s *securityContext) SetQuery(query interface{}) { } } +func (s *securityContext) GetData() interface{} { + return s.ctx.Data +} + +func (s *securityContext) SetData(data interface{}) { + s.ctx.Data = data +} + func (s *securityContext) GetResult() interface{} { return s.ctx.Result } diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go index 3bff822..2279306 100644 --- a/pkg/resolvemcp/tx_test.go +++ b/pkg/resolvemcp/tx_test.go @@ -293,6 +293,10 @@ func TestOnTxBeginErrorRollsBackEveryOperation(t *testing.T) { type stubProvider struct{ security.SecurityProvider } +func (stubProvider) GetColumnSecurity(context.Context, int, string, string) ([]security.ColumnSecurity, error) { + return nil, nil +} + func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { h, mock, ctx := newTxHarness(t) list, err := security.NewSecurityList(stubProvider{}) @@ -304,6 +308,7 @@ func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { }) RegisterSecurityHooks(h, list) ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx = context.WithValue(ctx, security.UserIDKey, 7) // Update opens two transactions; each must be stamped before any other SQL. cols := []string{"id", "name"} diff --git a/pkg/resolvespec/ops_tx_test.go b/pkg/resolvespec/ops_tx_test.go index b5da745..456cc4f 100644 --- a/pkg/resolvespec/ops_tx_test.go +++ b/pkg/resolvespec/ops_tx_test.go @@ -421,3 +421,65 @@ func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { t.Fatal(err) } } + +// Columns hidden or masked for the user cannot be written. Handlers are called +// directly here, so the BeforeHandle preload is done by hand. +func TestColumnSecurityDropsHiddenColumnOnWrite(t *testing.T) { + setup := func(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context, common.ResponseWriter, *httptest.ResponseRecorder) { + h, mock, _ := newDeleteHarness(t) + list, err := security.NewSecurityList(columnSecProvider{}) + if err != nil { + t.Fatal(err) + } + RegisterSecurityHooks(h, list) + + ctx, w, rec := opCtx(t) + ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx = context.WithValue(ctx, security.UserIDKey, 7) + pre := newSecurityContext(&HookContext{Context: ctx, Schema: "public", Entity: "items", Model: &delItem{}}) + for _, op := range []string{"create", "update"} { + if err := security.PreloadSecurityRules(pre, list, op); err != nil { + t.Fatal(err) + } + } + return h, mock, ctx, w, rec + } + + t.Run("create", func(t *testing.T) { + h, mock, ctx, w, rec := setup(t) + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WithArgs(7).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "")) + mock.ExpectCommit() + + payload := map[string]interface{}{"id": 7, "name": "secret"} + h.handleCreate(ctx, w, payload, common.RequestOptions{}) + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if _, ok := payload["name"]; ok { + t.Fatalf("hidden column must be dropped from the insert payload: %v", payload) + } + }) + + t.Run("update", func(t *testing.T) { + h, mock, ctx, w, rec := setup(t) + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WithArgs(float64(7), "a", "7").WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectCommit() + + payload := map[string]interface{}{"name": "secret"} + h.handleUpdate(ctx, w, "7", nil, payload, common.RequestOptions{}) + if _, ok := payload["name"]; ok { + t.Fatalf("hidden column must be dropped from the update payload: %v (status %d %s)", payload, rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatalf("update must keep the stored value: %v", err) + } + }) +} diff --git a/pkg/resolvespec/security_hooks.go b/pkg/resolvespec/security_hooks.go index 5db28a5..ae2d0cb 100644 --- a/pkg/resolvespec/security_hooks.go +++ b/pkg/resolvespec/security_hooks.go @@ -34,6 +34,15 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation) }) + // BeforeCreate/BeforeUpdate: drop columns hidden or masked for the user from the + // write payload, so they cannot be inserted or updated. + handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + // Hook 1: BeforeRead - Load security rules handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error { secCtx := newSecurityContext(hookCtx) @@ -133,6 +142,14 @@ func (s *securityContext) SetQuery(query interface{}) { } } +func (s *securityContext) GetData() interface{} { + return s.ctx.Data +} + +func (s *securityContext) SetData(data interface{}) { + s.ctx.Data = data +} + func (s *securityContext) GetResult() interface{} { return s.ctx.Result } diff --git a/pkg/restheadspec/security_hooks.go b/pkg/restheadspec/security_hooks.go index e1c18f8..ea722f5 100644 --- a/pkg/restheadspec/security_hooks.go +++ b/pkg/restheadspec/security_hooks.go @@ -33,6 +33,15 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation) }) + // BeforeCreate/BeforeUpdate: drop columns hidden or masked for the user from the + // write payload, so they cannot be inserted or updated. + handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + // Hook 1: BeforeRead - Load security rules handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error { secCtx := newSecurityContext(hookCtx) @@ -120,6 +129,14 @@ func (s *securityContext) SetQuery(query interface{}) { s.ctx.Query = query } +func (s *securityContext) GetData() interface{} { + return s.ctx.Data +} + +func (s *securityContext) SetData(data interface{}) { + s.ctx.Data = data +} + func (s *securityContext) GetResult() interface{} { return s.ctx.Result } diff --git a/pkg/security/README.md b/pkg/security/README.md index 7f11ace..1a51538 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -226,6 +226,12 @@ type ColumnSecurityProvider interface { } ``` +Write side (`RegisterSecurityHooks`, all specs except funcspec): +- Columns with a `hide` or `mask` rule (single-element `Path`) are removed from create/update payloads in `BeforeCreate`/`BeforeUpdate`; the write is not rejected. +- Match is case-insensitive on the rule path vs payload key, model field/JSON name or `gorm` column. +- Rules are preloaded in `BeforeHandle` (outside the tx); the hook only reads the cache and fails closed if rules were not loaded. +- Not covered: nested child records, nested `Path` (JSON sub-values), funcspec. + #### 3. RowSecurityProvider Manages row-level security (WHERE clause filtering): diff --git a/pkg/security/hooks.go b/pkg/security/hooks.go index 23809d1..697d769 100644 --- a/pkg/security/hooks.go +++ b/pkg/security/hooks.go @@ -245,13 +245,26 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error // cache for read operations. Call it from a BeforeHandle hook, i.e. before the // handler opens its transaction, so the provider queries do not need a second // pooled connection while the transaction holds one. Later LoadSecurityRules -// calls in the same request are then cache hits. Non-read operations and -// models with security disabled are skipped. +// calls in the same request are then cache hits. Reads load column and row +// rules; create/update load the column rules that ApplyWriteColumnSecurity +// reads. Other operations and models with security disabled are skipped. func PreloadSecurityRules(secCtx SecurityContext, securityList *SecurityList, operation string) error { - if operation != "read" || IsModelSecurityDisabled(secCtx) { + if IsModelSecurityDisabled(secCtx) { return nil } - return loadSecurityRules(secCtx, securityList) + switch { + case operation == "read": + return loadSecurityRules(secCtx, securityList) + case isWriteOperation(operation): + userID, ok := secCtx.GetUserID() + if !ok { + return nil + } + if err := securityList.LoadColumnSecurity(secCtx.GetContext(), userID, secCtx.GetSchema(), secCtx.GetEntity(), false); err != nil { + logger.Warn("Failed to load column security: %v", err) + } + } + return nil } // ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext diff --git a/pkg/security/writesecurity.go b/pkg/security/writesecurity.go new file mode 100644 index 0000000..44f48a7 --- /dev/null +++ b/pkg/security/writesecurity.go @@ -0,0 +1,159 @@ +package security + +import ( + "fmt" + "reflect" + "strings" + + "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/reflection" +) + +// WriteDataContext is implemented by security contexts that expose the +// create/update payload of the operation in flight. +type WriteDataContext interface { + GetData() interface{} + SetData(interface{}) +} + +// isWriteOperation reports whether the operation writes columns. +func isWriteOperation(operation string) bool { + return operation == "create" || operation == "update" +} + +// cachedColumnRules returns the cached column rules for the user/table and +// whether they were loaded. It never calls the provider, so it is safe inside +// a transaction. +func (m *SecurityList) cachedColumnRules(userID int, schema, table string) ([]ColumnSecurity, bool) { + m.ColumnSecurityMutex.RLock() + defer m.ColumnSecurityMutex.RUnlock() + rules, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", schema, table, userID)] + return rules, ok && rules != nil +} + +// blockedWriteColumns returns the lower-cased top-level column names that a +// "hide" or "mask" rule removes from write payloads. +func blockedWriteColumns(rules []ColumnSecurity) map[string]struct{} { + blocked := make(map[string]struct{}) + for i := range rules { + r := &rules[i] + if !strings.EqualFold(r.Accesstype, "hide") && !strings.EqualFold(r.Accesstype, "mask") { + continue + } + if len(r.Path) != 1 { + continue // nested paths address JSON sub-values, not columns + } + blocked[strings.ToLower(r.Path[0])] = struct{}{} + } + return blocked +} + +// modelColumnAliases maps each lower-cased field/column/JSON name of the model +// to all of its lower-cased names, so a rule on "name" also blocks its column. +func modelColumnAliases(model interface{}) map[string][]string { + aliases := make(map[string][]string) + if model == nil { + return aliases + } + v := reflect.ValueOf(model) + for v.Kind() == reflect.Pointer || v.Kind() == reflect.Slice || v.Kind() == reflect.Array { + if v.Kind() == reflect.Pointer && v.IsNil() { + v = reflect.New(v.Type().Elem()) + } + if v.Kind() == reflect.Slice || v.Kind() == reflect.Array { + v = reflect.New(v.Type().Elem()).Elem() + continue + } + v = v.Elem() + } + if v.Kind() != reflect.Struct { + return aliases + } + for _, c := range reflection.GetModelColumnDetail(v) { + names := []string{strings.ToLower(c.Name), strings.ToLower(c.SQLName)} + for _, n := range names { + if n != "" { + aliases[n] = names + } + } + } + return aliases +} + +// stripBlocked removes blocked keys from one payload map in place. +func stripBlocked(m map[string]interface{}, blocked map[string]struct{}, aliases map[string][]string) []string { + var dropped []string + for key := range m { + lk := strings.ToLower(key) + hit := false + if _, ok := blocked[lk]; ok { + hit = true + } else { + for _, a := range aliases[lk] { + if _, ok := blocked[a]; ok { + hit = true + break + } + } + } + if hit { + delete(m, key) + dropped = append(dropped, key) + } + } + return dropped +} + +// stripPayload strips blocked keys from a map, []map or []interface{} payload. +func stripPayload(data interface{}, blocked map[string]struct{}, aliases map[string][]string) (dropped []string) { + switch d := data.(type) { + case map[string]interface{}: + dropped = stripBlocked(d, blocked, aliases) + case []map[string]interface{}: + for _, m := range d { + dropped = append(dropped, stripBlocked(m, blocked, aliases)...) + } + case []interface{}: + for _, e := range d { + dropped = append(dropped, stripPayload(e, blocked, aliases)...) + } + } + return dropped +} + +// ApplyWriteColumnSecurity removes columns the user may not see (column +// security "hide" or "mask") from the create/update payload in place, so a +// hidden or masked column can never be written. It only reads the rules cache +// (see PreloadSecurityRules) and never queries the provider, so it is safe +// inside the transaction. Models with security disabled are skipped. Without +// a loaded rule set for a known user it fails closed. +func ApplyWriteColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error { + userID, ok := secCtx.GetUserID() + if !ok || securityList == nil || IsModelSecurityDisabled(secCtx) { + return nil + } + dc, ok := secCtx.(WriteDataContext) + if !ok { + return fmt.Errorf("column security: write payload not accessible for %s.%s", secCtx.GetSchema(), secCtx.GetEntity()) + } + data := dc.GetData() + if data == nil { + return nil + } + + rules, loaded := securityList.cachedColumnRules(userID, secCtx.GetSchema(), secCtx.GetEntity()) + if !loaded { + return fmt.Errorf("column security rules not loaded for %s.%s", secCtx.GetSchema(), secCtx.GetEntity()) + } + blocked := blockedWriteColumns(rules) + if len(blocked) == 0 { + return nil + } + + dropped := stripPayload(data, blocked, modelColumnAliases(secCtx.GetModel())) + if len(dropped) > 0 { + logger.Warn("Column security: dropped write to hidden/masked columns %v on %s.%s (user %d)", + dropped, secCtx.GetSchema(), secCtx.GetEntity(), userID) + } + return nil +} diff --git a/pkg/security/writesecurity_test.go b/pkg/security/writesecurity_test.go new file mode 100644 index 0000000..b3f5511 --- /dev/null +++ b/pkg/security/writesecurity_test.go @@ -0,0 +1,100 @@ +package security + +import ( + "context" + "reflect" + "testing" +) + +type wsCtx struct { + data interface{} + model interface{} + user bool +} + +func (c *wsCtx) GetContext() context.Context { + if c.user { + return context.WithValue(context.Background(), UserIDKey, 7) + } + return context.Background() +} +func (c *wsCtx) GetUserID() (int, bool) { return 7, c.user } +func (c *wsCtx) GetUserRef() (any, bool) { return 7, c.user } +func (c *wsCtx) GetSchema() string { return "public" } +func (c *wsCtx) GetEntity() string { return "items" } +func (c *wsCtx) GetModel() interface{} { return c.model } +func (c *wsCtx) GetQuery() interface{} { return nil } +func (c *wsCtx) SetQuery(interface{}) {} +func (c *wsCtx) GetResult() interface{} { return nil } +func (c *wsCtx) SetResult(interface{}) {} +func (c *wsCtx) GetData() interface{} { return c.data } +func (c *wsCtx) SetData(d interface{}) { c.data = d } + +type wsModel struct { + ID int `json:"id" bun:"id,pk"` + Name string `json:"name" bun:"name"` + Email string `json:"email" gorm:"column:email_addr"` + Other string `json:"other" bun:"other"` +} + +func wsList(rules ...ColumnSecurity) *SecurityList { + l := &SecurityList{ColumnSecurity: map[string][]ColumnSecurity{"public.items@7": rules}} + if rules == nil { + l.ColumnSecurity["public.items@7"] = []ColumnSecurity{} + } + return l +} + +func TestApplyWriteColumnSecurityStripsHiddenAndMasked(t *testing.T) { + list := wsList( + ColumnSecurity{Path: []string{"Name"}, Accesstype: "hide"}, + ColumnSecurity{Path: []string{"email_addr"}, Accesstype: "mask"}, + ColumnSecurity{Path: []string{"other"}, Accesstype: "allow"}, + ColumnSecurity{Path: []string{"id", "sub"}, Accesstype: "hide"}, + ) + tests := map[string]struct{ in, want interface{} }{ + "map": { + map[string]interface{}{"id": 1, "name": "a", "email": "e", "other": "o"}, + map[string]interface{}{"id": 1, "other": "o"}, + }, + "slice": { + []interface{}{map[string]interface{}{"NAME": "a", "id": 1}, map[string]interface{}{"email_addr": "e"}}, + []interface{}{map[string]interface{}{"id": 1}, map[string]interface{}{}}, + }, + } + for name, tc := range tests { + t.Run(name, func(t *testing.T) { + c := &wsCtx{data: tc.in, model: &wsModel{}, user: true} + if err := ApplyWriteColumnSecurity(c, list); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(c.data, tc.want) { + t.Fatalf("got %v want %v", c.data, tc.want) + } + }) + } +} + +func TestApplyWriteColumnSecurityNoRulesKeepsPayload(t *testing.T) { + c := &wsCtx{data: map[string]interface{}{"name": "a"}, model: &wsModel{}, user: true} + if err := ApplyWriteColumnSecurity(c, wsList()); err != nil { + t.Fatal(err) + } + if _, ok := c.data.(map[string]interface{})["name"]; !ok { + t.Fatal("payload changed without rules") + } +} + +func TestApplyWriteColumnSecurityFailsClosedWhenRulesNotLoaded(t *testing.T) { + c := &wsCtx{data: map[string]interface{}{"name": "a"}, model: &wsModel{}, user: true} + if err := ApplyWriteColumnSecurity(c, &SecurityList{}); err == nil { + t.Fatal("expected an error when rules are not loaded") + } +} + +func TestApplyWriteColumnSecuritySkipsWithoutUser(t *testing.T) { + c := &wsCtx{data: map[string]interface{}{"name": "a"}, model: &wsModel{}} + if err := ApplyWriteColumnSecurity(c, &SecurityList{}); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/websocketspec/security_hooks.go b/pkg/websocketspec/security_hooks.go index 41d5c21..061aab2 100644 --- a/pkg/websocketspec/security_hooks.go +++ b/pkg/websocketspec/security_hooks.go @@ -27,6 +27,21 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList return nil }) + // BeforeHandle: preload column rules for writes before the handler opens its + // transaction; the write hooks below only read the cache. + handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { + return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation) + }) + + // BeforeCreate/BeforeUpdate: drop columns hidden or masked for the user from the + // write payload, so they cannot be inserted or updated. + handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { + return security.ApplyWriteColumnSecurity(newSecurityContext(hookCtx), securityList) + }) + // Hook 1: BeforeRead - Load security rules handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error { secCtx := newSecurityContext(hookCtx) @@ -122,6 +137,14 @@ func (s *securityContext) SetQuery(query interface{}) { s.ctx.Metadata["query"] = query } +func (s *securityContext) GetData() interface{} { + return s.ctx.Data +} + +func (s *securityContext) SetData(data interface{}) { + s.ctx.Data = data +} + func (s *securityContext) GetResult() interface{} { return s.ctx.Result }