mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-13 15:32:24 -06:00
Compare commits
220 Commits
v1.6.1
...
stable/1.5
| Author | SHA1 | Date | |
|---|---|---|---|
| 4a38b835f5 | |||
| 415be00149 | |||
| 346ad2a6aa | |||
| 8003fcbebe | |||
| d4f3711d63 | |||
| cea229206b | |||
| b7b4dcc0df | |||
| 97080e1df9 | |||
| af0cbaaec3 | |||
| b9ce0d388e | |||
| d07d2242aa | |||
| 70514cc406 | |||
| 687a3367c0 | |||
| bfd99c6a81 | |||
| 1a813c8130 | |||
| d6aa85db6d | |||
| c0c7fda7f9 | |||
| 7fae2698d7 | |||
| 4bbe64755e | |||
| c1281b9721 | |||
| 3124dbe52f | |||
| a8a1e738ca | |||
| 27768abecc | |||
| 1e5017293a | |||
| b2c688f3b1 | |||
| 09b1f07e18 | |||
| 1de8f6b4c7 | |||
| 4f89c1c3f3 | |||
| 6898860f18 | |||
| 3d07cad272 | |||
| 3737ddf89f | |||
| 464450b9e2 | |||
| 1fba80a9b5 | |||
| 28a5914be0 | |||
| bb2515eddf | |||
| 19978bfd89 | |||
| 8eec44d809 | |||
| f1cf516eb6 | |||
| 1df1e739ef | |||
| fbbb21012a | |||
| 617de2488f | |||
| 1d5189bb8d | |||
| 89a282f86b | |||
| 59c943b83a | |||
| 14b7516b3f | |||
| 5d0ec99449 | |||
| 6dfd1b5c18 | |||
| 658c65aee8 | |||
| 641ce8e7f6 | |||
| d76c57e687 | |||
| 58bf811a0e | |||
| df942e375e | |||
| 84fd5dc859 | |||
| ba8b1d9126 | |||
| 685b1e3d9b | |||
| 91300060fa | |||
| 688f047ce1 | |||
| 87893aa4ab | |||
| 0988142303 | |||
| 40ecebf012 | |||
| c91869c7e5 | |||
| 53b52092f9 | |||
| 0d1a009a4c | |||
| e8352bd8e5 | |||
| de4cc568c4 | |||
| b477c85ddc | |||
| 00bd80a658 | |||
| 47df9d23c5 | |||
| 16fc7efce2 | |||
| e8eca2ec9b | |||
| 57563b0c12 | |||
| 43e622840b | |||
| 584437b98e | |||
| 3abe1c0058 | |||
| e24d8b9597 | |||
| e11e6f6b70 | |||
| b8d728fd79 | |||
| 2138c19821 | |||
| 627bf06ced | |||
| b67da0f48a | |||
| e05b6adc67 | |||
| 40355d8303 | |||
| 2cbd926b9d | |||
| 40c0ab5a58 | |||
| a5e8b8b17e | |||
| bed691c62a | |||
| c930078f3d | |||
| f148c4b423 | |||
| 8ce4c8e737 | |||
| 95e67dc768 | |||
| dc35cbc7bf | |||
| 79f4d0030d | |||
| 14e504db1f | |||
| 0abe0cb77d | |||
| 610513398b | |||
| 2d6519f9a8 | |||
| a99ce49311 | |||
| 46f3571c93 | |||
| 53cabe7e20 | |||
| 1d9fd94e23 | |||
| eb89ddab1e | |||
| eb92e61755 | |||
| 0e2ea122eb | |||
| eb9dd2402a | |||
| 0c58910c4b | |||
| 2e393d76b4 | |||
| dec175f176 | |||
| 592433b46d | |||
| da5321eb88 | |||
| baa2214f96 | |||
| 3b60a69e4f | |||
| 5d1213d3dc | |||
| 9ae2b376c7 | |||
| b368bdeecc | |||
| d767aca784 | |||
| 12bc580dee | |||
| aa1446364b | |||
| 21507dc02e | |||
| a0ed4b9897 | |||
| ab8ee0d759 | |||
| d34f6cd0b1 | |||
| b219c47ba8 | |||
| d770a811a8 | |||
| 912e9c57b0 | |||
| d7c6053441 | |||
| e7a17a20b0 | |||
| 931a1eca9d | |||
| 048285a423 | |||
| 481347eb17 | |||
| 31a554a4bd | |||
| af2c0ae13a | |||
| 0808dc0af0 | |||
| cfc8a6c8c0 | |||
| 266e3536aa | |||
| 42bf9aecaf | |||
| 191775dd7e | |||
| c41fd2be2e | |||
| 1787fb5c11 | |||
| 814c42763d | |||
| 242596ced3 | |||
| bde0913442 | |||
| 570b198f1b | |||
| 96d935f1f7 | |||
| 1a1043c4df | |||
| 55aab54774 | |||
| 0f8c8b38a3 | |||
| b0f7029ff1 | |||
| a4c335d7bf | |||
| 21663d1567 | |||
| c823156af5 | |||
| bace928477 | |||
| d16c911750 | |||
| b8fadad94f | |||
| 63aecdf2fa | |||
| cbe8940b30 | |||
| 1dcd1e2ec4 | |||
| 32e29ff255 | |||
| 3e2fe0bc9d | |||
| 5f5eee4aab | |||
| 6d532ed776 | |||
| c3d9cdae82 | |||
| 366d316941 | |||
| 3e87f4262e | |||
| f50b559792 | |||
| c6b3c0bc5f | |||
| cefb74a226 | |||
| 273d547f4e | |||
| bbb404c363 | |||
| 072113f7ca | |||
| 7f1b0acf7a | |||
| 6904bd8f39 | |||
| 4d677d1ebf | |||
| 35f462a46d | |||
| ec74334e74 | |||
| 2fd0c29a92 | |||
| 1207d27363 | |||
| 733c9818d4 | |||
| 28a2779c10 | |||
| 56364b0b5b | |||
| afb5804a7c | |||
| 4b508a1319 | |||
| 9c2cb185e1 | |||
| 0519b847bd | |||
| bbc8b99a9f | |||
| bd9f780b21 | |||
| 5d14b5f675 | |||
| 4693fa95f1 | |||
| c3423d6606 | |||
| 53f1222c22 | |||
| 802d87a57f | |||
| 8349d9994d | |||
| ac1fd67137 | |||
| 1b40ae79f9 | |||
| b078ddccf0 | |||
| 4b6c93a0e9 | |||
| 9d283e951f | |||
| 4e407e7d4f | |||
| 7ab24e500b | |||
| 5bcbcb73b9 | |||
| af6749421a | |||
| 4d6cb77075 | |||
| 5c225ef39b | |||
| f5a843f44a | |||
| cf44841624 | |||
| 7ffab6a272 | |||
| ba3bc9d989 | |||
| dbe023b4dd | |||
| 3d3a8b7367 | |||
| 99eff73a97 | |||
| 99fcd30299 | |||
| 423c2e80b7 | |||
| a0eb77360d | |||
| 3dd0e196fe | |||
| 19c3db5329 | |||
| 0bea72019e | |||
| b9ff52d582 | |||
| 961f999c93 | |||
| 8bdb916064 | |||
| 25fe4e728a | |||
| 0d1a32ff65 |
@@ -47,9 +47,9 @@ jobs:
|
||||
# explicit setup, that suite silently skips if the runner
|
||||
# image happens not to ship Node, masking regressions in
|
||||
# the browser-side renderer.
|
||||
- uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5
|
||||
- uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6
|
||||
with:
|
||||
node-version: "20"
|
||||
node-version: "24"
|
||||
- run: pip install -e ".[test]"
|
||||
- run: pytest tests/ -m "not live" --cov=turnstone --cov-report=term-missing --cov-report=xml -q
|
||||
- uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7
|
||||
@@ -79,9 +79,9 @@ jobs:
|
||||
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
|
||||
with:
|
||||
python-version: "3.14"
|
||||
- uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5
|
||||
- uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6
|
||||
with:
|
||||
node-version: "20"
|
||||
node-version: "24"
|
||||
- run: pip install -e ".[test,postgres]"
|
||||
- run: pytest tests/ -m "not live" --storage-backend=postgresql -q
|
||||
env:
|
||||
|
||||
+649
-3
@@ -8,13 +8,659 @@ version numbers (`X.Y.Z`, with `X.Y.ZaN` / `bN` / `rcN` for pre-releases).
|
||||
|
||||
Three release tracks are maintained:
|
||||
|
||||
- **`stable/1.0`** — patch-only (`v1.0.x`)
|
||||
- **`stable/1.3`** — patch-only (`v1.3.x`)
|
||||
- **`stable/1.4`** — patch-only (`v1.4.x`)
|
||||
- **`main`** — experimental (`v1.5.0aN`)
|
||||
- **`stable/1.5`** — patch-only (`v1.5.x`)
|
||||
- **`main`** — experimental (next major)
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [1.5.18]
|
||||
|
||||
Backports the `turnstone-admin` config-loading alignment from `main`
|
||||
plus the accompanying `load_config` permission-warning hardening. No
|
||||
schema changes.
|
||||
|
||||
### Added
|
||||
|
||||
- **`turnstone-admin` reads `config.toml`** — the admin CLI now honors
|
||||
the same `[database]` section that `turnstone-server` does, with the
|
||||
same precedence (`CLI / config.toml > TURNSTONE_DB_* env > defaults`).
|
||||
Operators with DB credentials in `config.toml` no longer need to
|
||||
re-export `TURNSTONE_DB_URL` before every admin invocation. Newly
|
||||
plumbed through to `init_storage`: `pool_size`, `sslmode`,
|
||||
`sslrootcert`, `sslcert`, `sslkey` — previously the admin CLI
|
||||
silently dropped these. A new `--config PATH` flag mirrors the
|
||||
one already on `turnstone-server`.
|
||||
|
||||
### Security
|
||||
|
||||
- **Permissive `config.toml` now warns** — `turnstone.core.config.load_config`
|
||||
logs a single warning when the resolved config file is group- or
|
||||
world-readable (any bit in `0o077`). DB password and TLS key paths
|
||||
live in `[database]`; operators usually want the file at `0600`.
|
||||
|
||||
## [1.5.17]
|
||||
|
||||
Backports a clutch of coordinator-tool clarity fixes plus a watch-delivery
|
||||
correctness fix from `main` to the `stable/1.5` track, plus a previously-
|
||||
latent intent-verdicts persistence bug exposed by the new heuristic-verdict
|
||||
INSERT paths. No schema changes.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **`intent_verdicts` PK collisions on every llm_fallback delivery** —
|
||||
async LLM-tier "llm_fallback" verdicts (`turnstone/core/judge.py` —
|
||||
`_deliver_fallbacks` and the in-loop fallback path) deliberately
|
||||
reuse the heuristic verdict's `verdict_id` so the row gets
|
||||
"upgraded in place" from `tier="heuristic"` → `tier="llm_fallback"`
|
||||
when the LLM judge times out, is cancelled, or returns no content.
|
||||
The consumer `_persist_intent_verdict` was doing a plain INSERT,
|
||||
hitting the `intent_verdicts_pkey` constraint on every fallback
|
||||
delivery; Postgres logged the duplicate-key error, the application
|
||||
try/except swallowed it at `log.debug`, and the row never actually
|
||||
got upgraded — the LLM judge's annotation
|
||||
(`"(LLM judge did not return a verdict)"`) was lost. The collision
|
||||
rate exploded on this release because the new heuristic-INSERT
|
||||
paths in the auto-approve early-return branches of `approve_tools`
|
||||
(introduced below) leave no gap for the fallback to land cleanly
|
||||
into. Fix: new `upsert_intent_verdict` storage method using
|
||||
`ON CONFLICT (verdict_id) DO UPDATE` that updates only `tier`,
|
||||
`reasoning`, `judge_model` — the three fields that genuinely
|
||||
change between heuristic and llm_fallback. Every other column
|
||||
(identity, carried-verbatim, and `user_decision`) is excluded;
|
||||
`user_decision` in particular would otherwise be clobbered back
|
||||
to `"pending"` when a fallback arrives after the operator has
|
||||
already resolved the approval. The bulk-INSERT path stays as
|
||||
plain INSERT — fresh UUIDs in `judge.evaluate` make in-turn dups
|
||||
impossible; the inverse race (fallback wins before bulk lands) is
|
||||
reachable but unchanged in observable behavior by this fix,
|
||||
documented at the bulk site for a future hardening pass.
|
||||
- **Coordinator LLM re-spawn loops on large fan-outs** — the spawn-tool
|
||||
return JSON used `ws_id` as its key, which primed the model's recency
|
||||
bias to feed the spawn result straight back into another
|
||||
`spawn_workstream(ws_id=...)` call instead of progressing to
|
||||
`wait_for_workstream(ws_ids=[...])`. On 10+ child fan-outs this cascaded
|
||||
into self-inflicted re-spawn loops. The LLM-facing tool result now emits
|
||||
`child_ws_id` (the storage column / HTTP API contract is unchanged); the
|
||||
field name is already an existing project term so the rename aligns
|
||||
rather than introduces new vocabulary. Also handles the silent
|
||||
upstream-omits-ws_id success-shape edge that previously emitted
|
||||
`{"child_ws_id": null}` to the LLM — now surfaces a tool error so the
|
||||
model retries rather than chasing a null id.
|
||||
- **`inspect_workstream` blowing the coordinator context budget** — a
|
||||
coord doing a fan-out wave against tool-heavy children could land
|
||||
>100 KB of raw output per inspect call, and the previous safety net
|
||||
(`_truncate_output`'s head+tail strategy) silently dropped *middle*
|
||||
messages — exactly the wrong shape for understanding a child's
|
||||
trajectory (the FIRST sets the brief, the LAST shows the conclusion,
|
||||
the middle is the connective tissue). Output now goes through a
|
||||
three-tier degradation ladder mirroring the search tool's
|
||||
`_format_search_results`: `_tier="full"` (every message verbatim) →
|
||||
`_tier="compact"` (per-message head/tail-snipped content + snipped
|
||||
`tool_calls.arguments`, falling through a `(20,30)` / `(10,20)` /
|
||||
`(5,10)` message-list trim ladder) → `_tier="skeleton"` (counts, role
|
||||
distribution, last-assistant preview). Budget 32 KiB matches the
|
||||
search tool's; the chosen tier is annotated on the response so the
|
||||
model can recall with a tighter `message_limit` if signal was lost.
|
||||
- **Auto-approved verdicts indistinguishable from pending review** —
|
||||
`intent_verdict` rows for auto-approved tool calls landed with
|
||||
`user_decision=""`, which read identically to "still waiting for the
|
||||
operator" in the audit trail and led to a real misdiagnosis incident.
|
||||
The column now carries an explicit vocabulary at insert: `pending` /
|
||||
`approved` / `denied` / `timeout` / `policy` / `blanket` / `skill` /
|
||||
`always` / `auto_approve_tools`. The auto-approve early-return
|
||||
branches in `approve_tools` now persist heuristic verdicts stamped
|
||||
with their reason (previously dropped on the floor), and late LLM-tier
|
||||
verdicts that arrive for an already-auto-approved call_id are stamped
|
||||
via a TTL-pruned lookup map — so the audit row carries the
|
||||
auto-approve reason even when the LLM judge daemon completes after
|
||||
the synchronous approval cycle finished. `resolve_approval` gains a
|
||||
`timeout` kwarg writing `"timeout"` (the previous shape collapsed
|
||||
passive timeouts and active denials into the same column).
|
||||
- **`list_skills` empty `allowed_tools` misread as "no tool access"** —
|
||||
the response previously emitted `"allowed_tools": []` for every skill
|
||||
that hadn't declared an auto-approve allowlist, which a coordinator
|
||||
model read as "this skill can't use any tools" (real misdiagnosis: a
|
||||
code-review child appeared to have been spawned with zero tool
|
||||
access). The field is now omitted entirely when empty — absence
|
||||
carries the unambiguous meaning "no tool is pre-approved for this
|
||||
skill", presence (non-empty list) keeps the standard Claude Code
|
||||
skill-spec shape. The tool description rewrite makes the
|
||||
auto-approve-allowlist semantics explicit so a future reader doesn't
|
||||
re-derive the gating misread.
|
||||
- **Watch terminal-fires silently dropped on backpressure** —
|
||||
delivery now routes terminal events through the same path as
|
||||
normal fires instead of being filtered out when the consumer was
|
||||
saturated.
|
||||
|
||||
### Documentation
|
||||
|
||||
- **Storage `LIKE_ESCAPE` contract** — clarify that callers passing
|
||||
`.like(escape=...)` must use the same escape character that the
|
||||
storage helper assumes; previous wording let a reader pass a
|
||||
different escape and silently produce no matches.
|
||||
|
||||
## [1.5.15]
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Admin console blank-page on MCP server rows with consented users** — a
|
||||
Phase 9 (1.5.14) regression in `admin.js` used double-quote string
|
||||
delimiters on the bulk-revoke button HTML literal, but the literal embeds
|
||||
a `"` mid-attribute. JS closed the string early, turned `bulk-revoke (`
|
||||
into bare tokens, and the resulting `SyntaxError` wiped out every global
|
||||
in `admin.js` — `showAdmin` and all other admin entry points became
|
||||
undefined, so the console UI was non-functional whenever the rendered MCP
|
||||
server list contained at least one row with `consented_users_count > 0`.
|
||||
Switch the literal to single-quote delimiters to match the surrounding
|
||||
block.
|
||||
|
||||
## [1.5.14]
|
||||
|
||||
Backports OAuth-MCP Phase 9 from `main` to the `stable/1.5` track.
|
||||
|
||||
### Added
|
||||
|
||||
- **OAuth-MCP Phase 9 — admin status, deferred-consent persistence, operator
|
||||
docs** — completes the per-(user, server) OAuth-MCP build-out. The sync pool
|
||||
dispatchers now upsert into a new `mcp_pending_consent` table on
|
||||
`mcp_consent_required` / `mcp_insufficient_scope`, so a non-interactive run
|
||||
(scheduled / channel) that hits an unconsented server surfaces the deferred
|
||||
prompt to the user on their next dashboard load via the gear-icon badge —
|
||||
rows are cleared automatically by the OAuth callback handler on consent
|
||||
completion, or via new DELETE endpoints for manual dismiss. The MCP Servers
|
||||
admin row gains a `consented_users_count` pill and a two-step-confirm
|
||||
bulk-revoke button for `auth_type=oauth_user` servers (upstream RFC 7009
|
||||
revoke is intentionally not attempted in bulk to avoid N synchronous
|
||||
round-trips against the provider). Operator-facing docs land at
|
||||
`docs/mcp-oauth.md` and `docs/operations/mcp-oauth-headless.md`.
|
||||
|
||||
Introduces forward-only migrations `054_mcp_pending_consent` and
|
||||
`055_mcp_user_tokens_server_index`.
|
||||
|
||||
## [1.5.13]
|
||||
|
||||
This release introduces one forward-only schema migration:
|
||||
`053_services_notify_trigger` — installs the `services_notify` PostgreSQL
|
||||
trigger that backs the new LISTEN/NOTIFY dispatcher (no-op on SQLite, where
|
||||
the dispatcher uses in-process fan-out).
|
||||
|
||||
### Added
|
||||
|
||||
- **Reactive node discovery via PG LISTEN/NOTIFY** — the console gains a
|
||||
`NotifyDispatcher` that holds a dedicated session-mode PostgreSQL `LISTEN`
|
||||
connection (bypasses pgbouncer transaction pooling) and fans wake-ups out to
|
||||
per-channel handlers on a separate dispatch thread. The cluster collector
|
||||
subscribes to a new `services` channel and reacts to node register /
|
||||
deregister within ~500 ms instead of waiting up to 60 s for the next discovery
|
||||
loop; the 60 s loop is retained as the backstop for crash-shaped loss
|
||||
(NOTIFY only fires on real writes). The storage layer also gains a uniform
|
||||
`notify` / `listen` API with an SQLite synthetic-sweep fallback so consumer
|
||||
code is identical across backends. `TURNSTONE_DB_LISTEN_URL` (or
|
||||
`[database] listen_url` in `config.toml`) points the dispatcher at a
|
||||
direct-to-Postgres URL; defaults to the main DB URL when unset.
|
||||
- **Event-driven `wait_for_workstream`** — coord's block-wait tool no longer
|
||||
polls storage every 500 ms. A new in-process `ChildEventBus` notifies waiters
|
||||
whenever a child state change is dispatched to the UI, and the wait loop
|
||||
blocks on `threading.Event.wait` with a 2 s heartbeat cap (matching the
|
||||
existing `wait_progress` SSE cadence). A 600 s wait that previously hit
|
||||
storage ~2400 times now wakes only on real state transitions, with ~4× lower
|
||||
SSE traffic in the quiescent case.
|
||||
- **Memory tool audit trail** — the memory tool now emits `memory.save`,
|
||||
`memory.update`, and `memory.delete` audit events (the admin-console DELETE
|
||||
route previously emitted only `memory.delete`, so tool-initiated mutations
|
||||
had no audit footprint). All emissions are best-effort and never break the
|
||||
tool call itself.
|
||||
- **`task_agent` per-call personas via `skill=`** — `task_agent` now accepts
|
||||
an optional `skill=<name>` argument that loads the named skill's content as
|
||||
the sub-agent's persona in place of the hardcoded identity statement. The
|
||||
fixed operating-guidance block (one-shot, tool-use over narration,
|
||||
no follow-up questions) is still layered on top of every persona. High- and
|
||||
critical-risk skills surface their risk tier in the approval header and
|
||||
emit a `task_agent.high_risk_skill` warning, matching the existing
|
||||
session-load gate.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Per-role plan / task model overrides could be bypassed by the LLM** — the
|
||||
back-compat `default` alias auto-synthesised by `load_model_registry`
|
||||
remained visible to the model even when an operator had configured
|
||||
`model.task_alias` / `model.plan_alias`, so `task_agent(model="default")`
|
||||
routed to whichever backend the synthesised alias was attached to at boot
|
||||
instead of the configured per-role default. The synthesised alias is now
|
||||
only added when neither the DB nor `[models.*]` populates the registry,
|
||||
filtered out of the LLM-visible alias list, and explicitly rejected at the
|
||||
validator chokepoint as defense-in-depth.
|
||||
- **Mermaid streaming parse errors + progressive `hljs`** — live-streamed
|
||||
mermaid blocks with bare `(`, `[`, `{` inside unquoted edge or rectangle
|
||||
node labels were re-entering the shape parser and producing
|
||||
`Parse error, got 'PS'` messages. The renderer now autoquotes the two
|
||||
affected label forms (`|content|` and `ID[content]`) before the SVG cache
|
||||
lookup; shapes whose syntax already nests delimiters (cylinders, subroutines,
|
||||
trapezoids, etc.) are intentionally left alone. The companion `hljs` change
|
||||
highlights code blocks progressively as they stream rather than only after
|
||||
completion.
|
||||
- **Re-auth from inside the proxy-prefixed UI** — on a proxied node page
|
||||
(`/node/{id}/...`), an expiring JWT triggered an in-page login modal whose
|
||||
POST went to `/v1/api/auth/login` and was rewritten to
|
||||
`/node/{id}/v1/api/auth/login`. Two latent bugs both blocked re-auth: the
|
||||
console's `AuthMiddleware` didn't recognise the `/node/{id}/` prefix over a
|
||||
public path, and `proxy_api` would have forwarded the login request to the
|
||||
upstream node (which mints `JWT_AUD_SERVER` tokens the console then rejects).
|
||||
Both fixed: proxied public paths stay public, and `proxy_api` now dispatches
|
||||
every entry in `_PROXY_AUTH_LOCAL_HANDLERS` (login, logout, setup, refresh,
|
||||
status, whoami, oidc/authorize, oidc/callback) to the console's own auth
|
||||
handlers. The dispatch table is a single `(method, path) → handler` mapping
|
||||
so the test parametrize list can't drift from the implementation.
|
||||
- **Appbar visibility + gear-icon dropdown on the dashboard** — the dashboard
|
||||
overlay was covering the entire appbar, hiding the proxy-injected node
|
||||
picker. The overlay now starts at `top: 48px` and the dashboard's role
|
||||
downgrades from `dialog+aria-modal` to `region` so the appbar above it
|
||||
remains reachable. The gear icon converts from a direct settings-panel
|
||||
click into a dropdown with "MCP connections" and "Logout" (the latter with
|
||||
`.destructive` styling). The settings-menu keydown handler is now attached
|
||||
synchronously so `Escape` can't fall through the brief window between the
|
||||
menu opening and its listeners being installed.
|
||||
- **PostgreSQL test backend on the notify dispatcher suite** — migration 053's
|
||||
`services_notify` trigger lives only in the alembic chain, but the test
|
||||
fixture creates tables via `metadata.create_all`. The trigger function +
|
||||
trigger are now declared in `_schema.py` and attached via
|
||||
`sa.event.listen(services, "after_create", ...)` DDL events gated on the
|
||||
PostgreSQL dialect, with the same SQL constants imported by migration 053
|
||||
so there's a single source of truth.
|
||||
|
||||
## [1.5.12]
|
||||
|
||||
### Added
|
||||
|
||||
- **Enriched backend error messages** — provider name and attempted URL are now
|
||||
included in session error responses, so operators can triage connectivity
|
||||
failures without enabling debug logging.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **`/rewind` always emits a `history` SSE event** — pre-fix, if the session
|
||||
had no messages remaining after a rewind the history event was skipped,
|
||||
leaving connected UIs with stale content and blocking edit-and-resend flows.
|
||||
|
||||
## [1.5.11]
|
||||
|
||||
This release introduces one forward-only schema migration:
|
||||
`052_model_reasoning_persistence` — `surface_persisted_reasoning` and
|
||||
`replay_reasoning_to_model` flag columns on `model_definitions`.
|
||||
|
||||
### Added
|
||||
|
||||
- **SSE refresh-resume** — clients that reload mid-stream (browser refresh, tab
|
||||
restore) now receive an `in_progress_snapshot` event carrying the buffered
|
||||
partial response, so the UI can resume rendering the in-flight turn without
|
||||
losing content. The snapshot is keyed by a monotonic `_ws_inflight_seq`
|
||||
counter so a reconnecting client can skip events it already saw.
|
||||
- **Reasoning persistence** (Phases 1–4) — model reasoning text can now be
|
||||
persisted to conversation history and optionally replayed to the model on
|
||||
subsequent turns. Phase 1 persists reasoning text on the history payload.
|
||||
Phase 2 wires a build-time shape filter and a per-model
|
||||
`replay_reasoning_to_model` flag. Phases 3+4 add full OpenAI Responses API
|
||||
(`include=["reasoning.encrypted_content"]`) and Chat Completions support;
|
||||
an `ANTHROPIC_VALID_BLOCK_TYPES` shape filter guards the Anthropic path. Two
|
||||
new per-model capability flags (`surface_persisted_reasoning`,
|
||||
`replay_reasoning_to_model`) both default `False` on unknown and
|
||||
local-server models.
|
||||
- **Console home composer: placeholders + toggle** — the console landing-page
|
||||
composer now shows context-aware placeholder text and a toggle component for
|
||||
advanced options; an admin polish pass tightened spacing and focus behaviour
|
||||
across the form.
|
||||
|
||||
### Changed
|
||||
|
||||
- **`judge.model` now requires a named alias** — raw provider model IDs on
|
||||
`judge.model` in config are no longer accepted; the judge must reference an
|
||||
alias registered in the model registry. The session-provider raw-model
|
||||
fallback is removed. Existing configs using an unregistered model ID need a
|
||||
corresponding alias entry.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **`replay_reasoning_to_model` AND-gated with model capability** — setting the
|
||||
flag for a model that does not declare reasoning-replay support now silently
|
||||
no-ops instead of forwarding reasoning blocks and triggering a provider error.
|
||||
- **Coordinator alias resolution unified across placeholder + factory** — a
|
||||
placeholder coordinator and the real coordinator factory could previously
|
||||
resolve to different model aliases, producing a visible mismatch in the model
|
||||
display. Both paths now share the same resolution logic.
|
||||
- **Console `cs=None` fallback in `/v1/api/models` placeholder** — an
|
||||
under-initialised coordinator state no longer 500s when the models endpoint
|
||||
is hit before the coordinator subsystem is fully bootstrapped.
|
||||
- **SSE `_ws_inflight_seq` always advances** — sequence numbers were previously
|
||||
skipped when an emit was past the buffer cap, leaving gaps in the monotonic
|
||||
counter that broke `state_change` / `in_progress_snapshot` ordering on
|
||||
reconnect.
|
||||
- **Reasoning persistence shape + replay fixes** — per-block
|
||||
`ANTHROPIC_VALID_BLOCK_TYPES` filter applied; `reasoning_text` is now
|
||||
synthesised alongside non-reasoning `provider_blocks` so both appear
|
||||
together in the history payload.
|
||||
|
||||
## [1.5.10]
|
||||
|
||||
This release introduces one forward-only schema migration:
|
||||
`051_skill_notify_on_complete_array_default` — backfills
|
||||
`prompt_templates.notify_on_complete` from `'{}'` to `'[]'`.
|
||||
|
||||
### Added
|
||||
|
||||
- **Skills unlock action** — operators can unlock an installed skill to allow
|
||||
local customisation. Once unlocked, the skill's resource content, system
|
||||
prompt additions, and notify configuration are editable through the admin UI.
|
||||
Skills shipped as part of a bundle remain locked (read-only) until explicitly
|
||||
unlocked; the unlock is logged to the audit trail. A lock icon in the
|
||||
top-right of the Skills detail pane doubles as the unlock trigger.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **`skills.sh` install endpoint** — the install script was targeting an
|
||||
endpoint removed in an earlier refactor; switched to `/api/download`.
|
||||
- **Skills `notify_on_complete` default** — the field defaulted to `{}`
|
||||
(object) instead of `[]` (array), causing notify configurations to be
|
||||
rejected at schema validation.
|
||||
- **Skills admin UI modal errors** — `.is-visible` class used consistently
|
||||
instead of inline `style.display`; stale error text is cleared on submit;
|
||||
designer-review lock-icon UX applied.
|
||||
|
||||
## [1.5.9]
|
||||
|
||||
### Fixed
|
||||
|
||||
- **`repair=False` on all display-read `load_messages` call sites** —
|
||||
passing `repair=True` on display paths was silently mutating the stored
|
||||
message list, causing divergence between what the UI showed and what the
|
||||
model received on the next turn.
|
||||
|
||||
## [1.5.8]
|
||||
|
||||
This release introduces two forward-only schema migrations:
|
||||
`049_mcp_oauth_schema` — OAuth token + consent tables for MCP servers;
|
||||
`050_conversations_source_and_reminders` — `_source` and `_reminders` columns
|
||||
on `conversations`.
|
||||
|
||||
### Added
|
||||
|
||||
- **MCP OAuth 2.1 + PKCE** — MCP servers that require OAuth can now be
|
||||
configured with a client ID and secret through the admin UI. The full token
|
||||
lifecycle (acquire → refresh → rotate) is managed automatically; tokens are
|
||||
stored encrypted at rest using a key derived from the JWT secret. The consent
|
||||
flow runs in-browser via a provider redirect. Rolled out in phases:
|
||||
|
||||
- Minimum admin form and OAuth schema (`21663d15`).
|
||||
- Token-at-rest AES-GCM encryption layer (`a4c335d7`).
|
||||
- Per-(user, server) OAuth 2.1 + PKCE flow (`b0f7029f`).
|
||||
- Per-(user, server) `ClientSession` pool with OAuth dispatch (`1a1043c4`).
|
||||
- SDK 401/403 introspection via httpx response hook (`bde09134`).
|
||||
- Phase 7 — per-user tool catalog scoping: each user sees only the tools
|
||||
their OAuth token is permitted to call (`cfc8a6c8`).
|
||||
- Phase 7b — per-user resource + prompt pool dispatch (`b368bdee`).
|
||||
- Phase 8 — per-user MCP consent UX: users see a consent dialog on first
|
||||
use of an OAuth-gated server and can revoke consent from their profile;
|
||||
admins see per-server consent counts in the MCP Servers tab (`61051339`).
|
||||
|
||||
- **Metacognition NudgeQueue** — all advisory channels (repeat-tool nudges,
|
||||
watch reminders, wake triggers) are unified into a pull-model `NudgeQueue`
|
||||
that delivers at most one nudge per turn, preventing multi-channel pile-ups
|
||||
that inflate context. Observable changes:
|
||||
|
||||
- Watch results carry metadata (watch ID, `valid_until`, trigger type)
|
||||
through to the system message so the model can reason about recency.
|
||||
- Coordinator idle-children observer: a coordinator with no in-flight
|
||||
children for longer than the configured idle threshold receives a nudge.
|
||||
- Wake trigger (`IdleNudgeWatcher`): sessions waiting on an external event
|
||||
can be unblocked via `ChatSession.deliver_wake_nudge_from_queue`.
|
||||
- Watch switchover: watch results are now enqueued on the `NudgeQueue`
|
||||
rather than the previous `_watch_pending` list, giving them the same
|
||||
delivery guarantees and priority handling as other advisories.
|
||||
|
||||
- **Structured watch-result card** — the UI renders watch results as a styled
|
||||
card with a system-nudge marker, distinct from the assistant message body.
|
||||
On history replay, system-nudge turns are visually distinguished from normal
|
||||
assistant turns.
|
||||
- **Side-channel persistence** — `_source` and `_reminders` side-channel
|
||||
fields are persisted to the `conversations` storage table and restored on
|
||||
session resume, so metacognitive context survives process restarts. A
|
||||
`REMINDER_TEXT_STORAGE_CAP` byte clamp prevents unbounded growth.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Replay consistency** — queued user messages captured mid-loop are now
|
||||
persisted and replayed in the correct order on a subsequent `events`
|
||||
subscription. Coordinator history replay fixed: blank assistant cards and
|
||||
out-of-order tool results on the coordinator tree no longer occur when the
|
||||
coordinator has mixed queued + delivered messages.
|
||||
- **Session reminder preservation on fork + resume** — `_source` and
|
||||
`_reminders` are carried through workstream fork and restored from storage
|
||||
on resume.
|
||||
- **NUL-byte sanitization in storage** — PostgreSQL rejects `\x00` in text
|
||||
columns; `_source` and `_reminders` now strip NUL bytes on write.
|
||||
- **Console coordinator subsystem bootstrap** — the coordinator subsystem is
|
||||
now committed atomically on first model add; startup teardown is offloaded
|
||||
to avoid blocking the event loop.
|
||||
- **MCP `asyncio.timeout` over `asyncio.wait_for`** — Python 3.11's
|
||||
`wait_for` wraps the coroutine in a fresh task, breaking anyio's `aclose`
|
||||
scope exit. Replaced with `async with asyncio.timeout(N)` for safe cleanup.
|
||||
- **MCP pool-reuse 401 recovery** — a reused `ClientSession` returning 401
|
||||
now replaces the pool entry with a fresh session; the carrier token is
|
||||
owned by the pool entry to prevent a race between the 401 handler and a
|
||||
concurrent request.
|
||||
- **OIDC hardening** — multiple security and correctness fixes:
|
||||
SSRF + plaintext credential exfil via discovery document (sec-1, sec-3);
|
||||
`TURNSTONE_OIDC_REDIRECT_BASE` now required, Host-header fallback removed
|
||||
(sec-2); atomic user + identity provisioning prevents orphan rows (bug-1);
|
||||
callback robustness — typed exceptions, shape checks, log sanitization, JS
|
||||
race (bug-4–6, sec-4); role-mapping concurrency serialized (bug-2, perf-1);
|
||||
stranded-user self-heal on role-mapping failure (cumulative bug-1).
|
||||
|
||||
## [1.5.7]
|
||||
|
||||
### Added
|
||||
|
||||
- **Inline node picker** — a compact node-switcher dropdown in the console
|
||||
header replaces the "← Back to console" banner, so operators can switch
|
||||
between nodes without a full navigation.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Queued user messages injected mid-loop** — messages queued while a
|
||||
generation was in progress were not being delivered at the correct seam and
|
||||
could be dropped or reordered when the worker consumed the queue.
|
||||
- **Search tool output bounded** — pathological inputs (very long lines with
|
||||
no whitespace) could produce search results exceeding the context budget.
|
||||
Output is now clamped before reaching the message.
|
||||
|
||||
## [1.5.6]
|
||||
|
||||
### Added
|
||||
|
||||
- **`api_surface` toggle** — model definitions gain an `api_surface` field
|
||||
(`"chat"` | `"responses"`) that selects which OpenAI-compatible API surface
|
||||
the provider client uses. Enables Mistral Medium reasoning via the Responses
|
||||
surface; Chat Completions remains the default for all other models.
|
||||
- **Healthy model aliases per node** — `GET /v1/api/cluster/nodes` now
|
||||
includes a `healthy_aliases` list per node, so the coordinator and operators
|
||||
can see which model aliases are currently reachable without a separate
|
||||
per-model health probe.
|
||||
- **Plan/task agent settings in Models → Roles** — the Models admin tab's
|
||||
Roles sub-tab gains `plan_agent` and `task_agent` rows so operators can
|
||||
configure per-kind reasoning effort and alias overrides from the UI rather
|
||||
than editing `config.toml`. Live-refresh dropdowns update in place when
|
||||
model definitions change.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Memory candidate selection** — recall now uses OR-of-terms BM25 with
|
||||
query-aware candidate-set selection, dramatically improving recall for
|
||||
queries whose terms span multiple stored entries.
|
||||
- **Workstream model + config preserved on rehydrate** — reopening a closed
|
||||
workstream no longer overwrites the model alias and per-workstream config
|
||||
with session defaults.
|
||||
- **Console home composer: attachments + user-message pills** — multipart
|
||||
attachments in the home composer were not forwarded correctly; user-message
|
||||
pills in the coordinator chat pane were missing.
|
||||
|
||||
## [1.5.5]
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Saved-workstream tool result rendering** — tool results in closed
|
||||
workstreams were not rendering on history replay. Audit-trail decoration for
|
||||
tool calls is now applied on the replay path.
|
||||
|
||||
## [1.5.4]
|
||||
|
||||
### Added
|
||||
|
||||
- **Stage 3 SessionManager Children primitive lift** — child workstreams are
|
||||
first-class citizens in the cluster event bus. `child_ws_state` events are
|
||||
pushed through the cluster SSE stream so the console tree view updates in
|
||||
real time without polling. `list_children` and `get_child` primitives on
|
||||
`SessionManager` provide a consistent cross-node view of the coordinator's
|
||||
spawn tree.
|
||||
- **Multi-select delete for Saved Coordinators** — the Saved Coordinators grid
|
||||
in the console admin panel now supports checkbox multi-select with a
|
||||
bulk-delete action.
|
||||
|
||||
## [1.5.3]
|
||||
|
||||
This release introduces one forward-only schema migration:
|
||||
`048_workstream_reaper_index` — partial composite index on `workstreams` for
|
||||
the orphan-reaper query.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Coordinator orphan reaping scoped by heartbeat** — the session manager's
|
||||
`close_idle` pass now scopes the DB-orphan reaper by
|
||||
`services.last_heartbeat` so workstreams belonging to a live node are not
|
||||
incorrectly reaped. `bulk_close_stale_orphans` and `touch_workstream`
|
||||
storage primitives added; a partial composite index keeps the reaper scan
|
||||
cheap.
|
||||
- **Coordinator pool idle cleanup** — a periodic task on the console now
|
||||
closes coordinator pool entries whose session has gone idle past the
|
||||
configurable threshold, preventing pool exhaustion on long-running consoles.
|
||||
|
||||
## [1.5.2]
|
||||
|
||||
### Added
|
||||
|
||||
- **Metacognition themed reminder bubble** — repeat-tool and user-reminder
|
||||
nudges are rendered as a distinct styled bubble rather than being injected
|
||||
inline into the assistant message, making it easier to distinguish model
|
||||
output from metacognitive annotations. The CLI REPL gains matching
|
||||
`on_user_reminder` / `on_tool_reminder` callbacks.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Metacog streak detector** — the N≥3 sequential-same-call streak detector
|
||||
now fires correctly on the third repetition; a write-success-clear that
|
||||
reset the counter after a successful tool call (preventing streaks across
|
||||
mixed-outcome sequences) was removed.
|
||||
- **Metacog reminders isolated to side-channel** — reminder text no longer
|
||||
appears in the user content turn; it flows through a dedicated side-channel
|
||||
the session injects into the system context, preventing the model from
|
||||
attributing it to the user.
|
||||
|
||||
## [1.5.1]
|
||||
|
||||
### Added
|
||||
|
||||
- **`pending_approval_detail` on child `ws_state` SSE events** — coordinators
|
||||
now receive the child's pending approval detail in `child_ws_state` events,
|
||||
enabling the coordinator to surface approval prompts without a separate poll.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Coordinator registry auto-refresh** — the console coordinator registry now
|
||||
refreshes when model definitions change, so a newly added alias is visible
|
||||
to coordinators without restarting.
|
||||
- **Coordinator fan-out default** — coordinators now fan out to independent
|
||||
child workstreams by default instead of serialising them, matching the
|
||||
documented contract for parallel-work patterns.
|
||||
- **`wait_for_workstream` message cap raised to 10 KiB** — large plan
|
||||
summaries and tool results from child workstreams were silently truncated at
|
||||
the previous 4 KiB cap.
|
||||
- **Coordinator SSE isolated on dedicated thread pool** — coordinator SSE
|
||||
polling now runs on a dedicated 200-thread executor, matching interactive's
|
||||
`sse_executor`, so coordinator long-poll blocking no longer contends with
|
||||
storage and routing workers on the default pool.
|
||||
|
||||
## [1.5.0]
|
||||
|
||||
User-visible additions: a unified workstream HTTP surface (interactive and
|
||||
coordinator under one URL family), inline child approvals, coordinator
|
||||
composer parity, progressive rendering, OIDC authentication, MCP OAuth
|
||||
foundations, and a redesigned UI built on the Design System v1 token layer.
|
||||
|
||||
This release removes the pre-1.5 body-keyed and query-keyed URL family.
|
||||
See **Removed (BREAKING)** below before upgrading from a 1.x stable line.
|
||||
|
||||
This release introduces the following forward-only schema migrations that the
|
||||
server applies automatically on first startup. All are additive; no data loss.
|
||||
|
||||
- `039_workstream_kind` — `kind` + `parent_ws_id` columns on `workstreams`.
|
||||
- `040_coord_cluster_admin_perms` — grants `admin.coordinator` +
|
||||
`admin.cluster.inspect` to the builtin-admin role.
|
||||
- `041_workstream_index_tuning` — refined indexes for the workstream query mix
|
||||
introduced by 039.
|
||||
- `042_coord_trust_send_perm` — adds `coordinator.trust.send` permission to
|
||||
builtin-admin.
|
||||
- `043_skill_description_required` — backfills empty `description` rows in
|
||||
`prompt_templates`.
|
||||
- `044_skill_kind` — adds `kind` classifier column to `prompt_templates`
|
||||
(`interactive` / `coordinator` / `any`).
|
||||
- `045_skill_risk_level_rename` — renames `prompt_templates.scan_status` →
|
||||
`risk_level`.
|
||||
- `046_drop_hash_ring_tables` — drops the hash-ring bucket tables superseded
|
||||
by rendezvous routing in 1.4.
|
||||
- `047_drop_coord_spawn_quota_settings` — removes the spawn-quota settings
|
||||
rows removed from the coordinator in 1.5.0a4.
|
||||
|
||||
### Added
|
||||
|
||||
- **Inline child approvals** — pending tool approvals on coordinator child
|
||||
workstreams surface directly in the coordinator tree view. A risk pill shows
|
||||
the judge verdict (or "pending" while the judge evaluates); Approve/Deny
|
||||
buttons appear inline so operators do not need to navigate to the child's
|
||||
workstream. `pending_approval_detail` is exposed on
|
||||
`GET /v1/api/dashboard` and passed through the cluster live-bulk SSE payload
|
||||
so all connected clients render approval prompts simultaneously. LLM judge
|
||||
verdicts are cached client-side and replayed on SSE reconnect.
|
||||
- **Coordinator composer parity** — the coordinator composer now supports
|
||||
Stop, Send-to-queue, and Attach (file upload), matching the interactive
|
||||
workstream composer feature set.
|
||||
- **Per-call model and judge override on coordinator composer** — operators
|
||||
can override the model alias and judge model for a single coordinator send
|
||||
from the composer, without changing the node-wide or role-wide defaults. Bad
|
||||
aliases return a corrective error listing available choices.
|
||||
- **Coordinator status bar + richer history replay** — each coordinator
|
||||
workstream gains a per-coordinator status bar showing active children, token
|
||||
spend, and generation state. History replay in the coordinator panel is
|
||||
extended to include tool results and thinking blocks.
|
||||
- **Coordinator child error surfacing + memory tool** — child workstream
|
||||
errors are surfaced as distinct error rows in the coordinator tree view
|
||||
rather than disappearing silently. The coordinator gains access to a
|
||||
`memory` tool (same interface as interactive) for retrieving stored facts.
|
||||
- **Coordinator inline tool-batch construct** — the coordinator tool approval
|
||||
UI replaces the separate approval dock with an inline batch construct that
|
||||
groups all pending tool calls for a given turn into a single review card.
|
||||
- **Node capability auto-detection** — nodes report kernel-level capabilities
|
||||
(available memory, CPU count, accelerator presence) via
|
||||
`/v1/api/node/capabilities` at startup, enabling the console to filter model
|
||||
aliases offered to coordinators routing to that node.
|
||||
- **Skills: paste `SKILL.md` to auto-fill the Create Skill modal** — pasting
|
||||
a `SKILL.md` file's content into the modal auto-populates the name,
|
||||
description, and configuration fields.
|
||||
- **Progressive mermaid rendering** — Mermaid diagrams begin rendering as
|
||||
soon as a complete diagram block is detected in the stream rather than
|
||||
waiting for the full response; the diagram re-renders in place as the model
|
||||
extends it.
|
||||
- **LaTeX and MathML delimiter support** — `\(…\)` inline and `\[…\]` block
|
||||
math delimiters are now recognised alongside the existing `$$` fences.
|
||||
|
||||
### Removed (BREAKING — 1.5.0)
|
||||
|
||||
- **Legacy body-keyed and query-keyed URL family for the workstream
|
||||
|
||||
+6
-3
@@ -8,14 +8,17 @@ FROM python:3.14-slim
|
||||
LABEL org.opencontainers.image.title="turnstone" \
|
||||
org.opencontainers.image.description="Multi-node AI orchestration platform"
|
||||
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.11.7 /uv /usr/local/bin/uv
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.11.14 /uv /usr/local/bin/uv
|
||||
|
||||
# Remove the slim image's man page exclusion so man-db has actual content
|
||||
RUN rm -f /etc/dpkg/dpkg.cfg.d/docker
|
||||
|
||||
# System dependencies: psycopg (libpq5), developer tooling for agent workflows
|
||||
# System dependencies: psycopg (libpq5), developer tooling for agent workflows.
|
||||
# ripgrep is the preferred backend for the search tool — natively bounds
|
||||
# per-line, per-file, and per-filesize so pathological inputs (minified
|
||||
# bundles, training-data JSONL with multi-MB single records) can't OOM us.
|
||||
RUN apt-get update && apt-get upgrade -y && apt-get install -y --no-install-recommends \
|
||||
libpq5 git curl jq man-db manpages procps file \
|
||||
libpq5 git curl jq man-db manpages procps file ripgrep \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Node.js LTS (for npx-based MCP servers like @modelcontextprotocol/server-github)
|
||||
|
||||
+46
-1
@@ -281,6 +281,7 @@ Each message in the `messages` array has:
|
||||
| `role` | string | `"user"`, `"assistant"`, or `"tool"` |
|
||||
| `content` | string or null | Text content of the message |
|
||||
| `tool_calls` | array or null | Present only on assistant messages with calls |
|
||||
| `reasoning` | string (optional) | Concatenated reasoning / chain-of-thought text on assistant turns whose `provider_data` carried reasoning-bearing blocks (Anthropic `thinking`, OpenAI Responses `reasoning`, or synthetic `reasoning_text` from local-model servers). Present only when the active model's `surface_persisted_reasoning` flag is True. |
|
||||
|
||||
Each entry in `tool_calls`:
|
||||
|
||||
@@ -325,6 +326,44 @@ finalize any in-progress assistant message.
|
||||
{"type": "stream_end"}
|
||||
```
|
||||
|
||||
**`state_change`** -- the worker thread transitioned to a new state. Drives
|
||||
the client's busy-mode (composer in send vs. stop, spinner indicators,
|
||||
auto-focus on idle). Sent live during normal operation AND on every fresh
|
||||
SSE subscribe (so a mid-stream page refresh restores the correct composer
|
||||
state without waiting for the next live transition).
|
||||
|
||||
```json
|
||||
{"type": "state_change", "state": "running"}
|
||||
```
|
||||
|
||||
| Field | Type | Description |
|
||||
|----------|--------|----------------------------------------------------------------------|
|
||||
| `state` | string | One of `"running"`, `"thinking"`, `"attention"`, `"idle"`, `"error"` |
|
||||
|
||||
**`in_progress_snapshot`** -- one-shot replay of the in-progress turn's
|
||||
content + reasoning text-so-far when this client connects mid-stream.
|
||||
Lets a refreshing browser tab restore partial assistant text immediately
|
||||
instead of waiting for the response to complete. Yielded once after the
|
||||
kind-specific replay phase (history + pending), only when at least one
|
||||
of `content` / `reasoning` is non-empty. Both halves render into the same
|
||||
assistant bubble the live `content` / `reasoning` events would target;
|
||||
clients should treat the snapshot as idempotent (skip overwrite if the
|
||||
current local buffer is already a superset prefix — covers EventSource
|
||||
auto-reconnect re-replays).
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "in_progress_snapshot",
|
||||
"content": "Here is the answer so far: it depends on ",
|
||||
"reasoning": "The user is asking about a comparison; let me think about..."
|
||||
}
|
||||
```
|
||||
|
||||
| Field | Type | Description |
|
||||
|--------------|--------|------------------------------------------------------------|
|
||||
| `content` | string | Joined assistant content text accumulated this turn |
|
||||
| `reasoning` | string | Joined reasoning / chain-of-thought text accumulated |
|
||||
|
||||
**`tool_info`** -- one or more tool calls that were auto-approved (no user
|
||||
action required).
|
||||
|
||||
@@ -522,7 +561,13 @@ Each SSE connection to a workstream receives its own delivery queue. Events
|
||||
produced by the worker thread are fanned out to all registered listener queues,
|
||||
so multiple consumers (browser, console proxy, SDK) can connect
|
||||
simultaneously and each receives every event. On reconnect the client receives
|
||||
a full history replay, so no catch-up mechanism is needed.
|
||||
the kind-specific replay (`connected` + `status` + `history` + pending
|
||||
approval / plan for interactive; `connected` + `status` + pending for coord)
|
||||
followed by a `state_change` carrying the current worker state and an
|
||||
optional `in_progress_snapshot` carrying any partial content / reasoning
|
||||
buffered for the in-progress turn — so a mid-stream refresh restores both
|
||||
the busy-mode UI and the partial assistant text without waiting for the
|
||||
response to complete.
|
||||
|
||||
---
|
||||
|
||||
|
||||
+51
-12
@@ -91,7 +91,7 @@ turnstone/
|
||||
discord/ Discord adapter (bot, cog, views, streaming, config)
|
||||
slack/ Slack adapter (Socket Mode bot, DM routing, approval buttons)
|
||||
shared_static/ Shared design system (base.css, auth.js, theme.js, toast.js, utils.js, kb.js)
|
||||
katex-0.16.45/ Vendored KaTeX math rendering library (MIT, woff2 fonts)
|
||||
katex-0.16.47/ Vendored KaTeX math rendering library (MIT, woff2 fonts)
|
||||
ui/
|
||||
colors.py ANSI color constants with NO_COLOR support
|
||||
markdown.py Streaming terminal markdown renderer (line-buffered)
|
||||
@@ -231,11 +231,13 @@ The engine emits state changes via `_emit_state()` which calls
|
||||
|
||||
> See also: [Core Engine Classes diagram](diagrams/png/03-core-engine-classes.png)
|
||||
|
||||
Defined in `turnstone.core.session.SessionUI` as a `typing.Protocol` with 14
|
||||
Defined in `turnstone.core.session.SessionUI` as a `typing.Protocol` with 16
|
||||
methods. Every frontend must implement all of them.
|
||||
|
||||
```python
|
||||
class SessionUI(Protocol):
|
||||
def on_turn_start(self) -> None: ...
|
||||
def on_turn_committed(self) -> None: ...
|
||||
def on_thinking_start(self) -> None: ...
|
||||
def on_thinking_stop(self) -> None: ...
|
||||
def on_reasoning_token(self, text: str) -> None: ...
|
||||
@@ -252,6 +254,14 @@ class SessionUI(Protocol):
|
||||
def on_rename(self, name: str) -> None: ... # propagate alias to tab/UI label
|
||||
```
|
||||
|
||||
`on_turn_start` fires at the top of each iteration of the send-loop;
|
||||
`on_turn_committed` fires immediately after `messages.append(assistant_msg)`.
|
||||
`SessionUIBase` uses both to reset the per-turn inflight buffers
|
||||
(`_ws_inflight_content` / `_ws_inflight_reasoning` / `_ws_inflight_seq`)
|
||||
that fuel the SSE refresh-resume `in_progress_snapshot` event — see
|
||||
the per-workstream events stream in
|
||||
[`docs/api-reference.md`](api-reference.md#get-v1apiworkstreamsws_idevents).
|
||||
|
||||
`on_rename` is called by the `/name` command (on success) and after a successful `/resume` (if the resumed session has an alias or title). `WebUI.on_rename` broadcasts a `ws_rename` event on the global SSE channel and updates the in-memory `Workstream.name`; `TerminalUI.on_rename` is a no-op.
|
||||
|
||||
### Three Implementations
|
||||
@@ -546,11 +556,9 @@ adds, removes, or reconnects servers as needed.
|
||||
6. `_exec_mcp_tool()` calls `call_tool_sync()` which dispatches to the async loop
|
||||
via `asyncio.run_coroutine_threadsafe()`
|
||||
|
||||
**Tool refresh:** Three mechanisms keep tools up-to-date without restart:
|
||||
**Tool refresh:** Two mechanisms keep tools up-to-date without restart:
|
||||
- **Push:** Servers declaring `tools.listChanged` send `ToolListChangedNotification`;
|
||||
the registered `message_handler` triggers immediate single-server refresh.
|
||||
- **Periodic:** Servers without push support are polled on a staggered interval
|
||||
(default 4 h, configurable via `[mcp] refresh_interval` or `--mcp-refresh-interval`).
|
||||
- **Manual:** `/mcp refresh [server]` calls `refresh_sync()` for on-demand refresh
|
||||
(also attempts reconnection for disconnected servers).
|
||||
|
||||
@@ -572,10 +580,11 @@ from a healthy connection do not trip the breaker. When the cooldown expires
|
||||
(`call_tool_sync`, `read_resource_sync`, `get_prompt_sync`, `refresh_sync`)
|
||||
cancel orphaned futures on timeout to prevent coroutine accumulation on the
|
||||
background event loop. Push notification refreshes are debounced (5 s per
|
||||
server) to protect against notification storms. The periodic refresh loop
|
||||
attempts reconnection for disconnected servers with exponential backoff
|
||||
(60 s–1 h). Transport stream references are pre-closed before stack teardown to
|
||||
work around the MCP SDK's anyio cancel-scope CPU busy-loop (SDK #2147).
|
||||
server) to protect against notification storms. Operators can force a
|
||||
catalog refresh or full reconnect from the admin panel; reconnects clear
|
||||
the circuit breaker and run a fresh handshake. Transport stream references
|
||||
are pre-closed before stack teardown to work around the MCP SDK's anyio
|
||||
cancel-scope CPU busy-loop (SDK #2147).
|
||||
|
||||
**Error isolation:** Per-server connection/refresh failures are caught and logged; other
|
||||
servers are unaffected. Tool execution errors return error strings to the LLM
|
||||
@@ -620,14 +629,15 @@ LLMProvider (protocol)
|
||||
| `get_capabilities()` | Per-model flags (`ModelCapabilities`) |
|
||||
| `convert_tools()` | Translate OpenAI tool schemas to provider format |
|
||||
| `retryable_error_names` | Exception class names that trigger retry |
|
||||
| `extract_reasoning_text()` | Walk stored `provider_blocks`, return concatenated reasoning text for UI rehydration (per-provider block-type knowledge: Anthropic `thinking`, OpenAI Responses `reasoning`, OpenAI Chat synthetic `reasoning_text`) |
|
||||
|
||||
**Normalized data types:**
|
||||
|
||||
| Type | Fields |
|
||||
|------|--------|
|
||||
| `StreamChunk` | `content_delta`, `reasoning_delta`, `tool_call_deltas`, `info_delta`, `usage`, `finish_reason` |
|
||||
| `CompletionResult` | `content`, `tool_calls`, `finish_reason`, `usage` |
|
||||
| `ModelCapabilities` | `context_window`, `max_output_tokens`, `supports_temperature`, `token_param`, `thinking_mode`, `supports_effort`, `supports_web_search`, `supports_tool_search`, `supports_vision` |
|
||||
| `StreamChunk` | `content_delta`, `reasoning_delta`, `tool_call_deltas`, `info_delta`, `usage`, `finish_reason`, `provider_blocks` |
|
||||
| `CompletionResult` | `content`, `tool_calls`, `finish_reason`, `usage`, `provider_blocks` |
|
||||
| `ModelCapabilities` | `context_window`, `max_output_tokens`, `supports_temperature`, `token_param`, `thinking_mode`, `supports_effort`, `supports_web_search`, `supports_tool_search`, `supports_vision`, `supports_reasoning_replay` |
|
||||
| `UsageInfo` | `prompt_tokens`, `completion_tokens`, `total_tokens`, `cache_creation_tokens`, `cache_read_tokens` |
|
||||
|
||||
**OpenAIProvider** (`_openai.py`): passes messages through unchanged (they are
|
||||
@@ -715,6 +725,35 @@ and `"openai-compatible"`.
|
||||
`max_tokens`, and `reasoning_effort` to override the global defaults from
|
||||
ConfigStore. When unset (`NULL`), the global default is used.
|
||||
|
||||
**Per-model reasoning persistence:** Two booleans on `model_definitions`
|
||||
(migration 052) control how reasoning text round-trips:
|
||||
|
||||
* `surface_persisted_reasoning` (default `True`) — gates whether stored
|
||||
reasoning text is surfaced on `/history` payloads for UI rehydration.
|
||||
**Storage of reasoning bytes happens regardless of this flag** — they
|
||||
ride in `provider_data` independently. Phase-1 admin UI label "Surface
|
||||
persisted reasoning."
|
||||
* `replay_reasoning_to_model` (default `False`) — gates whether stored
|
||||
reasoning blocks are sent back to the provider on subsequent turns.
|
||||
Capability-gated: `ModelCapabilities.supports_reasoning_replay` must
|
||||
also be `True` for the wire path to actually replay (canonical OpenAI
|
||||
gpt-5*/o-series and Anthropic Claude entries set it; unknown / local-
|
||||
server models default to `False`).
|
||||
|
||||
Three reasoning paths are recognised:
|
||||
|
||||
| Path | Provider | Capture | Persist | Replay |
|
||||
|------|----------|---------|---------|--------|
|
||||
| 1 | Anthropic Messages API | `thinking_delta` | `provider_blocks` (`type="thinking"`) | Verbatim via `_provider_content` |
|
||||
| 2 | OpenAI Responses (gpt-5*, o-series) | `response.reasoning_text.delta` events | `provider_blocks` (`type="reasoning"`) — only when `include=["reasoning.encrypted_content"]` | `ResponseReasoningItemParam` input items |
|
||||
| 3 | OpenAI Chat Completions (vLLM, llama.cpp, Gemini-compat) | `delta.reasoning_content` Pydantic extras | Synthetic `{type: "reasoning_text", text, source}` block stamped at end-of-stream | None — no API surface for replay on Chat Completions |
|
||||
|
||||
Cross-provider safety is enforced by `ANTHROPIC_VALID_BLOCK_TYPES` (a
|
||||
shape filter in `_anthropic.py:_convert_messages`): foreign blocks
|
||||
(OpenAI `reasoning`, synthetic `reasoning_text`) fall through to the
|
||||
text+tool_calls rebuild path rather than reaching Anthropic's input
|
||||
boundary as malformed content.
|
||||
|
||||
```toml
|
||||
[models.local]
|
||||
base_url = "http://localhost:8000/v1"
|
||||
|
||||
@@ -110,11 +110,18 @@ owns it; the node is just currently unreachable.
|
||||
|
||||
### Example — `spawn_batch`
|
||||
|
||||
This is the coordinator-tool result shape (the JSON the LLM receives),
|
||||
not an HTTP API response — the table above keys it under "model tool"
|
||||
to distinguish it from the `/v1/api/...` endpoints in the same table.
|
||||
The underlying HTTP spawn endpoint still returns `ws_id`; the tool
|
||||
result re-keys it to `child_ws_id` to defuse a coordinator-LLM recency
|
||||
bias (see `docs/coordinator-skills.md`).
|
||||
|
||||
```json
|
||||
{
|
||||
"results": {
|
||||
"0": {"ws_id": "d4e5f6...", "name": "csrf-audit", "node_id": "gpu-3"},
|
||||
"2": {"ws_id": "f1a2b3...", "name": "xss-audit", "node_id": "gpu-1"}
|
||||
"0": {"child_ws_id": "d4e5f6...", "name": "csrf-audit", "node_id": "gpu-3"},
|
||||
"2": {"child_ws_id": "f1a2b3...", "name": "xss-audit", "node_id": "gpu-1"}
|
||||
},
|
||||
"denied": [
|
||||
{"idx": 1, "reason": "skill not found: nonexistent-skill"}
|
||||
|
||||
@@ -115,7 +115,8 @@ with a `type` field. The recurring shapes a UI has to handle:
|
||||
| `approve_request` | One or more tool calls need operator approval | `items: [{call_id, header, preview, func_name, approval_label, needs_approval}]` |
|
||||
| `approval_resolved` | Operator answered the approval prompt | `approved`, `feedback` |
|
||||
| `state_change` | Worker-thread state transition (also re-emitted with the current state on every fresh subscribe so refresh-mid-stream restores composer mode) | `state` ∈ `running`, `thinking`, `attention`, `idle`, `error` |
|
||||
| `status` | Token usage + context-window snapshot (fires on every streaming tick) | `prompt_tokens`, `completion_tokens`, `total_tokens`, `context_window`, `pct`, `effort`, `cache_creation_tokens`, `cache_read_tokens` |
|
||||
| `in_progress_snapshot` | One-shot replay of the in-progress turn's content + reasoning when this client connects mid-stream | `content`, `reasoning` |
|
||||
| `status` | Token usage + context-window snapshot (fires on every streaming tick) | `prompt_tokens`, `completion_tokens`, `total_tokens`, `context_window`, `pct`, `effort`, `cache_creation_tokens`, `cache_read_tokens` |
|
||||
| `rename` | Session's display name changed | `name` |
|
||||
| `intent_verdict` | Intent judge produced a verdict on a pending tool call | `risk_level`, `recommendation`, `reasons` |
|
||||
| `output_warning` | Output guard flagged a tool result | `call_id`, `risk_level`, `flags` |
|
||||
@@ -130,9 +131,13 @@ with a `type` field. The recurring shapes a UI has to handle:
|
||||
**Reconnection contract:** a freshly-opened SSE connection receives
|
||||
the current snapshot of any pending tool approval (`approve_request`
|
||||
is re-sent if unresolved), any in-flight `wait_*` / `batch_*`
|
||||
indicator — so a tab refresh mid-approval doesn't strand the
|
||||
operator.
|
||||
|
||||
indicator, the worker's current `state_change`, and an
|
||||
`in_progress_snapshot` carrying any partial content / reasoning the
|
||||
model has produced for the in-progress turn — so a tab refresh
|
||||
mid-approval, mid-tool-execution, or mid-stream restores both the
|
||||
correct composer mode and the partial assistant text without waiting
|
||||
for the response to complete.
|
||||
|
||||
---
|
||||
|
||||
## 3. Send the first user message
|
||||
|
||||
@@ -169,14 +169,21 @@ validates ws_id against `parent_ws_id=coord_ws_id` AND
|
||||
the wait into reporting "complete".
|
||||
|
||||
Pattern: capture each spawn result in the next tool call's input.
|
||||
The JSON tool-result carries `{"ws_id": "...", "name": "...",
|
||||
The JSON tool-result carries `{"child_ws_id": "...", "name": "...",
|
||||
"node_id": "...", "routing_strategy": "..."}`; the model should
|
||||
extract the ws_id and pass it to `inspect_workstream` /
|
||||
`wait_for_workstream` / `send_to_workstream` / `close_workstream`
|
||||
verbatim.
|
||||
extract the `child_ws_id` and pass it as `ws_id` (or in the `ws_ids`
|
||||
list) to `inspect_workstream` / `wait_for_workstream` /
|
||||
`send_to_workstream` / `close_workstream` verbatim. The asymmetry
|
||||
— spawn returns `child_ws_id` but the other tools accept `ws_id` /
|
||||
`ws_ids` — is intentional: it defuses a coordinator-LLM recency
|
||||
bias where seeing `ws_id` in a spawn return primed re-spawn loops
|
||||
instead of progression to the wait phase.
|
||||
|
||||
A UI that wants human-readable identifiers should render the `name`
|
||||
field and keep the ws_id as the click-through key.
|
||||
field and keep the workstream id as the click-through key — note
|
||||
that the id *value* is the same regardless of whether it arrived
|
||||
under the `child_ws_id` key (spawn return) or the `ws_id` key
|
||||
(every other tool's input/output); only the field name differs.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ package "turnstone/core/" <<Rectangle>> {
|
||||
component [auth.py\nAuthentication] as auth <<core>>
|
||||
component [healthcheck.py\nBackendHealthMonitor] as healthcheck <<core>>
|
||||
component [ratelimit.py\nRateLimiter] as ratelimit <<core>>
|
||||
component [mcp_client.py\nMCPClientManager\n(push + periodic refresh)] as mcp <<core>>
|
||||
component [mcp_client.py\nMCPClientManager\n(push + manual refresh)] as mcp <<core>>
|
||||
component [tool_search.py\nToolSearchManager, BM25] as toolsearch <<core>>
|
||||
component [model_registry.py\nModelRegistry] as registry <<core>>
|
||||
}
|
||||
|
||||
@@ -69,9 +69,10 @@ class "NullUI" as NullUI {
|
||||
interface "LLMProvider" as LLMProvider <<Protocol>> {
|
||||
+ provider_name: str {property}
|
||||
+ get_capabilities(model) → ModelCapabilities
|
||||
+ create_streaming(client, model, messages, ...) → Iterator[StreamChunk]
|
||||
+ create_completion(client, model, messages, ...) → CompletionResult
|
||||
+ create_streaming(client, model, messages, ..., replay_reasoning_to_model) → Iterator[StreamChunk]
|
||||
+ create_completion(client, model, messages, ..., replay_reasoning_to_model) → CompletionResult
|
||||
+ convert_tools(tools) → list[dict]
|
||||
+ extract_reasoning_text(provider_blocks) → str
|
||||
+ retryable_error_names: frozenset[str] {property}
|
||||
--
|
||||
core/providers/_protocol.py
|
||||
@@ -126,6 +127,7 @@ class "ModelCapabilities" as ModelCaps <<frozen>> {
|
||||
+ supports_web_search: bool
|
||||
+ supports_tool_search: bool
|
||||
+ supports_vision: bool
|
||||
+ supports_reasoning_replay: bool
|
||||
}
|
||||
|
||||
' ChatSession
|
||||
@@ -253,7 +255,7 @@ class "MCPClientManager" as MCPMgr {
|
||||
Background asyncio event loop
|
||||
bridges async MCP SDK to
|
||||
sync ChatSession dispatch.
|
||||
Push + periodic + manual refresh.
|
||||
Push + manual refresh.
|
||||
Resources + prompts discovered
|
||||
alongside tools at startup.
|
||||
--
|
||||
|
||||
@@ -24,6 +24,14 @@ CS -> DB : save_message(ws_id, "user", input)
|
||||
|
||||
group loop [while tool_calls present]
|
||||
|
||||
CS -> UI : on_turn_start()
|
||||
note right of UI
|
||||
SessionUIBase resets the per-turn inflight
|
||||
buffers (_ws_inflight_content / reasoning /
|
||||
seq) that fuel the SSE in_progress_snapshot
|
||||
event for mid-stream refresh resume.
|
||||
end note
|
||||
|
||||
CS -> UI : on_state_change("thinking")
|
||||
CS -> UI : on_thinking_start()
|
||||
|
||||
@@ -73,6 +81,14 @@ group loop [while tool_calls present]
|
||||
|
||||
CS -> CS : _update_token_table()\ncalibrate chars_per_token ratio
|
||||
CS -> CS : messages.append(assistant_msg)
|
||||
CS -> UI : on_turn_committed()
|
||||
note right of UI
|
||||
Drops the per-turn inflight buffers — the
|
||||
assistant message is now in the history
|
||||
list, so the in_progress_snapshot must
|
||||
not re-render it during the next tool-
|
||||
execution window or the next streaming turn.
|
||||
end note
|
||||
CS -> DB : save_message(ws_id, "assistant", content)
|
||||
CS -> DB : save_message(ws_id, "tool_call", ...) ×N
|
||||
|
||||
|
||||
@@ -190,21 +190,25 @@ group Push Notifications (debounced 5s per server)
|
||||
MCPMgr -> Storage : sync_prompts_to_storage()
|
||||
end
|
||||
|
||||
group Periodic Polling (default 4h)
|
||||
MCPMgr -> MCPMgr : _periodic_refresh()
|
||||
group Manual Refresh
|
||||
Session -> MCPMgr : refresh_sync()
|
||||
note right
|
||||
Only polls capabilities
|
||||
without push support.
|
||||
Staggered per-server.
|
||||
Disconnected servers get
|
||||
reconnect attempts with
|
||||
exponential backoff (60s-1h).
|
||||
/mcp refresh [server] —
|
||||
re-fetches catalog and
|
||||
attempts reconnect for
|
||||
disconnected servers.
|
||||
end note
|
||||
end
|
||||
|
||||
group Manual Refresh
|
||||
Session -> MCPMgr : refresh_sync()
|
||||
note right: /mcp refresh [server]
|
||||
group Manual Reconnect
|
||||
Session -> MCPMgr : reconnect_sync(name)
|
||||
note right
|
||||
Operator-driven via the
|
||||
console admin panel —
|
||||
tears down session, clears
|
||||
circuit breaker, runs a
|
||||
fresh handshake.
|
||||
end note
|
||||
end
|
||||
|
||||
== Policy Evaluation ==
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a3b5c59403a6febd81667fc8fd2a7d22bc59da6130eba0dea5449c42668d0ede
|
||||
size 387044
|
||||
oid sha256:95dd5ebc899a1261d516686a5aa3319a7f45015d411302825fa28afbfc82e1ce
|
||||
size 326766
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:474b900448ec04d1117b48a2b55614524721b2f04ac4bda66170bd0a06aae0f2
|
||||
size 624573
|
||||
oid sha256:9857db23fe3c4316d492073aac69c7e7558b1abe3b95ad7756d4a5933bd0ece7
|
||||
size 620214
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:3aa8d972bba40d78152f9f0c762b9f5ec616d8052c45fa52b7dd1c679ed81d61
|
||||
size 325245
|
||||
oid sha256:c14dfbb2db8dcb22cd332b2cf0e53ba75141adb213dae47dd9dbfd389ed482fe
|
||||
size 354799
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7623df33be9baf7647ca1c2450640df57e1cd73e8be1f8168aae16e546ad683c
|
||||
size 459941
|
||||
oid sha256:d6aff446a062aa08f316985d00c2183148694f786d7f22172bc50b30046c728b
|
||||
size 379259
|
||||
|
||||
@@ -84,6 +84,7 @@ Auth is always enabled. `TURNSTONE_JWT_SECRET` is required.
|
||||
|----------|---------|-------------|
|
||||
| `TURNSTONE_DB_BACKEND` | `sqlite` | Storage backend: `sqlite` or `postgresql` |
|
||||
| `TURNSTONE_DB_URL` | — | Database URL (e.g. `postgresql+psycopg://user:pass@postgres:5432/turnstone`). For SQLite, defaults to `/data/.turnstone.db` |
|
||||
| `TURNSTONE_DB_LISTEN_URL` | (falls back to `TURNSTONE_DB_URL`) | Direct-to-PostgreSQL URL for the console's dedicated `LISTEN` connection. Set this when `TURNSTONE_DB_URL` points at PgBouncer in transaction pooling mode — LISTEN is session state and the transaction-pooled connection can't hold it. See [pgbouncer.md](pgbouncer.md). |
|
||||
| `TURNSTONE_DB_POOL_SIZE` | `2` | PostgreSQL connection pool size per process (default: 2 base + 3 overflow = 5 max) |
|
||||
| `POSTGRES_USER` | `turnstone` | PostgreSQL container username (used in default `TURNSTONE_DB_URL` for cluster/channel) |
|
||||
| `POSTGRES_PASSWORD` | — | PostgreSQL container password (required for production and cluster profiles) |
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
# MCP OAuth — per-user authorization for MCP servers
|
||||
|
||||
Turnstone supports **per-(user, MCP server) OAuth 2.1 + PKCE** delegation so each Turnstone user authorizes a remote MCP server with their own identity, rather than sharing a single bearer token across the deployment. This is the right shape for MCP servers that expose user-specific data (a personal CRM, an email inbox, a calendar) and for MCP servers that want per-user audit attribution.
|
||||
|
||||
Per-user OAuth is opt-in per `mcp_servers` row. Local-auth Turnstone installs with no `oauth_user` rows exercise zero new code paths — the entire feature is dark by default.
|
||||
|
||||
> **Note**: This is a separate authorization layer from Turnstone's own user authentication. A user who logs into Turnstone with a local username + password can still authorize a per-server OAuth MCP server. OIDC SSO and per-server OAuth are orthogonal.
|
||||
|
||||
---
|
||||
|
||||
## When to use which `auth_type`
|
||||
|
||||
The MCP server admin form exposes three authorization modes ("Multitenant Authorization"):
|
||||
|
||||
| `auth_type` | What it means | When to use |
|
||||
|---|---|---|
|
||||
| `none` | No headers attached. Open MCP server (or one gated by network policy only). | Internal MCP servers on a trusted network. |
|
||||
| `static` | One static bearer token, configured per server, sent on every request from every user. | Service-to-service MCP servers where per-user attribution doesn't matter, or single-tenant deployments. |
|
||||
| `oauth_user` *(recommended for user-data servers)* | Each user authorizes separately via OAuth 2.1 + PKCE; Turnstone stores per-user tokens encrypted at rest. | MCP servers that expose user-specific data or that want per-user audit attribution. |
|
||||
|
||||
Switching `auth_type` away from `oauth_user` orphans existing per-user tokens. Use the admin **bulk-revoke** affordance on the server row (Phase 9) to clear them, or let them expire naturally — they're inert without the matching `auth_type` value.
|
||||
|
||||
---
|
||||
|
||||
## Prerequisites for `auth_type=oauth_user`
|
||||
|
||||
1. **Encryption key**. Tokens are stored encrypted with Fernet. Set `[security] mcp_token_encryption_key` in `config.toml` (Turnstone won't start with an `oauth_user` row configured but no key installed). Rotate via `MultiFernet` — add the new key first, then later remove the old one once all rows have been re-encrypted.
|
||||
|
||||
2. **MCP server publishes RFC 9728 PRM and RFC 8414 AS metadata** *or* you configure the AS URL override on the server row. PKCE S256 is mandatory; Turnstone refuses to connect to authorization servers that don't advertise `code_challenge_methods_supported: ["S256"]`.
|
||||
|
||||
3. **OAuth client registration**. Two paths:
|
||||
- **Pre-registered** (most common): you create an OAuth client at the authorization server (manually, via admin console, or via Terraform), then paste the `client_id` / `client_secret` into the Turnstone admin form.
|
||||
- **Dynamic client registration** (RFC 7591): if the AS supports it and you select that mode in the admin form, Turnstone registers a client at first use and persists the `client_id` automatically.
|
||||
|
||||
4. **Redirect URI** registered at the authorization server: `https://your-turnstone-host/v1/api/mcp/oauth/callback`.
|
||||
|
||||
---
|
||||
|
||||
## Configuration
|
||||
|
||||
### Per-server fields (admin UI)
|
||||
|
||||
| Field | Required | Description |
|
||||
|---|---|---|
|
||||
| Server URL | Yes | The MCP server's `streamable-http` base URL. |
|
||||
| Multitenant Authorization | Yes | `none` / `static` / `oauth_user` (recommended). |
|
||||
| Authorization Server URL | No | Override for RFC 9728 PRM discovery. Set when your AS endpoint differs from the MCP server URL (e.g., corporate AS protecting a third-party MCP). When unset, Turnstone falls back to PRM discovery against the MCP server itself. |
|
||||
| Client Registration | Yes (oauth_user) | `preregistered` or `dynamic`. |
|
||||
| Client ID | Yes (preregistered) | OAuth 2.0 client ID. Stored unencrypted. |
|
||||
| Client Secret | Optional (write-only) | OAuth 2.0 client secret (confidential client). Encrypted at rest. Written but never re-read by the API; field stays masked. |
|
||||
| Scopes | No | Space-separated default scope set requested at the authorize endpoint. Per-tool step-up may union additional scopes from a server's `insufficient_scope` response. |
|
||||
| Audience | No | RFC 8707 `resource=` parameter sent on every authorize and token request. Defaults to the MCP server URL when unset. Validate against the `aud` claim in returned JWT tokens. |
|
||||
|
||||
### Encryption key
|
||||
|
||||
```toml
|
||||
[security]
|
||||
mcp_token_encryption_key = "base64-fernet-key"
|
||||
# For rotation, list the keys in priority order — first is used for new
|
||||
# writes, all are tried for reads.
|
||||
# mcp_token_encryption_keys = ["new-key", "old-key"]
|
||||
```
|
||||
|
||||
Keep this in `config.toml` rather than environment variables. An in-process LLM with shell-tool access can read the server's environment via `env` / `os.environ` and exfiltrate any secret stored there; secrets in `config.toml` are only loaded into the server at startup and never re-read on a tool-driven path, so a prompt-injection attack against the agent cannot reach them.
|
||||
|
||||
---
|
||||
|
||||
## Lifecycle
|
||||
|
||||
1. **First tool call** for a user against an `oauth_user` MCP server: pool dispatch finds no stored token, returns `mcp_consent_required` to the agent. Dashboard renders an inline "Connect" action card.
|
||||
|
||||
2. **User clicks Connect**: opens `/v1/api/mcp/oauth/start?server=<name>` in a popup. Browser redirects through the AS authorize endpoint, user grants consent, AS redirects back to `/v1/api/mcp/oauth/callback`. Turnstone exchanges code → tokens via PKCE, validates audience, encrypts, persists in `mcp_user_tokens`, redirects user back to the originating URL.
|
||||
|
||||
3. **Subsequent tool calls** by the same user against the same server reuse the persisted token via the per-(user, server) session pool. Tokens auto-refresh via the refresh-token grant when expired; failed refresh emits `mcp_consent_required` to drive re-consent.
|
||||
|
||||
4. **Step-up scope**: when a tool call hits `403` with `WWW-Authenticate: error="insufficient_scope"`, Turnstone emits `mcp_insufficient_scope` with the parsed scope set; the dashboard offers a "Connect with additional scopes" affordance that opens `/v1/api/mcp/oauth/start?server=<name>&scopes=<extra>` so the union of original + new scopes flows into the AS authorize request.
|
||||
|
||||
5. **User revoke** (settings modal): `DELETE /v1/api/mcp/oauth/connections/{server_name}` runs the authoritative local delete + best-effort RFC 7009 upstream revoke (fire-and-forget, capped at 256 concurrent in-flight tasks).
|
||||
|
||||
6. **Admin bulk-revoke** (Phase 9): `POST /v1/api/admin/mcp-servers/{name}/bulk-revoke` drops every user's token for the server. Upstream RFC 7009 revoke is intentionally **not** attempted in bulk (avoids N upstream HTTP calls per admin click); tokens at the AS expire naturally. Use the per-user revoke endpoint if you need guaranteed upstream invalidation.
|
||||
|
||||
---
|
||||
|
||||
## Admin status indicators
|
||||
|
||||
The MCP Servers admin tab shows per-server status pills (Phase 9):
|
||||
|
||||
- **Consented users count** — distinct users with a non-expired token for this server. Surfaced as a `bulk-revoke (N)` button when ≥1; clicking it opens a confirmation dialog. Hidden when 0.
|
||||
- **Last refresh** — timestamp + outcome (`ok` / `error:ClassName`) of the most recent manual or auto-reconnect refresh. Per node. Absent until at least one refresh has occurred (renders as "never" in the admin UI).
|
||||
|
||||
Additional indicators (circuit-breaker state, encryption-key mismatch) are exposed via `get_server_status` on the API but do not yet have a dedicated admin pill — operators see them today via the per-server status text + error tooltip and in audit logs. A future phase may surface these as discrete pills.
|
||||
|
||||
---
|
||||
|
||||
## Auth-type transitions
|
||||
|
||||
| From | To | What happens |
|
||||
|---|---|---|
|
||||
| `none` / `static` → `oauth_user` | — | New code path activates for this server. Existing static headers (if any) are no longer sent. Users must authorize on first use. |
|
||||
| `oauth_user` → `none` / `static` | — | Existing `mcp_user_tokens` rows are **orphaned** — inert without a matching `auth_type`. Use admin bulk-revoke to drop them, or let them expire. Switching back to `oauth_user` later re-activates the orphaned rows if they haven't been deleted. |
|
||||
| OAuth `client_id` or `client_secret` rotated | — | Existing tokens may stop refreshing if the AS treats them as bound to the previous client. Bulk-revoke after rotation. |
|
||||
|
||||
The orphan-by-default behavior is chosen so switching back to `oauth_user` is non-destructive. Bulk-revoke is the explicit cleanup path.
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Symptom | Likely cause | Action |
|
||||
|---|---|---|
|
||||
| `mcp_consent_required` even after consenting | Token persistence failed, or refresh-token rejected by AS | Check audit log for `mcp_server.oauth.persist_failed` or `mcp_server.oauth.token_revoked`. Re-consent via settings modal. |
|
||||
| `mcp_token_undecryptable_key_unknown` | Encryption key rotated without keeping the previous key in the keyring | Add the previous key back to `mcp_token_encryption_keys` until all rows have been re-encrypted, then drop. |
|
||||
| `mcp_oauth_url_insecure` | MCP server URL is `http://` (not `https://`) on a non-loopback host | Use `https://`. Per-user bearers must not transit cleartext. |
|
||||
| Tools fail in scheduled / Discord / Slack runs | OAuth-MCP requires browser-based consent | Users must pre-consent via the web UI. Phase 9 dashboard badge surfaces deferred consents from these runs on next login. |
|
||||
| Circuit breaker open repeatedly | Transport-level errors on the MCP server (DNS, TLS, 5xx) | Check the per-server error pill; auth errors do not trip the breaker. |
|
||||
|
||||
See also: `docs/operations/mcp-oauth-headless.md` for the cron / channel-driven run caveat.
|
||||
+82
-13
@@ -39,18 +39,19 @@ are set.
|
||||
| `TURNSTONE_OIDC_ROLE_CLAIM` | No | — | ID token claim containing role/group values (see [Role Mapping](#role-mapping)) |
|
||||
| `TURNSTONE_OIDC_ROLE_MAP` | No | — | Mapping from claim values to Turnstone role IDs (see [Role Mapping](#role-mapping)) |
|
||||
| `TURNSTONE_OIDC_PASSWORD_ENABLED` | No | `true` | Set to `false` to hide the password form and block all username/password logins (including admin). API tokens continue to work. |
|
||||
| `TURNSTONE_OIDC_REDIRECT_BASE` | No | — | Externally-reachable origin for the OIDC redirect URI (e.g. `https://app.example.com`). Recommended when running behind a reverse proxy. When unset, derived from the request Host header. |
|
||||
| `TURNSTONE_OIDC_REDIRECT_BASE` | Yes | — | Externally-reachable origin for the OIDC redirect URI (e.g. `https://app.example.com`). Without this, OIDC will refuse to start. The previous Host-header fallback was unsafe under permissive reverse proxies. |
|
||||
| `TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS` | No | — | Comma-separated list of additional hostnames whose endpoints the IdP discovery document is allowed to reference. See [Cross-host endpoints](#cross-host-endpoints). |
|
||||
|
||||
OIDC is enabled when all three required fields (issuer, client ID, client
|
||||
secret) are non-empty. If any is missing, OIDC is silently disabled and
|
||||
the login screen shows only the password form.
|
||||
All four required fields — issuer, client ID, client secret, and
|
||||
`TURNSTONE_OIDC_REDIRECT_BASE` — must be set. If any are missing OIDC
|
||||
is disabled at startup (an error is logged when only `redirect_base`
|
||||
is missing) and the login screen shows only the password form.
|
||||
|
||||
### Reverse Proxy / Load Balancer
|
||||
### Redirect base (required)
|
||||
|
||||
When Turnstone runs behind a reverse proxy, the internal `Host` header may
|
||||
not match the externally-reachable URL. Set `TURNSTONE_OIDC_REDIRECT_BASE`
|
||||
to the public origin so the redirect URI sent to the identity provider is
|
||||
correct:
|
||||
`TURNSTONE_OIDC_REDIRECT_BASE` pins the redirect URI sent to the identity
|
||||
provider to a known externally-visible origin. Set it to the public origin
|
||||
of your Turnstone deployment:
|
||||
|
||||
```bash
|
||||
TURNSTONE_OIDC_REDIRECT_BASE=https://app.example.com
|
||||
@@ -60,6 +61,44 @@ The resulting callback URL will be
|
||||
`https://app.example.com/v1/api/auth/oidc/callback` — register this as the
|
||||
authorized redirect URI in your identity provider.
|
||||
|
||||
OIDC will refuse to start when this variable is unset. There is no
|
||||
Host-header fallback: a permissive reverse proxy or direct backend access
|
||||
would otherwise let an attacker spoof `Host` and steer the IdP redirect
|
||||
to a callback origin they control.
|
||||
|
||||
### Cross-host endpoints
|
||||
|
||||
By default, every endpoint in the IdP discovery document
|
||||
(`token_endpoint`, `jwks_uri`, `userinfo_endpoint`) must share the
|
||||
issuer's `(scheme, host, port)`. This prevents a hostile or compromised
|
||||
IdP from redirecting the token-exchange POST (which carries
|
||||
`client_secret`) to an arbitrary host, and prevents JWKS fetches from
|
||||
being aimed at internal services.
|
||||
|
||||
A few public IdPs legitimately split endpoints across hostnames. Google
|
||||
is the canonical example:
|
||||
|
||||
| Field | Hostname |
|
||||
|-------|----------|
|
||||
| issuer | `accounts.google.com` |
|
||||
| token_endpoint | `oauth2.googleapis.com` |
|
||||
| jwks_uri | `www.googleapis.com` |
|
||||
| userinfo_endpoint | `openidconnect.googleapis.com` |
|
||||
|
||||
Google's set is built in — operators using `https://accounts.google.com`
|
||||
need no extra configuration.
|
||||
|
||||
For other IdPs whose discovery document references a non-issuer host,
|
||||
extend the allow-list explicitly:
|
||||
|
||||
```bash
|
||||
TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS=token.example.com,keys.example.com
|
||||
```
|
||||
|
||||
The same scheme / no-userinfo / SSRF rules apply to allow-listed hosts —
|
||||
this knob only relaxes the same-origin check, not the security gates.
|
||||
Each entry is a hostname (no scheme, no path).
|
||||
|
||||
### config.toml alternative
|
||||
|
||||
```toml
|
||||
@@ -198,6 +237,19 @@ TURNSTONE_OIDC_ROLE_MAP="admin:builtin-admin,engineering:builtin-operator,viewer
|
||||
the user authenticates via OIDC, so new group memberships are picked
|
||||
up on the next login.
|
||||
|
||||
### `assigned_by` markers
|
||||
|
||||
Role assignments record an `assigned_by` value that controls how the
|
||||
sync logic treats them. OIDC-driven flows use two distinct markers:
|
||||
|
||||
- `oidc` — set by claim-driven role mapping; revoked automatically on
|
||||
the next login when the corresponding claim value is no longer
|
||||
present.
|
||||
- `oidc-default` — applied to brand-new OIDC users who have no
|
||||
claim-mapped roles, as a safety net so they still get
|
||||
`builtin-viewer` access on first login. Survives subsequent logins
|
||||
regardless of claim contents and is never revoked by `apply_role_mapping`.
|
||||
|
||||
### Built-in Roles
|
||||
|
||||
| Role ID | Permissions |
|
||||
@@ -375,10 +427,27 @@ callback validation. Entries are automatically cleaned up after 5 minutes.
|
||||
|
||||
### "OIDC not configured"
|
||||
|
||||
All three required environment variables must be set:
|
||||
`TURNSTONE_OIDC_ISSUER`, `TURNSTONE_OIDC_CLIENT_ID`, and
|
||||
`TURNSTONE_OIDC_CLIENT_SECRET`. Check that none are empty or
|
||||
whitespace-only.
|
||||
All four required environment variables must be set:
|
||||
`TURNSTONE_OIDC_ISSUER`, `TURNSTONE_OIDC_CLIENT_ID`,
|
||||
`TURNSTONE_OIDC_CLIENT_SECRET`, and `TURNSTONE_OIDC_REDIRECT_BASE`.
|
||||
Check that none are empty or whitespace-only.
|
||||
|
||||
### "OIDC enabled but TURNSTONE_OIDC_REDIRECT_BASE is unset"
|
||||
|
||||
This error is logged when the three credential variables are set but
|
||||
`TURNSTONE_OIDC_REDIRECT_BASE` is missing. OIDC is disabled at startup
|
||||
to prevent Host-header-derived redirect URI spoofing. Set the variable
|
||||
to your service's externally-visible origin (e.g.
|
||||
`https://app.example.com`) and restart the server. See
|
||||
[Redirect base](#redirect-base-required) for the rationale.
|
||||
|
||||
### Discovery silently disables OIDC with "host does not match issuer"
|
||||
|
||||
The IdP discovery document points `token_endpoint`, `jwks_uri`, or
|
||||
`userinfo_endpoint` at a hostname that doesn't share the issuer's
|
||||
origin. If the IdP is legitimate, add the additional hostname(s) to
|
||||
`TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS`. Google is allow-listed
|
||||
automatically; see [Cross-host endpoints](#cross-host-endpoints).
|
||||
|
||||
### "Login session expired"
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
# MCP OAuth in headless / scheduled / channel-driven runs
|
||||
|
||||
**Constraint**: OAuth-MCP servers (`auth_type=oauth_user`) require browser-based user consent. Users must pre-consent via the web UI before any run that cannot drive a browser redirect.
|
||||
|
||||
**Affected surfaces**:
|
||||
|
||||
- Scheduled workstreams (`turnstone-console` task scheduler).
|
||||
- Discord adapter runs.
|
||||
- Slack adapter runs.
|
||||
- Any future channel adapter without an interactive browser session.
|
||||
|
||||
**What happens when consent is missing**:
|
||||
|
||||
A tool call against an `oauth_user` server returns a structured `mcp_consent_required` error to the agent. The agent surfaces the deferred work in its output. Turnstone persists a record to `mcp_pending_consent` so the dashboard badge surfaces the deferred consent need to the user on next login.
|
||||
|
||||
**Recovery**:
|
||||
|
||||
The user opens the dashboard, sees the gear-icon badge counting pending consents, opens the settings modal, clicks Connect for each affected server, and completes the OAuth dance. The pending-consent record is cleared by the OAuth callback handler on success. Subsequent scheduled / channel runs use the freshly-stored token.
|
||||
|
||||
**Pre-consent recipe**:
|
||||
|
||||
Before scheduling a workstream that depends on an `oauth_user` MCP server, the user should:
|
||||
|
||||
1. Open the dashboard.
|
||||
2. Open the settings modal (gear icon).
|
||||
3. Click Connect on each MCP server the schedule will use.
|
||||
4. Confirm consent in the popup.
|
||||
|
||||
This stores tokens that the scheduled run will reuse. Refresh-token rotation is handled transparently on the run side; only the first consent requires browser interaction.
|
||||
@@ -199,4 +199,28 @@ does not support prepared statements. Turnstone's SQLAlchemy layer does
|
||||
not use server-side prepared statements by default, so this is not an
|
||||
issue.
|
||||
|
||||
**LISTEN / NOTIFY not supported in transaction mode** — PgBouncer's
|
||||
transaction pooling assigns a real server connection only for the
|
||||
duration of each transaction, then returns it to the pool. PostgreSQL
|
||||
`LISTEN` is session state — a transaction-pooled client can't hold the
|
||||
multi-statement session a long-lived `LISTEN` needs. The console's
|
||||
`NotifyDispatcher` (reactive node discovery via the `services` channel)
|
||||
therefore opens a **dedicated, direct-to-Postgres** connection that
|
||||
bypasses PgBouncer.
|
||||
|
||||
Configure via `config.toml` `[database] listen_url` (preferred —
|
||||
co-located with the main `url`) or the `TURNSTONE_DB_LISTEN_URL` env var
|
||||
(config.toml wins when both are set). Defaults to the main DB URL when
|
||||
unset.
|
||||
|
||||
| Setting | Behaviour |
|
||||
|---|---|
|
||||
| unset | Listener uses `TURNSTONE_DB_URL` as-is. Fine when PgBouncer is in **session** mode, or when there's no pooler in front of Postgres. With transaction-mode PgBouncer the listener's `LISTEN` will fail and the dispatcher retries with exponential backoff (1 s → 30 s cap) without ever succeeding. Reactive NOTIFY-driven node discovery is silently lost; the cluster collector's 60 s `_discovery_loop` is the only remaining backstop. |
|
||||
| set to direct-to-PG URL (e.g. `postgresql://…/turnstone`) | Listener bypasses PgBouncer for its one dedicated connection. Reactive discovery latency drops from up-to-60 s to ~500 ms. The rest of the storage layer continues to go through PgBouncer in transaction mode. |
|
||||
|
||||
Set this whenever PgBouncer is in transaction mode (the recommended
|
||||
setting per this doc). The override only adds one long-lived PG
|
||||
connection per console process — sized into the cluster's
|
||||
`max_connections` budget alongside the pool.
|
||||
|
||||
See also: [Docker deployment](docker.md) · [Security](security.md)
|
||||
|
||||
@@ -138,6 +138,9 @@ SSE events are deserialized into typed dataclasses. Use `event.type` to discrimi
|
||||
| `error` | `ErrorEvent` | `message` |
|
||||
| `info` | `InfoEvent` | `message` |
|
||||
| `stream_end` | `StreamEndEvent` | — |
|
||||
| `state_change` | `StateChangeEvent` | `state` ∈ `running`/`thinking`/`attention`/`idle`/`error` |
|
||||
| `in_progress_snapshot` | `InProgressSnapshotEvent` | `content`, `reasoning` (one-shot mid-stream refresh resume) |
|
||||
| `approval_resolved` | `ApprovalResolvedEvent` | `approved`, `feedback` |
|
||||
| `cancelled` | `CancelledEvent` | — |
|
||||
|
||||
**Global events** (from `stream_global_events()`):
|
||||
|
||||
+16
-1
@@ -59,6 +59,21 @@ from ConfigStore. Model names and context windows are now configured per-model
|
||||
in the Models tab. A startup warning is logged if these keys appear in
|
||||
`config.toml`.
|
||||
|
||||
### Reasoning persistence (per-model)
|
||||
|
||||
Two boolean flags on `model_definitions` (migration 052) control how
|
||||
reasoning text round-trips per model:
|
||||
|
||||
| Flag | Default | Effect |
|
||||
|------|---------|--------|
|
||||
| `surface_persisted_reasoning` | `True` | Surface stored reasoning text on `/history` payloads so a page reload re-renders the reasoning bubble. **Storage of reasoning bytes is independent of this flag** — they ride in `provider_data` regardless. |
|
||||
| `replay_reasoning_to_model` | `False` | Send stored reasoning blocks back to the provider on subsequent turns. Capability-gated: only takes effect when the model's `ModelCapabilities.supports_reasoning_replay` is also `True`. Set on canonical OpenAI gpt-5*/o-series and Anthropic Claude entries; unknown / local-server models default to `False` so an operator who flips the flag on a model whose API doesn't understand reasoning replay silently no-ops rather than 400-ing. |
|
||||
|
||||
Edit both via the admin Models tab. See the architecture doc for the
|
||||
provider-side mechanics (Anthropic `thinking`, OpenAI Responses
|
||||
`reasoning` + `include=["reasoning.encrypted_content"]`, synthetic
|
||||
`reasoning_text` for Chat Completions / vLLM / llama.cpp / Gemini-compat).
|
||||
|
||||
### Plan / task agent overrides
|
||||
|
||||
`plan_agent` and `task_agent` sub-sessions resolve independently from the
|
||||
@@ -100,7 +115,7 @@ initialization:
|
||||
| `tools` | timeout, truncation, agent_max_turns, skip_permissions, search, search_threshold, search_max_results |
|
||||
| `server` | workstream_idle_timeout, max_workstreams |
|
||||
| `cluster` | node_fan_out_limit, mcp_max_servers |
|
||||
| `mcp` | config_path, refresh_interval, registry_url |
|
||||
| `mcp` | config_path, registry_url |
|
||||
| `ratelimit` | enabled, requests_per_second, burst, trusted_proxies |
|
||||
| `health` | backend_probe_interval, backend_probe_timeout, circuit_breaker_threshold, circuit_breaker_cooldown |
|
||||
| `judge` | enabled, model, provider, base_url, api_key, confidence_threshold, max_context_ratio, timeout, read_only_tools, output_guard, redact_secrets, cancel_on_approval |
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
---
|
||||
name: import-conversation-history
|
||||
description: Use this skill when the user wants to import or migrate conversation history from another LLM chat or coding tool (e.g. ChatGPT, Claude.ai, Cursor, Copilot Chat, Aider, Gemini, a custom JSON export) into Turnstone. The skill teaches Turnstone's destination contracts — workstream identity, the OpenAI-shaped message rows, tool-call/result pairing, provider-fidelity blobs, attachments, and archive-vs-resumable choice — so the agent can map any source format onto them. Trigger phrases: "import my chats", "migrate this transcript into Turnstone", "bring my Claude.ai history over", "load this export as a workstream".
|
||||
version: 1.0.0
|
||||
---
|
||||
|
||||
# Importing Conversation History into Turnstone
|
||||
|
||||
## Overview
|
||||
|
||||
Source formats vary; the destination does not. Your job is to translate whatever the user hands you (JSON dump, ZIP export, scraped HTML, screenshot OCR, raw transcript) into Turnstone's internal shape: **one workstream row** plus an ordered sequence of **conversation rows** in OpenAI message format. This skill documents the destination so you can write a correct mapper for any source.
|
||||
|
||||
Two questions to settle with the user before writing anything:
|
||||
|
||||
1. **Archive or resumable?** An archive ("saved" workstream — `state="closed"`) is read-only history. A resumable workstream (`state="idle"`) lets the user continue the conversation; this only works cleanly when the source LLM matches a Turnstone-supported provider/model and tool definitions still resolve.
|
||||
2. **One workstream per source thread, or merge?** Default to one-to-one unless the user explicitly asks to merge.
|
||||
|
||||
Default to **archive** when in doubt — resuming a foreign transcript with mismatched tool schemas or stale provider signatures will fail at the next turn.
|
||||
|
||||
## Turnstone Data Model (the destination)
|
||||
|
||||
Two tables carry the conversation:
|
||||
|
||||
### `workstreams` (one row per imported thread)
|
||||
|
||||
| Column | Required | Notes |
|
||||
|---|---|---|
|
||||
| `ws_id` | yes | 32-char lowercase hex. Auto-generate with `secrets.token_hex(16)` if you don't already have one. **First 4 hex chars are the routing bucket** — see "Identity & Routing" below. |
|
||||
| `name` | yes | Short title. Pull from source thread title; fall back to first ~60 chars of first user message. |
|
||||
| `state` | yes | `"closed"` for archive, `"idle"` for resumable. Never set `"running"` on import. |
|
||||
| `kind` | yes | `"interactive"` for normal threads. Do NOT use `"coordinator"` for imports — that's reserved for cluster-spawned coordinator workstreams. |
|
||||
| `parent_ws_id` | no | Leave NULL. Only set if you're importing a coordinator-spawned subtree and re-parenting it; rare. |
|
||||
| `user_id` | yes | Owner. Must exist in `users`; importer must know which Turnstone user owns the imported history. |
|
||||
| `node_id` | yes (multi-node) | Denormalized cache of the node that owns this `ws_id`'s bucket. Single-node deployments can leave it NULL or set it to the only node. |
|
||||
| `alias` | no | Human-typeable short name. Optional; must be unique cluster-wide if set. |
|
||||
| `title` | no | Auto-titled later by the LLM; safe to leave NULL on import. |
|
||||
| `skill_id`, `skill_version` | yes | Default `""` and `0` unless the source thread was scoped to a Turnstone skill. |
|
||||
| `created`, `updated` | yes | ISO8601 strings. Use the source's first/last message timestamps when available. |
|
||||
|
||||
### `conversations` (many rows per thread, ordered by `id`/`timestamp`)
|
||||
|
||||
| Column | Notes |
|
||||
|---|---|
|
||||
| `ws_id` | The workstream this row belongs to. |
|
||||
| `timestamp` | ISO8601 string. Preserve source timestamps; fall back to monotonically increasing values if unknown. **Order is canonical via `id` (autoincrement), not `timestamp`** — but always insert in conversational order so both agree. |
|
||||
| `role` | One of `system`, `user`, `assistant`, `tool`, `developer`. See role mapping below. |
|
||||
| `content` | Text. May be NULL for assistant rows that are *only* tool calls. |
|
||||
| `tool_name` | Set on `role="tool"` rows (the tool whose result this is). NULL otherwise. |
|
||||
| `tool_call_id` | Set on `role="tool"` rows (matches the assistant row's `tool_calls[].id`). NULL otherwise. |
|
||||
| `tool_calls` | JSON-encoded list, on `role="assistant"` rows that issued tool calls. OpenAI shape — see "Tool Calls" below. |
|
||||
| `provider_data` | JSON blob preserving provider-native content blocks (Anthropic `signature`, Gemini `thought_signature`, etc.). Optional; only matters for **resumable** imports against the same provider. Skip for archives. |
|
||||
|
||||
The internal format is **OpenAI-shaped**, even when the source was Anthropic or Gemini. Providers translate at their own API boundary; storage stays uniform.
|
||||
|
||||
## Identity & Routing (`ws_id`)
|
||||
|
||||
- `ws_id` is **32-char lowercase hex** (i.e. `secrets.token_hex(16)`).
|
||||
- The **routing bucket** is `int(ws_id[:4], 16)` — the first 4 hex chars place this workstream on a specific node via the consistent hash ring.
|
||||
- For multi-node imports: either insert through the console's routing proxy (which forwards to the owning node), or generate `ws_id`s and write directly to each node's database in batches grouped by bucket.
|
||||
- For single-node imports: bucket math is irrelevant; any `ws_id` works.
|
||||
- **Do not reuse the source platform's IDs as `ws_id`** unless they happen to be 32-char hex. Generate fresh; if you need the old ID for traceability, store it in `workstream_config` under a key like `import.source_id`.
|
||||
|
||||
## Recommended Import Path
|
||||
|
||||
Three options, in order of preference:
|
||||
|
||||
### 1. Storage protocol (recommended for full history)
|
||||
|
||||
Use `turnstone.core.storage.Storage.save_messages_bulk(rows)`. This is the canonical bulk-insert primitive and bypasses the LLM round-trip entirely.
|
||||
|
||||
```python
|
||||
from turnstone.core.storage import get_storage # construct via the same path the server uses
|
||||
|
||||
storage = get_storage(...) # see turnstone.core.storage.__init__ for the project's wiring
|
||||
|
||||
storage.create_workstream( # or whatever the project's exposed creator is — check turnstone/core/storage/_protocol.py
|
||||
ws_id=ws_id,
|
||||
user_id=user_id,
|
||||
name=name,
|
||||
state="closed",
|
||||
kind="interactive",
|
||||
...
|
||||
)
|
||||
|
||||
storage.save_messages_bulk([
|
||||
{"ws_id": ws_id, "role": "user", "content": "Hello"},
|
||||
{"ws_id": ws_id, "role": "assistant", "content": "Hi! What can I help with?"},
|
||||
{"ws_id": ws_id, "role": "assistant", "content": None,
|
||||
"tool_calls": json.dumps([{"id": "call_1", "type": "function",
|
||||
"function": {"name": "search", "arguments": "{\"q\":\"x\"}"}}])},
|
||||
{"ws_id": ws_id, "role": "tool", "tool_name": "search", "tool_call_id": "call_1",
|
||||
"content": "result text"},
|
||||
# ...
|
||||
])
|
||||
```
|
||||
|
||||
`save_messages_bulk` handles `timestamp` and the workstream's `updated` column internally, so you don't need to compute them per row. **Verify the exact creator signature** by reading `turnstone/core/storage/_protocol.py` — table layout has shifted across migrations and the Storage protocol is the source of truth.
|
||||
|
||||
### 2. SDK `create_workstream(resume_ws=...)` (when the source is already a Turnstone workstream)
|
||||
|
||||
Only useful for *Turnstone → Turnstone* re-parenting. Not relevant for foreign sources.
|
||||
|
||||
### 3. SDK `create_workstream(initial_message=...)` + `send()` per turn (last resort)
|
||||
|
||||
Only fits archives where the source had **no tool calls** and you don't care about preserving assistant turns verbatim. Each `send()` triggers a real LLM round-trip, which is expensive and rewrites assistant content. Don't use this for full history.
|
||||
|
||||
## Role Mapping
|
||||
|
||||
Common source-role conventions and how they map to Turnstone:
|
||||
|
||||
| Source role | Turnstone `role` | Notes |
|
||||
|---|---|---|
|
||||
| `user`, `human` | `user` | Direct map. |
|
||||
| `assistant`, `ai`, `model`, `bot` | `assistant` | Direct map. |
|
||||
| `system` | `system` | Preserve only if it's content the user wrote (custom instructions). Drop boilerplate provider preambles — Turnstone composes its own system message. |
|
||||
| `developer` (OpenAI o-series) | `developer` | Preserve. |
|
||||
| `tool`, `function`, `tool_result` | `tool` | Must carry `tool_name` and `tool_call_id` matching the prior assistant row's `tool_calls[].id`. |
|
||||
| `tool_use` (Anthropic) | `assistant` with `tool_calls` | Anthropic emits tool calls *inside* an assistant message; flatten to OpenAI shape. |
|
||||
| `human_feedback`, `revision` | `user` | Treat as a follow-up user turn. |
|
||||
|
||||
## Tool Calls (the most error-prone part)
|
||||
|
||||
Turnstone stores tool calls in OpenAI's nested-function shape on the assistant row, and matches them with `role="tool"` result rows by `tool_call_id`.
|
||||
|
||||
### Assistant row with tool calls
|
||||
|
||||
```json
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_web",
|
||||
"arguments": "{\"query\":\"turnstone import\"}"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`tool_calls[].function.arguments` is **a JSON-encoded string**, not an object. Source formats commonly get this wrong — Anthropic stores arguments as a parsed object, Gemini as a struct. Always re-serialize to a string.
|
||||
|
||||
### Tool result row
|
||||
|
||||
```json
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_name": "search_web",
|
||||
"tool_call_id": "call_abc123",
|
||||
"content": "..."
|
||||
}
|
||||
```
|
||||
|
||||
Pairing rules:
|
||||
- Every assistant `tool_calls[].id` MUST be followed by exactly one `role="tool"` row with the matching `tool_call_id`, before the next user/assistant turn.
|
||||
- If the source dropped the tool result (cut-off transcript), insert a synthetic `role="tool"` row with `content="[tool result missing in source]"` to keep the chain valid. An assistant row with an unanswered `tool_calls[].id` will break replay and any LLM round-trip.
|
||||
- Multi-tool assistant turns: one `role="tool"` row per call, in any order, all before the next non-tool row.
|
||||
|
||||
### Tool ID generation
|
||||
|
||||
If the source used opaque tool IDs that aren't unique within a thread (some platforms reuse them), regenerate with a stable scheme like `f"call_{i}"` where `i` is a per-thread counter. Update both the assistant and tool rows together.
|
||||
|
||||
## Provider Fidelity (`provider_data`)
|
||||
|
||||
Skip this entirely for **archive** imports.
|
||||
|
||||
For **resumable** imports against the same provider, populate `provider_data` to preserve provider-specific tool-call metadata that the next API round-trip will require:
|
||||
|
||||
- **Anthropic**: `signature` field on thinking blocks; required for round-tripping extended-thinking responses.
|
||||
- **Gemini**: `thought_signature` on tool calls; required for fidelity.
|
||||
- **OpenAI**: typically nothing to preserve.
|
||||
|
||||
The runtime-side dict key is `_provider_content` (a list of provider-native blocks); the persisted column is `provider_data` (the same list, JSON-encoded). If you don't have provider-native blocks from the source — and you usually won't, because a foreign export won't include them — leave `provider_data` NULL. The first new turn will succeed without it, but the previous assistant turn's reasoning won't replay back to the model.
|
||||
|
||||
## Attachments
|
||||
|
||||
If the source thread had image or file attachments:
|
||||
|
||||
- **Size limits**: images ≤ 4 MiB, text documents ≤ 512 KiB. Reject or downsample anything bigger.
|
||||
- **Allowed types**: server validates magic bytes for images and UTF-8-decodes for text. Binary blobs that aren't images won't pass.
|
||||
- **Lifecycle**: pending → reserved → consumed. For imports, the cleanest path is to upload as pending and immediately consume by attaching to the relevant `conversations.id`.
|
||||
|
||||
Two import paths:
|
||||
|
||||
1. **Bulk-insert + post-attach**: insert messages first, get back the assistant/user `conversations.id`, then write `workstream_attachments` rows linking the file to `message_id`.
|
||||
2. **SDK multipart create**: `create_workstream(attachments=[...], initial_message=...)` for the *first* turn only — the server reserves and consumes them onto that turn. Doesn't help for mid-thread attachments.
|
||||
|
||||
For full-history imports with multiple attachments at different turns, path (1) is the only option.
|
||||
|
||||
## Validation Checklist
|
||||
|
||||
Before declaring success, verify:
|
||||
|
||||
- [ ] `ws_id` is 32-char lowercase hex.
|
||||
- [ ] `workstreams` row exists with the right `user_id`, `state`, `kind`.
|
||||
- [ ] Conversation rows are inserted **in order** (autoincrement `id` will reflect insert order).
|
||||
- [ ] Every assistant `tool_calls[].id` has a matching `role="tool"` row with the same `tool_call_id`.
|
||||
- [ ] `tool_calls[].function.arguments` is a JSON-encoded **string**, not a parsed object.
|
||||
- [ ] First message is typically `role="user"` (not `system`) — Turnstone composes its own system prompt at runtime.
|
||||
- [ ] No empty assistant rows (`content=NULL` AND `tool_calls=NULL` is invalid).
|
||||
- [ ] If multi-node: the `ws_id`'s bucket maps to a node that exists; `workstreams.node_id` matches.
|
||||
- [ ] Round-trip test: run `Storage.load_messages(ws_id)` and confirm the reconstructed list matches what you inserted (modulo timestamps).
|
||||
|
||||
## Anti-patterns
|
||||
|
||||
- **Don't import the source provider's system prompt verbatim.** Provider boilerplate ("You are Claude...", "You are ChatGPT...") will conflict with Turnstone's composed system message and confuse the model on resume. Drop it; preserve only user-authored custom instructions.
|
||||
- **Don't preserve foreign tool definitions as Turnstone tools.** If the source had custom tools that don't exist in Turnstone, the assistant rows that called them are still valid history (archive), but the workstream is **not resumable** — mark `state="closed"`.
|
||||
- **Don't fabricate `tool_call_id`s without re-pairing.** Mismatched ids silently break the replay chain on the next turn.
|
||||
- **Don't skip the `tool_name` field on `role="tool"` rows.** Some load paths use it for display and audit; NULL there will render as "unknown tool".
|
||||
- **Don't write through the LLM (`send()` per turn) for full history.** It's expensive, rewrites assistant turns, and rate-limits will bite long imports.
|
||||
|
||||
## Quick Reference
|
||||
|
||||
| Task | Path |
|
||||
|---|---|
|
||||
| Generate ws_id | `secrets.token_hex(16)` |
|
||||
| Bulk insert messages | `Storage.save_messages_bulk(rows)` |
|
||||
| Archive (read-only) | `state="closed"`, skip `provider_data` |
|
||||
| Resumable | `state="idle"`, populate `provider_data` if same provider |
|
||||
| Tool call id | OpenAI shape: `{"id": ..., "type": "function", "function": {"name": ..., "arguments": "<json string>"}}` |
|
||||
| Tool result row | `role="tool"`, `tool_name`, `tool_call_id`, `content` |
|
||||
| Source role → Turnstone role | See "Role Mapping" table |
|
||||
| Per-thread metadata | Store source IDs in `workstream_config` under `import.*` keys |
|
||||
|
||||
## Files to read before writing the importer
|
||||
|
||||
- `turnstone/core/storage/_schema.py` — authoritative table definitions.
|
||||
- `turnstone/core/storage/_protocol.py` — `save_message`, `save_messages_bulk`, `load_messages` signatures.
|
||||
- `turnstone/core/session.py` (around the message-save section) — how the runtime constructs in-memory message dicts; mirror this shape on import to round-trip cleanly.
|
||||
- `turnstone/api/server_schemas.py` — Pydantic shapes for the SDK paths if you go through HTTP.
|
||||
+5
-14
@@ -758,22 +758,18 @@ MCP tools (3):
|
||||
|
||||
### Dynamic tool refresh
|
||||
|
||||
MCP tool lists stay up-to-date without restart through three mechanisms:
|
||||
MCP tool lists stay up-to-date without restart through two mechanisms:
|
||||
|
||||
1. **Push notifications** -- MCP servers that declare `tools.listChanged: true` in
|
||||
their capabilities send `notifications/tools/list_changed` when their tool list
|
||||
changes. `MCPClientManager` registers a `message_handler` on each `ClientSession`
|
||||
that triggers an immediate refresh for that server.
|
||||
|
||||
2. **Periodic timer** -- Servers that do *not* support push notifications are polled
|
||||
on a configurable interval (default 4 hours). The timer is staggered using a
|
||||
launch-time seed (`monotonic_ns ^ pid`) so cluster nodes don't all hit MCP
|
||||
servers simultaneously. Configure via `[mcp] refresh_interval` in `config.toml`
|
||||
or `--mcp-refresh-interval SECONDS` on the CLI. Set to `0` to disable.
|
||||
|
||||
3. **Manual** -- `/mcp refresh` re-fetches tools from all servers immediately.
|
||||
2. **Manual** -- `/mcp refresh` re-fetches tools from all servers immediately.
|
||||
`/mcp refresh <server>` targets a single server. If a server has disconnected,
|
||||
manual refresh attempts reconnection.
|
||||
manual refresh attempts reconnection. The console admin panel exposes the
|
||||
same controls (refresh / reconnect buttons per server) for cluster-wide
|
||||
fan-out.
|
||||
|
||||
When tools change, `MCPClientManager` rebuilds its merged tool list using copy-on-write
|
||||
(new list/dict objects assigned atomically) and notifies all active `ChatSession`
|
||||
@@ -781,11 +777,6 @@ instances via registered listener callbacks. Each session rebuilds its `_tools`,
|
||||
`_task_tools`, `_agent_tools`, and reconstructs its `ToolSearchManager` (if active),
|
||||
preserving the set of previously expanded (discovered) tools.
|
||||
|
||||
```toml
|
||||
[mcp]
|
||||
refresh_interval = 14400 # seconds (default 4h), 0 to disable
|
||||
```
|
||||
|
||||
```
|
||||
/mcp refresh
|
||||
MCP refresh complete:
|
||||
|
||||
+5
-4
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "turnstone"
|
||||
version = "1.5.0"
|
||||
version = "1.5.18"
|
||||
description = "Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."
|
||||
readme = "README.md"
|
||||
license = "BUSL-1.1"
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
dependencies = [
|
||||
"openai>=2.24",
|
||||
"httpx>=0.28",
|
||||
"mcp>=1.6",
|
||||
"mcp>=1.27",
|
||||
"starlette>=0.45",
|
||||
"uvicorn>=0.34",
|
||||
"sse-starlette>=2.0",
|
||||
@@ -35,6 +35,7 @@ dependencies = [
|
||||
"structlog>=24.1",
|
||||
"PyJWT>=2.8",
|
||||
"bcrypt>=4.0",
|
||||
"cryptography>=42",
|
||||
"python-frontmatter>=1.0",
|
||||
]
|
||||
|
||||
@@ -81,9 +82,9 @@ include = [
|
||||
"turnstone/console/static/coordinator/*.js",
|
||||
"turnstone/shared_static/*.css",
|
||||
"turnstone/shared_static/*.js",
|
||||
"turnstone/shared_static/katex-0.16.45/**/*",
|
||||
"turnstone/shared_static/katex-0.16.47/**/*",
|
||||
"turnstone/shared_static/hljs-11.11.1/**/*",
|
||||
"turnstone/shared_static/mermaid-11.14.0/**/*",
|
||||
"turnstone/shared_static/mermaid-11.15.0/**/*",
|
||||
"turnstone/shared_static/hls-1.6.16/**/*",
|
||||
"turnstone/sdk/py.typed",
|
||||
"turnstone/deploy/*.yaml",
|
||||
|
||||
@@ -50,11 +50,14 @@ update_refs() {
|
||||
local old_pattern="$1" # e.g. katex-0.16.38
|
||||
local new_pattern="$2" # e.g. katex-0.16.39
|
||||
|
||||
# Find all files with version references (excludes vendored JS and worktrees)
|
||||
# Find all files with version references. Excludes the old versioned vendor
|
||||
# directory itself (about to be rm -rf'd anyway) so we don't bother rewriting
|
||||
# self-references inside it — but does NOT exclude all of shared_static/,
|
||||
# because shared_static/renderer.js loads the vendored libs and needs the bump.
|
||||
local files
|
||||
files=$(grep -rl --include='*.toml' --include='*.html' --include='*.js' --include='*.md' \
|
||||
files=$(grep -rl --include='*.toml' --include='*.html' --include='*.js' --include='*.md' --include='*.py' \
|
||||
-F "$old_pattern" . \
|
||||
--exclude-dir='.claude' --exclude-dir='node_modules' --exclude-dir='shared_static' \
|
||||
--exclude-dir='.claude' --exclude-dir='node_modules' --exclude-dir="$old_pattern" \
|
||||
2>/dev/null || true)
|
||||
for f in $files; do
|
||||
sed -i "s|${old_pattern}|${new_pattern}|g" "$f"
|
||||
|
||||
Generated
+136
-136
@@ -74,9 +74,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@oxc-project/types": {
|
||||
"version": "0.127.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.127.0.tgz",
|
||||
"integrity": "sha512-aIYXQBo4lCbO4z0R3FHeucQHpF46l2LbMdxRvqvuRuW2OxdnSkcng5B8+K12spgLDj93rtN3+J2Vac/TIO+ciQ==",
|
||||
"version": "0.130.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.130.0.tgz",
|
||||
"integrity": "sha512-ibD2usx9JRu7f5pu2tMKMI4cpA4NgXJQoYRP4pQ7Pxmn1l6k/53qWtQWZayhYy3X4QZkt90Ot+mJEaeXouio6Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"funding": {
|
||||
@@ -84,9 +84,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-android-arm64": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-s70pVGhw4zqGeFnXWvAzJDlvxhlRollagdCCKRgOsgUOH3N1l0LIxf83AtGzmb5SiVM4Hjl5HyarMRfdfj3DaQ==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.0.1.tgz",
|
||||
"integrity": "sha512-fJI3I0r3C3Oj/zdBCpaCmBRZYf07xpaq4yCfDDoSFm+beWNzbIl26puW8RraUdugoJw/95zerNOn6jasAhzSmg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -101,9 +101,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-darwin-arm64": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-arm64/-/binding-darwin-arm64-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-4ksWc9n0mhlZpZ9PMZgTGjeOPRu8MB1Z3Tz0Mo02eWfWCHMW1zN82Qz/pL/rC+yQa+8ZnutMF0JjJe7PjwasYw==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-arm64/-/binding-darwin-arm64-1.0.1.tgz",
|
||||
"integrity": "sha512-cKnAhWEsV7TPcA/5EAteDp6KcJZBQ2G+BqE7zayMMi7kMvwRsbv7WT9aOnn0WNl4SKEIf43vjS31iUPu80nzXg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -118,9 +118,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-darwin-x64": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-x64/-/binding-darwin-x64-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-SUSDOI6WwUVNcWxd02QEBjLdY1VPHvlEkw6T/8nYG322iYWCTxRb1vzk4E+mWWYehTp7ERibq54LSJGjmouOsw==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-x64/-/binding-darwin-x64-1.0.1.tgz",
|
||||
"integrity": "sha512-YKrVwQjIRBPo+5G/u03wGjbdy4q7pyzCe93DK9VJ7zkVmeg8LJ7GbgsiHWdR4xSoe4CAXRD7Bcjgbtr64bkXNg==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -135,9 +135,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-freebsd-x64": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-freebsd-x64/-/binding-freebsd-x64-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-hwnz3nw9dbJ05EDO/PvcjaaewqqDy7Y1rn1UO81l8iIK1GjenME75dl16ajbvSSMfv66WXSRCYKIqfgq2KCfxw==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-freebsd-x64/-/binding-freebsd-x64-1.0.1.tgz",
|
||||
"integrity": "sha512-z/oBsREo46SsFqBwYtFe0kpJeBijAT48O/WXLI4suiCLBkr03RTtTJMCzSdDd2znlh8VJizL09XVkQgk8IZonw==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -152,9 +152,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm-gnueabihf": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm-gnueabihf/-/binding-linux-arm-gnueabihf-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-IS+W7epTcwANmFSQFrS1SivEXHtl1JtuQA9wlxrZTcNi6mx+FDOYrakGevvvTwgj2JvWiK8B29/qD9BELZPyXQ==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm-gnueabihf/-/binding-linux-arm-gnueabihf-1.0.1.tgz",
|
||||
"integrity": "sha512-ik8q7GM11zxvYxFc2PeDcT6TBvhCQMaUxfph/M5l9sKuTs/Sjg3L+Byw0F7w0ZVLBZmx30P+gG0ECzzN+MFcmQ==",
|
||||
"cpu": [
|
||||
"arm"
|
||||
],
|
||||
@@ -169,9 +169,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm64-gnu": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-gnu/-/binding-linux-arm64-gnu-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-e6usGaHKW5BMNZOymS1UcEYGowQMWcgZ71Z17Sl/h2+ZziNJ1a9n3Zvcz6LdRyIW5572wBCTH/Z+bKuZouGk9Q==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-gnu/-/binding-linux-arm64-gnu-1.0.1.tgz",
|
||||
"integrity": "sha512-QoSx2EkyrrdZ6kcyE8stqZ62t0Yra8Fs5ia9lOxJrh6TMQJK7gQKmscdTHf7pOXKREKrVwOtJcQG3qVSfc866A==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -189,9 +189,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-arm64-musl": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-musl/-/binding-linux-arm64-musl-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-b/CgbwAJpmrRLp02RPfhbudf5tZnN9nsPWK82znefso832etkem8H7FSZwxrOI9djcdTP7U6YfNhbRnh7djErg==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-musl/-/binding-linux-arm64-musl-1.0.1.tgz",
|
||||
"integrity": "sha512-uwNwFpwKeNiZawfAWBgg0VIztPTV3ihhh1vV334h9ivnNLorxnQMU6Fz8wG1Zb4Qh9LC1/MkcyT3YlDXG3Rsgg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -209,9 +209,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-ppc64-gnu": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-ppc64-gnu/-/binding-linux-ppc64-gnu-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-4EII1iNGRUN5WwGbF/kOh/EIkoDN9HsupgLQoXfY+D1oyJm7/F4t5PYU5n8SWZgG0FEwakyM8pGgwcBYruGTlA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-ppc64-gnu/-/binding-linux-ppc64-gnu-1.0.1.tgz",
|
||||
"integrity": "sha512-zY1bul7OWr7DFBiJ++wofXvnr8B45ce3QsQUhKrIhXsygAh7bTkwyeM1bi1a2g5C/yC/N8TZyGDEoMfm/l9mpg==",
|
||||
"cpu": [
|
||||
"ppc64"
|
||||
],
|
||||
@@ -229,9 +229,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-s390x-gnu": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-s390x-gnu/-/binding-linux-s390x-gnu-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-AH8oq3XqQo4IibpVXvPeLDI5pzkpYn0WiZAfT05kFzoJ6tQNzwRdDYQ45M8I/gslbodRZwW8uxLhbSBbkv96rA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-s390x-gnu/-/binding-linux-s390x-gnu-1.0.1.tgz",
|
||||
"integrity": "sha512-0frlsT/f4Ft6I7SMESTKnF3cZsdicQn1dCMkF/jT9wDLE+gGoiQfv1nmT9e+s7s/fekvvy6tZM2jHvI2tkbJDQ==",
|
||||
"cpu": [
|
||||
"s390x"
|
||||
],
|
||||
@@ -249,9 +249,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-x64-gnu": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-gnu/-/binding-linux-x64-gnu-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-cLnjV3xfo7KslbU41Z7z8BH/E1y5mzUYzAqih1d1MDaIGZRCMqTijqLv76/P7fyHuvUcfGsIpqCdddbxLLK9rA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-gnu/-/binding-linux-x64-gnu-1.0.1.tgz",
|
||||
"integrity": "sha512-XABVmGp9Tg0WspTVvwduTc4fpqy6JnAUrSQe6OuyqD/03nI7r0O9OWUkMIwFrjKAIqolvqoA4ZrJppgwE0Gxmw==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -269,9 +269,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-linux-x64-musl": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-musl/-/binding-linux-x64-musl-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-0phclDw1spsL7dUB37sIARuis2tAgomCJXAHZlpt8PXZ4Ba0dRP1e+66lsRqrfhISeN9bEGNjQs+T/Fbd7oYGw==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-musl/-/binding-linux-x64-musl-1.0.1.tgz",
|
||||
"integrity": "sha512-bV4fzswuzVcKD90o/VM6QqKxnxlDq0g2BISDLNVmxrnhpv1DDbyPhCIjYfvzYLV+MvkKKnQt2Q6AO86SEBULUQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -289,9 +289,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-openharmony-arm64": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-openharmony-arm64/-/binding-openharmony-arm64-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-0ag/hEgXOwgw4t8QyQvUCxvEg+V0KBcA6YuOx9g0r02MprutRF5dyljgm3EmR02O292UX7UeS6HzWHAl6KgyhA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-openharmony-arm64/-/binding-openharmony-arm64-1.0.1.tgz",
|
||||
"integrity": "sha512-/Mh0Zhq3OP7fVs0kcQHZP6lZEthMGTaSf8UBQYSFEZDWGXXlEC+nJ6EqenaK2t4LBXMe3A+K/G2BVXXdtOr4PQ==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -306,9 +306,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-wasm32-wasi": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-wasm32-wasi/-/binding-wasm32-wasi-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-LEXei6vo0E5wTGwpkJ4KoT3OZJRnglwldt5ziLzOlc6qqb55z4tWNq2A+PFqCJuvWWdP53CVhG1Z9NtToDPJrA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-wasm32-wasi/-/binding-wasm32-wasi-1.0.1.tgz",
|
||||
"integrity": "sha512-+1xc9X45l8ufsBAm6Gjvx2qDRIY9lTVt0cgWNcJ+1gdhXvkbxePA60yRTwSTuXL09CMhyJmjpV7E3NoyxbqFQQ==",
|
||||
"cpu": [
|
||||
"wasm32"
|
||||
],
|
||||
@@ -325,9 +325,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-win32-arm64-msvc": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-arm64-msvc/-/binding-win32-arm64-msvc-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-gUmyzBl3SPMa6hrqFUth9sVfcLBlYsbMzBx5PlexMroZStgzGqlZ26pYG89rBb45Mnia+oil6YAIFeEWGWhoZA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-arm64-msvc/-/binding-win32-arm64-msvc-1.0.1.tgz",
|
||||
"integrity": "sha512-1D+UqZdfnuR+Jy1GgMJwi85bD40H21uNmOPRWQhw4oRSuolZ/B5rixZ45DK2KXOTCvmVCecauWgEhbw8bI7tOw==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
@@ -342,9 +342,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/binding-win32-x64-msvc": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-x64-msvc/-/binding-win32-x64-msvc-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-3hkiolcUAvPB9FLb3UZdfjVVNWherN1f/skkGWJP/fgSQhYUZpSIRr0/I8ZK9TkF3F7kxvJAk0+IcKvPHk9qQg==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/binding-win32-x64-msvc/-/binding-win32-x64-msvc-1.0.1.tgz",
|
||||
"integrity": "sha512-INAycaWuhlOK3wk4mRHGsdgwYWmd9cChdPdE9bwWmy6rn9VqVNYNFGhOdXrofXUxwHIncSiPNb8tNm8knDVIeQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
@@ -359,9 +359,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@rolldown/pluginutils": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-n8iosDOt6Ig1UhJ2AYqoIhHWh/isz0xpicHTzpKBeotdVsTEcxsSA/i3EVM7gQAj0rU27OLAxCjzlj15IWY7bg==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.1.tgz",
|
||||
"integrity": "sha512-2j9bGt5Jh8hj+vPtgzPtl72j0yRxHAyumoo6TNfAjsLB04UtpSvPbPcDcBMxz7n+9CYB0c1GxQFxYRg2jimqGw==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
@@ -373,9 +373,9 @@
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@tybys/wasm-util": {
|
||||
"version": "0.10.1",
|
||||
"resolved": "https://registry.npmjs.org/@tybys/wasm-util/-/wasm-util-0.10.1.tgz",
|
||||
"integrity": "sha512-9tTaPJLSiejZKx+Bmog4uSubteqTvFrVrURwkmHixBo0G4seD0zUxp98E1DzUBJxLQ3NPwXrGKDiVjwx/DpPsg==",
|
||||
"version": "0.10.2",
|
||||
"resolved": "https://registry.npmjs.org/@tybys/wasm-util/-/wasm-util-0.10.2.tgz",
|
||||
"integrity": "sha512-RoBvJ2X0wuKlWFIjrwffGw1IqZHKQqzIchKaadZZfnNpsAYp2mM0h36JtPCjNDAHGgYez/15uMBpfGwchhiMgg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
@@ -402,23 +402,23 @@
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@types/estree": {
|
||||
"version": "1.0.8",
|
||||
"resolved": "https://registry.npmjs.org/@types/estree/-/estree-1.0.8.tgz",
|
||||
"integrity": "sha512-dWHzHa2WqEXI/O1E9OjrocMTKJl2mSrEolh1Iomrv6U+JuNwaHXsXx9bLu5gG7BUWFIN0skIQJQ/L1rIex4X6w==",
|
||||
"version": "1.0.9",
|
||||
"resolved": "https://registry.npmjs.org/@types/estree/-/estree-1.0.9.tgz",
|
||||
"integrity": "sha512-GhdPgy1el4/ImP05X05Uw4cw2/M93BCUmnEvWZNStlCzEKME4Fkk+YpoA5OiHNQmoS7Cafb8Xa3Pya8m1Qrzeg==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@vitest/expect": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.5.tgz",
|
||||
"integrity": "sha512-PWBaRY5JoKuRnHlUHfpV/KohFylaDZTupcXN1H9vYryNLOnitSw60Mw9IAE2r67NbwwzBw/Cc/8q9BK3kIX8Kw==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.6.tgz",
|
||||
"integrity": "sha512-7EHDquPthALSV0jhhjgEW8FXaviMx7rSqu8W6oqCoAuOhKov814P99QDV1pxMA3QPv21YudvJngIhjrNI4opLg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@standard-schema/spec": "^1.1.0",
|
||||
"@types/chai": "^5.2.2",
|
||||
"@vitest/spy": "4.1.5",
|
||||
"@vitest/utils": "4.1.5",
|
||||
"@vitest/spy": "4.1.6",
|
||||
"@vitest/utils": "4.1.6",
|
||||
"chai": "^6.2.2",
|
||||
"tinyrainbow": "^3.1.0"
|
||||
},
|
||||
@@ -427,13 +427,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/mocker": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.5.tgz",
|
||||
"integrity": "sha512-/x2EmFC4mT4NNzqvC3fmesuV97w5FC903KPmey4gsnJiMQ3Be1IlDKVaDaG8iqaLFHqJ2FVEkxZk5VmeLjIItw==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.6.tgz",
|
||||
"integrity": "sha512-MCFc63czMjEInOlcY2cpQCvCN+KgbAn+60xu9cMgP4sKaLC5JNAKw7JH8QdAnoAC88hW1IiSNZ+GgVXlN1UcMQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/spy": "4.1.5",
|
||||
"@vitest/spy": "4.1.6",
|
||||
"estree-walker": "^3.0.3",
|
||||
"magic-string": "^0.30.21"
|
||||
},
|
||||
@@ -454,9 +454,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/pretty-format": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.5.tgz",
|
||||
"integrity": "sha512-7I3q6l5qr03dVfMX2wCo9FxwSJbPdwKjy2uu/YPpU3wfHvIL4QHwVRp57OfGrDFeUJ8/8QdfBKIV12FTtLn00g==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.6.tgz",
|
||||
"integrity": "sha512-h5SxD/IzNhZYnrSZRsUZQIC+vD0GY8cUvq0iwsmkFKixRCKLLWqCXa/FIQ4S1R+sI+PGoojkHsdNrbZiM9Qpgw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
@@ -467,13 +467,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/runner": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.5.tgz",
|
||||
"integrity": "sha512-2D+o7Pr82IEO46YPpoA/YU0neeyr6FTerQb5Ro7BUnBuv6NQtT/kmVnczngiMEBhzgqz2UZYl5gArejsyERDSQ==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.6.tgz",
|
||||
"integrity": "sha512-nOPCmn2+yD0ZNmKdsXGv/UxMMWbMuKeD6GyYncNwdkYDxpQvrPSKYj2rWuDjC2Y4b6w6hjip5dBKFzEUuZe3vA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/utils": "4.1.5",
|
||||
"@vitest/utils": "4.1.6",
|
||||
"pathe": "^2.0.3"
|
||||
},
|
||||
"funding": {
|
||||
@@ -481,14 +481,14 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/snapshot": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.5.tgz",
|
||||
"integrity": "sha512-zypXEt4KH/XgKGPUz4eC2AvErYx0My5hfL8oDb1HzGFpEk1P62bxSohdyOmvz+d9UJwanI68MKwr2EquOaOgMQ==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.6.tgz",
|
||||
"integrity": "sha512-YhsdE6xAVfTDmzjxL2ZDUvjj+ZsgyOKe+TdQzqkD72wIOmHka8NuGQ6NpTNZv9D2Z63fbwWKJPeVpEw4EQgYxw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/pretty-format": "4.1.5",
|
||||
"@vitest/utils": "4.1.5",
|
||||
"@vitest/pretty-format": "4.1.6",
|
||||
"@vitest/utils": "4.1.6",
|
||||
"magic-string": "^0.30.21",
|
||||
"pathe": "^2.0.3"
|
||||
},
|
||||
@@ -497,9 +497,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/spy": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.5.tgz",
|
||||
"integrity": "sha512-2lNOsh6+R2Idnf1TCZqSwYlKN2E/iDlD8sgU59kYVl+OMDmvldO1VDk39smRfpUNwYpNRVn3w4YfuC7KfbBnkQ==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.6.tgz",
|
||||
"integrity": "sha512-JFKxMx6udhwKh/Ldo270e17QX710vgunMkuPAvXjHSvC6oqLWAHhVhjg/I71q0u0CBSErIODV1Kjv0FQNSWjdg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"funding": {
|
||||
@@ -507,13 +507,13 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@vitest/utils": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.5.tgz",
|
||||
"integrity": "sha512-76wdkrmfXfqGjueGgnb45ITPyUi1ycZ4IHgC2bhPDUfWHklY/q3MdLOAB+TF1e6xfl8NxNY0ZYaPCFNWSsw3Ug==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.6.tgz",
|
||||
"integrity": "sha512-FxIY+U81R3LGKCxaHHFRQ5+g6/iRgGLmeHWdp2Amj4ljQRrEIWHmZyDfDYBRZlpyqA7qKxtS9DD1dhk8RnRIVQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/pretty-format": "4.1.5",
|
||||
"@vitest/pretty-format": "4.1.6",
|
||||
"convert-source-map": "^2.0.0",
|
||||
"tinyrainbow": "^3.1.0"
|
||||
},
|
||||
@@ -902,9 +902,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/nanoid": {
|
||||
"version": "3.3.11",
|
||||
"resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.11.tgz",
|
||||
"integrity": "sha512-N8SpfPUnUp1bK+PMYW8qSWdl9U+wwNWI4QKxOYDy9JAro3WMX7p2OeVRF9v+347pnakNevPmiHhNmZ2HbFA76w==",
|
||||
"version": "3.3.12",
|
||||
"resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.12.tgz",
|
||||
"integrity": "sha512-ZB9RH/39qpq5Vu6Y+NmUaFhQR6pp+M2Xt76XBnEwDaGcVAqhlvxrl3B2bKS5D3NH3QR76v3aSrKaF/Kiy7lEtQ==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
@@ -959,9 +959,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/postcss": {
|
||||
"version": "8.5.12",
|
||||
"resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.12.tgz",
|
||||
"integrity": "sha512-W62t/Se6rA0Az3DfCL0AqJwXuKwBeYg6nOaIgzP+xZ7N5BFCI7DYi1qs6ygUYT6rvfi6t9k65UMLJC+PHZpDAA==",
|
||||
"version": "8.5.14",
|
||||
"resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.14.tgz",
|
||||
"integrity": "sha512-SoSL4+OSEtR99LHFZQiJLkT59C5B1amGO1NzTwj7TT1qCUgUO6hxOvzkOYxD+vMrXBM3XJIKzokoERdqQq/Zmg==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
@@ -988,14 +988,14 @@
|
||||
}
|
||||
},
|
||||
"node_modules/rolldown": {
|
||||
"version": "1.0.0-rc.17",
|
||||
"resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.0-rc.17.tgz",
|
||||
"integrity": "sha512-ZrT53oAKrtA4+YtBWPQbtPOxIbVDbxT0orcYERKd63VJTF13zPcgXTvD4843L8pcsI7M6MErt8QtON6lrB9tyA==",
|
||||
"version": "1.0.1",
|
||||
"resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.1.tgz",
|
||||
"integrity": "sha512-X0KQHljNnEkWNqqiz9zJrGunh1B0HgOxLXvnFpCOcadzcy5qohZ3tqMEUg00vncoRovXuK3ZqCT9KnnKzoInFQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@oxc-project/types": "=0.127.0",
|
||||
"@rolldown/pluginutils": "1.0.0-rc.17"
|
||||
"@oxc-project/types": "=0.130.0",
|
||||
"@rolldown/pluginutils": "^1.0.0"
|
||||
},
|
||||
"bin": {
|
||||
"rolldown": "bin/cli.mjs"
|
||||
@@ -1004,21 +1004,21 @@
|
||||
"node": "^20.19.0 || >=22.12.0"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
"@rolldown/binding-android-arm64": "1.0.0-rc.17",
|
||||
"@rolldown/binding-darwin-arm64": "1.0.0-rc.17",
|
||||
"@rolldown/binding-darwin-x64": "1.0.0-rc.17",
|
||||
"@rolldown/binding-freebsd-x64": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-arm-gnueabihf": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-arm64-gnu": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-arm64-musl": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-ppc64-gnu": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-s390x-gnu": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-x64-gnu": "1.0.0-rc.17",
|
||||
"@rolldown/binding-linux-x64-musl": "1.0.0-rc.17",
|
||||
"@rolldown/binding-openharmony-arm64": "1.0.0-rc.17",
|
||||
"@rolldown/binding-wasm32-wasi": "1.0.0-rc.17",
|
||||
"@rolldown/binding-win32-arm64-msvc": "1.0.0-rc.17",
|
||||
"@rolldown/binding-win32-x64-msvc": "1.0.0-rc.17"
|
||||
"@rolldown/binding-android-arm64": "1.0.1",
|
||||
"@rolldown/binding-darwin-arm64": "1.0.1",
|
||||
"@rolldown/binding-darwin-x64": "1.0.1",
|
||||
"@rolldown/binding-freebsd-x64": "1.0.1",
|
||||
"@rolldown/binding-linux-arm-gnueabihf": "1.0.1",
|
||||
"@rolldown/binding-linux-arm64-gnu": "1.0.1",
|
||||
"@rolldown/binding-linux-arm64-musl": "1.0.1",
|
||||
"@rolldown/binding-linux-ppc64-gnu": "1.0.1",
|
||||
"@rolldown/binding-linux-s390x-gnu": "1.0.1",
|
||||
"@rolldown/binding-linux-x64-gnu": "1.0.1",
|
||||
"@rolldown/binding-linux-x64-musl": "1.0.1",
|
||||
"@rolldown/binding-openharmony-arm64": "1.0.1",
|
||||
"@rolldown/binding-wasm32-wasi": "1.0.1",
|
||||
"@rolldown/binding-win32-arm64-msvc": "1.0.1",
|
||||
"@rolldown/binding-win32-x64-msvc": "1.0.1"
|
||||
}
|
||||
},
|
||||
"node_modules/siginfo": {
|
||||
@@ -1060,9 +1060,9 @@
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/tinyexec": {
|
||||
"version": "1.1.1",
|
||||
"resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-1.1.1.tgz",
|
||||
"integrity": "sha512-VKS/ZaQhhkKFMANmAOhhXVoIfBXblQxGX1myCQ2faQrfmobMftXeJPcZGp0gS07ocvGJWDLZGyOZDadDBqYIJg==",
|
||||
"version": "1.1.2",
|
||||
"resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-1.1.2.tgz",
|
||||
"integrity": "sha512-dAqSqE/RabpBKI8+h26GfLq6Vb3JVXs30XYQjdMjaj/c2tS8IYYMbIzP599KtRj7c57/wYApb3QjgRgXmrCukA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
@@ -1119,16 +1119,16 @@
|
||||
}
|
||||
},
|
||||
"node_modules/vite": {
|
||||
"version": "8.0.10",
|
||||
"resolved": "https://registry.npmjs.org/vite/-/vite-8.0.10.tgz",
|
||||
"integrity": "sha512-rZuUu9j6J5uotLDs+cAA4O5H4K1SfPliUlQwqa6YEwSrWDZzP4rhm00oJR5snMewjxF5V/K3D4kctsUTsIU9Mw==",
|
||||
"version": "8.0.13",
|
||||
"resolved": "https://registry.npmjs.org/vite/-/vite-8.0.13.tgz",
|
||||
"integrity": "sha512-MFtjBYgzmSxmgA4RAfjIyXWpGe1oALnjgUTzzV7QLx/TKxCzjtMH6Fd9/eVK+5Fg1qNoz5VAwsmMs/NofrmJvw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"lightningcss": "^1.32.0",
|
||||
"picomatch": "^4.0.4",
|
||||
"postcss": "^8.5.10",
|
||||
"rolldown": "1.0.0-rc.17",
|
||||
"postcss": "^8.5.14",
|
||||
"rolldown": "1.0.1",
|
||||
"tinyglobby": "^0.2.16"
|
||||
},
|
||||
"bin": {
|
||||
@@ -1145,7 +1145,7 @@
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@types/node": "^20.19.0 || >=22.12.0",
|
||||
"@vitejs/devtools": "^0.1.0",
|
||||
"@vitejs/devtools": "^0.1.18",
|
||||
"esbuild": "^0.27.0 || ^0.28.0",
|
||||
"jiti": ">=1.21.0",
|
||||
"less": "^4.0.0",
|
||||
@@ -1197,19 +1197,19 @@
|
||||
}
|
||||
},
|
||||
"node_modules/vitest": {
|
||||
"version": "4.1.5",
|
||||
"resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.5.tgz",
|
||||
"integrity": "sha512-9Xx1v3/ih3m9hN+SbfkUyy0JAs72ap3r7joc87XL6jwF0jGg6mFBvQ1SrwaX+h8BlkX6Hz9shdd1uo6AF+ZGpg==",
|
||||
"version": "4.1.6",
|
||||
"resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.6.tgz",
|
||||
"integrity": "sha512-6lvjbS3p9b4CrdCmguzbh2/4uoXhGE2q71R4OX5sqF9R1bo9Xd6fGrMAfvp5wnCzlBnFVdCOp6onuTQVbo8iUQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@vitest/expect": "4.1.5",
|
||||
"@vitest/mocker": "4.1.5",
|
||||
"@vitest/pretty-format": "4.1.5",
|
||||
"@vitest/runner": "4.1.5",
|
||||
"@vitest/snapshot": "4.1.5",
|
||||
"@vitest/spy": "4.1.5",
|
||||
"@vitest/utils": "4.1.5",
|
||||
"@vitest/expect": "4.1.6",
|
||||
"@vitest/mocker": "4.1.6",
|
||||
"@vitest/pretty-format": "4.1.6",
|
||||
"@vitest/runner": "4.1.6",
|
||||
"@vitest/snapshot": "4.1.6",
|
||||
"@vitest/spy": "4.1.6",
|
||||
"@vitest/utils": "4.1.6",
|
||||
"es-module-lexer": "^2.0.0",
|
||||
"expect-type": "^1.3.0",
|
||||
"magic-string": "^0.30.21",
|
||||
@@ -1237,12 +1237,12 @@
|
||||
"@edge-runtime/vm": "*",
|
||||
"@opentelemetry/api": "^1.9.0",
|
||||
"@types/node": "^20.0.0 || ^22.0.0 || >=24.0.0",
|
||||
"@vitest/browser-playwright": "4.1.5",
|
||||
"@vitest/browser-preview": "4.1.5",
|
||||
"@vitest/browser-webdriverio": "4.1.5",
|
||||
"@vitest/coverage-istanbul": "4.1.5",
|
||||
"@vitest/coverage-v8": "4.1.5",
|
||||
"@vitest/ui": "4.1.5",
|
||||
"@vitest/browser-playwright": "4.1.6",
|
||||
"@vitest/browser-preview": "4.1.6",
|
||||
"@vitest/browser-webdriverio": "4.1.6",
|
||||
"@vitest/coverage-istanbul": "4.1.6",
|
||||
"@vitest/coverage-v8": "4.1.6",
|
||||
"@vitest/ui": "4.1.6",
|
||||
"happy-dom": "*",
|
||||
"jsdom": "*",
|
||||
"vite": "^6.0.0 || ^7.0.0 || ^8.0.0"
|
||||
|
||||
@@ -13,6 +13,20 @@ export interface ConnectedEvent {
|
||||
|
||||
export interface HistoryEvent {
|
||||
type: "history";
|
||||
/**
|
||||
* Per-message dicts the frontend consumes directly. Common optional keys:
|
||||
* - `role`: "user" | "assistant" | "tool"
|
||||
* - `content`: string or list (image/document parts)
|
||||
* - `tool_calls`: assistant turns — list of `{id, name, arguments, verdict?, output_assessment?}`
|
||||
* - `tool_call_id`: tool turns — id of the originating call
|
||||
* - `reminders`: metacognitive nudge bubbles (user/tool channels)
|
||||
* - `advisories`: extracted `UserInterjection` payloads on tool turns
|
||||
* - `reasoning`: concatenated reasoning text for assistant turns whose
|
||||
* `provider_data` carried reasoning-bearing blocks (Anthropic
|
||||
* `thinking`, OpenAI Responses `reasoning`, or synthetic
|
||||
* `reasoning_text` from path-3 servers). Present only when the
|
||||
* active model's `surface_persisted_reasoning` flag is true.
|
||||
*/
|
||||
messages: Array<Record<string, unknown>>;
|
||||
}
|
||||
|
||||
@@ -38,6 +52,14 @@ export interface StreamEndEvent {
|
||||
type: "stream_end";
|
||||
}
|
||||
|
||||
/** One-shot replay of the in-progress turn's content + reasoning emitted
|
||||
* by the events SSE handler when a fresh subscriber connects mid-stream. */
|
||||
export interface InProgressSnapshotEvent {
|
||||
type: "in_progress_snapshot";
|
||||
content: string;
|
||||
reasoning: string;
|
||||
}
|
||||
|
||||
export interface StateChangeEvent {
|
||||
type: "state_change";
|
||||
state: "idle" | "thinking" | "running" | "attention" | "error";
|
||||
@@ -162,6 +184,7 @@ export type ServerEvent =
|
||||
| ContentEvent
|
||||
| ReasoningEvent
|
||||
| StreamEndEvent
|
||||
| InProgressSnapshotEvent
|
||||
| StateChangeEvent
|
||||
| ToolInfoEvent
|
||||
| ApproveRequestEvent
|
||||
@@ -261,6 +284,12 @@ export function isStreamEndEvent(e: ServerEvent): e is StreamEndEvent {
|
||||
return e.type === "stream_end";
|
||||
}
|
||||
|
||||
export function isInProgressSnapshotEvent(
|
||||
e: ServerEvent,
|
||||
): e is InProgressSnapshotEvent {
|
||||
return e.type === "in_progress_snapshot";
|
||||
}
|
||||
|
||||
export function isStateChangeEvent(e: ServerEvent): e is StateChangeEvent {
|
||||
return e.type === "state_change";
|
||||
}
|
||||
|
||||
@@ -35,11 +35,10 @@ def _seed_children(
|
||||
The production path populates the registry via the cluster-event
|
||||
fan-out thread observing ``ws_created`` events. These tests just
|
||||
need a known-children set for the endpoint handlers to iterate —
|
||||
inject directly under ``_children_lock`` rather than spinning up
|
||||
the collector + fan-out plumbing.
|
||||
inject directly via the registry's bulk-merge surface rather than
|
||||
spinning up the collector + fan-out plumbing.
|
||||
"""
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked(coord_ws_id, child_ws_ids)
|
||||
adapter._registry.merge_children(coord_ws_id, child_ws_ids)
|
||||
|
||||
|
||||
class _AuthMiddleware(BaseHTTPMiddleware):
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Shared test helpers — kept out of conftest.py since these are factories,
|
||||
not fixtures, and several test files want to import them directly."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def make_chat_session(**overrides: Any) -> Any:
|
||||
"""Build a minimal ``ChatSession`` with sane test defaults.
|
||||
|
||||
Caller passes any constructor arg as a kwarg to override the default —
|
||||
e.g. ``make_chat_session(memory_config=MemoryConfig(fetch_limit=5))``.
|
||||
"""
|
||||
from turnstone.core.session import ChatSession
|
||||
|
||||
defaults: dict[str, Any] = {
|
||||
"client": MagicMock(),
|
||||
"model": "test-model",
|
||||
"ui": MagicMock(),
|
||||
"instructions": None,
|
||||
"temperature": 0.5,
|
||||
"max_tokens": 4096,
|
||||
"tool_timeout": 30,
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return ChatSession(**defaults)
|
||||
|
||||
|
||||
def patch_session_storage(
|
||||
monkeypatch: Any,
|
||||
*,
|
||||
active: bool = True,
|
||||
raise_on_is_active: bool = False,
|
||||
) -> list[str]:
|
||||
"""Patch ``session.get_storage`` to a stub whose ``is_watch_active``
|
||||
returns *active* (or raises if *raise_on_is_active*). Returns the
|
||||
list of ``watch_id``s the predicate was called with.
|
||||
"""
|
||||
from turnstone.core import session as session_mod
|
||||
|
||||
calls: list[str] = []
|
||||
|
||||
class _Stub:
|
||||
def is_watch_active(self, watch_id: str) -> bool:
|
||||
calls.append(watch_id)
|
||||
if raise_on_is_active:
|
||||
raise RuntimeError("storage down")
|
||||
return active
|
||||
|
||||
monkeypatch.setattr(session_mod, "get_storage", lambda: _Stub())
|
||||
return calls
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Shared session-test helpers.
|
||||
|
||||
Two reasoning-test modules (``test_session_replay_reasoning.py`` and
|
||||
``test_session_synth_reasoning_block.py``) need the same minimal
|
||||
``ChatSession`` factory + a ``SessionUIBase`` no-op subclass. Hoisting
|
||||
keeps a future third caller from drifting on the defaults — the third
|
||||
existing ``_make_session`` (``test_model_registry.py``) deliberately
|
||||
takes a different signature (registry / model_alias / reasoning_effort
|
||||
+ ``_FakeUI``) and is NOT a candidate for sharing this helper.
|
||||
|
||||
Module is named with a leading underscore so pytest doesn't try to
|
||||
collect it as a test file — it's an importable utility, not a test.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.session_ui_base import SessionUIBase
|
||||
|
||||
|
||||
class NullUI(SessionUIBase):
|
||||
"""Bare-bones UI satisfying the SessionUIBase contract for tests
|
||||
that don't care about UI side effects."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
|
||||
def make_session(**kwargs: Any) -> ChatSession:
|
||||
"""Build a ChatSession with minimal defaults; tests override
|
||||
individual fields via kwargs."""
|
||||
defaults: dict[str, Any] = {
|
||||
"client": MagicMock(),
|
||||
"model": "test-model",
|
||||
"ui": NullUI(),
|
||||
"instructions": None,
|
||||
"temperature": 0.5,
|
||||
"max_tokens": 4096,
|
||||
"tool_timeout": 30,
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return ChatSession(**defaults)
|
||||
@@ -1,10 +1,79 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.mcp_client import MCPClientManager, StaticServerState
|
||||
from turnstone.core.mcp_crypto import MCPTokenCipher
|
||||
from turnstone.core.oidc import OIDCConfig
|
||||
|
||||
|
||||
def make_mcp_token_cipher() -> MCPTokenCipher:
|
||||
"""Build a single-key MCP token cipher for tests.
|
||||
|
||||
Used by test files that need to exercise ``MCPTokenStore`` round-
|
||||
trips without the lifespan-side configuration loader; centralised
|
||||
here so the key/material defaults stay aligned across files.
|
||||
"""
|
||||
import base64
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from turnstone.core.mcp_crypto import MCPTokenCipher, MCPTokenCipherConfig
|
||||
|
||||
raw = base64.urlsafe_b64decode(Fernet.generate_key())
|
||||
return MCPTokenCipher(MCPTokenCipherConfig(keys=(raw,)))
|
||||
|
||||
|
||||
def _seed_static_state(mgr: MCPClientManager, name: str, **overrides: Any) -> StaticServerState:
|
||||
"""Get-or-create a ``StaticServerState`` on ``mgr`` and apply ``overrides``.
|
||||
|
||||
Shared across MCP test files so the helper stays in one place. Imported
|
||||
where needed; ``StaticServerState`` is constructed lazily so non-MCP
|
||||
tests don't pay the import cost.
|
||||
"""
|
||||
from turnstone.core.mcp_client import StaticServerState
|
||||
|
||||
state = mgr._static_servers.get(name)
|
||||
if state is None:
|
||||
state = StaticServerState(name=name)
|
||||
mgr._static_servers[name] = state
|
||||
for k, v in overrides.items():
|
||||
setattr(state, k, v)
|
||||
return state
|
||||
|
||||
|
||||
def make_oidc_test_config(**overrides: Any) -> OIDCConfig:
|
||||
"""Build a test ``OIDCConfig`` with sensible defaults.
|
||||
|
||||
Shared between ``test_oidc.py`` and ``test_oidc_handlers.py`` so the
|
||||
defaults (including the now-required ``redirect_base``) stay aligned.
|
||||
"""
|
||||
from turnstone.core.oidc import OIDCConfig
|
||||
|
||||
defaults: dict[str, Any] = {
|
||||
"enabled": True,
|
||||
"issuer": "https://idp.example.com",
|
||||
"client_id": "my-client",
|
||||
"client_secret": "my-secret",
|
||||
"scopes": "openid email profile",
|
||||
"provider_name": "TestIDP",
|
||||
"role_claim": "",
|
||||
"role_map": {},
|
||||
"password_enabled": True,
|
||||
"redirect_base": "https://app.example.com",
|
||||
"authorization_endpoint": "https://idp.example.com/authorize",
|
||||
"token_endpoint": "https://idp.example.com/token",
|
||||
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
||||
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return OIDCConfig(**defaults)
|
||||
|
||||
|
||||
def pytest_addoption(parser: pytest.Parser) -> None:
|
||||
parser.addoption(
|
||||
|
||||
@@ -0,0 +1,391 @@
|
||||
"""Spike 1 — validate MCP SDK behavior for the per-(user, server) session pool.
|
||||
|
||||
Three scenarios:
|
||||
|
||||
1. N=20 concurrent ClientSession instances to the same URL.
|
||||
Verifies: no FD blow-up, no shared transport state, each session's
|
||||
tools/list returns independently.
|
||||
|
||||
2. Two concurrent tools/call on a shared ClientSession with interleaving
|
||||
payloads. Verifies: request_id demux works under contention.
|
||||
|
||||
3. Per-session Authorization header isolation. Verifies: different Bearer
|
||||
tokens per ClientSession reach the server with the expected
|
||||
Authorization header — i.e. httpx connection pooling does not cross
|
||||
headers between sessions.
|
||||
|
||||
Run: uv run python tests/spike_sdk_concurrency.py
|
||||
|
||||
Outcome gates Phase 5's pool architecture; if any scenario fails, fall
|
||||
back to per-call header injection (Alternative F in the OAuth-MCP RFC).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import uvicorn
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
# Reduce uvicorn / mcp log noise so spike output is readable.
|
||||
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("mcp").setLevel(logging.WARNING)
|
||||
|
||||
# Records (auth_header, tool_name) per request — populated by the
|
||||
# AuthHeaderRecorder middleware below. Indexed by call sequence.
|
||||
SERVER_OBSERVATIONS: list[tuple[str | None, str | None]] = []
|
||||
# Tool-call payloads observed (for request_id demux verification).
|
||||
TOOL_CALL_PAYLOADS: list[str] = []
|
||||
|
||||
|
||||
class AuthHeaderRecorder(BaseHTTPMiddleware):
|
||||
"""Records the Authorization header on every request the server sees."""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
auth = request.headers.get("authorization")
|
||||
# We only record the auth header here; tool name comes from the
|
||||
# body payload which we can't read non-destructively. The tool
|
||||
# handler logs the payload it received.
|
||||
SERVER_OBSERVATIONS.append((auth, None))
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def find_free_port() -> int:
|
||||
"""Bind to port 0, return the assigned port."""
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def build_server(port: int) -> uvicorn.Server:
|
||||
"""Create a minimal FastMCP server with one echo tool."""
|
||||
mcp = FastMCP(name="spike-target", streamable_http_path="/mcp")
|
||||
|
||||
@mcp.tool()
|
||||
async def echo(payload: str) -> str:
|
||||
"""Echo the payload back. Records the payload server-side."""
|
||||
TOOL_CALL_PAYLOADS.append(payload)
|
||||
# Add a small await so two concurrent calls can interleave
|
||||
# on the wire if the SDK pools the requests.
|
||||
await asyncio.sleep(0.05)
|
||||
return f"echoed:{payload}"
|
||||
|
||||
app = mcp.streamable_http_app()
|
||||
app.add_middleware(AuthHeaderRecorder)
|
||||
|
||||
config = uvicorn.Config(
|
||||
app,
|
||||
host="127.0.0.1",
|
||||
port=port,
|
||||
log_level="warning",
|
||||
access_log=False,
|
||||
)
|
||||
return uvicorn.Server(config)
|
||||
|
||||
|
||||
def run_server_in_thread(server: uvicorn.Server) -> threading.Thread:
|
||||
"""Boot the server in a background thread on its own asyncio loop."""
|
||||
|
||||
def _run() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.serve())
|
||||
|
||||
t = threading.Thread(target=_run, daemon=True, name="spike-server")
|
||||
t.start()
|
||||
return t
|
||||
|
||||
|
||||
async def wait_for_server_ready(url: str, timeout: float = 5.0) -> None:
|
||||
"""Poll the server until it accepts connections."""
|
||||
import urllib.parse
|
||||
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
reader, writer = await asyncio.open_connection(parsed.hostname, parsed.port)
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
return
|
||||
except OSError:
|
||||
await asyncio.sleep(0.05)
|
||||
raise TimeoutError(f"server at {url} not ready within {timeout}s")
|
||||
|
||||
|
||||
def fd_count() -> int:
|
||||
"""Count open file descriptors for the current process."""
|
||||
try:
|
||||
return len(os.listdir(f"/proc/{os.getpid()}/fd"))
|
||||
except OSError:
|
||||
return -1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scenario 1: N=20 concurrent ClientSession instances
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def scenario_1_concurrent_sessions(url: str, n: int = 20) -> dict:
|
||||
"""Open N concurrent ClientSession instances and call tools/list on each."""
|
||||
print(f"\n=== Scenario 1: {n} concurrent ClientSession instances ===")
|
||||
fd_before = fd_count()
|
||||
|
||||
async def one_session(idx: int) -> dict:
|
||||
headers = {"Authorization": f"Bearer test-token-{idx}"}
|
||||
async with (
|
||||
streamablehttp_client(url=url, headers=headers) as (read, write, _),
|
||||
ClientSession(read, write) as session,
|
||||
):
|
||||
await session.initialize()
|
||||
tools = await session.list_tools()
|
||||
return {
|
||||
"idx": idx,
|
||||
"tool_count": len(tools.tools),
|
||||
"tool_names": [t.name for t in tools.tools],
|
||||
}
|
||||
|
||||
start = time.monotonic()
|
||||
results = await asyncio.gather(*[one_session(i) for i in range(n)], return_exceptions=True)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
fd_after = fd_count()
|
||||
# Allow some settling time for FDs to release.
|
||||
await asyncio.sleep(0.5)
|
||||
fd_settled = fd_count()
|
||||
|
||||
successes = [r for r in results if isinstance(r, dict)]
|
||||
failures = [r for r in results if isinstance(r, Exception)]
|
||||
|
||||
# Verify every session got the same tool catalog.
|
||||
catalog_consistent = (
|
||||
len(successes) == n and len({tuple(r["tool_names"]) for r in successes}) == 1
|
||||
)
|
||||
|
||||
return {
|
||||
"scenario": "concurrent_sessions",
|
||||
"n": n,
|
||||
"successes": len(successes),
|
||||
"failures": len(failures),
|
||||
"elapsed_seconds": round(elapsed, 3),
|
||||
"fd_before": fd_before,
|
||||
"fd_during_peak": fd_after,
|
||||
"fd_settled": fd_settled,
|
||||
"fd_growth_during": fd_after - fd_before,
|
||||
"fd_growth_settled": fd_settled - fd_before,
|
||||
"catalog_consistent": catalog_consistent,
|
||||
"first_failure": str(failures[0]) if failures else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scenario 2: 2 concurrent tools/call on a shared session
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def scenario_2_concurrent_calls_shared_session(url: str) -> dict:
|
||||
"""Two concurrent tools/call on one ClientSession with interleaving payloads.
|
||||
|
||||
The echo tool sleeps 50ms, so concurrent calls overlap on the wire.
|
||||
Each call passes a distinct payload (~10KB) to make request bodies
|
||||
spannable across multiple stream frames.
|
||||
"""
|
||||
print("\n=== Scenario 2: 2 concurrent tools/call on shared session ===")
|
||||
|
||||
# Generous-size payloads so both bodies live during the await.
|
||||
payload_a = "A" * 10000
|
||||
payload_b = "B" * 10000
|
||||
|
||||
headers = {"Authorization": "Bearer shared-session-token"}
|
||||
async with (
|
||||
streamablehttp_client(url=url, headers=headers) as (read, write, _),
|
||||
ClientSession(read, write) as session,
|
||||
):
|
||||
await session.initialize()
|
||||
|
||||
TOOL_CALL_PAYLOADS.clear()
|
||||
|
||||
start = time.monotonic()
|
||||
results = await asyncio.gather(
|
||||
session.call_tool("echo", {"payload": payload_a}),
|
||||
session.call_tool("echo", {"payload": payload_b}),
|
||||
return_exceptions=True,
|
||||
)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
successes = [r for r in results if not isinstance(r, Exception)]
|
||||
failures = [r for r in results if isinstance(r, Exception)]
|
||||
|
||||
# Each result.content[0].text should be "echoed:{payload}".
|
||||
response_payloads: list[str] = []
|
||||
if len(successes) == 2:
|
||||
for r in successes:
|
||||
text = r.content[0].text if r.content else ""
|
||||
response_payloads.append(text)
|
||||
|
||||
# Order may not match call order — what matters is both payloads echo.
|
||||
expected = {f"echoed:{payload_a}", f"echoed:{payload_b}"}
|
||||
received = set(response_payloads)
|
||||
demux_ok = received == expected
|
||||
|
||||
# Did both calls actually overlap? If sequential, elapsed ~= 0.1+s;
|
||||
# if concurrent, ~0.05s.
|
||||
concurrent_observed = elapsed < 0.09
|
||||
|
||||
return {
|
||||
"scenario": "concurrent_calls_shared_session",
|
||||
"successes": len(successes),
|
||||
"failures": len(failures),
|
||||
"elapsed_seconds": round(elapsed, 3),
|
||||
"demux_ok": demux_ok,
|
||||
"expected_payloads_received": list(received) if demux_ok else None,
|
||||
"actual_payloads_received_count": len(received),
|
||||
"appears_concurrent_on_wire": concurrent_observed,
|
||||
"first_failure": str(failures[0]) if failures else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scenario 3: per-session header isolation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def scenario_3_header_isolation(url: str, n: int = 5) -> dict:
|
||||
"""Open N sessions with distinct Authorization headers, call echo on each.
|
||||
|
||||
Verifies the server sees each session's own header — i.e. httpx
|
||||
connection pooling does not cross headers between concurrent
|
||||
ClientSession instances against the same URL.
|
||||
"""
|
||||
print(f"\n=== Scenario 3: {n}-session Authorization-header isolation ===")
|
||||
SERVER_OBSERVATIONS.clear()
|
||||
|
||||
async def one_session(idx: int) -> str | None:
|
||||
headers = {"Authorization": f"Bearer iso-token-{idx}"}
|
||||
async with (
|
||||
streamablehttp_client(url=url, headers=headers) as (read, write, _),
|
||||
ClientSession(read, write) as session,
|
||||
):
|
||||
await session.initialize()
|
||||
# One call per session.
|
||||
await session.call_tool("echo", {"payload": f"session-{idx}"})
|
||||
return f"Bearer iso-token-{idx}"
|
||||
|
||||
start = time.monotonic()
|
||||
expected_tokens = await asyncio.gather(*[one_session(i) for i in range(n)])
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
# Tally observed Authorization headers, ignoring None entries (initial
|
||||
# handshake sometimes lacks auth).
|
||||
observed_auth = [auth for auth, _ in SERVER_OBSERVATIONS if auth]
|
||||
expected_set = set(expected_tokens)
|
||||
observed_set = set(observed_auth)
|
||||
|
||||
# Every expected token must show up at least once on the server.
|
||||
all_present = expected_set.issubset(observed_set)
|
||||
# No spurious tokens.
|
||||
no_extras = observed_set.issubset(expected_set)
|
||||
# Frequency: at least one observation per token.
|
||||
counts = defaultdict(int)
|
||||
for a in observed_auth:
|
||||
counts[a] += 1
|
||||
each_seen = all(counts[t] >= 1 for t in expected_tokens)
|
||||
|
||||
return {
|
||||
"scenario": "header_isolation",
|
||||
"n": n,
|
||||
"elapsed_seconds": round(elapsed, 3),
|
||||
"expected_tokens": sorted(expected_set),
|
||||
"observed_tokens": sorted(observed_set),
|
||||
"all_expected_present": all_present,
|
||||
"no_extra_tokens_observed": no_extras,
|
||||
"each_token_seen_at_least_once": each_seen,
|
||||
"header_counts_per_token": dict(counts),
|
||||
"total_requests_observed": len(observed_auth),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Driver
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
port = find_free_port()
|
||||
url = f"http://127.0.0.1:{port}/mcp"
|
||||
|
||||
server = build_server(port)
|
||||
server_thread = run_server_in_thread(server)
|
||||
try:
|
||||
await wait_for_server_ready(url)
|
||||
print(f"server up at {url}\n")
|
||||
|
||||
result_1 = await scenario_1_concurrent_sessions(url, n=20)
|
||||
print_scenario_result(result_1)
|
||||
|
||||
result_2 = await scenario_2_concurrent_calls_shared_session(url)
|
||||
print_scenario_result(result_2)
|
||||
|
||||
result_3 = await scenario_3_header_isolation(url, n=5)
|
||||
print_scenario_result(result_3)
|
||||
|
||||
# Final verdict
|
||||
verdict_1 = (
|
||||
result_1["successes"] == result_1["n"]
|
||||
and result_1["catalog_consistent"]
|
||||
and result_1["fd_growth_settled"] < 30 # 20 sessions, generous bound
|
||||
)
|
||||
verdict_2 = result_2["demux_ok"] and result_2["successes"] == 2
|
||||
verdict_3 = (
|
||||
result_3["all_expected_present"]
|
||||
and result_3["no_extra_tokens_observed"]
|
||||
and result_3["each_token_seen_at_least_once"]
|
||||
)
|
||||
|
||||
print("\n=== VERDICT ===")
|
||||
print(f" Scenario 1 (concurrent sessions): {'PASS' if verdict_1 else 'FAIL'}")
|
||||
print(f" Scenario 2 (concurrent calls shared): {'PASS' if verdict_2 else 'FAIL'}")
|
||||
print(f" Scenario 3 (header isolation): {'PASS' if verdict_3 else 'FAIL'}")
|
||||
all_pass = verdict_1 and verdict_2 and verdict_3
|
||||
print(
|
||||
f"\n Phase 5 per-(user, server) pool architecture: "
|
||||
f"{'VIABLE' if all_pass else 'NEEDS REWORK (Alternative F fallback)'}"
|
||||
)
|
||||
sys.exit(0 if all_pass else 1)
|
||||
finally:
|
||||
server.should_exit = True
|
||||
server_thread.join(timeout=5)
|
||||
|
||||
|
||||
def print_scenario_result(result: dict) -> None:
|
||||
print(f"\nresult[{result['scenario']}]:")
|
||||
for k, v in result.items():
|
||||
if k == "scenario":
|
||||
continue
|
||||
print(f" {k}: {v}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
with contextlib.suppress(KeyboardInterrupt):
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,220 @@
|
||||
"""Tests for turnstone-admin DB configuration precedence.
|
||||
|
||||
Locks in the alignment with turnstone-server:
|
||||
CLI / config.toml [database] > TURNSTONE_DB_* env > hardcoded default
|
||||
|
||||
The motivation is to keep DB secrets in config.toml (see
|
||||
feedback_secrets_not_in_env) rather than forcing operators to export
|
||||
TURNSTONE_DB_URL before every admin invocation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from typing import TYPE_CHECKING
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
|
||||
import turnstone.core.config as config_mod
|
||||
from turnstone.admin import _get_storage
|
||||
|
||||
|
||||
def _reset_cache() -> None:
|
||||
config_mod._cache = None
|
||||
config_mod._config_path = None
|
||||
|
||||
|
||||
def _build_args(config_path: str | None) -> argparse.Namespace:
|
||||
"""Build an args namespace the way admin.main() does.
|
||||
|
||||
Skips ``add_config_arg`` (which reads ``sys.argv``) — the test
|
||||
constructs the args programmatically instead.
|
||||
"""
|
||||
config_mod.set_config_path(config_path or "/nonexistent/turnstone-admin-test.toml")
|
||||
parser = argparse.ArgumentParser()
|
||||
config_mod.apply_config(parser, ["database"])
|
||||
sub = parser.add_subparsers(dest="command")
|
||||
sub.add_parser("list-users")
|
||||
return parser.parse_args(["list-users"])
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_db_env(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
||||
"""Clean slate: no TURNSTONE_DB_* env vars unless a test sets them."""
|
||||
for var in (
|
||||
"TURNSTONE_DB_BACKEND",
|
||||
"TURNSTONE_DB_URL",
|
||||
"TURNSTONE_DB_PATH",
|
||||
"TURNSTONE_DB_POOL_SIZE",
|
||||
"TURNSTONE_DB_SSLMODE",
|
||||
"TURNSTONE_DB_SSLROOTCERT",
|
||||
"TURNSTONE_DB_SSLCERT",
|
||||
"TURNSTONE_DB_SSLKEY",
|
||||
"TURNSTONE_CONFIG",
|
||||
):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
_reset_cache()
|
||||
yield
|
||||
_reset_cache()
|
||||
|
||||
|
||||
def test_defaults_to_sqlite_when_neither_config_nor_env_set() -> None:
|
||||
args = _build_args(None)
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
assert init.call_args.args == ("sqlite",)
|
||||
assert init.call_args.kwargs["url"] == ""
|
||||
assert init.call_args.kwargs["path"] == ""
|
||||
assert init.call_args.kwargs["pool_size"] == 2
|
||||
|
||||
|
||||
def test_config_toml_database_section_drives_init_storage(tmp_path: Path) -> None:
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(
|
||||
"[database]\n"
|
||||
'backend = "postgresql"\n'
|
||||
'url = "postgresql+psycopg://fromconfig:x@host/db"\n'
|
||||
"pool_size = 5\n"
|
||||
'sslmode = "verify-full"\n'
|
||||
'sslrootcert = "/etc/ssl/ca.pem"\n'
|
||||
'sslcert = "/etc/ssl/client.pem"\n'
|
||||
'sslkey = "/etc/ssl/client.key"\n'
|
||||
)
|
||||
args = _build_args(str(cfg))
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
assert init.call_args.args == ("postgresql",)
|
||||
kw = init.call_args.kwargs
|
||||
assert kw["url"] == "postgresql+psycopg://fromconfig:x@host/db"
|
||||
assert kw["pool_size"] == 5
|
||||
assert kw["sslmode"] == "verify-full"
|
||||
assert kw["sslrootcert"] == "/etc/ssl/ca.pem"
|
||||
assert kw["sslcert"] == "/etc/ssl/client.pem"
|
||||
assert kw["sslkey"] == "/etc/ssl/client.key"
|
||||
|
||||
|
||||
def test_env_used_as_fallback_when_config_absent(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("TURNSTONE_DB_BACKEND", "postgresql")
|
||||
monkeypatch.setenv("TURNSTONE_DB_URL", "postgresql+psycopg://fromenv:x@host/db")
|
||||
monkeypatch.setenv("TURNSTONE_DB_POOL_SIZE", "7")
|
||||
monkeypatch.setenv("TURNSTONE_DB_SSLMODE", "require")
|
||||
|
||||
args = _build_args(None)
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
assert init.call_args.args == ("postgresql",)
|
||||
kw = init.call_args.kwargs
|
||||
assert kw["url"] == "postgresql+psycopg://fromenv:x@host/db"
|
||||
assert kw["pool_size"] == 7
|
||||
assert kw["sslmode"] == "require"
|
||||
|
||||
|
||||
def test_config_toml_wins_over_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
"""config.toml beats env — operators should put secrets in TOML."""
|
||||
monkeypatch.setenv("TURNSTONE_DB_BACKEND", "sqlite")
|
||||
monkeypatch.setenv("TURNSTONE_DB_URL", "postgresql+psycopg://fromenv:x@host/db")
|
||||
monkeypatch.setenv("TURNSTONE_DB_SSLMODE", "require")
|
||||
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(
|
||||
"[database]\n"
|
||||
'backend = "postgresql"\n'
|
||||
'url = "postgresql+psycopg://fromconfig:x@host/db"\n'
|
||||
'sslmode = "verify-full"\n'
|
||||
)
|
||||
args = _build_args(str(cfg))
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
assert init.call_args.args == ("postgresql",)
|
||||
kw = init.call_args.kwargs
|
||||
assert kw["url"] == "postgresql+psycopg://fromconfig:x@host/db"
|
||||
assert kw["sslmode"] == "verify-full"
|
||||
|
||||
|
||||
def test_partial_config_falls_through_to_env_per_key(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
"""A key missing from [database] should fall back to its env var."""
|
||||
monkeypatch.setenv("TURNSTONE_DB_SSLMODE", "require")
|
||||
monkeypatch.setenv("TURNSTONE_DB_POOL_SIZE", "9")
|
||||
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(
|
||||
'[database]\nbackend = "postgresql"\nurl = "postgresql+psycopg://fromconfig:x@host/db"\n'
|
||||
)
|
||||
args = _build_args(str(cfg))
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
kw = init.call_args.kwargs
|
||||
assert kw["url"] == "postgresql+psycopg://fromconfig:x@host/db"
|
||||
assert kw["sslmode"] == "require"
|
||||
assert kw["pool_size"] == 9
|
||||
|
||||
|
||||
def test_empty_string_in_config_beats_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
"""`url = ""` in config.toml beats an env var.
|
||||
|
||||
Locks in the `is not None` guard — a falsy-but-present TOML value
|
||||
should NOT silently fall through to the env fallback.
|
||||
"""
|
||||
monkeypatch.setenv("TURNSTONE_DB_URL", "postgresql+psycopg://fromenv:x@host/db")
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[database]\nbackend = "sqlite"\nurl = ""\n')
|
||||
args = _build_args(str(cfg))
|
||||
with patch("turnstone.core.storage.init_storage") as init:
|
||||
_get_storage(args)
|
||||
assert init.call_args.kwargs["url"] == ""
|
||||
|
||||
|
||||
def test_main_threads_config_toml_through_real_argv(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
"""End-to-end: ``turnstone-admin --config <toml> list-users`` honors TOML.
|
||||
|
||||
Covers the ``add_config_arg`` -> ``apply_config`` -> ``_get_storage``
|
||||
chain that the programmatic ``_build_args`` helper skips.
|
||||
"""
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(
|
||||
'[database]\nbackend = "postgresql"\nurl = "postgresql+psycopg://fromcli:x@host/db"\n'
|
||||
)
|
||||
monkeypatch.setattr("sys.argv", ["turnstone-admin", "--config", str(cfg), "list-users"])
|
||||
|
||||
fake_storage = patch("turnstone.core.storage.init_storage").start()
|
||||
fake_storage.return_value.list_users.return_value = []
|
||||
try:
|
||||
from turnstone.admin import main
|
||||
|
||||
main()
|
||||
finally:
|
||||
patch.stopall()
|
||||
|
||||
assert fake_storage.call_args.args == ("postgresql",)
|
||||
assert fake_storage.call_args.kwargs["url"] == "postgresql+psycopg://fromcli:x@host/db"
|
||||
|
||||
|
||||
def test_get_storage_initializes_real_sqlite_backend(tmp_path: Path) -> None:
|
||||
"""Drives the real ``init_storage`` boundary on a fresh sqlite file.
|
||||
|
||||
Mock-only tests would miss a kwarg-name typo (sslmode -> ssl_mode).
|
||||
This test trips on any such drift because Alembic + the backend
|
||||
actually run.
|
||||
"""
|
||||
from turnstone.core.storage import reset_storage
|
||||
|
||||
db_file = tmp_path / "admin.db"
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text(f'[database]\nbackend = "sqlite"\npath = "{db_file}"\n')
|
||||
args = _build_args(str(cfg))
|
||||
|
||||
reset_storage()
|
||||
try:
|
||||
storage = _get_storage(args)
|
||||
assert storage.list_users() == []
|
||||
finally:
|
||||
reset_storage()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -92,3 +92,437 @@ def test_tool_error_does_not_overwrite_approval_badge() -> None:
|
||||
"badge instead so the approval verdict stays visible alongside "
|
||||
"the error."
|
||||
)
|
||||
|
||||
|
||||
def test_replay_history_renders_content_before_tool_block() -> None:
|
||||
"""In ``replayHistory``'s ``role === "assistant"`` branch, the
|
||||
``msg.content`` render must precede the ``msg.tool_calls`` render.
|
||||
|
||||
Two reasons, both load-bearing:
|
||||
|
||||
1. **Structural** — the next loop iteration's ``role === "tool"``
|
||||
message anchors to ``lastToolBlock``. The tool-block branch sets
|
||||
that anchor; the content branch clears it. If content runs after
|
||||
the tool block, the clear silently drops the upcoming tool
|
||||
result. Pre-fix, every interactive tool result was missing from
|
||||
saved-workstream replays whenever the assistant turn carried
|
||||
both narration and tool calls (very common output shape).
|
||||
|
||||
2. **Visual** — the live SSE path renders content first
|
||||
(``stream_text`` streams before ``tool_info`` /
|
||||
``approve_request``), so replay should match.
|
||||
|
||||
The test pins the order via the offsets of the ``msg.content`` and
|
||||
``msg.tool_calls`` branch headers inside the function body."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("Pane.prototype.replayHistory = function")
|
||||
end = body.index("Pane.prototype._attachRetryToLastAssistant", start)
|
||||
fn = body[start:end]
|
||||
# Locate the assistant branch and bound the search to its body —
|
||||
# the function also handles user / tool roles which would otherwise
|
||||
# confuse the offset comparison.
|
||||
asst_start = fn.index('msg.role === "assistant"')
|
||||
asst_end = fn.index('msg.role === "tool"', asst_start)
|
||||
asst = fn[asst_start:asst_end]
|
||||
# ``if (msg.content && msg.content.trim())`` guards against a
|
||||
# whitespace-only content row (Qwen-style "\n\n" left over after a
|
||||
# reasoning-parser model strips ``<think>…</think>`` and emits
|
||||
# nothing else before the tool call). Pre-trim guard, those rows
|
||||
# rendered as a visible-but-empty ``.msg.assistant`` card on
|
||||
# replay. Match the substring up to ``msg.content`` so the test
|
||||
# tolerates either guard shape without locking the trim() in.
|
||||
content_idx = asst.index("if (msg.content")
|
||||
tool_calls_idx = asst.index("if (msg.tool_calls && msg.tool_calls.length)")
|
||||
assert content_idx < tool_calls_idx, (
|
||||
"replayHistory must render msg.content BEFORE msg.tool_calls "
|
||||
"inside the assistant branch — otherwise the lastToolBlock "
|
||||
"anchor is clobbered before the next iteration's tool result "
|
||||
"can attach to it (and the visual order also drifts from the "
|
||||
"live SSE flow)."
|
||||
)
|
||||
|
||||
|
||||
def test_replay_history_renders_persisted_verdict_badge() -> None:
|
||||
"""Saved-workstream replays must paint the persisted intent verdict
|
||||
next to each tool div, using the same ``renderVerdictBadge`` helper
|
||||
the live ``showInlineToolBlock`` path uses. Pre-fix the audit trail
|
||||
was complete in storage (``intent_verdicts`` table) but never
|
||||
surfaced on replay — operators reviewing a saved workstream
|
||||
couldn't see what the heuristic / LLM judge thought of any tool
|
||||
call. This test pins the call site so a refactor that drops the
|
||||
decoration regresses the audit surface."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("Pane.prototype.replayHistory = function")
|
||||
end = body.index("Pane.prototype._attachRetryToLastAssistant", start)
|
||||
fn = body[start:end]
|
||||
# Match a `renderVerdictBadge(<something>.verdict, ...)` call inside
|
||||
# the replay loop. Loose on whitespace + identifier so a future
|
||||
# rename of the iteration variable doesn't trip CI.
|
||||
badge_call_re = re.compile(
|
||||
r"renderVerdictBadge\(\s*\w+\.verdict\b",
|
||||
)
|
||||
assert badge_call_re.search(fn), (
|
||||
"replayHistory must call renderVerdictBadge(tc.verdict, ...) "
|
||||
"when a persisted verdict is attached to a tool_call entry — "
|
||||
"otherwise the audit-trail data persisted to intent_verdicts "
|
||||
"doesn't surface on saved-workstream replays."
|
||||
)
|
||||
|
||||
|
||||
def test_shared_utils_defines_replay_advisories_after_tool() -> None:
|
||||
"""The shared ``replayAdvisoriesAfterTool`` helper in
|
||||
``shared_static/utils.js`` is the single source of advisory-walk +
|
||||
type-filter logic for both ``app.js`` (interactive) and
|
||||
``coordinator.js`` (coord). A refactor that drops the helper
|
||||
breaks both surfaces, so guard its definition + filter shape here.
|
||||
"""
|
||||
utils_js = Path(__file__).resolve().parent.parent / "turnstone/shared_static/utils.js"
|
||||
body = utils_js.read_text(encoding="utf-8")
|
||||
assert "function replayAdvisoriesAfterTool" in body, (
|
||||
"shared/utils.js must define replayAdvisoriesAfterTool — "
|
||||
"interactive and coord both invoke it."
|
||||
)
|
||||
# The type filter — ``adv.type !== 'user_interjection'`` — must
|
||||
# remain in the helper so a future advisory shape (output_guard,
|
||||
# metacognitive nudge, etc.) doesn't silently render as a user
|
||||
# bubble.
|
||||
assert 'adv.type !== "user_interjection"' in body, (
|
||||
"replayAdvisoriesAfterTool must filter by advisory type so a "
|
||||
"future non-user_interjection advisory shape doesn't silently "
|
||||
"render as a user bubble."
|
||||
)
|
||||
|
||||
|
||||
def test_replay_renders_user_interjection_advisory_after_tool_block() -> None:
|
||||
"""Queued user messages spliced into the last tool-result envelope
|
||||
of a batch (Seam 1) persist on the tool DB row as a wrapped
|
||||
``<tool_output>`` envelope. ``decorate_history_messages`` extracts
|
||||
the advisory back out and the wire layer projects it onto
|
||||
``msg.advisories``; ``replayHistory`` must invoke the shared
|
||||
``replayAdvisoriesAfterTool`` helper (defined in
|
||||
``shared/utils.js``) so each ``user_interjection`` renders through
|
||||
``addUserMessage`` and the bubble looks identical to a Seam 2/3
|
||||
user row.
|
||||
|
||||
This test pins the call site so a refactor that drops the helper
|
||||
invocation regresses the queued-during-batch replay shape
|
||||
silently."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("Pane.prototype.replayHistory = function")
|
||||
end = body.index("Pane.prototype._attachRetryToLastAssistant", start)
|
||||
fn = body[start:end]
|
||||
# The replay loop must invoke the shared helper, passing
|
||||
# ``msg.advisories`` and a renderer that routes through
|
||||
# ``addUserMessage``. The helper itself filters on
|
||||
# ``adv.type !== "user_interjection"``; that branch lives in
|
||||
# ``shared/utils.js`` (test_shared_utils_js or runtime smoke covers
|
||||
# the helper's body).
|
||||
assert "replayAdvisoriesAfterTool(msg.advisories" in fn, (
|
||||
"replayHistory must invoke replayAdvisoriesAfterTool with "
|
||||
"msg.advisories so queued messages spliced into the tool "
|
||||
"envelope render as user bubbles after the tool block."
|
||||
)
|
||||
assert "addUserMessage(text" in fn, (
|
||||
"replayHistory's renderer callback must route the extracted "
|
||||
"advisory text through addUserMessage so the rendered bubble "
|
||||
"matches a normal user-row replay."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 8 — Chunk D: MCP error embed + settings panel UX
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_INDEX_HTML = Path(__file__).resolve().parent.parent / "turnstone/ui/static/index.html"
|
||||
_STYLE_CSS = Path(__file__).resolve().parent.parent / "turnstone/ui/static/style.css"
|
||||
|
||||
# The Phase-8 D-chunk pins the absence of an unsafe DOM-write API
|
||||
# in two regions of app.js. Spell the property name out of literal
|
||||
# concatenation so the tooling that flags occurrences in code
|
||||
# strings doesn't false-positive on the test source.
|
||||
_UNSAFE_DOM_WRITE_RE = re.compile(r"\.inner" + r"HTML\s*=")
|
||||
|
||||
|
||||
def test_phase8_mcp_error_helpers_defined_in_app_js() -> None:
|
||||
"""The Phase 8 dashboard renderer adds three load-bearing helpers
|
||||
next to the existing media-embed pattern: ``tryParseMcpError``
|
||||
(envelope detector), ``buildMcpErrorEmbed`` (interactive card),
|
||||
and the ``_pendingConsentServers`` set that drives the gear-icon
|
||||
badge. A regression that drops any of them silently degrades the
|
||||
OAuth consent UX to a plain JSON dump, so guard their existence
|
||||
here."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
assert "function tryParseMcpError" in body, (
|
||||
"tryParseMcpError must remain defined — appendToolOutput's "
|
||||
"error branch depends on it to detect the MCP error envelope."
|
||||
)
|
||||
assert "function buildMcpErrorEmbed" in body, (
|
||||
"buildMcpErrorEmbed must remain defined — it renders the "
|
||||
"interactive consent / forbidden / operator card."
|
||||
)
|
||||
assert "_pendingConsentServers" in body, (
|
||||
"_pendingConsentServers state must remain — it backs the "
|
||||
"gear-icon badge so a user who scrolls past a consent prompt "
|
||||
"still has a stable signal that consent is pending."
|
||||
)
|
||||
# The buildMcpErrorEmbed pattern must also wire the "actionable"
|
||||
# branch (consent_required / insufficient_scope) into the badge
|
||||
# via _onConsentDetected; pin the helper name.
|
||||
assert "_onConsentDetected" in body, (
|
||||
"_onConsentDetected must remain — buildMcpErrorEmbed calls it "
|
||||
"for the actionable category to surface the gear-icon badge."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_settings_panel_handlers_defined() -> None:
|
||||
"""The settings modal exposes four entry points that the inline
|
||||
``onclick`` attributes in index.html depend on. Renaming or
|
||||
deleting any of them breaks the modal silently (the buttons are
|
||||
still rendered but click-to-action is dead). Catch that here."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
for name in [
|
||||
"function openSettingsPanel",
|
||||
"function closeSettingsPanel",
|
||||
"function confirmRevokeMcp",
|
||||
"function cancelRevokeMcp",
|
||||
]:
|
||||
assert name in body, f"Missing required handler: {name}"
|
||||
# The connections list is fetched against the Phase-7 endpoint —
|
||||
# pin the URL so a server-side rename forces an explicit UI bump.
|
||||
assert "/v1/api/mcp/oauth/connections" in body, (
|
||||
"Settings panel must fetch /v1/api/mcp/oauth/connections — "
|
||||
"a server-side rename needs an explicit UI update."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_appendtooloutput_dispatches_mcp_error_before_renderer() -> None:
|
||||
"""``appendToolOutput`` must call ``tryParseMcpError`` inside its
|
||||
``isError`` branch BEFORE falling through to the plain
|
||||
``renderToolOutput`` path. The ordering is what makes the
|
||||
interactive consent card replace the JSON dump; reverse the calls
|
||||
and the user sees the raw error envelope as text again."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("Pane.prototype.appendToolOutput = function")
|
||||
end = body.index("Pane.prototype.", start + 10)
|
||||
fn = body[start:end]
|
||||
parse_idx = fn.find("tryParseMcpError(")
|
||||
render_idx = fn.find("renderToolOutput(")
|
||||
assert parse_idx >= 0, (
|
||||
"appendToolOutput must call tryParseMcpError on the error path "
|
||||
"before renderToolOutput, otherwise the consent card never "
|
||||
"replaces the plain JSON output."
|
||||
)
|
||||
assert render_idx >= 0, "renderToolOutput call must remain present"
|
||||
assert parse_idx < render_idx, (
|
||||
"tryParseMcpError must run BEFORE renderToolOutput so the "
|
||||
"interactive card path takes precedence over plain rendering."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_no_unsafe_dom_write_in_settings_panel() -> None:
|
||||
"""Defensive XSS guard: the settings panel renders user-controlled
|
||||
server names, scope strings, and timestamp values into the DOM.
|
||||
The whole section MUST go through ``textContent``-style APIs; an
|
||||
unsafe-DOM-write assignment would be a regression vector. Bound
|
||||
the check to the section 15 body to avoid false positives
|
||||
elsewhere."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("// 15. MCP server connections settings panel")
|
||||
# Bound to the full settings section (terminates at the next
|
||||
# top-level keydown handler block).
|
||||
end = body.index('document.addEventListener("keydown"', start)
|
||||
section = body[start:end]
|
||||
assert not _UNSAFE_DOM_WRITE_RE.search(section), (
|
||||
"Section 15 must not assign to the unsafe DOM-write property — "
|
||||
"server names and scope values flow through here and would be "
|
||||
"XSS-injectable. Use textContent / DOM APIs instead."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_settings_button_in_index_html() -> None:
|
||||
"""The gear-icon entry-point for the settings menu must remain
|
||||
in the appbar's actions span. The console proxy IIFE prepends a
|
||||
node pill to ``header.firstChild`` (turnstone/console/server.py:
|
||||
202); our button is appended inside ``<span class='appbar-actions'>``
|
||||
on the right, so they don't collide. Pin both shape constraints
|
||||
here so a future appbar refactor keeps them disjoint."""
|
||||
body = _INDEX_HTML.read_text(encoding="utf-8")
|
||||
assert 'id="settings-btn"' in body, (
|
||||
"index.html must keep the #settings-btn — onclick handlers "
|
||||
"and the consent badge target it by id."
|
||||
)
|
||||
assert 'onclick="toggleSettingsMenu(this)"' in body, (
|
||||
"settings-btn must wire onclick=toggleSettingsMenu(this) — "
|
||||
"the gear opens a dropdown with MCP connections + Logout; "
|
||||
"losing the binding leaves the menu unreachable."
|
||||
)
|
||||
# The button must live inside <span class="appbar-actions"> so the
|
||||
# console proxy's header.insertBefore(pill, header.firstChild)
|
||||
# leaves it untouched.
|
||||
actions_open = body.index('class="appbar-actions"')
|
||||
actions_close = body.index("</span>", actions_open)
|
||||
assert 'id="settings-btn"' in body[actions_open:actions_close], (
|
||||
"settings-btn must be inside <span class='appbar-actions'> "
|
||||
"so the console proxy's firstChild prepend doesn't shift it."
|
||||
)
|
||||
|
||||
|
||||
def test_settings_menu_handlers_defined() -> None:
|
||||
"""The gear-icon dropdown exposes a toggle/open/close trio that the
|
||||
inline ``onclick="toggleSettingsMenu(this)"`` in index.html depends
|
||||
on, plus the menu items themselves must wire to existing entry
|
||||
points (``openSettingsPanel`` for MCP connections, ``logout`` for
|
||||
sign-out). Pin all four so a rename or deletion fails loudly here
|
||||
instead of silently leaving the gear's menu broken or wired to a
|
||||
stale function."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
for name in [
|
||||
"function toggleSettingsMenu",
|
||||
"function openSettingsMenu",
|
||||
"function closeSettingsMenu",
|
||||
]:
|
||||
assert name in body, f"Missing required handler: {name}"
|
||||
# Bound to the settings-menu region so we don't accidentally match
|
||||
# an unrelated openSettingsPanel/logout call elsewhere in the file.
|
||||
start = body.index("function openSettingsMenu(")
|
||||
end = body.index("function closeSettingsMenu(", start)
|
||||
section = body[start:end]
|
||||
assert "openSettingsPanel()" in section, (
|
||||
"Settings menu's MCP-connections item must call openSettingsPanel() "
|
||||
"— otherwise the existing settings overlay is unreachable from the "
|
||||
"new dropdown."
|
||||
)
|
||||
assert "logout()" in section, (
|
||||
"Settings menu's Logout item must call logout() — that's the "
|
||||
"shared auth.js entry point that clears the cookie + session state."
|
||||
)
|
||||
|
||||
|
||||
def test_dashboard_overlay_is_region_not_dialog() -> None:
|
||||
"""The dashboard overlay must be role='region' (not role='dialog' +
|
||||
aria-modal='true'). The role downgrade is what allows ui-header to
|
||||
stay interactive while the dashboard is open — see the comment at
|
||||
showDashboard() in app.js. A revert to role='dialog' + aria-modal
|
||||
would re-trap focus and break the gear/theme buttons + the console
|
||||
proxy's node-picker pill while the dashboard is open."""
|
||||
body = _INDEX_HTML.read_text(encoding="utf-8")
|
||||
idx = body.index('id="dashboard"')
|
||||
# Bound to ~600 chars after the tag so we only check this element's
|
||||
# attributes — same shape as test_phase8_settings_modal_in_index_html.
|
||||
chunk = body[idx : idx + 600]
|
||||
assert 'role="region"' in chunk, (
|
||||
"dashboard must be role='region' — see showDashboard() comment."
|
||||
)
|
||||
assert "aria-modal" not in chunk, (
|
||||
"dashboard must NOT be aria-modal — re-trapping focus breaks "
|
||||
"the appbar's interactive controls (theme toggle, settings menu, "
|
||||
"proxy node-picker pill) while the dashboard is open."
|
||||
)
|
||||
|
||||
|
||||
def test_close_settings_menu_resets_aria() -> None:
|
||||
"""closeSettingsMenu must reset aria-expanded='false' AND remove
|
||||
aria-controls from the gear trigger. Without the reset the gear
|
||||
keeps reporting 'expanded' to assistive tech after the menu closes;
|
||||
without the removal aria-controls points at a dead DOM id."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("function closeSettingsMenu(")
|
||||
# Bound to ~600 chars so we don't catch unrelated handlers.
|
||||
section = body[start : start + 600]
|
||||
assert 'setAttribute("aria-expanded", "false")' in section, (
|
||||
"closeSettingsMenu must set aria-expanded='false' on the gear."
|
||||
)
|
||||
assert 'removeAttribute("aria-controls")' in section, (
|
||||
"closeSettingsMenu must remove aria-controls from the gear."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_settings_modal_in_index_html() -> None:
|
||||
"""Both the settings overlay and the revoke-confirmation overlay
|
||||
must remain in the modal area. The Escape-key deferral list in
|
||||
app.js targets these ids, so removing them silently breaks the
|
||||
handler chain."""
|
||||
body = _INDEX_HTML.read_text(encoding="utf-8")
|
||||
assert 'id="settings-overlay"' in body
|
||||
assert 'id="revoke-mcp-overlay"' in body
|
||||
# Each overlay must have role="dialog" + aria-modal="true" so
|
||||
# screen readers and the existing modal-deferral handlers can
|
||||
# treat them like the rest of the modal stack.
|
||||
for overlay_id in ("settings-overlay", "revoke-mcp-overlay"):
|
||||
idx = body.index(f'id="{overlay_id}"')
|
||||
# Bound to ~600 chars after the open tag so we only check this
|
||||
# overlay's attributes.
|
||||
chunk = body[idx : idx + 600]
|
||||
assert 'role="dialog"' in chunk, f"{overlay_id} missing role=dialog"
|
||||
assert 'aria-modal="true"' in chunk, f"{overlay_id} missing aria-modal=true"
|
||||
|
||||
|
||||
def test_phase8_xss_safe_render_in_build_mcp_error_embed() -> None:
|
||||
"""Adversarial input — the renderer for an MCP error envelope
|
||||
must use ``textContent`` (not the unsafe DOM-write API) for every
|
||||
field that flows from the server: ``err.detail``, ``err.server``,
|
||||
scopes list. The card builder uses createElement + textContent
|
||||
throughout so a script-tag server name renders harmlessly. Pin
|
||||
the absence of the unsafe-write inside ``buildMcpErrorEmbed``."""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
start = body.index("function buildMcpErrorEmbed(")
|
||||
# Bound to the function body — find its closing brace at column 0.
|
||||
rest = body[start:]
|
||||
# Closing function brace at line start (matches existing functions)
|
||||
end_match = re.search(r"\n}\n", rest)
|
||||
assert end_match is not None
|
||||
fn = rest[: end_match.end()]
|
||||
assert not _UNSAFE_DOM_WRITE_RE.search(fn), (
|
||||
"buildMcpErrorEmbed must not use the unsafe-DOM-write API — "
|
||||
"server names and detail strings flow through here. An "
|
||||
"adversarial server name must render harmlessly via "
|
||||
"textContent."
|
||||
)
|
||||
|
||||
|
||||
def test_phase8_css_classes_present_in_stylesheet() -> None:
|
||||
"""The card / badge / modal classes referenced from app.js must
|
||||
have CSS rules. Without them the DOM still works but the visual
|
||||
treatment is gone, which would silently degrade the consent UX."""
|
||||
css = _STYLE_CSS.read_text(encoding="utf-8")
|
||||
for selector in [
|
||||
".mcp-error-card",
|
||||
".mcp-error-icon",
|
||||
".mcp-error-action-btn",
|
||||
".mcp-scope-pill",
|
||||
"#settings-overlay",
|
||||
"#settings-box",
|
||||
".settings-revoke-btn",
|
||||
".settings-consent-badge",
|
||||
"#revoke-mcp-overlay",
|
||||
]:
|
||||
assert selector in css, f"Missing CSS rule for {selector}"
|
||||
|
||||
|
||||
def test_phase8_consent_url_prefix_check_in_click_handler() -> None:
|
||||
"""Defence-in-depth: the consent button's click handler must reject
|
||||
any ``consent_url`` that doesn't start with the dispatcher's known
|
||||
prefix (``/v1/api/mcp/oauth/start``). ``_build_consent_url`` always
|
||||
emits a path-relative URL with that exact prefix; a non-prefix
|
||||
value implies the producer drifted (or was compromised) and a
|
||||
``window.open("javascript:...")`` would be catastrophic.
|
||||
|
||||
The renderer is the last line of defence before ``window.open`` and
|
||||
must not rely on the producer-side guarantee alone. Pin the prefix
|
||||
string and the ``startsWith`` form so a future refactor can't
|
||||
silently weaken the guard.
|
||||
"""
|
||||
body = _APP_JS.read_text(encoding="utf-8")
|
||||
# Bound the search to the click handler region (between the
|
||||
# ``buildMcpErrorEmbed`` function and the next top-level helper) to
|
||||
# avoid false positives from unrelated string occurrences.
|
||||
start = body.index("function buildMcpErrorEmbed(")
|
||||
end = body.index("\n}\n", start) + 1
|
||||
fn = body[start:end]
|
||||
assert 'consentUrl.startsWith("/v1/api/mcp/oauth/start")' in fn, (
|
||||
"Click handler must guard window.open with "
|
||||
'consentUrl.startsWith("/v1/api/mcp/oauth/start"). Without it '
|
||||
"a future producer drift to a non-path-relative URL (or a "
|
||||
'"javascript:" injection) would be passed straight to '
|
||||
"window.open."
|
||||
)
|
||||
|
||||
@@ -83,6 +83,39 @@ class TestIsPublicPath:
|
||||
def test_shared_static_public(self):
|
||||
assert is_public_path("/shared/base.css") is True
|
||||
|
||||
# Console proxy: a public proxied path must still be public, otherwise
|
||||
# the login modal can never re-authenticate from inside a ``/node/{id}/``
|
||||
# proxied page once the cookie expires.
|
||||
def test_proxy_v1_login_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/login") is True
|
||||
|
||||
def test_proxy_no_v1_login_public(self):
|
||||
assert is_public_path("/node/node-a/api/auth/login") is True
|
||||
|
||||
def test_proxy_v1_status_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/status") is True
|
||||
|
||||
def test_proxy_v1_setup_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/setup") is True
|
||||
|
||||
def test_proxy_v1_logout_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/logout") is True
|
||||
|
||||
def test_proxy_v1_oidc_authorize_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/oidc/authorize") is True
|
||||
|
||||
def test_proxy_v1_oidc_callback_public(self):
|
||||
assert is_public_path("/node/node-a/v1/api/auth/oidc/callback") is True
|
||||
|
||||
def test_proxy_v1_workstreams_still_not_public(self):
|
||||
"""Proxy prefix must not turn protected paths into public ones."""
|
||||
assert is_public_path("/node/node-a/v1/api/workstreams") is False
|
||||
|
||||
def test_proxy_v1_refresh_still_requires_auth(self):
|
||||
"""Refresh isn't in PUBLIC_PATHS — the caller must already have
|
||||
a valid cookie. Proxy-prefix shouldn't change that."""
|
||||
assert is_public_path("/node/node-a/v1/api/auth/refresh") is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TestRequiredRole
|
||||
@@ -207,6 +240,30 @@ class TestRequiredScope:
|
||||
"""Only POST is elevated — GET falls through to read."""
|
||||
assert required_scope("GET", "/api/_internal/mcp-reload") == "read"
|
||||
|
||||
def test_internal_mcp_refresh_one_needs_approve(self):
|
||||
assert required_scope("POST", "/api/_internal/mcp-refresh/srv") == "approve"
|
||||
|
||||
def test_v1_internal_mcp_refresh_one_needs_approve(self):
|
||||
assert required_scope("POST", "/v1/api/_internal/mcp-refresh/srv") == "approve"
|
||||
|
||||
def test_proxy_internal_mcp_refresh_one_needs_approve(self):
|
||||
assert required_scope("POST", "/node/n1/v1/api/_internal/mcp-refresh/srv") == "approve"
|
||||
|
||||
def test_proxy_no_v1_internal_mcp_refresh_one_needs_approve(self):
|
||||
assert required_scope("POST", "/node/n1/api/_internal/mcp-refresh/srv") == "approve"
|
||||
|
||||
def test_internal_mcp_reconnect_one_needs_approve(self):
|
||||
assert required_scope("POST", "/api/_internal/mcp-reconnect/srv") == "approve"
|
||||
|
||||
def test_v1_internal_mcp_reconnect_one_needs_approve(self):
|
||||
assert required_scope("POST", "/v1/api/_internal/mcp-reconnect/srv") == "approve"
|
||||
|
||||
def test_proxy_internal_mcp_reconnect_one_needs_approve(self):
|
||||
assert required_scope("POST", "/node/n1/v1/api/_internal/mcp-reconnect/srv") == "approve"
|
||||
|
||||
def test_proxy_no_v1_internal_mcp_reconnect_one_needs_approve(self):
|
||||
assert required_scope("POST", "/node/n1/api/_internal/mcp-reconnect/srv") == "approve"
|
||||
|
||||
# Workstream sub-resource mutations (parametric paths)
|
||||
def test_ws_delete_needs_write(self):
|
||||
assert required_scope("POST", "/api/workstreams/abc123/delete") == "write"
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
"""Tests for ``turnstone.server._build_history`` reminder + source surfacing.
|
||||
|
||||
The replay path (``_build_history``) projects the ``_source`` and
|
||||
``_reminders`` side-channels onto the wire entry the frontend
|
||||
consumes. Persisted via migration 050 (Commit 1) so multi-tab /
|
||||
multi-device replay sees the same metacognitive bubble shape the
|
||||
originating tab saw live.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
from turnstone.server import _build_history
|
||||
|
||||
|
||||
def _make_stub_session(messages: list[dict[str, Any]]) -> Any:
|
||||
"""Minimal ChatSession-shaped stub. ``_build_history`` only reads
|
||||
``session.messages`` plus calls ``_load_verdict_indexes(ws_id)`` —
|
||||
the latter we patch out below.
|
||||
"""
|
||||
return SimpleNamespace(messages=messages, _ws_id="ws-test")
|
||||
|
||||
|
||||
def _build(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Run ``_build_history`` against a stub session, bypassing the
|
||||
verdicts / output-assessment storage round-trip (no tool_calls in
|
||||
these tests, so the indexes are unused anyway).
|
||||
"""
|
||||
session = _make_stub_session(messages)
|
||||
with patch(
|
||||
"turnstone.server._load_verdict_indexes",
|
||||
return_value=({}, {}),
|
||||
):
|
||||
return _build_history(session)
|
||||
|
||||
|
||||
class TestSourceSurfacing:
|
||||
def test_source_surfaces_when_set(self) -> None:
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": "",
|
||||
"_source": "system_nudge",
|
||||
}
|
||||
history = _build([msg])
|
||||
assert len(history) == 1
|
||||
assert history[0]["source"] == "system_nudge"
|
||||
|
||||
def test_source_absent_when_unset(self) -> None:
|
||||
msg = {"role": "user", "content": "hello"}
|
||||
history = _build([msg])
|
||||
assert "source" not in history[0]
|
||||
|
||||
|
||||
class TestRemindersWidening:
|
||||
def test_watch_triggered_optional_fields_propagate(self) -> None:
|
||||
"""The widened payload (Commit 2) carries watch_name / command /
|
||||
poll_count / max_polls / is_final on each ``watch_triggered``
|
||||
reminder so the frontend renders ``.msg.watch-result``.
|
||||
"""
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": "",
|
||||
"_source": "system_nudge",
|
||||
"_reminders": [
|
||||
{
|
||||
"type": "watch_triggered",
|
||||
"text": "$ ls\nfile.txt",
|
||||
"watch_name": "w1",
|
||||
"command": "ls",
|
||||
"poll_count": 2,
|
||||
"max_polls": 100,
|
||||
"is_final": False,
|
||||
}
|
||||
],
|
||||
}
|
||||
history = _build([msg])
|
||||
assert history[0]["source"] == "system_nudge"
|
||||
assert history[0]["reminders"] == [
|
||||
{
|
||||
"type": "watch_triggered",
|
||||
"text": "$ ls\nfile.txt",
|
||||
"watch_name": "w1",
|
||||
"command": "ls",
|
||||
"poll_count": 2,
|
||||
"max_polls": 100,
|
||||
"is_final": False,
|
||||
}
|
||||
]
|
||||
|
||||
def test_legacy_two_field_reminders_still_work(self) -> None:
|
||||
"""Producers without optional fields (correction / denial /
|
||||
idle_children) keep the legacy ``{type, text}`` shape — the
|
||||
widened filter just doesn't add anything beyond that."""
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": "noted",
|
||||
"_reminders": [{"type": "correction", "text": "watch out"}],
|
||||
}
|
||||
history = _build([msg])
|
||||
assert history[0]["reminders"] == [{"type": "correction", "text": "watch out"}]
|
||||
|
||||
def test_unknown_keys_are_dropped(self) -> None:
|
||||
"""The wire-layer filter projects on a known set of keys so a
|
||||
future producer accidentally stuffing arbitrary fields can't
|
||||
leak them through replay.
|
||||
"""
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": "x",
|
||||
"_reminders": [
|
||||
{
|
||||
"type": "correction",
|
||||
"text": "hi",
|
||||
"secret": "leak-me",
|
||||
"internal_id": 42,
|
||||
}
|
||||
],
|
||||
}
|
||||
history = _build([msg])
|
||||
clean = history[0]["reminders"][0]
|
||||
assert "secret" not in clean
|
||||
assert "internal_id" not in clean
|
||||
assert clean == {"type": "correction", "text": "hi"}
|
||||
|
||||
def test_malformed_reminder_skipped(self) -> None:
|
||||
"""A non-dict / empty entry is filtered out instead of breaking
|
||||
the rest of the list (mirrors the defensive filter in
|
||||
``_apply_reminders_for_provider``).
|
||||
"""
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": "x",
|
||||
"_reminders": [
|
||||
"garbage string",
|
||||
{"type": "", "text": ""}, # empty type + text → drop
|
||||
{"type": "denial", "text": "ok"},
|
||||
],
|
||||
}
|
||||
history = _build([msg])
|
||||
assert history[0]["reminders"] == [{"type": "denial", "text": "ok"}]
|
||||
|
||||
|
||||
class _StubRegistry:
|
||||
"""Minimal model registry — only ``get_config`` is read by
|
||||
``_build_history``."""
|
||||
|
||||
def __init__(self, surface_persisted_reasoning: bool = True) -> None:
|
||||
self._cfg = SimpleNamespace(surface_persisted_reasoning=surface_persisted_reasoning)
|
||||
|
||||
def get_config(self, alias: str) -> Any:
|
||||
return self._cfg
|
||||
|
||||
|
||||
def _build_with_registry(
|
||||
messages: list[dict[str, Any]],
|
||||
surface_persisted_reasoning: bool = True,
|
||||
) -> list[dict[str, Any]]:
|
||||
session = SimpleNamespace(
|
||||
messages=messages,
|
||||
_ws_id="ws-test",
|
||||
_registry=_StubRegistry(surface_persisted_reasoning=surface_persisted_reasoning),
|
||||
_model_alias="claude-opus-4-7",
|
||||
)
|
||||
with patch(
|
||||
"turnstone.server._load_verdict_indexes",
|
||||
return_value=({}, {}),
|
||||
):
|
||||
return _build_history(session)
|
||||
|
||||
|
||||
class TestReasoningSurfacing:
|
||||
"""Phase 1 — surface stored Anthropic thinking blocks on the
|
||||
history payload so refresh-the-page rehydrates the reasoning bubble.
|
||||
Drives through the real ``AnthropicProvider`` extractor (no mock-of-
|
||||
extractor) — only the model registry is stubbed.
|
||||
"""
|
||||
|
||||
def test_reasoning_surfaces_for_anthropic_thinking_msg(self) -> None:
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": "Final answer.",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": "let me think", "signature": "s"},
|
||||
{"type": "text", "text": "Final answer."},
|
||||
],
|
||||
}
|
||||
history = _build_with_registry([msg], surface_persisted_reasoning=True)
|
||||
assert len(history) == 1
|
||||
assert history[0]["reasoning"] == "let me think"
|
||||
|
||||
def test_reasoning_empty_when_persist_flag_false(self) -> None:
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": "Final answer.",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": "hidden", "signature": "s"},
|
||||
],
|
||||
}
|
||||
history = _build_with_registry([msg], surface_persisted_reasoning=False)
|
||||
assert "reasoning" not in history[0]
|
||||
|
||||
def test_provider_content_never_in_wire_entry(self) -> None:
|
||||
# The build path does not copy ``_provider_content`` into the
|
||||
# entry dict regardless of flag — wire payload stays tight.
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": "Final answer.",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": "x", "signature": "s"},
|
||||
],
|
||||
}
|
||||
history = _build_with_registry([msg], surface_persisted_reasoning=True)
|
||||
assert "_provider_content" not in history[0]
|
||||
|
||||
def test_no_reasoning_field_when_provider_content_missing(self) -> None:
|
||||
msg = {"role": "assistant", "content": "plain answer"}
|
||||
history = _build_with_registry([msg], surface_persisted_reasoning=True)
|
||||
assert "reasoning" not in history[0]
|
||||
|
||||
def test_no_reasoning_field_for_non_assistant_messages(self) -> None:
|
||||
# Defensive — user/tool messages with a stray _provider_content
|
||||
# do not get the reasoning field stamped.
|
||||
msgs: list[dict[str, Any]] = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "c1",
|
||||
"content": "out",
|
||||
"_provider_content": [{"type": "thinking", "thinking": "leak", "signature": "s"}],
|
||||
},
|
||||
]
|
||||
history = _build_with_registry(msgs, surface_persisted_reasoning=True)
|
||||
assert "reasoning" not in history[0]
|
||||
assert "reasoning" not in history[1]
|
||||
|
||||
def test_default_true_when_registry_lookup_raises(self) -> None:
|
||||
# Conservative default — Phase 1 spec mandates rehydration on
|
||||
# refresh. A registry/alias mismatch must not silently kill the
|
||||
# bubble.
|
||||
class BrokenRegistry:
|
||||
def get_config(self, alias: str) -> Any:
|
||||
raise KeyError(alias)
|
||||
|
||||
session = SimpleNamespace(
|
||||
messages=[
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": "still works", "signature": "s"}
|
||||
],
|
||||
}
|
||||
],
|
||||
_ws_id="ws-test",
|
||||
_registry=BrokenRegistry(),
|
||||
_model_alias="missing-alias",
|
||||
)
|
||||
with patch(
|
||||
"turnstone.server._load_verdict_indexes",
|
||||
return_value=({}, {}),
|
||||
):
|
||||
history = _build_history(session)
|
||||
assert history[0]["reasoning"] == "still works"
|
||||
@@ -19,6 +19,12 @@ class NullUI:
|
||||
self.infos = []
|
||||
self.stream_ends = 0
|
||||
|
||||
def on_turn_start(self):
|
||||
pass
|
||||
|
||||
def on_turn_committed(self):
|
||||
pass
|
||||
|
||||
def on_thinking_start(self):
|
||||
pass
|
||||
|
||||
@@ -820,3 +826,109 @@ class TestForceCancelThreaded:
|
||||
assert "idle" in ui.states
|
||||
assistant_msgs = [m for m in session.messages if m["role"] == "assistant"]
|
||||
assert any("Fresh response" in m.get("content", "") for m in assistant_msgs)
|
||||
|
||||
|
||||
class TestSynthesizeCancelledResults:
|
||||
"""Regression coverage for ``_synthesize_cancelled_results`` — must
|
||||
fire ``on_tool_result`` for each synthesized cancellation so live
|
||||
SSE listeners (e.g. coord's ``--running`` indicator added by
|
||||
tool_info) can complete the in-DOM tool batch. Without this, the
|
||||
coord JS would spin the running indicator forever on cancelled
|
||||
batches because ``state_change`` doesn't strip ``--running`` from
|
||||
individual batches."""
|
||||
|
||||
def _ui_with_tool_result_tracking(self):
|
||||
class _TrackingUI(NullUI):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.tool_results: list[tuple[str, str, str, bool]] = []
|
||||
|
||||
def on_tool_result(self, call_id, name, output, **kwargs):
|
||||
self.tool_results.append(
|
||||
(call_id, name, output, bool(kwargs.get("is_error", False))),
|
||||
)
|
||||
|
||||
return _TrackingUI()
|
||||
|
||||
def test_synthesizes_tool_result_for_unanswered_calls(self, tmp_db):
|
||||
ui = self._ui_with_tool_result_tracking()
|
||||
session = _make_session(ui=ui)
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "calling tools",
|
||||
"tool_calls": [
|
||||
{"id": "call_a", "function": {"name": "search", "arguments": "{}"}},
|
||||
{"id": "call_b", "function": {"name": "compute", "arguments": "{}"}},
|
||||
],
|
||||
},
|
||||
)
|
||||
session._msg_tokens.append(1)
|
||||
|
||||
session._synthesize_cancelled_results("Cancelled by user.")
|
||||
|
||||
# Both unanswered calls fired ``on_tool_result``.
|
||||
assert len(ui.tool_results) == 2
|
||||
ids = {tr[0] for tr in ui.tool_results}
|
||||
assert ids == {"call_a", "call_b"}
|
||||
# All emitted as errors so the live UI renders them as
|
||||
# ``coord-tool-row-result--error``.
|
||||
assert all(tr[3] is True for tr in ui.tool_results)
|
||||
# Reason text propagates as the synthetic tool output.
|
||||
assert all(tr[2] == "Cancelled by user." for tr in ui.tool_results)
|
||||
# And the message list has the synthesized tool entries
|
||||
# (preserves the prior contract).
|
||||
tool_msgs = [m for m in session.messages if m.get("role") == "tool"]
|
||||
assert len(tool_msgs) == 2
|
||||
|
||||
def test_skips_calls_already_answered(self, tmp_db):
|
||||
ui = self._ui_with_tool_result_tracking()
|
||||
session = _make_session(ui=ui)
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{"id": "call_a", "function": {"name": "search", "arguments": "{}"}},
|
||||
{"id": "call_b", "function": {"name": "compute", "arguments": "{}"}},
|
||||
],
|
||||
},
|
||||
)
|
||||
session._msg_tokens.append(1)
|
||||
# call_a already answered.
|
||||
session.messages.append(
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": "result"},
|
||||
)
|
||||
session._msg_tokens.append(1)
|
||||
|
||||
session._synthesize_cancelled_results("Cancelled by user.")
|
||||
|
||||
# Only call_b synthesized.
|
||||
assert len(ui.tool_results) == 1
|
||||
assert ui.tool_results[0][0] == "call_b"
|
||||
|
||||
def test_ui_emit_failure_does_not_break_synthesis(self, tmp_db):
|
||||
"""The UI hook is wrapped in try/except — a hook failure
|
||||
during cancel must NOT compound the problem. Synthesis still
|
||||
appends to messages + storage."""
|
||||
|
||||
class _ExplodingUI(NullUI):
|
||||
def on_tool_result(self, call_id, name, output, **kwargs):
|
||||
raise RuntimeError("ui hook blew up")
|
||||
|
||||
ui = _ExplodingUI()
|
||||
session = _make_session(ui=ui)
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{"id": "call_a", "function": {"name": "search", "arguments": "{}"}},
|
||||
],
|
||||
},
|
||||
)
|
||||
session._msg_tokens.append(1)
|
||||
|
||||
# Must not raise.
|
||||
session._synthesize_cancelled_results("Cancelled by user.")
|
||||
|
||||
tool_msgs = [m for m in session.messages if m.get("role") == "tool"]
|
||||
assert len(tool_msgs) == 1
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
"""ChatSession interactivity flag tests (Phase 9).
|
||||
|
||||
Validates that ``ChatSession._is_interactive_for_consent`` is computed
|
||||
correctly from ``client_type`` on construction. This is the front of
|
||||
the Phase 9 plumb-through: the flag flows from here to
|
||||
``_dispatch_pool_sync`` to the structured-error → pending-consent
|
||||
write path.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from tests._session_helpers import make_session
|
||||
from turnstone.prompts import INTERACTIVE_CONSENT_CLIENT_TYPES, ClientType
|
||||
|
||||
|
||||
def test_web_is_interactive() -> None:
|
||||
s = make_session(client_type=ClientType.WEB)
|
||||
assert s._is_interactive_for_consent is True
|
||||
|
||||
|
||||
def test_cli_is_interactive() -> None:
|
||||
s = make_session(client_type=ClientType.CLI)
|
||||
assert s._is_interactive_for_consent is True
|
||||
|
||||
|
||||
def test_chat_is_not_interactive() -> None:
|
||||
# Discord / Slack adapters cannot drive a browser redirect from
|
||||
# inside the channel — consent prompts must be deferred to the
|
||||
# dashboard badge.
|
||||
s = make_session(client_type=ClientType.CHAT)
|
||||
assert s._is_interactive_for_consent is False
|
||||
|
||||
|
||||
def test_scheduled_is_not_interactive() -> None:
|
||||
# The scheduler runs autonomously; the user isn't online to
|
||||
# complete the OAuth redirect.
|
||||
s = make_session(client_type=ClientType.SCHEDULED)
|
||||
assert s._is_interactive_for_consent is False
|
||||
|
||||
|
||||
def test_interactive_set_matches_module_constant() -> None:
|
||||
# Pin the module-level frozenset against the flag computation —
|
||||
# a future reorganisation that drifts the set vs the per-session
|
||||
# logic would silently break the gating.
|
||||
for ct in ClientType:
|
||||
s = make_session(client_type=ct)
|
||||
assert s._is_interactive_for_consent == (ct in INTERACTIVE_CONSENT_CLIENT_TYPES), ct
|
||||
|
||||
|
||||
def test_default_client_type_is_cli_interactive() -> None:
|
||||
# Defaults preserved — make_session uses ChatSession's default
|
||||
# which is CLI. Sanity check that the default user experience
|
||||
# stays interactive-for-consent.
|
||||
s = make_session()
|
||||
assert s._client_type == ClientType.CLI
|
||||
assert s._is_interactive_for_consent is True
|
||||
|
||||
|
||||
def test_scheduled_env_file_exists() -> None:
|
||||
"""The SCHEDULED env module must exist; otherwise
|
||||
``compose_system_message`` for a scheduled session would 500."""
|
||||
from turnstone.prompts import _load
|
||||
|
||||
text = _load("env/scheduled.md")
|
||||
assert "Output Environment" in text
|
||||
assert "consent" in text.lower()
|
||||
@@ -0,0 +1,215 @@
|
||||
"""Unit tests for :class:`turnstone.core.child_event_bus.ChildEventBus`.
|
||||
|
||||
The bus is the in-process wakeup primitive for ``wait_for_workstream``
|
||||
(see :mod:`turnstone.console.coordinator_client`). It's a small dict
|
||||
of ws_id → set[threading.Event] under a lock — focused tests for
|
||||
register/notify symmetry, no-subscriber notify, multi-waiter fan-out,
|
||||
multi-child waiter, and concurrent register/notify (smoke). End-to-end
|
||||
integration with the dispatch sink lives in
|
||||
``test_coordinator_adapter.py`` and ``test_coordinator_client.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
|
||||
def test_register_returns_event_that_starts_unset() -> None:
|
||||
"""A waiter must not see leftover state from before it registered —
|
||||
a fresh wait should always block until the first notify."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
assert isinstance(event, threading.Event)
|
||||
assert not event.is_set()
|
||||
|
||||
|
||||
def test_notify_wakes_waiter_on_matching_ws_id() -> None:
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
bus.notify("ws-1")
|
||||
assert event.is_set()
|
||||
|
||||
|
||||
def test_notify_does_not_wake_waiter_on_unrelated_ws_id() -> None:
|
||||
"""Different ws_ids must keep independent waiter sets — a notify on
|
||||
a stranger ws can't wake the wait or the bus stops being keyed."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
bus.notify("ws-other")
|
||||
assert not event.is_set()
|
||||
|
||||
|
||||
def test_notify_with_no_subscribers_is_noop() -> None:
|
||||
"""The dispatch sink calls notify on every translated event; the
|
||||
steady state has no wait tool active. Must not raise."""
|
||||
bus = ChildEventBus()
|
||||
bus.notify("ws-nobody-cares") # no exception
|
||||
|
||||
|
||||
def test_multi_waiter_each_gets_independent_event() -> None:
|
||||
"""Two waits on the same ws_id must wake independently — clearing
|
||||
one Event must not silence the other."""
|
||||
bus = ChildEventBus()
|
||||
e1 = bus.register_waiter(["ws-1"])
|
||||
e2 = bus.register_waiter(["ws-1"])
|
||||
assert e1 is not e2
|
||||
bus.notify("ws-1")
|
||||
assert e1.is_set()
|
||||
assert e2.is_set()
|
||||
|
||||
|
||||
def test_multi_child_waiter_fires_on_any_listed_ws_id() -> None:
|
||||
"""A wait on [A, B, C] returns a single Event registered against
|
||||
all three. Notify on ANY of A/B/C must wake the wait — the
|
||||
caller's snapshot re-read disambiguates which one changed."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-a", "ws-b", "ws-c"])
|
||||
bus.notify("ws-b")
|
||||
assert event.is_set()
|
||||
|
||||
|
||||
def test_unregister_removes_event_from_all_listed_ws_ids() -> None:
|
||||
"""After unregister, notify on any of the previously-watched ws_ids
|
||||
must NOT wake the Event — leaks would mean every future notify on
|
||||
that ws_id wakes a long-dead wait."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-a", "ws-b"])
|
||||
bus.unregister_waiter(["ws-a", "ws-b"], event)
|
||||
bus.notify("ws-a")
|
||||
bus.notify("ws-b")
|
||||
assert not event.is_set()
|
||||
|
||||
|
||||
def test_unregister_is_idempotent() -> None:
|
||||
"""A double-unregister must silently no-op — finally blocks may
|
||||
run twice in odd shutdown paths, the bus must not raise."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
bus.unregister_waiter(["ws-1"], event)
|
||||
bus.unregister_waiter(["ws-1"], event) # no exception
|
||||
|
||||
|
||||
def test_unregister_pops_empty_buckets() -> None:
|
||||
"""Empty per-ws_id buckets must be popped so a long-lived bus
|
||||
doesn't accumulate dead keys after many waits have churned through.
|
||||
Reaches into the private state — the property is structural, not
|
||||
behavioral, so the assertion is also."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
assert "ws-1" in bus._waiters
|
||||
bus.unregister_waiter(["ws-1"], event)
|
||||
assert "ws-1" not in bus._waiters
|
||||
|
||||
|
||||
def test_unregister_keeps_bucket_with_remaining_waiters() -> None:
|
||||
"""Removing one waiter from a multi-waiter bucket must not drop
|
||||
the others — popping the bucket would silently disable notifies
|
||||
for every concurrent wait on the same ws_id."""
|
||||
bus = ChildEventBus()
|
||||
e1 = bus.register_waiter(["ws-1"])
|
||||
e2 = bus.register_waiter(["ws-1"])
|
||||
bus.unregister_waiter(["ws-1"], e1)
|
||||
bus.notify("ws-1")
|
||||
assert not e1.is_set()
|
||||
assert e2.is_set()
|
||||
|
||||
|
||||
def test_empty_and_falsy_ws_ids_are_skipped_on_register() -> None:
|
||||
"""Defensive: ``wait_for_workstream`` cleans its inputs but the bus
|
||||
is reachable from other callers in future use; falsy ids should be
|
||||
silently dropped, not registered against an empty-string key."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["", "ws-1", ""])
|
||||
# Only the real ws_id should bucket the waiter.
|
||||
assert list(bus._waiters.keys()) == ["ws-1"]
|
||||
bus.notify("") # no crash, no spurious wake
|
||||
assert not event.is_set()
|
||||
bus.notify("ws-1")
|
||||
assert event.is_set()
|
||||
|
||||
|
||||
def test_notify_wakes_waiter_blocking_on_event_wait() -> None:
|
||||
"""End-to-end wake-up latency: a wait blocked on ``Event.wait``
|
||||
must return promptly after a notify on a watched ws_id. This is
|
||||
the property that retires the 0.5s polling cadence."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
woken_at = [0.0]
|
||||
|
||||
def _waiter() -> None:
|
||||
event.wait(timeout=2.0)
|
||||
woken_at[0] = time.monotonic()
|
||||
|
||||
t = threading.Thread(target=_waiter, daemon=True)
|
||||
t.start()
|
||||
# Give the waiter a beat to enter Event.wait, then notify.
|
||||
time.sleep(0.05)
|
||||
notified_at = time.monotonic()
|
||||
bus.notify("ws-1")
|
||||
t.join(timeout=1.0)
|
||||
assert not t.is_alive(), "waiter did not wake within 1s of notify"
|
||||
# Latency budget is generous; the contract is "well under the legacy
|
||||
# 0.5s poll cadence", not microsecond timing.
|
||||
assert woken_at[0] - notified_at < 0.2
|
||||
|
||||
|
||||
def test_clear_before_check_race_does_not_lose_wake() -> None:
|
||||
"""The wait-loop pattern is ``clear(); snapshot(); ...; wait()``.
|
||||
A notify between clear and wait must leave the Event set, so the
|
||||
next wait returns immediately and the loop re-snapshots. Same
|
||||
standard subscribe/check race the wait loop guards against."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
# Simulate wait-loop ordering: clear, then notify "between" clear
|
||||
# and the next wait.
|
||||
event.clear()
|
||||
bus.notify("ws-1")
|
||||
# The next wait must return True immediately (set is sticky until
|
||||
# the next clear).
|
||||
assert event.wait(timeout=0.1) is True
|
||||
|
||||
|
||||
def test_concurrent_register_and_notify_is_safe() -> None:
|
||||
"""Smoke test: many threads registering / notifying / unregistering
|
||||
in parallel must not raise or deadlock. Doesn't assert specific
|
||||
interleavings — only structural safety of the lock discipline."""
|
||||
bus = ChildEventBus()
|
||||
stop = threading.Event()
|
||||
errors: list[BaseException] = []
|
||||
|
||||
def _worker(ws_id: str) -> None:
|
||||
try:
|
||||
for _ in range(200):
|
||||
if stop.is_set():
|
||||
return
|
||||
ev = bus.register_waiter([ws_id])
|
||||
bus.notify(ws_id)
|
||||
bus.unregister_waiter([ws_id], ev)
|
||||
except BaseException as e: # noqa: BLE001
|
||||
errors.append(e)
|
||||
|
||||
threads = [threading.Thread(target=_worker, args=(f"ws-{i}",), daemon=True) for i in range(8)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=5.0)
|
||||
stop.set()
|
||||
assert not errors, f"worker threads raised: {errors!r}"
|
||||
# All buckets should have been popped (every register paired with
|
||||
# unregister).
|
||||
assert bus._waiters == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ws_id", ["", None])
|
||||
def test_notify_silently_ignores_falsy_ws_id(ws_id: object) -> None:
|
||||
"""Defensive: the dispatch sink already guards against empty
|
||||
ws_ids, but a falsy slip-through must not raise."""
|
||||
bus = ChildEventBus()
|
||||
event = bus.register_waiter(["ws-1"])
|
||||
bus.notify(ws_id) # type: ignore[arg-type]
|
||||
assert not event.is_set()
|
||||
@@ -0,0 +1,270 @@
|
||||
"""Unit tests for :mod:`turnstone.core.child_source`.
|
||||
|
||||
Covers both strategies in isolation against fakes — no live collector,
|
||||
no live SessionManager. Adapter-level integration coverage continues to
|
||||
live in ``test_coordinator_adapter.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core.child_source import ClusterChildSource, SameNodeChildSource
|
||||
from turnstone.core.children_registry import ChildrenRegistry
|
||||
from turnstone.core.workstream import WorkstreamState
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import queue
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SameNodeChildSource
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeManager:
|
||||
"""Minimal SessionManager stand-in implementing the subscribe API."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.subscribers: list[Any] = []
|
||||
|
||||
def subscribe_to_state(self, callback: Any) -> None:
|
||||
self.subscribers.append(callback)
|
||||
|
||||
def unsubscribe_from_state(self, callback: Any) -> None:
|
||||
with contextlib.suppress(ValueError):
|
||||
self.subscribers.remove(callback)
|
||||
|
||||
def fire(self, ws_id: str, state: WorkstreamState) -> None:
|
||||
for cb in self.subscribers:
|
||||
cb(ws_id, state)
|
||||
|
||||
|
||||
class TestSameNodeChildSource:
|
||||
def test_start_subscribes_to_manager(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
sink_calls: list[dict[str, Any]] = []
|
||||
src.start(sink=sink_calls.append)
|
||||
assert len(mgr.subscribers) == 1
|
||||
|
||||
def test_state_change_for_known_child_pushes_to_sink(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
registry.install("p1", object())
|
||||
registry.add_child("p1", "c1")
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
sink_calls: list[dict[str, Any]] = []
|
||||
src.start(sink=sink_calls.append)
|
||||
|
||||
mgr.fire("c1", WorkstreamState.RUNNING)
|
||||
|
||||
assert len(sink_calls) == 1
|
||||
ev = sink_calls[0]
|
||||
assert ev["type"] == "cluster_state"
|
||||
assert ev["ws_id"] == "c1"
|
||||
assert ev["state"] == "running"
|
||||
|
||||
def test_state_change_for_unknown_workstream_is_dropped(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
sink_calls: list[dict[str, Any]] = []
|
||||
src.start(sink=sink_calls.append)
|
||||
|
||||
# No registry entry — pre-filter drops the event without
|
||||
# invoking the sink.
|
||||
mgr.fire("ws-unknown", WorkstreamState.IDLE)
|
||||
assert sink_calls == []
|
||||
|
||||
def test_shutdown_unsubscribes(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
src.start(sink=lambda ev: None)
|
||||
assert len(mgr.subscribers) == 1
|
||||
src.shutdown()
|
||||
assert mgr.subscribers == []
|
||||
|
||||
def test_start_is_idempotent(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
src.start(sink=lambda ev: None)
|
||||
src.start(sink=lambda ev: None)
|
||||
# Second start is a no-op; only one subscription.
|
||||
assert len(mgr.subscribers) == 1
|
||||
|
||||
def test_sink_exception_does_not_propagate(self) -> None:
|
||||
mgr = _FakeManager()
|
||||
registry = ChildrenRegistry()
|
||||
registry.install("p1", object())
|
||||
registry.add_child("p1", "c1")
|
||||
src = SameNodeChildSource(mgr, registry)
|
||||
|
||||
def bad_sink(ev: dict[str, Any]) -> None:
|
||||
raise RuntimeError("sink boom")
|
||||
|
||||
src.start(sink=bad_sink)
|
||||
# Should not raise — the strategy catches sink failures and logs.
|
||||
mgr.fire("c1", WorkstreamState.RUNNING)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ClusterChildSource
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeCollector:
|
||||
"""Minimal ClusterCollector stand-in providing the listener API."""
|
||||
|
||||
def __init__(self, snapshot: dict[str, Any] | None = None) -> None:
|
||||
self._snapshot = snapshot or {"nodes": []}
|
||||
self.queues: list[queue.Queue[dict[str, Any]]] = []
|
||||
self.unregistered: list[queue.Queue[dict[str, Any]]] = []
|
||||
|
||||
def get_snapshot_and_register(self, q: queue.Queue[dict[str, Any]]) -> dict[str, Any]:
|
||||
self.queues.append(q)
|
||||
return self._snapshot
|
||||
|
||||
def unregister_listener(self, q: queue.Queue[dict[str, Any]]) -> None:
|
||||
self.unregistered.append(q)
|
||||
|
||||
def emit(self, event: dict[str, Any]) -> None:
|
||||
"""Push an event to all registered listener queues."""
|
||||
for q in self.queues:
|
||||
q.put(event)
|
||||
|
||||
|
||||
class TestClusterChildSource:
|
||||
def test_start_subscribes_to_collector(self) -> None:
|
||||
coll = _FakeCollector()
|
||||
registry = ChildrenRegistry()
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=list,
|
||||
)
|
||||
try:
|
||||
src.start(sink=lambda ev: None)
|
||||
assert len(coll.queues) == 1
|
||||
finally:
|
||||
src.shutdown()
|
||||
|
||||
def test_start_primes_registry_from_snapshot(self) -> None:
|
||||
snapshot = {
|
||||
"nodes": [
|
||||
{
|
||||
"workstreams": [
|
||||
{"id": "c1", "parent_ws_id": "p1"},
|
||||
{"id": "c2", "parent_ws_id": "p1"},
|
||||
# Unknown parent — dropped
|
||||
{"id": "x", "parent_ws_id": "p-unknown"},
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
coll = _FakeCollector(snapshot)
|
||||
registry = ChildrenRegistry()
|
||||
registry.install("p1", object())
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=lambda: ["p1"],
|
||||
)
|
||||
try:
|
||||
src.start(sink=lambda ev: None)
|
||||
assert set(registry.children_of("p1")) == {"c1", "c2"}
|
||||
assert registry.parent_for("x") is None
|
||||
finally:
|
||||
src.shutdown()
|
||||
|
||||
def test_event_dispatched_to_sink(self) -> None:
|
||||
coll = _FakeCollector()
|
||||
registry = ChildrenRegistry()
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=list,
|
||||
)
|
||||
sink_calls: list[dict[str, Any]] = []
|
||||
|
||||
try:
|
||||
src.start(sink=sink_calls.append)
|
||||
coll.emit({"type": "cluster_state", "ws_id": "c1", "state": "running"})
|
||||
# Daemon thread loop has 1.0s queue timeout; poll briefly.
|
||||
for _ in range(20):
|
||||
if sink_calls:
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert len(sink_calls) == 1
|
||||
assert sink_calls[0]["ws_id"] == "c1"
|
||||
finally:
|
||||
src.shutdown()
|
||||
|
||||
def test_shutdown_unregisters_and_joins_thread(self) -> None:
|
||||
coll = _FakeCollector()
|
||||
registry = ChildrenRegistry()
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=list,
|
||||
)
|
||||
src.start(sink=lambda ev: None)
|
||||
src.shutdown()
|
||||
assert coll.unregistered == coll.queues
|
||||
# Second shutdown is a no-op (idempotent).
|
||||
src.shutdown()
|
||||
|
||||
def test_start_is_idempotent(self) -> None:
|
||||
coll = _FakeCollector()
|
||||
registry = ChildrenRegistry()
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=list,
|
||||
)
|
||||
try:
|
||||
src.start(sink=lambda ev: None)
|
||||
src.start(sink=lambda ev: None)
|
||||
assert len(coll.queues) == 1
|
||||
finally:
|
||||
src.shutdown()
|
||||
|
||||
def test_sink_exception_does_not_kill_thread(self) -> None:
|
||||
coll = _FakeCollector()
|
||||
registry = ChildrenRegistry()
|
||||
src = ClusterChildSource(
|
||||
collector=coll,
|
||||
registry=registry,
|
||||
parents_provider=list,
|
||||
)
|
||||
survived_calls: list[dict[str, Any]] = []
|
||||
call_count = [0]
|
||||
|
||||
def flaky_sink(ev: dict[str, Any]) -> None:
|
||||
call_count[0] += 1
|
||||
if call_count[0] == 1:
|
||||
raise RuntimeError("first one boom")
|
||||
survived_calls.append(ev)
|
||||
|
||||
try:
|
||||
src.start(sink=flaky_sink)
|
||||
coll.emit({"type": "cluster_state", "ws_id": "c1", "state": "x"})
|
||||
coll.emit({"type": "cluster_state", "ws_id": "c2", "state": "y"})
|
||||
for _ in range(40):
|
||||
if survived_calls:
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert len(survived_calls) == 1
|
||||
assert survived_calls[0]["ws_id"] == "c2"
|
||||
finally:
|
||||
src.shutdown()
|
||||
|
||||
|
||||
# Multi-subscriber observer tests for ``SessionManager.subscribe_to_state``
|
||||
# / ``unsubscribe_from_state`` live in ``test_session_manager.py`` where
|
||||
# the proper FakeAdapter / FakeStorage construction helpers already exist.
|
||||
@@ -0,0 +1,239 @@
|
||||
"""Unit tests for :class:`turnstone.core.children_registry.ChildrenRegistry`.
|
||||
|
||||
The registry was lifted from ``CoordinatorAdapter`` in Stage 3 Step 1.
|
||||
Adapter-level coverage for the integrated behavior already lives in
|
||||
``test_coordinator_adapter.py``; this file pins the data structure
|
||||
invariants in isolation so the registry can be reused by future
|
||||
``ChildSource`` strategies (Step 2) without re-deriving the behavior
|
||||
from the adapter test surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.children_registry import ChildrenRegistry
|
||||
|
||||
|
||||
class _Sentinel:
|
||||
"""Lightweight UI stand-in; identity-comparable, no behavior."""
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registry() -> ChildrenRegistry:
|
||||
return ChildrenRegistry()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# install / uninstall
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInstallUninstall:
|
||||
def test_install_seeds_empty_child_set_and_presence(self, registry: ChildrenRegistry) -> None:
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
assert registry.children_of("p1") == []
|
||||
assert registry.ui_for("p1") is ui
|
||||
assert registry.parents() == ["p1"]
|
||||
|
||||
def test_install_is_idempotent_repoints_ui_keeps_children(
|
||||
self, registry: ChildrenRegistry
|
||||
) -> None:
|
||||
ui_a = _Sentinel()
|
||||
ui_b = _Sentinel()
|
||||
registry.install("p1", ui_a)
|
||||
registry.merge_children("p1", ["c1", "c2"])
|
||||
registry.install("p1", ui_b)
|
||||
assert registry.ui_for("p1") is ui_b
|
||||
assert set(registry.children_of("p1")) == {"c1", "c2"}
|
||||
|
||||
def test_uninstall_clears_forward_reverse_and_presence(
|
||||
self, registry: ChildrenRegistry
|
||||
) -> None:
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
registry.merge_children("p1", ["c1", "c2"])
|
||||
registry.uninstall("p1")
|
||||
assert registry.children_of("p1") == []
|
||||
assert registry.ui_for("p1") is None
|
||||
assert registry.parents() == []
|
||||
assert registry.parent_for("c1") is None
|
||||
assert registry.parent_for("c2") is None
|
||||
|
||||
def test_uninstall_unknown_parent_is_noop(self, registry: ChildrenRegistry) -> None:
|
||||
registry.uninstall("never-installed") # must not raise
|
||||
|
||||
def test_uninstall_does_not_clobber_other_parents(self, registry: ChildrenRegistry) -> None:
|
||||
registry.install("p1", _Sentinel())
|
||||
registry.install("p2", _Sentinel())
|
||||
registry.merge_children("p1", ["c1"])
|
||||
registry.merge_children("p2", ["c2"])
|
||||
registry.uninstall("p1")
|
||||
assert registry.parent_for("c1") is None
|
||||
assert registry.parent_for("c2") == "p2"
|
||||
assert registry.parents() == ["p2"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# add_child — atomic check-and-route
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAddChild:
|
||||
def test_add_child_returns_ui_on_success(self, registry: ChildrenRegistry) -> None:
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
assert registry.add_child("p1", "c1") is ui
|
||||
assert registry.parent_for("c1") == "p1"
|
||||
assert registry.children_of("p1") == ["c1"]
|
||||
|
||||
def test_add_child_returns_none_when_parent_not_installed(
|
||||
self, registry: ChildrenRegistry
|
||||
) -> None:
|
||||
assert registry.add_child("absent", "c1") is None
|
||||
assert registry.parent_for("c1") is None
|
||||
|
||||
def test_add_child_returns_none_on_duplicate(self, registry: ChildrenRegistry) -> None:
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
assert registry.add_child("p1", "c1") is ui
|
||||
# second add for same child returns None — caller must not
|
||||
# double-dispatch.
|
||||
assert registry.add_child("p1", "c1") is None
|
||||
assert registry.children_of("p1") == ["c1"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# merge_children — bulk seeding
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMergeChildren:
|
||||
def test_merge_seeds_forward_and_reverse(self, registry: ChildrenRegistry) -> None:
|
||||
registry.merge_children("p1", ["c1", "c2", "c3"])
|
||||
assert set(registry.children_of("p1")) == {"c1", "c2", "c3"}
|
||||
for cid in ("c1", "c2", "c3"):
|
||||
assert registry.parent_for(cid) == "p1"
|
||||
|
||||
def test_merge_is_idempotent(self, registry: ChildrenRegistry) -> None:
|
||||
registry.merge_children("p1", ["c1"])
|
||||
registry.merge_children("p1", ["c1"])
|
||||
assert registry.children_of("p1") == ["c1"]
|
||||
|
||||
def test_merge_skips_empty_or_falsy_ids(self, registry: ChildrenRegistry) -> None:
|
||||
registry.merge_children("p1", ["", "c1", "", "c2"])
|
||||
assert set(registry.children_of("p1")) == {"c1", "c2"}
|
||||
|
||||
def test_merge_does_not_require_install(self, registry: ChildrenRegistry) -> None:
|
||||
# Snapshot-priming may run before the parent's install fires —
|
||||
# the merge still seeds the forward set so the install picks
|
||||
# the children up. (Storage-seeded rebuild relies on this.)
|
||||
registry.merge_children("p1", ["c1"])
|
||||
assert registry.children_of("p1") == ["c1"]
|
||||
# ui_for is still None because install hasn't run
|
||||
assert registry.ui_for("p1") is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lookups — return copies, not live refs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLookups:
|
||||
def test_children_of_returns_copy(self, registry: ChildrenRegistry) -> None:
|
||||
registry.install("p1", _Sentinel())
|
||||
registry.merge_children("p1", ["c1", "c2"])
|
||||
snap = registry.children_of("p1")
|
||||
snap.append("c3-injected")
|
||||
assert "c3-injected" not in registry.children_of("p1")
|
||||
|
||||
def test_children_of_unknown_parent_returns_empty(self, registry: ChildrenRegistry) -> None:
|
||||
assert registry.children_of("absent") == []
|
||||
|
||||
def test_parent_for_unknown_child_returns_none(self, registry: ChildrenRegistry) -> None:
|
||||
assert registry.parent_for("absent") is None
|
||||
|
||||
def test_parents_returns_copy(self, registry: ChildrenRegistry) -> None:
|
||||
registry.install("p1", _Sentinel())
|
||||
snap = registry.parents()
|
||||
snap.append("p2-injected")
|
||||
assert "p2-injected" not in registry.parents()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Concurrency — concurrent add_child must not exceed the unique-set
|
||||
# invariant or leave a half-installed reverse-index entry.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConcurrency:
|
||||
def test_concurrent_add_child_returns_ui_exactly_once_per_unique(
|
||||
self, registry: ChildrenRegistry
|
||||
) -> None:
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
results: list[object] = []
|
||||
results_lock = threading.Lock()
|
||||
|
||||
def attempt_add(child_id: str) -> None:
|
||||
r = registry.add_child("p1", child_id)
|
||||
with results_lock:
|
||||
results.append(r)
|
||||
|
||||
threads = [threading.Thread(target=attempt_add, args=("c1",)) for _ in range(20)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
# Exactly one thread sees the UI; the remaining 19 see None
|
||||
# (duplicate). The forward + reverse indexes carry exactly one
|
||||
# entry for c1.
|
||||
successes = [r for r in results if r is ui]
|
||||
nones = [r for r in results if r is None]
|
||||
assert len(successes) == 1
|
||||
assert len(nones) == 19
|
||||
assert registry.children_of("p1") == ["c1"]
|
||||
assert registry.parent_for("c1") == "p1"
|
||||
|
||||
def test_concurrent_install_and_add_child_no_resurrect(
|
||||
self, registry: ChildrenRegistry
|
||||
) -> None:
|
||||
# add_child racing with uninstall: either lands first (registry
|
||||
# populated) or the parent is gone (returns None). Must NOT
|
||||
# leave a forward-set entry without presence — that would be
|
||||
# the "resurrected after close" leak the locked dispatch path
|
||||
# was guarding against.
|
||||
ui = _Sentinel()
|
||||
registry.install("p1", ui)
|
||||
|
||||
outcomes: list[object] = []
|
||||
|
||||
def adder() -> None:
|
||||
outcomes.append(registry.add_child("p1", "c1"))
|
||||
|
||||
def uninstaller() -> None:
|
||||
registry.uninstall("p1")
|
||||
|
||||
threads = [
|
||||
threading.Thread(target=adder),
|
||||
threading.Thread(target=uninstaller),
|
||||
]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
# If add_child landed first: c1 is in the forward set, then
|
||||
# uninstall clears everything. End state: nothing.
|
||||
# If uninstall landed first: add_child sees no presence,
|
||||
# returns None, no entry added. End state: nothing.
|
||||
# Either way, the leak invariant holds: child set is empty or
|
||||
# parent is gone, never "child set populated but no presence".
|
||||
children = registry.children_of("p1")
|
||||
ui_present = registry.ui_for("p1") is not None
|
||||
if children:
|
||||
assert ui_present, "registry leaked: children set without presence"
|
||||
@@ -30,9 +30,25 @@ def _full_hdr() -> dict[str, str]:
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_metrics(monkeypatch):
|
||||
"""Swap ``turnstone.server._metrics`` for a fresh collector
|
||||
per-test, with auto-restore.
|
||||
|
||||
Bare ``srv_mod._metrics = MetricsCollector()`` (the prior
|
||||
pattern) leaks into any test file that already bound the name
|
||||
via ``from turnstone.server import _metrics`` at import time —
|
||||
those tests' patches then operate on a different instance from
|
||||
the one the live ``_publish_models_metadata`` reads, and the
|
||||
monkeypatch silently no-ops. ``monkeypatch.setattr`` restores
|
||||
after the test, so the leak is contained.
|
||||
"""
|
||||
fresh = MetricsCollector()
|
||||
fresh.model = "test-model"
|
||||
monkeypatch.setattr(srv_mod, "_metrics", fresh)
|
||||
|
||||
|
||||
def _make_app(storage: Any) -> TestClient:
|
||||
srv_mod._metrics = MetricsCollector()
|
||||
srv_mod._metrics.model = "test-model"
|
||||
mock_session = MagicMock()
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = "ws-target"
|
||||
|
||||
@@ -50,6 +50,41 @@ def test_load_config_invalid_toml(tmp_path):
|
||||
assert load_config() == {}
|
||||
|
||||
|
||||
def test_load_config_warns_when_world_readable(tmp_path, caplog):
|
||||
"""Secrets in config.toml — warn if anyone but the owner can read it."""
|
||||
import logging
|
||||
import os
|
||||
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[database]\nurl = "postgresql+psycopg://u:secret@h/d"\n')
|
||||
os.chmod(cfg, 0o644)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="turnstone.core.config"):
|
||||
load_config()
|
||||
|
||||
messages = [r.getMessage() for r in caplog.records]
|
||||
assert any("group/world-readable" in m for m in messages)
|
||||
|
||||
|
||||
def test_load_config_quiet_when_mode_0600(tmp_path, caplog):
|
||||
import logging
|
||||
import os
|
||||
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
cfg.write_text('[database]\nurl = "postgresql+psycopg://u:secret@h/d"\n')
|
||||
os.chmod(cfg, 0o600)
|
||||
set_config_path(str(cfg))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="turnstone.core.config"):
|
||||
load_config()
|
||||
|
||||
messages = [r.getMessage() for r in caplog.records]
|
||||
assert not any("group/world-readable" in m for m in messages)
|
||||
|
||||
|
||||
def test_load_config_caches(tmp_path):
|
||||
_reset_cache()
|
||||
cfg = tmp_path / "config.toml"
|
||||
|
||||
+489
-19
@@ -3,11 +3,13 @@
|
||||
import asyncio
|
||||
import json
|
||||
import queue
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.console.collector import ClusterCollector, NodeSnapshot
|
||||
from turnstone.console.server import _PROXY_AUTH_LOCAL_HANDLERS
|
||||
|
||||
# Shared test auth — JWT-based
|
||||
_TEST_JWT_SECRET = "test-jwt-secret-minimum-32-chars!"
|
||||
@@ -149,6 +151,78 @@ class TestCollectorDiscovery:
|
||||
assert c._nodes["node-a"].started == 1234567890.0
|
||||
|
||||
|
||||
class TestCollectorNotifyWireIn:
|
||||
"""NotifyDispatcher-driven discovery — reactive node visibility."""
|
||||
|
||||
def test_start_subscribes_to_services_channel(self):
|
||||
# Stub dispatcher records subscriptions without spawning threads.
|
||||
class _StubDispatcher:
|
||||
def __init__(self):
|
||||
self.subscriptions: list[tuple[str, Any]] = []
|
||||
|
||||
def subscribe(self, channel, handler):
|
||||
self.subscriptions.append((channel, handler))
|
||||
return lambda: None
|
||||
|
||||
stub = _StubDispatcher()
|
||||
storage = MockStorage()
|
||||
c = ClusterCollector(
|
||||
storage=storage,
|
||||
discovery_interval=999,
|
||||
notify_dispatcher=stub,
|
||||
)
|
||||
try:
|
||||
c.start()
|
||||
assert len(stub.subscriptions) == 1
|
||||
channel, handler = stub.subscriptions[0]
|
||||
assert channel == "services"
|
||||
assert handler == c._on_services_notify
|
||||
finally:
|
||||
c.stop()
|
||||
|
||||
def test_no_dispatcher_means_no_subscribe(self):
|
||||
# Collector without a dispatcher (single-node / SQLite dev) just
|
||||
# falls back to the 60 s discovery-loop polling — no error.
|
||||
c = _make_collector(MockStorage())
|
||||
try:
|
||||
c.start()
|
||||
assert c._notify_unsubscribe is None
|
||||
finally:
|
||||
c.stop()
|
||||
|
||||
def test_on_notify_runs_discovery(self):
|
||||
# Construct a synthetic Notify and invoke the handler directly —
|
||||
# asserts the wire-in delegates back to ``_discover_nodes``.
|
||||
from turnstone.core.storage._notify import Notify
|
||||
|
||||
storage = MockStorage()
|
||||
c = _make_collector(storage)
|
||||
c._running = True # bypass start() so we don't spawn threads
|
||||
q: queue.Queue[dict[str, Any]] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
storage.services = [
|
||||
{"service_id": "node-z", "url": "http://z:8080", "metadata": "{}"},
|
||||
]
|
||||
c._on_services_notify(Notify(channel="services", payload="{}", pid=0))
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "node_joined"
|
||||
assert event["node_id"] == "node-z"
|
||||
|
||||
def test_on_notify_when_not_running_is_noop(self):
|
||||
# If a stray notify arrives after stop, the handler doesn't run
|
||||
# discovery on a half-torn-down collector.
|
||||
from turnstone.core.storage._notify import Notify
|
||||
|
||||
storage = MockStorage()
|
||||
storage.services = [{"service_id": "node-y", "url": "http://y:8080", "metadata": "{}"}]
|
||||
c = _make_collector(storage)
|
||||
# _running stays False (never called start()).
|
||||
c._on_services_notify(Notify(channel="services", payload="{}", pid=0))
|
||||
assert c.get_overview()["nodes"] == 0
|
||||
|
||||
|
||||
class TestCollectorSnapshot:
|
||||
"""Applying node_snapshot SSE events."""
|
||||
|
||||
@@ -314,6 +388,46 @@ class TestCollectorSnapshot:
|
||||
assert event["ws_id"] == "ws1"
|
||||
assert event["state"] == "running"
|
||||
|
||||
def test_apply_snapshot_state_change_does_not_carry_pending_approval_detail(self):
|
||||
"""Stage 3 cleanup — the snapshot-resync cluster_state event no
|
||||
longer piggybacks ``pending_approval_detail`` (the field is
|
||||
gone from cluster_state entirely). On reconnect the browser's
|
||||
bulk fetch — triggered by the ``activity_state="approval"``
|
||||
transition in the reducer — pulls the items directly from
|
||||
``ui.serialize_pending_approval_detail()`` via the dashboard
|
||||
endpoint."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "same", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_snapshot(
|
||||
"node-a",
|
||||
{
|
||||
"type": "node_snapshot",
|
||||
"node_id": "node-a",
|
||||
"workstreams": [
|
||||
{
|
||||
"id": "ws1",
|
||||
"name": "same",
|
||||
"state": "running",
|
||||
"activity_state": "approval",
|
||||
}
|
||||
],
|
||||
"health": {},
|
||||
"aggregate": {},
|
||||
},
|
||||
)
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "cluster_state"
|
||||
assert event["activity_state"] == "approval"
|
||||
assert "pending_approval_detail" not in event
|
||||
|
||||
def test_apply_snapshot_skips_empty_id_workstream(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
@@ -359,6 +473,36 @@ class TestCollectorDelta:
|
||||
# Verify in-memory state was updated
|
||||
assert c._nodes["node-a"].workstreams["ws1"]["state"] == "running"
|
||||
|
||||
def test_apply_delta_ws_state_does_not_carry_pending_approval_detail(self):
|
||||
"""Stage 3 cleanup — ``cluster_state`` no longer carries the
|
||||
``pending_approval_detail`` piggyback. Approval items now arrive
|
||||
via bulk fetch on activity_state transition; verdicts via the
|
||||
explicit ``intent_verdict`` event class. Symmetric event flow,
|
||||
no piggyback to dedupe against."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_delta(
|
||||
"node-a",
|
||||
{
|
||||
"type": "ws_state",
|
||||
"ws_id": "ws1",
|
||||
"state": "running",
|
||||
"activity_state": "approval",
|
||||
},
|
||||
)
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "cluster_state"
|
||||
assert event["activity_state"] == "approval"
|
||||
assert "pending_approval_detail" not in event
|
||||
|
||||
def test_apply_delta_ws_created(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
@@ -404,6 +548,139 @@ class TestCollectorDelta:
|
||||
assert event["name"] == "new-name"
|
||||
assert c._nodes["node-a"].workstreams["ws1"]["name"] == "new-name"
|
||||
|
||||
def test_apply_delta_intent_verdict_forwards_verbatim(self):
|
||||
"""Stage 3 Step 5 — node-emitted intent_verdict events flow
|
||||
through _apply_delta to cluster fan-out so coord adapters can
|
||||
re-emit as child_ws_intent_verdict on the parent's SSE."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
verdict = {
|
||||
"call_id": "c1",
|
||||
"risk_level": "low",
|
||||
"confidence": 0.9,
|
||||
"recommendation": "approve",
|
||||
}
|
||||
c._apply_delta(
|
||||
"node-a",
|
||||
{"type": "intent_verdict", "ws_id": "ws1", "verdict": verdict},
|
||||
)
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "intent_verdict"
|
||||
assert event["ws_id"] == "ws1"
|
||||
assert event["node_id"] == "node-a"
|
||||
assert event["verdict"] == verdict
|
||||
|
||||
def test_apply_delta_intent_verdict_drops_when_ws_id_missing(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_delta("node-a", {"type": "intent_verdict", "verdict": {}})
|
||||
|
||||
assert q.empty()
|
||||
|
||||
def test_apply_delta_approval_resolved_forwards_verbatim(self):
|
||||
"""Stage 3 Step 5 — paired with intent_verdict; clears the
|
||||
coord tree's pending-approval pill in lockstep with the
|
||||
actual decision rather than waiting for the state-change
|
||||
piggyback."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_delta(
|
||||
"node-a",
|
||||
{
|
||||
"type": "approval_resolved",
|
||||
"ws_id": "ws1",
|
||||
"approved": True,
|
||||
"feedback": "lgtm",
|
||||
"always": False,
|
||||
},
|
||||
)
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "approval_resolved"
|
||||
assert event["ws_id"] == "ws1"
|
||||
assert event["node_id"] == "node-a"
|
||||
assert event["approved"] is True
|
||||
assert event["feedback"] == "lgtm"
|
||||
assert event["always"] is False
|
||||
|
||||
def test_apply_delta_approve_request_forwards_detail(self):
|
||||
"""Push path for the initial approval items — eliminates the
|
||||
bulk-fetch race that left the coord row stuck on a loading
|
||||
placeholder when the bulk fetch landed in the gap between
|
||||
_emit_state(ATTENTION) and approve_tools setting _pending_approval."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
detail = {
|
||||
"type": "approve_request",
|
||||
"items": [{"call_id": "c1", "header": "tool x"}],
|
||||
"judge_pending": True,
|
||||
}
|
||||
c._apply_delta(
|
||||
"node-a",
|
||||
{"type": "approve_request", "ws_id": "ws1", "detail": detail},
|
||||
)
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["type"] == "approve_request"
|
||||
assert event["ws_id"] == "ws1"
|
||||
assert event["node_id"] == "node-a"
|
||||
assert event["detail"] == detail
|
||||
|
||||
def test_apply_delta_approve_request_drops_when_ws_id_missing(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080")
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_delta("node-a", {"type": "approve_request", "detail": {}})
|
||||
|
||||
assert q.empty()
|
||||
|
||||
def test_apply_delta_approval_resolved_coerces_missing_fields(self):
|
||||
"""Defensive: ``approved`` / ``always`` / ``feedback`` may be
|
||||
omitted by older nodes mid-rolling-upgrade; collector coerces
|
||||
to safe defaults."""
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
node_id="node-a",
|
||||
server_url="http://a:8080",
|
||||
workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle"}},
|
||||
)
|
||||
q: queue.Queue[dict] = queue.Queue()
|
||||
c.register_listener(q)
|
||||
|
||||
c._apply_delta("node-a", {"type": "approval_resolved", "ws_id": "ws1"})
|
||||
|
||||
event = q.get_nowait()
|
||||
assert event["approved"] is False
|
||||
assert event["feedback"] == ""
|
||||
assert event["always"] is False
|
||||
|
||||
def test_apply_delta_health_changed(self):
|
||||
c = _make_collector()
|
||||
c._nodes["node-a"] = NodeSnapshot(
|
||||
@@ -1286,6 +1563,142 @@ class TestConsoleProxy:
|
||||
assert sse_mock.await_count == 1
|
||||
assert sse_mock.await_args.kwargs.get("use_service_auth") is False
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Proxied auth endpoints — handled locally by the console, not
|
||||
# forwarded to the upstream node. Cases derive directly from
|
||||
# ``_PROXY_AUTH_LOCAL_HANDLERS`` so a new dispatch entry can't be
|
||||
# added without a matching test (or vice versa). See proxy_api's
|
||||
# docstring for the JWT-audience reasoning.
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("method", "path", "handler_name"),
|
||||
[
|
||||
(method, path, handler_name)
|
||||
for (method, path), handler_name in sorted(_PROXY_AUTH_LOCAL_HANDLERS.items())
|
||||
],
|
||||
)
|
||||
def test_proxy_auth_endpoint_dispatches_to_local_handler(
|
||||
self, client, method, path, handler_name
|
||||
):
|
||||
"""Every entry in ``_PROXY_AUTH_LOCAL_HANDLERS`` must route to its
|
||||
local console handler and never reach the upstream proxy. The
|
||||
lockout class of bug this dispatch was added to fix is exactly
|
||||
what a regression here would reintroduce silently — covering all
|
||||
eight branches keeps each path tied to its handler."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from starlette.responses import JSONResponse
|
||||
|
||||
with (
|
||||
patch(
|
||||
f"turnstone.console.server.{handler_name}",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "ok"}),
|
||||
) as local_mock,
|
||||
patch(
|
||||
"turnstone.console.server._proxy_post",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "should-not-be-called"}),
|
||||
) as post_mock,
|
||||
patch(
|
||||
"turnstone.console.server._proxy_get",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "should-not-be-called"}),
|
||||
) as get_mock,
|
||||
):
|
||||
resp = client.request(method, f"/node/node-a/v1/api/{path}")
|
||||
assert resp.status_code == 200
|
||||
assert local_mock.await_count == 1
|
||||
assert post_mock.await_count == 0
|
||||
assert get_mock.await_count == 0
|
||||
|
||||
def test_proxy_auth_login_works_without_cookie(self, mock_collector):
|
||||
"""Without this fix the AuthMiddleware 401s before any handler
|
||||
runs — the user is locked out of the proxied UI once the cookie
|
||||
expires. Test bypasses _TEST_AUTH_HEADERS to reproduce."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from turnstone.console.server import _load_static, create_app
|
||||
|
||||
_load_static()
|
||||
app = create_app(collector=mock_collector, jwt_secret=_TEST_JWT_SECRET)
|
||||
unauth_client = TestClient(app, raise_server_exceptions=False)
|
||||
try:
|
||||
with patch(
|
||||
"turnstone.console.server.auth_login",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "ok"}),
|
||||
) as local_mock:
|
||||
resp = unauth_client.post(
|
||||
"/node/node-a/v1/api/auth/login",
|
||||
json={"username": "x", "password": "y"},
|
||||
)
|
||||
# AuthMiddleware must classify the proxied login path as
|
||||
# public (is_public_path change) AND proxy_api must
|
||||
# dispatch to the local handler (proxy_api change).
|
||||
assert resp.status_code == 200, (
|
||||
f"login locked out: got {resp.status_code}, body={resp.text}"
|
||||
)
|
||||
assert local_mock.await_count == 1
|
||||
finally:
|
||||
unauth_client.close()
|
||||
|
||||
def test_proxy_auth_wrong_method_returns_405_not_forwarded(self, client):
|
||||
"""A non-canonical method on an auth path (e.g. PUT on auth/login)
|
||||
must short-circuit with 405 instead of falling through to the
|
||||
upstream proxy — falling through would forward the request
|
||||
authenticated as the console's service token (``_proxy_auth_headers``
|
||||
fallback)."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from starlette.responses import JSONResponse
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.console.server._proxy_post",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "should-not-be-called"}),
|
||||
) as post_mock,
|
||||
patch(
|
||||
"turnstone.console.server._proxy_get",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"status": "should-not-be-called"}),
|
||||
) as get_mock,
|
||||
):
|
||||
# PUT on a POST-only auth path → 405
|
||||
put_resp = client.put("/node/node-a/v1/api/auth/login")
|
||||
assert put_resp.status_code == 405
|
||||
# POST on a GET-only auth path → 405
|
||||
post_resp = client.post("/node/node-a/v1/api/auth/status")
|
||||
assert post_resp.status_code == 405
|
||||
assert post_mock.await_count == 0
|
||||
assert get_mock.await_count == 0
|
||||
|
||||
def test_proxy_non_auth_endpoint_still_forwarded(self, client, mock_collector):
|
||||
"""Sanity: only auth/* paths intercept. Other API paths still
|
||||
forward to the upstream node."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from starlette.responses import JSONResponse
|
||||
|
||||
mock_collector.get_node_detail.return_value = {
|
||||
"node_id": "node-a",
|
||||
"server_url": "http://a:8080",
|
||||
"reachable": True,
|
||||
}
|
||||
with patch(
|
||||
"turnstone.console.server._proxy_get",
|
||||
new_callable=AsyncMock,
|
||||
return_value=JSONResponse({"ok": True}),
|
||||
) as proxy_mock:
|
||||
resp = client.get("/node/node-a/v1/api/workstreams")
|
||||
assert resp.status_code == 200
|
||||
assert proxy_mock.await_count == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Proxy URL rewriting unit tests (no HTTP needed)
|
||||
@@ -1309,11 +1722,62 @@ class TestProxyRewriting:
|
||||
assert "window.fetch" in _JS_PROXY_SHIM
|
||||
assert "window.EventSource" in _JS_PROXY_SHIM
|
||||
|
||||
def test_console_banner_contains_placeholder(self):
|
||||
from turnstone.console.server import _CONSOLE_BANNER_TEMPLATE
|
||||
def test_js_shim_carries_node_id_placeholder(self):
|
||||
"""The picker reads the current node_id from the shim's _nodeId
|
||||
closure variable; the placeholder must be present and substitutable."""
|
||||
from turnstone.console.server import _JS_PROXY_SHIM
|
||||
|
||||
assert "NODE_ID_PLACEHOLDER" in _CONSOLE_BANNER_TEMPLATE
|
||||
assert "Console" in _CONSOLE_BANNER_TEMPLATE
|
||||
assert "NODE_ID_PLACEHOLDER" in _JS_PROXY_SHIM
|
||||
replaced = _JS_PROXY_SHIM.replace("NODE_ID_PLACEHOLDER", "node-a")
|
||||
assert "node-a" in replaced
|
||||
assert "NODE_ID_PLACEHOLDER" not in replaced
|
||||
|
||||
def test_js_shim_includes_picker_pieces(self):
|
||||
"""Picker logic ships in the same IIFE as the prefix shim — verify
|
||||
the moving parts are present so a future refactor doesn't silently
|
||||
drop them. /v1/api/cluster/nodes is the lazy-fetch target;
|
||||
#ui-header is the DOM anchor; console-node-pill is the trigger
|
||||
class; ws-tab-dropdown is the menu shell we share with the
|
||||
workstream chevron menu (style + behaviour parity); ArrowDown is
|
||||
the keyboard-nav primitive that disambiguates this from a plain
|
||||
click-only menu."""
|
||||
from turnstone.console.server import _JS_PROXY_SHIM
|
||||
|
||||
# limit=1000 matches the collector's hard cap; without it the
|
||||
# picker would silently drop nodes past the 100-default in
|
||||
# clusters with >100 nodes.
|
||||
assert "/v1/api/cluster/nodes?limit=1000" in _JS_PROXY_SHIM
|
||||
assert "ui-header" in _JS_PROXY_SHIM
|
||||
assert "console-node-pill" in _JS_PROXY_SHIM
|
||||
assert "ws-tab-dropdown" in _JS_PROXY_SHIM
|
||||
assert "ArrowDown" in _JS_PROXY_SHIM
|
||||
assert "DOMContentLoaded" in _JS_PROXY_SHIM
|
||||
|
||||
def test_proxy_style_drops_banner_styles(self):
|
||||
"""The legacy banner CSS classes (.console-banner, .ts-header-back-link
|
||||
offsets, .dashboard-overlay top:32px hack) should be gone — the new
|
||||
picker lives inside #ui-header and doesn't need overlay offsets."""
|
||||
from turnstone.console.server import _CONSOLE_PROXY_STYLE
|
||||
|
||||
assert ".console-banner" not in _CONSOLE_PROXY_STYLE
|
||||
assert "dashboard-overlay" not in _CONSOLE_PROXY_STYLE
|
||||
assert ".console-node-pill" in _CONSOLE_PROXY_STYLE
|
||||
assert ".console-node-menu" in _CONSOLE_PROXY_STYLE
|
||||
|
||||
def test_proxy_style_uses_canonical_degraded_color(self):
|
||||
"""Degraded health dot must use --accent (the canonical "needs
|
||||
attention" token used by the cluster-overview node table at
|
||||
console/static/style.css:548) and not --yellow. Yellow is reserved
|
||||
for the dash-state attention dot, a stronger signal."""
|
||||
from turnstone.console.server import _CONSOLE_PROXY_STYLE
|
||||
|
||||
assert "console-node-menu-item-dot--degraded" in _CONSOLE_PROXY_STYLE
|
||||
# The degraded rule sits on its own line; assert it uses --accent
|
||||
# by checking the CSS substring has --accent and not --yellow.
|
||||
idx = _CONSOLE_PROXY_STYLE.find("console-node-menu-item-dot--degraded")
|
||||
rule = _CONSOLE_PROXY_STYLE[idx : idx + 200]
|
||||
assert "var(--accent)" in rule
|
||||
assert "var(--yellow)" not in rule
|
||||
|
||||
def test_html_rewriting_changes_static_paths(self):
|
||||
"""Simulate the proxy_index rewriting logic."""
|
||||
@@ -1330,16 +1794,24 @@ class TestProxyRewriting:
|
||||
assert 'href="/static/' not in rewritten
|
||||
assert 'src="/static/' not in rewritten
|
||||
|
||||
def test_banner_injection_after_body(self):
|
||||
"""Simulate the banner injection logic."""
|
||||
from turnstone.console.server import _CONSOLE_BANNER_TEMPLATE
|
||||
def test_shim_injection_after_body(self):
|
||||
"""Simulate the proxy shim injection — the shim ships the node-id
|
||||
and prefix as JS literals and renders the picker at runtime, so
|
||||
we assert the substituted JS literals land in the page."""
|
||||
from turnstone.console.server import _CONSOLE_PROXY_STYLE, _JS_PROXY_SHIM
|
||||
|
||||
sample_html = "<html><body><div>content</div></body></html>"
|
||||
banner = _CONSOLE_BANNER_TEMPLATE.replace("NODE_ID_PLACEHOLDER", "node-a")
|
||||
result = sample_html.replace("<body>", "<body>" + banner, 1)
|
||||
assert "node-a" in result
|
||||
assert "Console" in result
|
||||
assert result.startswith("<html><body><div")
|
||||
prefix = "/node/node-a"
|
||||
shim_js = _JS_PROXY_SHIM.replace('"PREFIX_PLACEHOLDER"', json.dumps(prefix)).replace(
|
||||
'"NODE_ID_PLACEHOLDER"', json.dumps("node-a")
|
||||
)
|
||||
injection = _CONSOLE_PROXY_STYLE + "<script>" + shim_js + "</script>"
|
||||
result = sample_html.replace("<body>", "<body>" + injection, 1)
|
||||
assert '"node-a"' in result
|
||||
assert '"/node/node-a"' in result
|
||||
assert "PREFIX_PLACEHOLDER" not in result
|
||||
assert "NODE_ID_PLACEHOLDER" not in result
|
||||
assert result.startswith("<html><body><style>")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1574,17 +2046,15 @@ class TestProxySharedStatic:
|
||||
def test_proxy_shim_injected_in_html(self):
|
||||
"""Verify shim is injected as inline script in proxied HTML."""
|
||||
|
||||
from turnstone.console.server import _CONSOLE_BANNER_TEMPLATE, _JS_PROXY_SHIM
|
||||
from turnstone.console.server import _JS_PROXY_SHIM
|
||||
|
||||
sample_html = "<html><body><div>content</div></body></html>"
|
||||
prefix = "/node/test-node"
|
||||
banner = _CONSOLE_BANNER_TEMPLATE.replace("NODE_ID_PLACEHOLDER", "test-node")
|
||||
shim = (
|
||||
"<script>"
|
||||
+ _JS_PROXY_SHIM.replace('"PREFIX_PLACEHOLDER"', json.dumps(prefix))
|
||||
+ "</script>"
|
||||
shim_js = _JS_PROXY_SHIM.replace('"PREFIX_PLACEHOLDER"', json.dumps(prefix)).replace(
|
||||
'"NODE_ID_PLACEHOLDER"', json.dumps("test-node")
|
||||
)
|
||||
result = sample_html.replace("<body>", "<body>" + banner + shim, 1)
|
||||
shim = "<script>" + shim_js + "</script>"
|
||||
result = sample_html.replace("<body>", "<body>" + shim, 1)
|
||||
assert "<script>" in result
|
||||
assert "/node/test-node" in result
|
||||
assert "window.fetch" in result
|
||||
|
||||
@@ -0,0 +1,345 @@
|
||||
"""``GET /v1/api/models`` resolution-chain coverage.
|
||||
|
||||
The console handler resolves four defaults from settings + the enabled
|
||||
model list:
|
||||
|
||||
* ``default_alias`` ← ``model.default_alias``
|
||||
* ``channel_default_alias`` ← ``channels.default_model_alias``
|
||||
* ``coordinator_default_alias`` ← ``coordinator.model_alias``, falling
|
||||
back to ``default_alias`` when empty *or* pointing at a disabled /
|
||||
removed alias (mirrors :mod:`turnstone.console.session_factory`).
|
||||
* ``judge_default_alias`` ← ``judge.model``, falling back to the
|
||||
resolved coordinator alias when empty *or* pointing at a value that
|
||||
isn't an enabled alias. ``judge.model`` is alias-only — same
|
||||
contract as the other model roles — and
|
||||
:class:`turnstone.core.judge.IntentJudge` silently inherits the
|
||||
session model when an unknown value is configured, so the API
|
||||
surfaces the resolved coordinator alias rather than echoing the
|
||||
misconfigured string.
|
||||
|
||||
These tests pin each branch so the home composer's resolved-alias
|
||||
placeholder stays correct as the precedence rules evolve.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.routing import Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from tests._coord_test_helpers import _AuthMiddleware, _FakeConfigStore
|
||||
from turnstone.console.server import list_available_models
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "available_models.db"))
|
||||
|
||||
|
||||
def _seed_model(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
definition_id: str,
|
||||
alias: str,
|
||||
model: str = "model-x",
|
||||
enabled: bool = True,
|
||||
) -> None:
|
||||
storage.create_model_definition(
|
||||
definition_id=definition_id,
|
||||
alias=alias,
|
||||
model=model,
|
||||
provider="openai-compatible",
|
||||
base_url="http://localhost:8000/v1",
|
||||
api_key="sk-test",
|
||||
context_window=8192,
|
||||
capabilities="{}",
|
||||
enabled=enabled,
|
||||
created_by="admin",
|
||||
)
|
||||
|
||||
|
||||
class _StubRegistry:
|
||||
"""Mimics the surface ``resolve_coordinator_alias`` reads from
|
||||
``coord_registry``: ``.default`` and ``.has_alias()``.
|
||||
|
||||
Production wires this through ``ModelRegistry``, which in turn
|
||||
pulls aliases from both DB rows and config.toml. The fixture
|
||||
mirrors the storage's enabled-row set so ``has_alias()`` agrees
|
||||
with what the placeholder's enabled-row filter would accept —
|
||||
without that alignment the helper rejects every tier-2 candidate
|
||||
and the placeholder goes blank in cases that production handles
|
||||
fine."""
|
||||
|
||||
def __init__(self, *, default: str, known: set[str]) -> None:
|
||||
self.default = default
|
||||
self._known = known
|
||||
|
||||
def has_alias(self, alias: str) -> bool:
|
||||
return alias in self._known
|
||||
|
||||
|
||||
def _make_client(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
settings: dict[str, str] | None = None,
|
||||
registry_default: str = "",
|
||||
config_store: bool = True,
|
||||
) -> TestClient:
|
||||
app = Starlette(
|
||||
routes=[Route("/v1/api/models", list_available_models)],
|
||||
middleware=[Middleware(_AuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
if config_store:
|
||||
app.state.config_store = _FakeConfigStore(dict(settings or {}))
|
||||
# ``coord_registry`` is always set in production after lifespan
|
||||
# startup; mirror that here. ``has_alias`` answers from the same
|
||||
# enabled-rows set the handler filters against.
|
||||
enabled = {r["alias"] for r in storage.list_model_definitions(enabled_only=True)}
|
||||
app.state.coord_registry = _StubRegistry(default=registry_default, known=enabled)
|
||||
client = TestClient(app)
|
||||
client.headers.update({"X-Test-User": "admin", "X-Test-Perms": ""})
|
||||
return client
|
||||
|
||||
|
||||
def _get_models(client: TestClient) -> dict[str, Any]:
|
||||
resp = client.get("/v1/api/models")
|
||||
assert resp.status_code == 200, resp.text
|
||||
return resp.json()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Coordinator resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_no_settings_leaves_all_defaults_blank(storage: SQLiteBackend) -> None:
|
||||
"""No model.default_alias, no per-role overrides → every default
|
||||
field is empty and ``models`` is an empty list."""
|
||||
body = _get_models(_make_client(storage))
|
||||
assert body == {
|
||||
"models": [],
|
||||
"default_alias": "",
|
||||
"channel_default_alias": "",
|
||||
"coordinator_default_alias": "",
|
||||
"judge_default_alias": "",
|
||||
}
|
||||
|
||||
|
||||
def test_coordinator_inherits_default_alias_when_unset(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(_make_client(storage, settings={"model.default_alias": "primary"}))
|
||||
assert body["default_alias"] == "primary"
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
|
||||
|
||||
def test_coordinator_explicit_enabled_alias_passes_through(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
_seed_model(storage, definition_id="m2", alias="fast")
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"coordinator.model_alias": "fast",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["coordinator_default_alias"] == "fast"
|
||||
|
||||
|
||||
def test_coordinator_set_to_disabled_alias_falls_back_to_default(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Operator disabled the alias the coordinator was pinned to —
|
||||
fall back to the registry default rather than advertising a model
|
||||
that workstream creation would refuse to use."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
_seed_model(storage, definition_id="m2", alias="legacy", enabled=False)
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"coordinator.model_alias": "legacy",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
|
||||
|
||||
def test_coordinator_set_to_unknown_alias_falls_back_to_default(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"coordinator.model_alias": "ghost",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
|
||||
|
||||
def test_coordinator_falls_back_to_registry_default_when_config_store_empty(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Match ``console/session_factory.py:109-110``: when both
|
||||
``coordinator.model_alias`` and ``model.default_alias`` are unset, new
|
||||
coordinator sessions run on ``registry.default`` (loaded from
|
||||
config.toml ``[model].default``). The placeholder must report the
|
||||
same alias rather than going blank — otherwise the home composer
|
||||
advertises "Default model" while sessions actually launch on a
|
||||
concrete alias."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(_make_client(storage, registry_default="primary"))
|
||||
assert body["default_alias"] == ""
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
assert body["judge_default_alias"] == "primary"
|
||||
|
||||
|
||||
def test_coordinator_skips_registry_default_when_alias_disabled(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Registry default points at an alias that's been disabled in the DB
|
||||
— the placeholder stays blank rather than advertising a model that
|
||||
workstream creation would refuse to use."""
|
||||
_seed_model(storage, definition_id="m1", alias="legacy", enabled=False)
|
||||
body = _get_models(_make_client(storage, registry_default="legacy"))
|
||||
assert body["coordinator_default_alias"] == ""
|
||||
|
||||
|
||||
def test_coordinator_falls_back_to_registry_default_when_config_store_missing(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Edge case from PR #500 review: lifespan can leave
|
||||
``app.state.config_store`` as None (e.g. a startup exception) while
|
||||
``coord_registry`` still binds successfully. The placeholder must
|
||||
still advertise ``registry.default`` (filtered against enabled rows)
|
||||
rather than going blank — otherwise the home composer is uselessly
|
||||
empty in a degraded-but-recoverable state."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(_make_client(storage, registry_default="primary", config_store=False))
|
||||
assert body["default_alias"] == ""
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
assert body["judge_default_alias"] == "primary"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Judge resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_judge_empty_inherits_resolved_coordinator_alias(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
_seed_model(storage, definition_id="m2", alias="fast")
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"coordinator.model_alias": "fast",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["coordinator_default_alias"] == "fast"
|
||||
assert body["judge_default_alias"] == "fast"
|
||||
|
||||
|
||||
def test_judge_explicit_enabled_alias_passes_through(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
_seed_model(storage, definition_id="m2", alias="judge-fast")
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"judge.model": "judge-fast",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["judge_default_alias"] == "judge-fast"
|
||||
|
||||
|
||||
def test_judge_set_to_unknown_value_inherits_coordinator(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""``judge.model`` is alias-only — same contract as the other model
|
||||
roles. An unknown value silently inherits the session model in
|
||||
:class:`IntentJudge`, so the API surfaces the resolved coordinator
|
||||
alias rather than echoing the misconfigured string."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"judge.model": "anthropic/claude-haiku-4-5", # raw, not an alias
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["coordinator_default_alias"] == "primary"
|
||||
assert body["judge_default_alias"] == "primary"
|
||||
|
||||
|
||||
def test_judge_set_to_disabled_alias_inherits_coordinator(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Disabled-alias case is handled identically to the unknown-value
|
||||
case — both trip the alias-not-resolved path."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
_seed_model(storage, definition_id="m2", alias="judge-old", enabled=False)
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={
|
||||
"model.default_alias": "primary",
|
||||
"judge.model": "judge-old",
|
||||
},
|
||||
)
|
||||
)
|
||||
assert body["judge_default_alias"] == "primary"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-existing fields stay correct under the new resolution code
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_channel_default_alias_blanked_when_disabled(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
_seed_model(storage, definition_id="m1", alias="primary", enabled=False)
|
||||
body = _get_models(
|
||||
_make_client(
|
||||
storage,
|
||||
settings={"channels.default_model_alias": "primary"},
|
||||
)
|
||||
)
|
||||
assert body["channel_default_alias"] == ""
|
||||
|
||||
|
||||
def test_models_payload_strips_secret_fields(storage: SQLiteBackend) -> None:
|
||||
"""Regression guard: only alias/model/provider land in the response,
|
||||
never api_key / base_url / context_window / capabilities."""
|
||||
_seed_model(storage, definition_id="m1", alias="primary")
|
||||
body = _get_models(_make_client(storage))
|
||||
assert body["models"] == [
|
||||
{"alias": "primary", "model": "model-x", "provider": "openai-compatible"}
|
||||
]
|
||||
@@ -0,0 +1,306 @@
|
||||
"""Tests for the console's coordinator idle-cleanup thread helper.
|
||||
|
||||
The helper itself is a tiny loop wrapping ``mgr.close_idle``; the heavy
|
||||
lifting is in ``SessionManager.close_idle`` (covered in
|
||||
``test_session_manager.py``) and ``bulk_close_stale_orphans`` (covered
|
||||
in ``test_storage_sqlite.py``). These tests verify the glue:
|
||||
|
||||
- the helper runs an initial sweep BEFORE its first wait (cold-start
|
||||
cleanup without blocking the lifespan),
|
||||
- the helper swallows exceptions so a transient DB blip can't kill the
|
||||
daemon thread,
|
||||
- the helper exits cleanly when ``stop_event`` is set,
|
||||
- the helper subscribes to ``mgr.subscribe_to_state`` and a state-change
|
||||
event wakes the next sweep early (event-driven, not polling),
|
||||
- the helper unsubscribes when the thread exits so the subscriber
|
||||
doesn't leak past one cleanup-thread lifetime.
|
||||
|
||||
The ``stop_event`` parameter is exclusively for tests — production
|
||||
callers pass ``None`` and the daemon runs for process lifetime.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import threading
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from turnstone.console.server import _coord_idle_cleanup_thread
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
class _StubMgr:
|
||||
"""Minimal SessionManager substitute exposing only what the cleanup
|
||||
thread touches: ``close_idle``, ``subscribe_to_state``,
|
||||
``unsubscribe_from_state``. Records call ordering for assertions
|
||||
and lets the test fire state-change events manually via
|
||||
:meth:`fire_state_change`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, *, stop_event: threading.Event, expected_calls: int, raise_after: int = -1
|
||||
) -> None:
|
||||
self.calls: list[float] = []
|
||||
self._stop_event = stop_event
|
||||
self._expected = expected_calls
|
||||
self._raise_after = raise_after
|
||||
self._subscribers: list[Callable[[str, object], None]] = []
|
||||
self._sub_lock = threading.Lock()
|
||||
|
||||
def close_idle(self, timeout_sec: float) -> list[str]:
|
||||
self.calls.append(timeout_sec)
|
||||
try:
|
||||
if 0 <= self._raise_after < len(self.calls):
|
||||
raise RuntimeError("simulated DB blip")
|
||||
finally:
|
||||
# Set stop after the helper has been exercised enough,
|
||||
# regardless of whether this call raised.
|
||||
if len(self.calls) >= self._expected:
|
||||
self._stop_event.set()
|
||||
return []
|
||||
|
||||
def subscribe_to_state(self, callback: Callable[[str, object], None]) -> None:
|
||||
with self._sub_lock:
|
||||
self._subscribers.append(callback)
|
||||
|
||||
def unsubscribe_from_state(self, callback: Callable[[str, object], None]) -> None:
|
||||
with self._sub_lock, contextlib.suppress(ValueError):
|
||||
self._subscribers.remove(callback)
|
||||
|
||||
@property
|
||||
def subscribers_count(self) -> int:
|
||||
with self._sub_lock:
|
||||
return len(self._subscribers)
|
||||
|
||||
def fire_state_change(self, ws_id: str = "ws-x", state: object = "idle") -> None:
|
||||
with self._sub_lock:
|
||||
snapshot = list(self._subscribers)
|
||||
for cb in snapshot:
|
||||
cb(ws_id, state)
|
||||
|
||||
|
||||
def _run_until_done(mgr: _StubMgr, stop_event: threading.Event, timeout_sec: float) -> None:
|
||||
# ``min_sweep_interval=0.0`` disables the production cadence floor
|
||||
# (default 5 s) so tests can fire many close_idle calls back-to-back
|
||||
# without waiting real time between them. The floor is exercised
|
||||
# in its own dedicated test below.
|
||||
thread = threading.Thread(
|
||||
target=_coord_idle_cleanup_thread,
|
||||
args=(mgr, timeout_sec, stop_event),
|
||||
kwargs={"min_sweep_interval": 0.0},
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
thread.join(timeout=2.0)
|
||||
assert not thread.is_alive(), "helper failed to exit on stop_event"
|
||||
|
||||
|
||||
def test_coord_idle_cleanup_runs_initial_sweep_before_wait() -> None:
|
||||
"""The first close_idle call must happen BEFORE the first wait —
|
||||
otherwise cold-start orphans wait one ``check_every`` interval (~30 min
|
||||
on default 2h timeout) for the first reap. Crucial because the
|
||||
lifespan no longer does a synchronous initial sweep.
|
||||
|
||||
Verified structurally: a single ``expected_calls=1`` run completes
|
||||
in well under one ``check_every`` (here 0.04 s timeout → 0.01 s
|
||||
check_every), so the initial sweep must have happened before any
|
||||
real wait could have blocked it.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=1)
|
||||
started = time.monotonic()
|
||||
_run_until_done(mgr, stop_event, timeout_sec=0.04)
|
||||
elapsed = time.monotonic() - started
|
||||
assert len(mgr.calls) == 1
|
||||
# check_every = min(300.0, 0.04/4) = 0.01 s. An initial sweep
|
||||
# gated behind one full wait would have taken ~0.01+ s anyway, so
|
||||
# the upper bound here is "much less than one check_every plus
|
||||
# process noise" — the explicit 1.0 s gives generous CI headroom
|
||||
# while still asserting the test is testing the right thing.
|
||||
assert elapsed < 1.0
|
||||
|
||||
|
||||
def test_coord_idle_cleanup_calls_close_idle_each_tick() -> None:
|
||||
"""Heartbeat path: with no state-change events, close_idle fires
|
||||
each ``check_every`` interval. Test uses a tiny timeout so the
|
||||
test runs fast — the contract under test is "the loop iterates",
|
||||
not the production cadence.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=3)
|
||||
_run_until_done(mgr, stop_event, timeout_sec=0.04)
|
||||
assert len(mgr.calls) == 3
|
||||
assert all(t == 0.04 for t in mgr.calls)
|
||||
|
||||
|
||||
def test_coord_idle_cleanup_survives_close_idle_exceptions() -> None:
|
||||
"""A transient DB error must not kill the daemon thread — the next
|
||||
tick should still fire close_idle. Without the try/except, a single
|
||||
blip would silently leak orphans forever."""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=4, raise_after=1)
|
||||
_run_until_done(mgr, stop_event, timeout_sec=0.04)
|
||||
# All four calls must have fired despite calls 2-4 raising.
|
||||
assert len(mgr.calls) == 4
|
||||
|
||||
|
||||
def test_coord_idle_cleanup_exits_cleanly_on_stop_event() -> None:
|
||||
"""The stop_event mechanism is the test contract; verify the thread
|
||||
actually exits when the event is set, without needing exceptions or
|
||||
daemon-process termination."""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=2)
|
||||
_run_until_done(mgr, stop_event, timeout_sec=0.04)
|
||||
assert stop_event.is_set()
|
||||
|
||||
|
||||
def test_state_change_wakes_close_idle_before_heartbeat() -> None:
|
||||
"""The event-driven path is the whole point of the refactor: a
|
||||
workstream state-change must wake the cleanup sweep without
|
||||
waiting one ``check_every`` interval. Tested with a long
|
||||
timeout_sec so the heartbeat would NOT have fired in the test
|
||||
window — the close_idle call past the initial sweep must come
|
||||
from a state-change wake.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=2)
|
||||
# check_every = min(300.0, 120.0/4) = 30 s — well outside the test
|
||||
# window. Any close_idle call past the initial sweep must come
|
||||
# from a fire_state_change-driven wake-up.
|
||||
thread = threading.Thread(
|
||||
target=_coord_idle_cleanup_thread,
|
||||
args=(mgr, 120.0, stop_event),
|
||||
kwargs={"min_sweep_interval": 0.0},
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
# Wait for the initial sweep to complete AND the thread to enter
|
||||
# its first ``tick_now.wait`` (signalled here by the subscriber
|
||||
# being registered + calls advancing to 1).
|
||||
deadline = time.monotonic() + 1.0
|
||||
while time.monotonic() < deadline:
|
||||
if mgr.subscribers_count == 1 and len(mgr.calls) >= 1:
|
||||
break
|
||||
time.sleep(0.01)
|
||||
assert mgr.subscribers_count == 1, "thread didn't subscribe to state"
|
||||
assert len(mgr.calls) == 1, "initial sweep didn't fire"
|
||||
# One state-change fire wakes the first ``wait`` → close_idle runs
|
||||
# again → stop_event is set (expected_calls=2) → thread exits.
|
||||
mgr.fire_state_change()
|
||||
thread.join(timeout=2.0)
|
||||
assert not thread.is_alive(), "thread didn't exit after state-change-driven sweep"
|
||||
# 2 = initial + state-change-driven. If the state change weren't
|
||||
# being honoured, close_idle would have stalled on the 30 s wait
|
||||
# and the thread.join would have timed out.
|
||||
assert len(mgr.calls) == 2
|
||||
|
||||
|
||||
def test_subscriber_unregisters_when_thread_exits() -> None:
|
||||
"""The cleanup thread's state-change subscriber must be removed
|
||||
when the thread exits — otherwise long-running processes that
|
||||
restart their cleanup threads (admin model-CRUD path, tests) leak
|
||||
subscribers and every state change fires N stale callbacks.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=1)
|
||||
_run_until_done(mgr, stop_event, timeout_sec=0.04)
|
||||
assert mgr.subscribers_count == 0, "subscriber leaked past thread exit"
|
||||
|
||||
|
||||
def test_state_change_during_close_idle_triggers_followup_sweep() -> None:
|
||||
"""A state-change fired during the initial sweep (e.g. close_idle's
|
||||
own ``close()`` calls firing subscribers) must wake the next
|
||||
``tick_now.wait`` rather than being lost to the clear-before-sweep
|
||||
ordering. The clear runs INSIDE the loop just before close_idle,
|
||||
so a fire during the initial sweep — which precedes the loop —
|
||||
arrives at an already-set event that the first wait sees set and
|
||||
returns on immediately.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=2)
|
||||
|
||||
real_close_idle = mgr.close_idle
|
||||
|
||||
# One-shot fire during the initial sweep, mirroring what
|
||||
# close_idle's own close() calls do in production (set_state →
|
||||
# state-change subscribers).
|
||||
fired = [False]
|
||||
|
||||
def _instrumented_close_idle(timeout_sec: float) -> list[str]:
|
||||
result = real_close_idle(timeout_sec)
|
||||
if not fired[0]:
|
||||
fired[0] = True
|
||||
mgr.fire_state_change()
|
||||
return result
|
||||
|
||||
mgr.close_idle = _instrumented_close_idle # type: ignore[method-assign]
|
||||
|
||||
thread = threading.Thread(
|
||||
target=_coord_idle_cleanup_thread,
|
||||
args=(mgr, 120.0, stop_event),
|
||||
kwargs={"min_sweep_interval": 0.0},
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
thread.join(timeout=2.0)
|
||||
assert not thread.is_alive(), "thread blocked on the next wait — mid-sweep wake was lost"
|
||||
# 2 = initial sweep + state-change-driven follow-up. Without the
|
||||
# event surviving the clear-before-sweep ordering, the thread
|
||||
# would have blocked on the 30 s ``wait`` and the test would have
|
||||
# timed out at thread.join.
|
||||
assert len(mgr.calls) == 2
|
||||
|
||||
|
||||
def test_min_sweep_interval_floors_close_idle_cadence_under_sustained_wakes() -> None:
|
||||
"""Cadence floor: even when state-change events keep firing
|
||||
``tick_now.set()``, ``close_idle`` must not run more often than
|
||||
``min_sweep_interval`` — otherwise the loop tight-spins close_idle
|
||||
at the rate of its own DB latency, doing 600-1500x more DB work
|
||||
than the pre-refactor fixed-30 s cadence.
|
||||
|
||||
Wires a state-change subscriber that fires another state change
|
||||
from inside close_idle, so the bus would tick forever if not
|
||||
floored. Asserts the elapsed-between-sweeps is at least
|
||||
``min_sweep_interval`` modulo small wall-clock noise.
|
||||
"""
|
||||
stop_event = threading.Event()
|
||||
mgr = _StubMgr(stop_event=stop_event, expected_calls=3)
|
||||
|
||||
real_close_idle = mgr.close_idle
|
||||
sweep_times: list[float] = []
|
||||
|
||||
def _instrumented_close_idle(timeout_sec: float) -> list[str]:
|
||||
sweep_times.append(time.monotonic())
|
||||
result = real_close_idle(timeout_sec)
|
||||
# Always fire another state-change to simulate sustained
|
||||
# activity (each turn fires thinking/running/attention/idle).
|
||||
# If the floor were absent, the next wake would race the next
|
||||
# close_idle immediately and ``sweep_times`` deltas would be
|
||||
# bounded by close_idle latency (microseconds), not the floor.
|
||||
mgr.fire_state_change()
|
||||
return result
|
||||
|
||||
mgr.close_idle = _instrumented_close_idle # type: ignore[method-assign]
|
||||
|
||||
# 0.15 s floor keeps the test fast (~0.3 s total) while still
|
||||
# representing a meaningful gap relative to close_idle's
|
||||
# near-zero stub latency.
|
||||
thread = threading.Thread(
|
||||
target=_coord_idle_cleanup_thread,
|
||||
args=(mgr, 120.0, stop_event),
|
||||
kwargs={"min_sweep_interval": 0.15},
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
thread.join(timeout=3.0)
|
||||
assert not thread.is_alive(), "thread didn't exit"
|
||||
assert len(sweep_times) >= 2, "fewer than two sweeps fired"
|
||||
# Gap between sweep 1 (post-initial) and sweep 2 must respect
|
||||
# the floor. Initial sweep at sweep_times[0] is unfloored
|
||||
# (no prior sweep to compare against), so the meaningful
|
||||
# assertion is on sweep_times[1] - sweep_times[0].
|
||||
gap = sweep_times[1] - sweep_times[0]
|
||||
assert gap >= 0.12, f"floor breached: gap {gap:.3f}s < min_sweep_interval 0.15s"
|
||||
@@ -0,0 +1,219 @@
|
||||
"""``console/session_factory.py`` alias-resolution coverage.
|
||||
|
||||
The console session factory resolves the coordinator alias through a
|
||||
three-tier chain that must stay in lockstep with the placeholder logic
|
||||
in ``console/server.py:list_available_models`` — otherwise the home
|
||||
composer advertises one alias while sessions launch on another.
|
||||
|
||||
Tier order (highest priority first):
|
||||
|
||||
1. Per-call ``model_alias`` arg, or the ``coordinator.model_alias``
|
||||
ConfigStore setting (admin-pinned coordinator-specific override).
|
||||
2. ``model.default_alias`` ConfigStore setting (admin-managed system
|
||||
default surfaced in the Models tab).
|
||||
3. ``registry.default`` (config.toml ``[model].default``, the boot-time
|
||||
fallback).
|
||||
|
||||
These tests pin each branch by intercepting ``registry.resolve`` —
|
||||
they short-circuit before ChatSession construction so the test never
|
||||
has to satisfy ChatSession's full kwarg contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tests._coord_test_helpers import _FakeConfigStore
|
||||
from turnstone.console.session_factory import build_console_session_factory
|
||||
|
||||
|
||||
class _StopBeforeChatSessionError(Exception):
|
||||
"""Sentinel raised by the capturing registry to short-circuit
|
||||
factory execution after alias resolution but before ChatSession is
|
||||
built. The factory's outer code path is irrelevant to alias
|
||||
resolution and would force the test to satisfy a long kwarg
|
||||
contract for no extra coverage."""
|
||||
|
||||
|
||||
class _CapturingRegistry:
|
||||
"""Records the alias passed to ``resolve()`` and short-circuits.
|
||||
|
||||
``has_alias`` answers from the configured known set so the
|
||||
``model.default_alias`` validation tier behaves realistically.
|
||||
Mirrors the public surface ``ModelRegistry`` exposes to
|
||||
session_factory: ``has_alias``, ``resolve``, and ``default``.
|
||||
"""
|
||||
|
||||
def __init__(self, *, default: str, known: set[str]) -> None:
|
||||
self.default = default
|
||||
self._known = known
|
||||
self.captured_alias: str | None = None
|
||||
|
||||
def has_alias(self, alias: str) -> bool:
|
||||
return alias in self._known
|
||||
|
||||
def resolve(self, alias: str) -> Any:
|
||||
self.captured_alias = alias
|
||||
raise _StopBeforeChatSessionError()
|
||||
|
||||
|
||||
def _build_factory(
|
||||
*,
|
||||
registry_default: str = "registry-default",
|
||||
known_aliases: set[str] | None = None,
|
||||
settings: dict[str, Any] | None = None,
|
||||
) -> tuple[Any, _CapturingRegistry]:
|
||||
"""Construct the factory with stub deps. Returns ``(factory_callable,
|
||||
registry)`` so tests can read back ``registry.captured_alias``."""
|
||||
|
||||
registry = _CapturingRegistry(
|
||||
default=registry_default,
|
||||
known=known_aliases if known_aliases is not None else {registry_default},
|
||||
)
|
||||
config_store = _FakeConfigStore(dict(settings or {}))
|
||||
factory = build_console_session_factory(
|
||||
registry=registry, # type: ignore[arg-type]
|
||||
config_store=config_store, # type: ignore[arg-type]
|
||||
node_id="console",
|
||||
coord_client_factory=lambda ws_id, uid: MagicMock(),
|
||||
)
|
||||
return factory, registry
|
||||
|
||||
|
||||
def _invoke(factory: Any, **factory_kwargs: Any) -> None:
|
||||
"""Call the factory with a stub UI and absorb the sentinel.
|
||||
|
||||
Forwards ``factory_kwargs`` to the factory so per-call overrides
|
||||
(e.g. ``model_alias``) can flow through. Raises if any other
|
||||
exception comes out — the test should fail loudly when alias
|
||||
resolution itself errors rather than swallowing it.
|
||||
"""
|
||||
ui = MagicMock()
|
||||
ui._user_id = "" # skip storage-backed username lookup branch
|
||||
with pytest.raises(_StopBeforeChatSessionError):
|
||||
factory(ui, **factory_kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tier 1 — explicit pin (per-call arg or coordinator.model_alias)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_per_call_model_alias_arg_wins_over_everything() -> None:
|
||||
"""The ``model_alias`` kwarg on the factory call (e.g. body field on
|
||||
POST /workstreams/new) wins over both ConfigStore tiers and the
|
||||
registry default."""
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"per-call", "coord-pin", "admin-default", "registry-default"},
|
||||
settings={
|
||||
"coordinator.model_alias": "coord-pin",
|
||||
"model.default_alias": "admin-default",
|
||||
},
|
||||
)
|
||||
_invoke(factory, model_alias="per-call")
|
||||
assert registry.captured_alias == "per-call"
|
||||
|
||||
|
||||
def test_coordinator_model_alias_wins_when_no_per_call_override() -> None:
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"coord-pin", "admin-default", "registry-default"},
|
||||
settings={
|
||||
"coordinator.model_alias": "coord-pin",
|
||||
"model.default_alias": "admin-default",
|
||||
},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "coord-pin"
|
||||
|
||||
|
||||
def test_coordinator_model_alias_passed_through_unvalidated() -> None:
|
||||
"""Tier 1 is an *explicit* operator pin — when it's stale or typoed
|
||||
we deliberately pass it through to ``registry.resolve`` so the
|
||||
request layer turns it into a 503 with the alias surfaced in the
|
||||
error. Falling through silently would mask the misconfiguration."""
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"admin-default", "registry-default"},
|
||||
settings={
|
||||
"coordinator.model_alias": "ghost", # unknown
|
||||
"model.default_alias": "admin-default",
|
||||
},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "ghost"
|
||||
|
||||
|
||||
def test_per_call_model_alias_arg_passed_through_unvalidated() -> None:
|
||||
"""The per-call ``model_alias`` kwarg (POST body field — the more
|
||||
common production trigger) is the same kind of explicit pin as the
|
||||
ConfigStore setting, so a stale value passes through to
|
||||
``registry.resolve`` rather than silently falling through to the
|
||||
system default."""
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"registry-default"},
|
||||
settings={"model.default_alias": "registry-default"},
|
||||
)
|
||||
_invoke(factory, model_alias="ghost")
|
||||
assert registry.captured_alias == "ghost"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tier 2 — model.default_alias (admin-managed system default)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_model_default_alias_used_when_coordinator_unset() -> None:
|
||||
"""Regression for the historical drift: admin sets the system
|
||||
default in the Models tab, the home composer advertises it, and new
|
||||
coordinator sessions must launch on the same alias rather than
|
||||
silently falling through to ``registry.default``."""
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"admin-default", "registry-default"},
|
||||
settings={"model.default_alias": "admin-default"},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "admin-default"
|
||||
|
||||
|
||||
def test_unknown_model_default_alias_falls_through_to_registry_default() -> None:
|
||||
"""Tier 2 is *not* an explicit pin — operators set
|
||||
``model.default_alias`` once in the UI and forget about it; an alias
|
||||
that's later disabled or typo'd should not 503 the coordinator,
|
||||
since tier 3 (``registry.default``) is guaranteed to resolve."""
|
||||
factory, registry = _build_factory(
|
||||
known_aliases={"registry-default"}, # admin-default got removed
|
||||
settings={"model.default_alias": "admin-default"},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "registry-default"
|
||||
|
||||
|
||||
def test_blank_model_default_alias_falls_through_to_registry_default() -> None:
|
||||
factory, registry = _build_factory(
|
||||
settings={"model.default_alias": ""},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "registry-default"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tier 3 — registry.default (config.toml [model].default)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_no_settings_uses_registry_default() -> None:
|
||||
factory, registry = _build_factory()
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "registry-default"
|
||||
|
||||
|
||||
def test_whitespace_only_coord_alias_falls_through() -> None:
|
||||
"""``" "`` is not an explicit pin — ``.strip()`` reduces it to
|
||||
"", which the chain should treat as unset."""
|
||||
factory, registry = _build_factory(
|
||||
settings={"coordinator.model_alias": " "},
|
||||
)
|
||||
_invoke(factory)
|
||||
assert registry.captured_alias == "registry-default"
|
||||
@@ -464,3 +464,148 @@ def test_coord_budget_override_survives_wildcard_allow_policy() -> None:
|
||||
assert approved is True
|
||||
types = [e.get("type") for e in captured_events]
|
||||
assert "approve_request" in types, "Wildcard allow must not strip the budget-override prompt"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cluster-bus broadcast hooks — _broadcast_intent_verdict / _approval_resolved
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBroadcastIntentVerdict:
|
||||
"""``ConsoleCoordinatorUI._broadcast_intent_verdict`` overrides the
|
||||
no-op base hook to push the verdict onto the cluster bus via
|
||||
``ClusterCollector.emit_console_ws_intent_verdict``. The far more
|
||||
common path is the per-node ``WebUI`` override (covered in
|
||||
test_webui_content.py); this lights up the rare coord-self path
|
||||
(a coord that runs its own LLM judge).
|
||||
"""
|
||||
|
||||
def test_calls_collector_emit_with_ws_id_and_verdict(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
verdict = {
|
||||
"call_id": "c1",
|
||||
"risk_level": "high",
|
||||
"confidence": 0.91,
|
||||
}
|
||||
ui._broadcast_intent_verdict(verdict)
|
||||
collector.emit_console_ws_intent_verdict.assert_called_once_with(
|
||||
"coord-a",
|
||||
verdict,
|
||||
)
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
def test_no_op_when_collector_unset(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
# Doesn't raise.
|
||||
ui._broadcast_intent_verdict({"call_id": "c1"})
|
||||
|
||||
def test_collector_exception_swallowed(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
collector.emit_console_ws_intent_verdict.side_effect = RuntimeError("boom")
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
# Doesn't raise — collector failures are observational only.
|
||||
ui._broadcast_intent_verdict({"call_id": "c1"})
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
|
||||
class TestBroadcastApprovalResolved:
|
||||
"""``ConsoleCoordinatorUI._broadcast_approval_resolved`` overrides
|
||||
the base hook to push the resolution onto the cluster bus via
|
||||
``ClusterCollector.emit_console_ws_approval_resolved``."""
|
||||
|
||||
def test_calls_collector_with_decision_fields(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
ui._broadcast_approval_resolved(True, "lgtm", always=True)
|
||||
collector.emit_console_ws_approval_resolved.assert_called_once_with(
|
||||
"coord-a",
|
||||
approved=True,
|
||||
feedback="lgtm",
|
||||
always=True,
|
||||
)
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
def test_normalises_none_feedback_to_empty_string(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
ui._broadcast_approval_resolved(False, None)
|
||||
collector.emit_console_ws_approval_resolved.assert_called_once_with(
|
||||
"coord-a",
|
||||
approved=False,
|
||||
feedback="",
|
||||
always=False,
|
||||
)
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
def test_no_op_when_collector_unset(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
# Doesn't raise.
|
||||
ui._broadcast_approval_resolved(True, None)
|
||||
|
||||
def test_collector_exception_swallowed(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
collector.emit_console_ws_approval_resolved.side_effect = RuntimeError("boom")
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
# Doesn't raise.
|
||||
ui._broadcast_approval_resolved(True, "ok")
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
|
||||
class TestBroadcastApproveRequest:
|
||||
"""Coord-side override for the approve_request push. Same rationale
|
||||
as the WebUI override — the coord-self path is rare today, but
|
||||
parity keeps the override symmetric with the rest of the broadcast
|
||||
family."""
|
||||
|
||||
def test_calls_collector_emit_with_ws_id_and_detail(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
detail = {
|
||||
"type": "approve_request",
|
||||
"items": [{"call_id": "c1", "header": "tool x"}],
|
||||
"judge_pending": True,
|
||||
}
|
||||
ui._broadcast_approve_request(detail)
|
||||
collector.emit_console_ws_approve_request.assert_called_once_with(
|
||||
"coord-a",
|
||||
detail,
|
||||
)
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
def test_no_op_when_collector_unset(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
# Doesn't raise.
|
||||
ui._broadcast_approve_request({"items": []})
|
||||
|
||||
def test_collector_exception_swallowed(self) -> None:
|
||||
ui = ConsoleCoordinatorUI(ws_id="coord-a", user_id="u1")
|
||||
collector = MagicMock()
|
||||
collector.emit_console_ws_approve_request.side_effect = RuntimeError("boom")
|
||||
ConsoleCoordinatorUI._collector = collector
|
||||
try:
|
||||
# Doesn't raise.
|
||||
ui._broadcast_approve_request({"items": []})
|
||||
finally:
|
||||
ConsoleCoordinatorUI._collector = None
|
||||
|
||||
@@ -378,13 +378,21 @@ class TestCoordinatorAdapterWorkerDispatch:
|
||||
|
||||
|
||||
class TestCoordinatorAdapterChildrenRegistry:
|
||||
def test_emit_created_seeds_empty_children_set(self) -> None:
|
||||
"""Adapter-level integration with :class:`ChildrenRegistry`.
|
||||
|
||||
Pure-registry invariants (forward/reverse consistency, idempotent
|
||||
merge, locking) live in ``test_children_registry.py``. These
|
||||
tests cover the adapter's wiring: that ``emit_*`` paths drive the
|
||||
registry correctly and that the snapshot-priming bridge between
|
||||
a collector snapshot and the registry preserves merge semantics.
|
||||
"""
|
||||
|
||||
def test_emit_created_installs_parent(self) -> None:
|
||||
adapter, _ = _make_adapter()
|
||||
ws = _make_ws()
|
||||
adapter.emit_created(ws)
|
||||
assert ws.id in adapter._children
|
||||
assert adapter._children[ws.id] == set()
|
||||
assert adapter._active_coords[ws.id] is ws.ui
|
||||
assert adapter._registry.children_of(ws.id) == []
|
||||
assert adapter._registry.ui_for(ws.id) is ws.ui
|
||||
|
||||
def test_emit_rehydrated_calls_rebuild(self) -> None:
|
||||
adapter, _ = _make_adapter()
|
||||
@@ -398,42 +406,36 @@ class TestCoordinatorAdapterChildrenRegistry:
|
||||
adapter.emit_rehydrated(ws)
|
||||
assert calls == [ws.id]
|
||||
|
||||
def test_emit_closed_clears_forward_and_reverse_indexes(self) -> None:
|
||||
def test_emit_closed_uninstalls_parent_and_clears_children(self) -> None:
|
||||
adapter, _ = _make_adapter()
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-a1", "child-a2"])
|
||||
adapter._merge_child_ids_locked("coord-b", ["child-b1"])
|
||||
adapter._active_coords["coord-a"] = object()
|
||||
adapter._active_coords["coord-b"] = object()
|
||||
adapter._registry.install("coord-a", object())
|
||||
adapter._registry.install("coord-b", object())
|
||||
adapter._registry.merge_children("coord-a", ["child-a1", "child-a2"])
|
||||
adapter._registry.merge_children("coord-b", ["child-b1"])
|
||||
|
||||
adapter.emit_closed("coord-a")
|
||||
|
||||
assert "coord-a" not in adapter._children
|
||||
assert "coord-a" not in adapter._active_coords
|
||||
assert "child-a1" not in adapter._child_to_coord
|
||||
assert "child-a2" not in adapter._child_to_coord
|
||||
assert adapter._registry.ui_for("coord-a") is None
|
||||
assert adapter._registry.children_of("coord-a") == []
|
||||
assert adapter._registry.parent_for("child-a1") is None
|
||||
assert adapter._registry.parent_for("child-a2") is None
|
||||
# coord-b untouched
|
||||
assert adapter._child_to_coord["child-b1"] == "coord-b"
|
||||
assert "coord-b" in adapter._children
|
||||
|
||||
def test_merge_child_ids_locked_is_idempotent(self) -> None:
|
||||
adapter, _ = _make_adapter()
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-1"])
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-1"])
|
||||
assert adapter._children["coord-a"] == {"child-1"}
|
||||
assert adapter._child_to_coord == {"child-1": "coord-a"}
|
||||
assert adapter._registry.parent_for("child-b1") == "coord-b"
|
||||
assert adapter._registry.ui_for("coord-b") is not None
|
||||
|
||||
def test_prime_children_from_snapshot_merges_without_overwriting(self) -> None:
|
||||
# Snapshot priming now lives on ClusterChildSource (production
|
||||
# path). The adapter no longer carries its own duplicate copy.
|
||||
from turnstone.core.child_source import ClusterChildSource
|
||||
|
||||
adapter, _ = _make_adapter()
|
||||
# Seed one in-memory coord + one existing child
|
||||
coord_ws = _make_ws()
|
||||
coord_ws.id = "coord-a"
|
||||
mgr = MagicMock()
|
||||
mgr.list_all.return_value = [coord_ws]
|
||||
adapter.attach(mgr)
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-a1"])
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
|
||||
source = ClusterChildSource(
|
||||
collector=MagicMock(),
|
||||
registry=adapter._registry,
|
||||
parents_provider=lambda: ["coord-a"],
|
||||
)
|
||||
|
||||
snapshot = {
|
||||
"nodes": [
|
||||
@@ -448,10 +450,13 @@ class TestCoordinatorAdapterChildrenRegistry:
|
||||
},
|
||||
],
|
||||
}
|
||||
adapter._prime_children_from_snapshot(snapshot)
|
||||
assert adapter._children["coord-a"] == {"child-a1", "child-a2"}
|
||||
assert adapter._child_to_coord["child-a2"] == "coord-a"
|
||||
assert "child-x" not in adapter._child_to_coord
|
||||
source._prime_from_snapshot(snapshot)
|
||||
assert set(adapter._registry.children_of("coord-a")) == {
|
||||
"child-a1",
|
||||
"child-a2",
|
||||
}
|
||||
assert adapter._registry.parent_for("child-a2") == "coord-a"
|
||||
assert adapter._registry.parent_for("child-x") is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -478,9 +483,7 @@ class TestCoordinatorAdapterDispatchChildEvent:
|
||||
coord_ws.id = coord_id
|
||||
recorder = _UIRecorder()
|
||||
coord_ws.ui = recorder # type: ignore[assignment]
|
||||
with adapter._children_lock:
|
||||
adapter._children.setdefault(coord_id, set())
|
||||
adapter._active_coords[coord_id] = recorder
|
||||
adapter._registry.install(coord_id, recorder)
|
||||
adapter.attach(_StubManager(coord_ws)) # type: ignore[arg-type]
|
||||
return adapter, recorder, coord_ws
|
||||
|
||||
@@ -510,12 +513,11 @@ class TestCoordinatorAdapterDispatchChildEvent:
|
||||
assert payload["child_ws_id"] == "child-a1"
|
||||
assert payload["parent_ws_id"] == "coord-a"
|
||||
# Reverse index updated for subsequent cluster_state events.
|
||||
assert adapter._child_to_coord["child-a1"] == "coord-a"
|
||||
assert adapter._registry.parent_for("child-a1") == "coord-a"
|
||||
|
||||
def test_dispatch_cluster_state_routes_via_reverse_index(self) -> None:
|
||||
adapter, recorder, _ = self._setup()
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-a1"])
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
@@ -533,8 +535,7 @@ class TestCoordinatorAdapterDispatchChildEvent:
|
||||
|
||||
def test_dispatch_ws_closed_routes_to_parent_coord(self) -> None:
|
||||
adapter, recorder, _ = self._setup()
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-a1"])
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{"type": "ws_closed", "ws_id": "child-a1", "reason": "evicted"}
|
||||
)
|
||||
@@ -548,8 +549,7 @@ class TestCoordinatorAdapterDispatchChildEvent:
|
||||
"""perf-6: _enqueue_on_ui mutates the payload dict in place with
|
||||
the coord's ws_id so the browser can discriminate child events."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
with adapter._children_lock:
|
||||
adapter._merge_child_ids_locked("coord-a", ["child-a1"])
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
@@ -558,3 +558,240 @@ class TestCoordinatorAdapterDispatchChildEvent:
|
||||
}
|
||||
)
|
||||
assert recorder.enqueued[0]["ws_id"] == "coord-a"
|
||||
|
||||
def test_dispatch_cluster_state_does_not_carry_pending_approval_detail(
|
||||
self,
|
||||
) -> None:
|
||||
"""Stage 3 cleanup — the ``pending_approval_detail`` piggyback
|
||||
on ``cluster_state`` is gone. Approval items now arrive via
|
||||
bulk fetch (triggered by ``activity_state="approval"`` in the
|
||||
browser); verdicts via ``child_ws_intent_verdict``; resolution
|
||||
via ``child_ws_approval_resolved``. The state event carries
|
||||
only state + activity_state — no detail field."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
"ws_id": "child-a1",
|
||||
"state": "running",
|
||||
"activity_state": "approval",
|
||||
}
|
||||
)
|
||||
assert len(recorder.enqueued) == 1
|
||||
payload = recorder.enqueued[0]
|
||||
assert payload["type"] == "child_ws_state"
|
||||
assert payload["activity_state"] == "approval"
|
||||
assert "pending_approval_detail" not in payload
|
||||
|
||||
def test_dispatch_intent_verdict_emits_child_ws_intent_verdict(self) -> None:
|
||||
"""Stage 3 Step 6 — explicit verdict events are re-emitted as
|
||||
child_ws_intent_verdict on the parent's SSE so the tree UI
|
||||
renders the risk pill without polling."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
verdict = {
|
||||
"call_id": "c1",
|
||||
"risk_level": "low",
|
||||
"confidence": 0.92,
|
||||
"recommendation": "approve",
|
||||
}
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "intent_verdict",
|
||||
"ws_id": "child-a1",
|
||||
"node_id": "node-1",
|
||||
"verdict": verdict,
|
||||
}
|
||||
)
|
||||
assert len(recorder.enqueued) == 1
|
||||
payload = recorder.enqueued[0]
|
||||
assert payload["type"] == "child_ws_intent_verdict"
|
||||
assert payload["child_ws_id"] == "child-a1"
|
||||
assert payload["parent_ws_id"] == "coord-a"
|
||||
assert payload["node_id"] == "node-1"
|
||||
assert payload["verdict"] == verdict
|
||||
|
||||
def test_dispatch_intent_verdict_unknown_child_drops(self) -> None:
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "intent_verdict",
|
||||
"ws_id": "ws-orphan",
|
||||
"verdict": {"call_id": "c1"},
|
||||
}
|
||||
)
|
||||
assert recorder.enqueued == []
|
||||
|
||||
def test_dispatch_approval_resolved_emits_child_ws_approval_resolved(
|
||||
self,
|
||||
) -> None:
|
||||
"""Stage 3 Step 6 — paired with intent_verdict; clears the
|
||||
pending-approval pill on the parent's tree UI in lockstep
|
||||
with the actual decision."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "approval_resolved",
|
||||
"ws_id": "child-a1",
|
||||
"node_id": "node-1",
|
||||
"approved": True,
|
||||
"feedback": "lgtm",
|
||||
"always": False,
|
||||
}
|
||||
)
|
||||
assert len(recorder.enqueued) == 1
|
||||
payload = recorder.enqueued[0]
|
||||
assert payload["type"] == "child_ws_approval_resolved"
|
||||
assert payload["child_ws_id"] == "child-a1"
|
||||
assert payload["parent_ws_id"] == "coord-a"
|
||||
assert payload["approved"] is True
|
||||
assert payload["feedback"] == "lgtm"
|
||||
assert payload["always"] is False
|
||||
|
||||
def test_dispatch_approval_resolved_coerces_missing_fields(self) -> None:
|
||||
"""Older nodes mid-rolling-upgrade may omit approved / always /
|
||||
feedback; dispatch coerces to safe defaults."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
adapter._dispatch_child_event({"type": "approval_resolved", "ws_id": "child-a1"})
|
||||
assert len(recorder.enqueued) == 1
|
||||
payload = recorder.enqueued[0]
|
||||
assert payload["approved"] is False
|
||||
assert payload["feedback"] == ""
|
||||
assert payload["always"] is False
|
||||
|
||||
def test_dispatch_approval_resolved_unknown_child_drops(self) -> None:
|
||||
"""Symmetric to the intent_verdict drop test — events for
|
||||
ws_ids the registry doesn't know about silently drop instead
|
||||
of fanning out to a parent that has no business seeing them."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "approval_resolved",
|
||||
"ws_id": "ws-orphan",
|
||||
"approved": True,
|
||||
},
|
||||
)
|
||||
assert recorder.enqueued == []
|
||||
|
||||
def test_dispatch_approve_request_emits_child_ws_approve_request(
|
||||
self,
|
||||
) -> None:
|
||||
"""Push path for the initial approval items — eliminates the
|
||||
bulk-fetch race that left the coord row stuck on a loading
|
||||
placeholder when the bulk fetch landed in the gap between
|
||||
_emit_state(ATTENTION) and approve_tools setting _pending_approval."""
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
detail = {
|
||||
"type": "approve_request",
|
||||
"items": [{"call_id": "c1", "header": "tool x"}],
|
||||
"judge_pending": True,
|
||||
}
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "approve_request",
|
||||
"ws_id": "child-a1",
|
||||
"node_id": "node-1",
|
||||
"detail": detail,
|
||||
},
|
||||
)
|
||||
assert len(recorder.enqueued) == 1
|
||||
payload = recorder.enqueued[0]
|
||||
assert payload["type"] == "child_ws_approve_request"
|
||||
assert payload["child_ws_id"] == "child-a1"
|
||||
assert payload["parent_ws_id"] == "coord-a"
|
||||
assert payload["node_id"] == "node-1"
|
||||
assert payload["detail"] == detail
|
||||
|
||||
def test_dispatch_approve_request_unknown_child_drops(self) -> None:
|
||||
adapter, recorder, _ = self._setup()
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "approve_request",
|
||||
"ws_id": "ws-orphan",
|
||||
"detail": {"items": []},
|
||||
},
|
||||
)
|
||||
assert recorder.enqueued == []
|
||||
|
||||
def test_dispatch_notifies_child_event_bus_on_state_event(self) -> None:
|
||||
"""Every translated state-class event must call
|
||||
``ChildEventBus.notify(ws_id)`` so a registered
|
||||
``wait_for_workstream`` waiter wakes promptly. Notify fires
|
||||
AFTER the UI enqueue so the SSE fan-out keeps priority — the
|
||||
order assertion here is structural (one notify call, matching
|
||||
ws_id) since the bus side-effect lookup is what guards against
|
||||
regressions, not the relative event ordering.
|
||||
"""
|
||||
adapter, _, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
bus = adapter.child_event_bus
|
||||
event = bus.register_waiter(["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
"ws_id": "child-a1",
|
||||
"state": "idle",
|
||||
}
|
||||
)
|
||||
assert event.is_set(), "bus notify did not fire on cluster_state dispatch"
|
||||
|
||||
def test_dispatch_notifies_for_all_state_class_event_types(self) -> None:
|
||||
"""The dispatch sink translates six event types into the
|
||||
``child_ws_*`` SSE shape; all six must also fire the bus so
|
||||
a wait on any of them wakes. ``ws_created`` is intentionally
|
||||
NOT in this set — waiters register against ws_ids they already
|
||||
know exist (the wait tool takes a pre-known list)."""
|
||||
for etype, extra in [
|
||||
("cluster_state", {"state": "running"}),
|
||||
("ws_closed", {"reason": "evicted"}),
|
||||
("ws_rename", {"name": "renamed"}),
|
||||
("intent_verdict", {"verdict": {"call_id": "c1"}}),
|
||||
("approval_resolved", {"approved": True}),
|
||||
("approve_request", {"detail": {}}),
|
||||
]:
|
||||
adapter, _, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
bus = adapter.child_event_bus
|
||||
event = bus.register_waiter(["child-a1"])
|
||||
adapter._dispatch_child_event(
|
||||
{"type": etype, "ws_id": "child-a1", **extra},
|
||||
)
|
||||
assert event.is_set(), f"bus notify did not fire on {etype} dispatch"
|
||||
|
||||
def test_dispatch_does_not_notify_for_unrelated_ws_id(self) -> None:
|
||||
"""Bus is keyed by ws_id — a dispatch for ws X must not wake a
|
||||
waiter registered against ws Y, or every state change anywhere
|
||||
in the system would shake every concurrent wait."""
|
||||
adapter, _, _ = self._setup()
|
||||
adapter._registry.merge_children("coord-a", ["child-a1"])
|
||||
bus = adapter.child_event_bus
|
||||
event = bus.register_waiter(["child-other"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
"ws_id": "child-a1",
|
||||
"state": "idle",
|
||||
}
|
||||
)
|
||||
assert not event.is_set(), "bus notify spuriously fired on unrelated ws_id"
|
||||
|
||||
def test_dispatch_does_not_notify_for_unknown_child(self) -> None:
|
||||
"""Events whose ws_id isn't in any coord's registry are dropped
|
||||
BEFORE the bus notify (early return at ``coord_id is None``).
|
||||
Notify only fires for events the dispatch sink fully translated,
|
||||
keeping the bus side-effect aligned with the UI enqueue."""
|
||||
adapter, _, _ = self._setup()
|
||||
bus = adapter.child_event_bus
|
||||
event = bus.register_waiter(["ws-orphan"])
|
||||
adapter._dispatch_child_event(
|
||||
{
|
||||
"type": "cluster_state",
|
||||
"ws_id": "ws-orphan",
|
||||
"state": "idle",
|
||||
}
|
||||
)
|
||||
assert not event.is_set(), "bus notify fired for ws_id the dispatch dropped"
|
||||
|
||||
@@ -9,6 +9,7 @@ storage-call path.
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
@@ -20,6 +21,7 @@ from turnstone.console.coordinator_client import (
|
||||
CoordinatorTokenManager,
|
||||
)
|
||||
from turnstone.core.auth import JWT_AUD_CONSOLE, validate_jwt
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -145,6 +147,7 @@ def _mock_client(
|
||||
coord_ws_id="coord-1",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=ChildEventBus(),
|
||||
)
|
||||
return client, captured
|
||||
|
||||
@@ -470,6 +473,7 @@ def _make_read_client(storage: SQLiteBackend) -> CoordinatorClient:
|
||||
coord_ws_id="coord-1",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=ChildEventBus(),
|
||||
)
|
||||
|
||||
|
||||
@@ -663,6 +667,7 @@ def _make_client_with_cluster_response(
|
||||
coord_ws_id="coord-1",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=ChildEventBus(),
|
||||
)
|
||||
|
||||
|
||||
@@ -949,6 +954,137 @@ def test_list_nodes_empty_on_no_matching_filters(storage_with_nodes):
|
||||
assert result["truncated"] is False
|
||||
|
||||
|
||||
def test_list_nodes_surfaces_healthy_model_aliases(tmp_path):
|
||||
"""The node's heartbeat loop projects its registry into a ``models``
|
||||
metadata entry shaped like ``[{alias, provider, healthy}, ...]``.
|
||||
``list_nodes`` flattens that to the healthy-alias list at the top
|
||||
level (under ``model_aliases``) so a coordinator can pass aliases
|
||||
straight to ``spawn_workstream(model=)`` without having to
|
||||
introspect the metadata blob. The provider-side model identifier
|
||||
(``cfg.model``) is intentionally NOT in the payload — coords kept
|
||||
reaching for it when they should pass the local alias."""
|
||||
st = SQLiteBackend(str(tmp_path / "nodes.db"))
|
||||
_set_meta(
|
||||
st,
|
||||
"node-x",
|
||||
[
|
||||
("arch", "x86_64", "auto"),
|
||||
(
|
||||
"models",
|
||||
[
|
||||
{"alias": "gpt5", "provider": "openai", "healthy": True},
|
||||
{"alias": "claude-opus-47", "provider": "anthropic", "healthy": True},
|
||||
{"alias": "broken", "provider": "openai", "healthy": False},
|
||||
],
|
||||
"auto",
|
||||
),
|
||||
],
|
||||
)
|
||||
_register_service(st, "node-x")
|
||||
client = _make_read_client(st)
|
||||
result = client.list_nodes()
|
||||
node = result["nodes"][0]
|
||||
assert node["model_aliases"] == ["gpt5", "claude-opus-47"]
|
||||
# Full per-alias info still available under metadata for callers
|
||||
# that want provider / healthy detail (e.g. surfacing degraded
|
||||
# aliases in a UI).
|
||||
full = node["metadata"]["models"]["value"]
|
||||
assert {row["alias"] for row in full} == {"gpt5", "claude-opus-47", "broken"}
|
||||
# ``model`` (the provider-side identifier) is intentionally absent
|
||||
# — keep the payload to the three values a coord actually uses.
|
||||
for row in full:
|
||||
assert "model" not in row
|
||||
|
||||
|
||||
def test_list_nodes_model_aliases_distinct_from_metadata_models(tmp_path):
|
||||
"""Pin the naming distinction explicitly: the top-level shortlist
|
||||
(``model_aliases``, list of strings) and the rich metadata blob
|
||||
(``metadata.models.value``, list of dicts) live under different
|
||||
keys so a caller that confuses them gets a clear KeyError rather
|
||||
than a silent shape mismatch."""
|
||||
st = SQLiteBackend(str(tmp_path / "nodes.db"))
|
||||
_set_meta(
|
||||
st,
|
||||
"node-x",
|
||||
[
|
||||
(
|
||||
"models",
|
||||
[{"alias": "a", "provider": "openai", "healthy": True}],
|
||||
"auto",
|
||||
),
|
||||
],
|
||||
)
|
||||
_register_service(st, "node-x")
|
||||
client = _make_read_client(st)
|
||||
node = client.list_nodes()["nodes"][0]
|
||||
# No top-level ``models`` field — only ``model_aliases``.
|
||||
assert "models" not in node
|
||||
assert node["model_aliases"] == ["a"]
|
||||
# Rich shape stays under metadata.
|
||||
assert isinstance(node["metadata"]["models"]["value"], list)
|
||||
assert isinstance(node["metadata"]["models"]["value"][0], dict)
|
||||
|
||||
|
||||
def test_list_nodes_model_aliases_empty_when_node_has_not_published(tmp_path):
|
||||
"""Nodes from older builds — or a node mid-startup before its first
|
||||
metadata write — won't have a ``models`` entry. The top-level
|
||||
``model_aliases`` field defaults to ``[]`` rather than being
|
||||
omitted so coordinators can rely on the key being present."""
|
||||
st = SQLiteBackend(str(tmp_path / "nodes.db"))
|
||||
_set_meta(st, "node-y", [("arch", "x86_64", "auto")])
|
||||
_register_service(st, "node-y")
|
||||
client = _make_read_client(st)
|
||||
result = client.list_nodes()
|
||||
assert result["nodes"][0]["model_aliases"] == []
|
||||
|
||||
|
||||
def test_list_nodes_models_tolerates_malformed_entries(tmp_path):
|
||||
"""If a node ever stores a malformed ``models`` entry (wrong outer
|
||||
type, missing alias, non-bool healthy), the projection drops the
|
||||
bad rows rather than raising — the rest of the response should
|
||||
still be useful."""
|
||||
st = SQLiteBackend(str(tmp_path / "nodes.db"))
|
||||
_set_meta(
|
||||
st,
|
||||
"node-z",
|
||||
[
|
||||
(
|
||||
"models",
|
||||
[
|
||||
{"alias": "ok", "provider": "p", "healthy": True},
|
||||
"not-a-dict",
|
||||
{"provider": "p", "healthy": True}, # missing alias
|
||||
{"alias": "", "healthy": True}, # empty alias
|
||||
{"alias": "degraded", "healthy": False},
|
||||
{"alias": 42, "healthy": True}, # non-string alias
|
||||
],
|
||||
"auto",
|
||||
),
|
||||
],
|
||||
)
|
||||
_register_service(st, "node-z")
|
||||
client = _make_read_client(st)
|
||||
result = client.list_nodes()
|
||||
assert result["nodes"][0]["model_aliases"] == ["ok"]
|
||||
|
||||
|
||||
def test_list_nodes_models_handles_non_list_payload(tmp_path):
|
||||
"""A node with a corrupted models entry (dict, scalar, null) shouldn't
|
||||
blow up the whole list_nodes call. ``model_aliases`` falls back to ``[]``."""
|
||||
st = SQLiteBackend(str(tmp_path / "nodes.db"))
|
||||
_set_meta(
|
||||
st,
|
||||
"node-w",
|
||||
[
|
||||
("models", {"oops": "not a list"}, "auto"),
|
||||
],
|
||||
)
|
||||
_register_service(st, "node-w")
|
||||
client = _make_read_client(st)
|
||||
result = client.list_nodes()
|
||||
assert result["nodes"][0]["model_aliases"] == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# list_skills
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1092,6 +1228,49 @@ def test_list_skills_hides_interactive_only_skills(tmp_path):
|
||||
assert skill["kind"] in {"coordinator", "any"}
|
||||
|
||||
|
||||
def test_list_skills_omits_allowed_tools_when_empty(tmp_path):
|
||||
"""``allowed_tools`` is the auto-approve allowlist (tools exempt
|
||||
from the operator approval gate), NOT the set of tools the skill
|
||||
can use. An empty list reads as "no tool access" to a model
|
||||
that doesn't know the semantics — real misdiagnosis source: a
|
||||
code-review skill with no auto-approve allowlist looked like it
|
||||
had been spawned with zero tools. Dropping the key when empty
|
||||
removes the ambiguity at the source; absence of the field carries
|
||||
the unambiguous meaning "no tool is pre-approved for this skill"
|
||||
while a tool list reads as "these specific tools bypass the prompt".
|
||||
"""
|
||||
st = SQLiteBackend(str(tmp_path / "skills_empty.db"))
|
||||
st.create_prompt_template(
|
||||
template_id="s-empty",
|
||||
name="empty-skill",
|
||||
category="ops",
|
||||
content="",
|
||||
variables="[]",
|
||||
is_default=False,
|
||||
org_id="",
|
||||
created_by="test",
|
||||
tags="[]",
|
||||
allowed_tools="[]",
|
||||
)
|
||||
st.create_prompt_template(
|
||||
template_id="s-nonempty",
|
||||
name="nonempty-skill",
|
||||
category="ops",
|
||||
content="",
|
||||
variables="[]",
|
||||
is_default=False,
|
||||
org_id="",
|
||||
created_by="test",
|
||||
tags="[]",
|
||||
allowed_tools='["read_file"]',
|
||||
)
|
||||
client = _make_read_client(st)
|
||||
result = client.list_skills()
|
||||
by_name = {s["name"]: s for s in result["skills"]}
|
||||
assert "allowed_tools" not in by_name["empty-skill"]
|
||||
assert by_name["nonempty-skill"]["allowed_tools"] == ["read_file"]
|
||||
|
||||
|
||||
def test_list_skills_projects_allowed_tools_capped_with_sentinel(tmp_path):
|
||||
"""Each row carries the skill's allowed_tools (capped at the projection
|
||||
cap with a +N more sentinel) so coordinators can pick a skill without
|
||||
@@ -1312,6 +1491,23 @@ def test_wait_for_workstream_denies_foreign_ws_id(populated_storage):
|
||||
assert result["elapsed"] < 1.0
|
||||
|
||||
|
||||
def test_wait_for_workstream_denies_cross_tenant_child(populated_storage):
|
||||
"""Defense-in-depth (Copilot #506): a row whose ``parent_ws_id``
|
||||
matches the coordinator but whose ``user_id`` belongs to a
|
||||
different tenant must collapse to ``denied`` — otherwise a
|
||||
forged / migration-era / pre-tenant-gate row would let a
|
||||
coordinator's LLM observe foreign-tenant state through
|
||||
``wait_for_workstream``. The ``populated_storage`` fixture's
|
||||
``cross-tenant-child`` row has exactly this shape
|
||||
(parent_ws_id="coord-1", user_id="user-2").
|
||||
"""
|
||||
client = _make_read_client(populated_storage)
|
||||
result = client.wait_for_workstream(["cross-tenant-child"], timeout=5, mode="any")
|
||||
assert result["results"]["cross-tenant-child"]["state"] == "denied"
|
||||
assert result["complete"] is False
|
||||
assert result["elapsed"] < 1.0
|
||||
|
||||
|
||||
def test_wait_for_workstream_missing_ws_id_indistinguishable_from_denied(populated_storage):
|
||||
"""A ws_id that doesn't exist collapses into the same 'denied'
|
||||
shape as a foreign ws_id so wait can't be used as an existence
|
||||
@@ -1400,10 +1596,22 @@ def test_wait_for_workstream_dedupes_ws_ids(populated_storage):
|
||||
assert list(result["results"].keys()) == ["child-a"]
|
||||
|
||||
|
||||
def test_wait_for_workstream_uses_batched_storage_calls(populated_storage, monkeypatch):
|
||||
"""Per-tick polling must issue batched storage calls — at the
|
||||
documented cap (32 ws_ids over a 600s wait) the naive per-id
|
||||
shape produced ~38k row reads. Guard against regression."""
|
||||
def test_wait_for_workstream_never_falls_back_to_per_id_storage_calls(
|
||||
populated_storage, monkeypatch
|
||||
):
|
||||
"""All storage reads issued by ``wait_for_workstream`` must go
|
||||
through the batched paths. At the documented cap (32 ws_ids over
|
||||
a 600 s wait) the naive per-id shape produced ~38k row reads, so
|
||||
a regression to per-id is the meaningful failure mode this test
|
||||
guards against.
|
||||
|
||||
The primary safety net is the ``pytest.fail`` mock on the per-id
|
||||
``get_workstream`` / ``sum_workstream_tokens`` paths — any call
|
||||
there blows up loudly with the regression message. The
|
||||
additional ``batch_calls`` / ``sum_calls`` assertions cover the
|
||||
subtler regression where the call IS batched but only covers a
|
||||
subset of ws_ids (e.g. one ws_id per call in a loop).
|
||||
"""
|
||||
client = _make_read_client(populated_storage)
|
||||
batch_calls: list[list[str]] = []
|
||||
sum_calls: list[list[str]] = []
|
||||
@@ -1434,11 +1642,16 @@ def test_wait_for_workstream_uses_batched_storage_calls(populated_storage, monke
|
||||
|
||||
result = client.wait_for_workstream(["child-a", "child-b"], timeout=5, mode="any")
|
||||
assert result["complete"] is True
|
||||
# One tick is enough since child-a is already idle (terminal).
|
||||
assert len(batch_calls) == 1
|
||||
assert len(sum_calls) == 1
|
||||
assert set(batch_calls[0]) == {"child-a", "child-b"}
|
||||
assert set(sum_calls[0]) == {"child-a", "child-b"}
|
||||
# Every batched call carried the full ws_id set. The exact count
|
||||
# (currently 2: one pre-loop ownership filter + one snapshot tick)
|
||||
# is incidental; if either gains another batched read it stays
|
||||
# batched, which is the property under test.
|
||||
assert batch_calls, "no batched get_workstreams_batch call observed"
|
||||
assert sum_calls, "no batched sum_workstream_tokens_batch call observed"
|
||||
first_batch = set(batch_calls[0])
|
||||
first_sum = set(sum_calls[0])
|
||||
assert first_batch == {"child-a", "child-b"}
|
||||
assert first_sum == {"child-a", "child-b"}
|
||||
|
||||
|
||||
def test_wait_for_workstream_handles_non_string_mode(populated_storage):
|
||||
@@ -1450,6 +1663,205 @@ def test_wait_for_workstream_handles_non_string_mode(populated_storage):
|
||||
assert "invalid mode" in result["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# wait_for_workstream — event-driven (ChildEventBus wired in)
|
||||
# ---------------------------------------------------------------------------
|
||||
#
|
||||
# When the coord adapter wires its ``child_event_bus`` into the client,
|
||||
# the wait loop blocks on a per-call ``threading.Event`` keyed by ws_id
|
||||
# and only re-snapshots storage on state-change wakes or the heartbeat
|
||||
# cap. The legacy ``time.sleep`` poll path remains intact for tests
|
||||
# that don't wire the bus (above), so this section adds focused
|
||||
# coverage of the bus-driven behaviour without re-running the full
|
||||
# matrix of mode / since / cross-tenant cases.
|
||||
|
||||
|
||||
def _make_read_client_with_bus(storage, bus) -> CoordinatorClient:
|
||||
"""Like ``_make_read_client`` but wires a real ``ChildEventBus``.
|
||||
|
||||
Caller owns the bus so the test can call ``bus.notify(ws_id)`` to
|
||||
simulate the dispatch-sink wake-up.
|
||||
"""
|
||||
transport = httpx.MockTransport(lambda r: httpx.Response(200))
|
||||
http = httpx.Client(transport=transport)
|
||||
return CoordinatorClient(
|
||||
console_base_url="http://x",
|
||||
storage=storage,
|
||||
token_factory=lambda: "t",
|
||||
coord_ws_id="coord-1",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=bus,
|
||||
)
|
||||
|
||||
|
||||
def test_wait_with_bus_returns_immediately_when_already_terminal(populated_storage):
|
||||
"""Subscribe-after-terminal race: the wait registers its waiter
|
||||
BEFORE the first snapshot, then re-snapshots — an already-terminal
|
||||
child must return at once without spinning the heartbeat cap.
|
||||
"""
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
result = client.wait_for_workstream(["child-a"], timeout=5, mode="any")
|
||||
assert result["complete"] is True
|
||||
assert result["results"]["child-a"]["state"] == "idle"
|
||||
assert result["elapsed"] < 1.0
|
||||
# Waiter must be unregistered on exit so a long-lived bus doesn't
|
||||
# accumulate dead keys across many waits.
|
||||
assert "child-a" not in bus._waiters
|
||||
|
||||
|
||||
def test_wait_with_bus_wakes_on_notify(populated_storage):
|
||||
"""The core property of the refactor: a state-change ``notify``
|
||||
must wake the wait promptly — well under the legacy 0.5 s poll
|
||||
cadence AND the 2 s heartbeat cap. Test fires a state update
|
||||
+ notify after a short delay and asserts the wait returns quickly.
|
||||
"""
|
||||
import threading as _t
|
||||
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
# child-b starts running; flip to idle + notify after the wait
|
||||
# blocks. 100 ms is enough that the wait is parked in event.wait()
|
||||
# but short enough that the test runs fast.
|
||||
timer = _t.Timer(
|
||||
0.1,
|
||||
lambda: (
|
||||
populated_storage.update_workstream_state("child-b", "idle"),
|
||||
bus.notify("child-b"),
|
||||
),
|
||||
)
|
||||
timer.start()
|
||||
start = time.monotonic()
|
||||
result = client.wait_for_workstream(["child-b"], timeout=5.0, mode="any")
|
||||
elapsed = time.monotonic() - start
|
||||
assert result["complete"] is True
|
||||
assert result["results"]["child-b"]["state"] == "idle"
|
||||
# Bus-driven wake should fire well under 1 s; legacy poll would
|
||||
# take ~0.5 s but bus-driven should be ~0.1 s (the timer delay)
|
||||
# plus a few ms. Generous 0.6 s budget for CI noise.
|
||||
assert elapsed < 0.6, f"wake-up too slow: {elapsed}s"
|
||||
|
||||
|
||||
def test_wait_with_bus_unrelated_notify_does_not_wake(populated_storage):
|
||||
"""A notify on a ws_id the wait isn't watching must NOT wake it —
|
||||
otherwise every state change anywhere on the system would shake
|
||||
every concurrent wait into a redundant storage snapshot.
|
||||
"""
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
# child-b is running indefinitely; mode='all' will time out unless
|
||||
# a relevant notify fires. Fire only unrelated notifies — wait
|
||||
# should still hit the full timeout.
|
||||
import threading as _t
|
||||
|
||||
def _fire_unrelated() -> None:
|
||||
for _ in range(5):
|
||||
bus.notify("ws-unrelated-1")
|
||||
bus.notify("ws-unrelated-2")
|
||||
time.sleep(0.05)
|
||||
|
||||
t = _t.Thread(target=_fire_unrelated, daemon=True)
|
||||
t.start()
|
||||
start = time.monotonic()
|
||||
result = client.wait_for_workstream(["child-b"], timeout=0.5, mode="all")
|
||||
elapsed = time.monotonic() - start
|
||||
assert result["complete"] is False, "unrelated notify falsely satisfied wait"
|
||||
# Wait should burn its full timeout (give or take heartbeat
|
||||
# granularity). The bus path doesn't have a 0.5 s poll, so the
|
||||
# bound is "approximately timeout".
|
||||
assert elapsed >= 0.5
|
||||
t.join(timeout=1.0)
|
||||
|
||||
|
||||
def test_wait_with_bus_heartbeat_still_progresses_without_notify(populated_storage):
|
||||
"""Without any notify, the wait must still progress through ticks
|
||||
via the heartbeat cap so ``progress_callback`` keeps firing for
|
||||
the sidebar UI. Verified by counting callback firings over an
|
||||
interval longer than the heartbeat.
|
||||
"""
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
# Shrink the heartbeat for test speed via the ClassVar seam —
|
||||
# instance attribute shadows the class-level default. Production
|
||||
# stays at 2.0 s; the test exercises the heartbeat-fires-without-
|
||||
# notify property in well under 1 s.
|
||||
client._WAIT_HEARTBEAT_INTERVAL = 0.1 # type: ignore[misc]
|
||||
snapshots: list[dict[str, dict[str, object]]] = []
|
||||
|
||||
def _cb(snap: dict[str, dict[str, object]], _elapsed: float) -> None:
|
||||
snapshots.append(snap)
|
||||
|
||||
# child-b is running indefinitely; wait will time out at 0.4 s.
|
||||
# With heartbeat = 0.1 s, we expect ~3-5 callback firings
|
||||
# (initial tick + ~3-4 heartbeats). Loose lower bound to avoid
|
||||
# CI flakiness.
|
||||
start = time.monotonic()
|
||||
result = client.wait_for_workstream(["child-b"], timeout=0.4, mode="all", progress_callback=_cb)
|
||||
elapsed = time.monotonic() - start
|
||||
assert result["complete"] is False
|
||||
assert elapsed >= 0.4
|
||||
# At least 2 callback firings: the initial snapshot plus at least
|
||||
# one heartbeat-driven re-tick. Tight upper bound would be
|
||||
# ~ceil(0.4/0.1) + 1 = 5 firings.
|
||||
assert len(snapshots) >= 2, f"heartbeat didn't fire: {len(snapshots)} snapshots"
|
||||
|
||||
|
||||
def test_wait_with_bus_unregisters_waiter_on_exit(populated_storage):
|
||||
"""Both the success path and the timeout path must unregister the
|
||||
waiter — otherwise a long-lived bus accumulates dead
|
||||
``threading.Event`` instances forever.
|
||||
"""
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
# Success path (already-terminal child).
|
||||
client.wait_for_workstream(["child-a"], timeout=5, mode="any")
|
||||
assert bus._waiters == {}, "success path leaked waiter"
|
||||
# Timeout path (running child, mode='all' that times out).
|
||||
client.wait_for_workstream(["child-a", "child-b"], timeout=0.3, mode="all")
|
||||
assert bus._waiters == {}, "timeout path leaked waiter"
|
||||
|
||||
|
||||
def test_wait_with_bus_multi_waiter_independence(populated_storage):
|
||||
"""Two concurrent waits on the same ws_id must be independent —
|
||||
one wait completing must not affect the other's wake-up state.
|
||||
Smoke-tests the multi-Event-per-bucket bus behaviour against the
|
||||
real wait-loop.
|
||||
"""
|
||||
import threading as _t
|
||||
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
|
||||
bus = ChildEventBus()
|
||||
client = _make_read_client_with_bus(populated_storage, bus)
|
||||
|
||||
results: dict[str, dict[str, object]] = {}
|
||||
|
||||
def _do_wait(label: str) -> None:
|
||||
results[label] = client.wait_for_workstream(["child-a"], timeout=5, mode="any")
|
||||
|
||||
threads = [_t.Thread(target=_do_wait, args=(f"t{i}",), daemon=True) for i in range(3)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=5.0)
|
||||
for label in ("t0", "t1", "t2"):
|
||||
assert results[label]["complete"] is True
|
||||
assert results[label]["results"]["child-a"]["state"] == "idle"
|
||||
# All waiters must be unregistered after exit.
|
||||
assert bus._waiters == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# wait_for_workstream — last-message bundling
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -2247,3 +2659,362 @@ def test_cleanup_dead_task_child_refs_storage_batch_failure_swallows(populated_s
|
||||
|
||||
populated_storage.get_workstreams_batch = _boom # type: ignore[method-assign]
|
||||
assert client.cleanup_dead_task_child_refs("coord-1") == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# inspect_workstream — three-tier output compression
|
||||
# ---------------------------------------------------------------------------
|
||||
#
|
||||
# A coord doing a fan-out wave against tool-heavy children would
|
||||
# otherwise blow the context budget on raw output alone. Mirrors the
|
||||
# search tool's Tier-1/Tier-2/Tier-3 ladder.
|
||||
|
||||
|
||||
def _make_inspect_result(
|
||||
*, ws_id: str = "ws-test", state: str = "running", n_messages: int = 5
|
||||
) -> dict[str, Any]:
|
||||
"""Build an inspect-result dict shaped like ``coordinator_client.inspect()``.
|
||||
|
||||
Production output keys (``ws_id``, ``skill_id``) mirror the storage
|
||||
row that ``inspect()`` spreads from ``get_workstream``. Tests that
|
||||
synthesize an inspect result must match these keys — otherwise a
|
||||
formatter that looks at the production keys silently emits null
|
||||
values against a fixture that uses different ones (real bug-1
|
||||
regression source: skeleton tier read ``skill`` from a fixture
|
||||
that wrote ``skill`` while production wrote ``skill_id``).
|
||||
"""
|
||||
return {
|
||||
"ws_id": ws_id,
|
||||
"state": state,
|
||||
"title": "test workstream",
|
||||
"skill_id": "researcher",
|
||||
"messages": [
|
||||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"msg {i} content"}
|
||||
for i in range(n_messages)
|
||||
],
|
||||
"verdicts": [],
|
||||
}
|
||||
|
||||
|
||||
def test_format_inspect_tiered_full_fits_returns_full_tier():
|
||||
"""Small payloads pass through with `_tier='full'` — no compression."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
result = _make_inspect_result(n_messages=3)
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "full"
|
||||
# Every message verbatim.
|
||||
assert len(parsed["messages"]) == 3
|
||||
assert parsed["messages"][0]["content"] == "msg 0 content"
|
||||
|
||||
|
||||
def test_format_inspect_tiered_compact_when_full_exceeds_budget():
|
||||
"""Large messages trigger the compact tier — head/tail-snipped
|
||||
content with the rest of the row intact."""
|
||||
from turnstone.console.coordinator_client import (
|
||||
_INSPECT_MSG_CONTENT_HEAD,
|
||||
_INSPECT_MSG_CONTENT_TAIL,
|
||||
_INSPECT_OUTPUT_BUDGET,
|
||||
_format_inspect_tiered,
|
||||
)
|
||||
|
||||
# Each message ~5KB; with 20 messages, full tier blows the 32KB budget.
|
||||
fat = "X" * 5000
|
||||
result = {
|
||||
"id": "ws-fat",
|
||||
"state": "running",
|
||||
"messages": [{"role": "assistant", "content": fat} for _ in range(20)],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "compact"
|
||||
# Every message preserved (compact keeps the count, just snips content).
|
||||
assert len(parsed["messages"]) == 20
|
||||
# Head/tail snip kicked in.
|
||||
msg_content = parsed["messages"][0]["content"]
|
||||
assert msg_content.startswith("X" * _INSPECT_MSG_CONTENT_HEAD)
|
||||
assert msg_content.endswith("X" * _INSPECT_MSG_CONTENT_TAIL)
|
||||
assert "chars elided" in msg_content
|
||||
# Budget invariant — the load-bearing contract of the formatter.
|
||||
# Without this assertion, a future change to ``_tier_note`` or
|
||||
# ``_compact_message`` could push the output over budget and the
|
||||
# ``_truncate_output`` head+tail safety net would silently mask
|
||||
# the regression, re-introducing the middle-message-drop pathology.
|
||||
assert len(out) <= _INSPECT_OUTPUT_BUDGET
|
||||
|
||||
|
||||
def test_format_inspect_tiered_compact_when_content_below_snip_threshold():
|
||||
"""When per-message content is below the snip threshold but the
|
||||
message COUNT alone overflows the budget, compact tier must still
|
||||
stay within budget — by trimming the message list (head + tail of
|
||||
messages) rather than degrading straight to skeleton. Bug-3
|
||||
regression cover: with 400 × 100-char messages, the original
|
||||
formatter fell through to skeleton because adding ``_tier_note``
|
||||
to an un-snipped tier-2 produced output strictly larger than
|
||||
tier-1 (both over budget). The fix preserves messages from both
|
||||
ends of the list and inserts an ``_omitted`` sentinel."""
|
||||
from turnstone.console.coordinator_client import (
|
||||
_INSPECT_OUTPUT_BUDGET,
|
||||
_format_inspect_tiered,
|
||||
)
|
||||
|
||||
# 400 × ~100 chars → Tier-1 ~53 KB (over budget), per-message
|
||||
# content under the 964-char snip threshold so content-snipping
|
||||
# saves nothing. Without the list-trim rung the formatter would
|
||||
# fall to skeleton and drop all 400 messages.
|
||||
smallish = "S" * 100
|
||||
result = {
|
||||
"ws_id": "ws-many-small",
|
||||
"state": "running",
|
||||
"messages": [
|
||||
{"role": "assistant" if i % 2 == 0 else "user", "content": smallish} for i in range(400)
|
||||
],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
# Should NOT fall through to skeleton — message-list trim preserves
|
||||
# head + tail of the conversation.
|
||||
assert parsed["_tier"] == "compact"
|
||||
assert "messages" in parsed
|
||||
# Some messages must survive; the trim shape is head + tail with an
|
||||
# ``_omitted`` sentinel between them.
|
||||
assert len(parsed["messages"]) > 0
|
||||
assert len(parsed["messages"]) < 400
|
||||
# Budget invariant.
|
||||
assert len(out) <= _INSPECT_OUTPUT_BUDGET
|
||||
|
||||
|
||||
def test_format_inspect_tiered_skeleton_when_compact_also_exceeds_budget():
|
||||
"""Tier 3 fallback: counts + last assistant preview only. Trigger by
|
||||
flooding with messages whose content is a multi-block list — the
|
||||
snipper correctly leaves non-string content unchanged (mirrors
|
||||
Anthropic/OpenAI multi-block content shape), so even after the
|
||||
(5, 10) message-list trim the surviving 15 messages don't fit in
|
||||
the 32 KB budget."""
|
||||
from turnstone.console.coordinator_client import (
|
||||
_INSPECT_OUTPUT_BUDGET,
|
||||
_format_inspect_tiered,
|
||||
)
|
||||
|
||||
# 50 messages × multi-block content (~30 KB each — list-shape
|
||||
# content bypasses the head/tail string snipper because lists
|
||||
# aren't strings). Even (5, 10) trim leaves 15 × 30 KB which
|
||||
# blows the 32 KB budget — forces skeleton.
|
||||
fat_block = {"type": "text", "text": "Y" * 3000}
|
||||
result = {
|
||||
"ws_id": "ws-flood",
|
||||
"state": "running",
|
||||
"title": "flood",
|
||||
"skill_id": "researcher",
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant" if i % 2 == 0 else "user",
|
||||
"content": [fat_block] * 10,
|
||||
}
|
||||
for i in range(50)
|
||||
],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "skeleton"
|
||||
assert parsed["message_count"] == 50
|
||||
# Role distribution surfaces — the "what shape of activity" signal.
|
||||
assert parsed["roles"]["assistant"] == 25
|
||||
assert parsed["roles"]["user"] == 25
|
||||
# No `messages` field at skeleton tier — only the aggregate signal.
|
||||
assert "messages" not in parsed
|
||||
# Budget invariant.
|
||||
assert len(out) <= _INSPECT_OUTPUT_BUDGET
|
||||
|
||||
|
||||
def test_format_inspect_tiered_skeleton_keeps_terminal_state_fields():
|
||||
"""``close_reason`` / ``last_error`` survive the skeleton fall — they're
|
||||
small, load-bearing, and the operator needs them to understand WHY
|
||||
a terminal child landed in its state."""
|
||||
from turnstone.console.coordinator_client import (
|
||||
_INSPECT_OUTPUT_BUDGET,
|
||||
_format_inspect_tiered,
|
||||
)
|
||||
|
||||
# Same flood pattern as the bare-skeleton test (multi-block content
|
||||
# bypasses the string snipper) — paired with terminal-state fields
|
||||
# that must survive the skeleton fall.
|
||||
fat_block = {"type": "text", "text": "Z" * 3000}
|
||||
result = {
|
||||
"ws_id": "ws-closed",
|
||||
"state": "closed",
|
||||
"title": "done",
|
||||
"skill_id": "researcher",
|
||||
"messages": [{"role": "user", "content": [fat_block] * 10} for _ in range(50)],
|
||||
"verdicts": [],
|
||||
"close_reason": "task complete: report attached",
|
||||
"live": None, # filtered by truthy check
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "skeleton"
|
||||
assert parsed["close_reason"] == "task complete: report attached"
|
||||
# Falsy ``live`` doesn't bleed through.
|
||||
assert "live" not in parsed
|
||||
assert len(out) <= _INSPECT_OUTPUT_BUDGET
|
||||
|
||||
|
||||
def test_format_inspect_tiered_error_shapes_bypass_tiering():
|
||||
"""Cross-tenant / not-found responses keep their original shape — they
|
||||
carry no messages, are already tiny, and changing them would break
|
||||
callers that key on the ``error`` field."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
result = {"error": "workstream not found", "ws_id": "ws-foreign"}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed == {"error": "workstream not found", "ws_id": "ws-foreign"}
|
||||
# No `_tier` annotation — error shapes are self-describing.
|
||||
assert "_tier" not in parsed
|
||||
|
||||
|
||||
def test_format_inspect_tiered_compact_preserves_tool_call_linkage():
|
||||
"""Compact tier keeps ``tool_name`` / ``tool_call_id`` / ``name`` so a
|
||||
model reading the snipped trace can still pair a tool call to its
|
||||
response — the linkage is load-bearing for "what happened" signal."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
fat = "Q" * 5000
|
||||
result = {
|
||||
"ws_id": "ws-tools",
|
||||
"state": "running",
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": fat,
|
||||
"tool_name": "bash",
|
||||
"tool_call_id": "call-1",
|
||||
}
|
||||
for _ in range(20)
|
||||
],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "compact"
|
||||
first = parsed["messages"][0]
|
||||
assert first["tool_name"] == "bash"
|
||||
assert first["tool_call_id"] == "call-1"
|
||||
|
||||
|
||||
def test_format_inspect_tiered_compact_preserves_assistant_tool_calls():
|
||||
"""Compact tier must preserve the assistant-side ``tool_calls`` list
|
||||
(OpenAI shape: ``[{id, type, function: {name, arguments}}]``) so a
|
||||
model reading the snipped trace can see WHICH tool was called and
|
||||
pair it with the corresponding result row via ``id`` ↔ ``tool_call_id``.
|
||||
Bug-2 regression cover: the pre-fix compactor stripped ``tool_calls``,
|
||||
leaving the audit reader with a tool-result orphan against an
|
||||
invisible call.
|
||||
|
||||
``function.arguments`` strings are snipped head/tail (analogous to
|
||||
content) because they can be multi-KB JSON; ``id`` and
|
||||
``function.name`` are preserved verbatim — they're the linkage."""
|
||||
from turnstone.console.coordinator_client import (
|
||||
_INSPECT_TOOL_ARG_HEAD,
|
||||
_INSPECT_TOOL_ARG_TAIL,
|
||||
_format_inspect_tiered,
|
||||
)
|
||||
|
||||
fat_content = "C" * 5000 # forces compact tier
|
||||
fat_args = "A" * 5000 # forces argument snipping
|
||||
tool_calls = [
|
||||
{
|
||||
"id": "call-abc-123",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": fat_args},
|
||||
},
|
||||
{
|
||||
"id": "call-def-456",
|
||||
"type": "function",
|
||||
"function": {"name": "read_file", "arguments": fat_args},
|
||||
},
|
||||
]
|
||||
result = {
|
||||
"ws_id": "ws-tool-calls",
|
||||
"state": "running",
|
||||
"messages": [
|
||||
{"role": "assistant", "content": fat_content, "tool_calls": tool_calls}
|
||||
for _ in range(20)
|
||||
],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "compact"
|
||||
first = parsed["messages"][0]
|
||||
# tool_calls survives compaction.
|
||||
assert "tool_calls" in first
|
||||
assert len(first["tool_calls"]) == 2
|
||||
# Linkage fields verbatim.
|
||||
assert first["tool_calls"][0]["id"] == "call-abc-123"
|
||||
assert first["tool_calls"][0]["function"]["name"] == "bash"
|
||||
assert first["tool_calls"][1]["id"] == "call-def-456"
|
||||
assert first["tool_calls"][1]["function"]["name"] == "read_file"
|
||||
# arguments snipped head/tail — both prefix and suffix preserved.
|
||||
snipped_args = first["tool_calls"][0]["function"]["arguments"]
|
||||
assert snipped_args.startswith("A" * _INSPECT_TOOL_ARG_HEAD)
|
||||
assert snipped_args.endswith("A" * _INSPECT_TOOL_ARG_TAIL)
|
||||
assert "chars elided" in snipped_args
|
||||
|
||||
|
||||
def test_format_inspect_tiered_compact_passes_small_messages_through_unsnipped():
|
||||
"""Messages under the snip threshold pass through verbatim at compact
|
||||
tier — snipping a 100-byte message costs more bytes (the elision
|
||||
marker) than it saves."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
# Mix: a few large messages force compact tier; small messages must
|
||||
# not be snipped.
|
||||
big = "B" * 5000
|
||||
small = "S" * 50
|
||||
result = {
|
||||
"id": "ws-mixed",
|
||||
"state": "running",
|
||||
"messages": [{"role": "assistant", "content": big} for _ in range(15)]
|
||||
+ [{"role": "user", "content": small}],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "compact"
|
||||
# The trailing small message is exact, not snipped.
|
||||
assert parsed["messages"][-1]["content"] == small
|
||||
|
||||
|
||||
def test_format_inspect_tiered_emits_tier_note_when_compressed():
|
||||
"""The ``_tier_note`` advisory tells the LLM how to ask for a tighter
|
||||
or fuller view next time — actionable feedback rather than a bare
|
||||
"we compressed your output" signal."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
fat = "F" * 5000
|
||||
result = {
|
||||
"id": "ws-noted",
|
||||
"state": "running",
|
||||
"messages": [{"role": "assistant", "content": fat} for _ in range(20)],
|
||||
"verdicts": [],
|
||||
}
|
||||
out = _format_inspect_tiered(result)
|
||||
parsed = json.loads(out)
|
||||
assert "_tier_note" in parsed
|
||||
assert "message_limit" in parsed["_tier_note"]
|
||||
|
||||
|
||||
def test_format_inspect_tiered_full_tier_omits_tier_note():
|
||||
"""When the full tier fits, no note is emitted — the absence of a
|
||||
note is the signal that nothing was compressed."""
|
||||
from turnstone.console.coordinator_client import _format_inspect_tiered
|
||||
|
||||
out = _format_inspect_tiered(_make_inspect_result(n_messages=2))
|
||||
parsed = json.loads(out)
|
||||
assert parsed["_tier"] == "full"
|
||||
assert "_tier_note" not in parsed
|
||||
|
||||
@@ -40,6 +40,7 @@ from turnstone.console.server import (
|
||||
_require_coord_mgr,
|
||||
)
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.child_event_bus import ChildEventBus
|
||||
from turnstone.core.session_manager import SessionManager
|
||||
from turnstone.core.session_routes import (
|
||||
SessionEndpointConfig,
|
||||
@@ -286,6 +287,7 @@ def test_coordinator_client_spawn_close_delete(tmp_path):
|
||||
coord_ws_id="coord-42",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=ChildEventBus(),
|
||||
)
|
||||
|
||||
# spawn ---------------------------------------------------------------
|
||||
@@ -387,6 +389,7 @@ def _read_client(storage: SQLiteBackend) -> CoordinatorClient:
|
||||
coord_ws_id="coord-root",
|
||||
user_id="user-1",
|
||||
http_client=http,
|
||||
child_event_bus=ChildEventBus(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,489 @@
|
||||
"""Unit tests for :class:`CoordinatorIdleObserver`.
|
||||
|
||||
Drives a fake :class:`SessionManager` that mirrors the real one's
|
||||
``subscribe_to_state`` / ``get`` contract, plus a fake storage with the
|
||||
``list_workstreams`` slice the observer queries.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import threading
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.console.coordinator_idle_observer import CoordinatorIdleObserver
|
||||
from turnstone.core.nudge_queue import NudgeQueue
|
||||
from turnstone.core.workstream import WorkstreamKind, WorkstreamState
|
||||
|
||||
|
||||
class _FakeRow:
|
||||
"""SQLAlchemy-Row-like wrapper exposing ``_mapping``."""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self._mapping = kwargs
|
||||
|
||||
|
||||
class _FakeStorage:
|
||||
def __init__(self) -> None:
|
||||
self.children: list[dict[str, Any]] = []
|
||||
self.list_calls: list[dict[str, Any]] = []
|
||||
self.count_calls: list[dict[str, Any]] = []
|
||||
self.list_raises: bool = False
|
||||
self.count_raises: bool = False
|
||||
|
||||
def list_workstreams(
|
||||
self,
|
||||
node_id: str | None = None,
|
||||
limit: int = 100,
|
||||
*,
|
||||
parent_ws_id: str | None = None,
|
||||
kind: WorkstreamKind | str | None = None,
|
||||
user_id: str | None = None,
|
||||
) -> list[Any]:
|
||||
self.list_calls.append(
|
||||
{
|
||||
"limit": limit,
|
||||
"parent_ws_id": parent_ws_id,
|
||||
"kind": kind,
|
||||
"user_id": user_id,
|
||||
}
|
||||
)
|
||||
if self.list_raises:
|
||||
raise RuntimeError("storage forced failure")
|
||||
return [_FakeRow(**c) for c in self.children]
|
||||
|
||||
def count_workstreams_by_state(
|
||||
self,
|
||||
*,
|
||||
parent_ws_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, int]:
|
||||
self.count_calls.append({"parent_ws_id": parent_ws_id, "user_id": user_id})
|
||||
if self.count_raises:
|
||||
raise RuntimeError("count forced failure")
|
||||
counts: dict[str, int] = {}
|
||||
for c in self.children:
|
||||
counts[c["state"]] = counts.get(c["state"], 0) + 1
|
||||
return counts
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self) -> None:
|
||||
self._nudge_queue = NudgeQueue()
|
||||
self.messages: list[dict[str, Any]] = []
|
||||
self._wake_source_tag: str = ""
|
||||
self._metacog_state: dict[str, float] = {}
|
||||
self._mem_cfg = MagicMock(nudge_cooldown=300)
|
||||
|
||||
def _visible_memory_count(self) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
class _FakeWorkstream:
|
||||
def __init__(
|
||||
self,
|
||||
ws_id: str = "ws-coord",
|
||||
kind: WorkstreamKind = WorkstreamKind.COORDINATOR,
|
||||
user_id: str = "u1",
|
||||
) -> None:
|
||||
self.id = ws_id
|
||||
self.kind = kind
|
||||
self.user_id = user_id
|
||||
self.session: _FakeSession | None = _FakeSession()
|
||||
|
||||
|
||||
class _FakeManager:
|
||||
def __init__(self) -> None:
|
||||
self._workstreams: dict[str, _FakeWorkstream] = {}
|
||||
self._subscribers: list[Any] = []
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def add_ws(self, ws: _FakeWorkstream) -> None:
|
||||
self._workstreams[ws.id] = ws
|
||||
|
||||
def remove_ws(self, ws_id: str) -> None:
|
||||
self._workstreams.pop(ws_id, None)
|
||||
|
||||
def get(self, ws_id: str) -> _FakeWorkstream | None:
|
||||
return self._workstreams.get(ws_id)
|
||||
|
||||
def subscribe_to_state(self, callback: Any) -> None:
|
||||
with self._lock:
|
||||
self._subscribers.append(callback)
|
||||
|
||||
def unsubscribe_from_state(self, callback: Any) -> None:
|
||||
with self._lock, contextlib.suppress(ValueError):
|
||||
self._subscribers.remove(callback)
|
||||
|
||||
def fire_state(self, ws_id: str, state: WorkstreamState) -> None:
|
||||
with self._lock:
|
||||
subs = list(self._subscribers)
|
||||
for cb in subs:
|
||||
with contextlib.suppress(Exception):
|
||||
cb(ws_id, state)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def coord_setup() -> tuple[_FakeManager, _FakeStorage, _FakeWorkstream]:
|
||||
mgr = _FakeManager()
|
||||
storage = _FakeStorage()
|
||||
ws = _FakeWorkstream()
|
||||
mgr.add_ws(ws)
|
||||
return mgr, storage, ws
|
||||
|
||||
|
||||
def _add_active_child(storage: _FakeStorage, **overrides: Any) -> None:
|
||||
storage.children.append(
|
||||
{
|
||||
"ws_id": overrides.get("ws_id", "child-1"),
|
||||
"name": overrides.get("name", "research"),
|
||||
"state": overrides.get("state", "running"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class TestEnqueueOnIdle:
|
||||
def test_idle_with_active_children_enqueues(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage, ws_id="child-a", state="running")
|
||||
_add_active_child(storage, ws_id="child-b", state="thinking")
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
|
||||
snap = ws.session._nudge_queue.pending("any")
|
||||
assert len(snap) == 1
|
||||
nudge_type, text = snap[0]
|
||||
assert nudge_type == "idle_children"
|
||||
assert "child-a" in text
|
||||
assert "child-b" in text
|
||||
|
||||
def test_idle_with_no_active_children_no_enqueue(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
# storage.children is empty
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
def test_idle_only_idle_state_children_no_enqueue(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
# All children "idle" — terminal-from-coord-perspective; not active.
|
||||
_add_active_child(storage, state="idle")
|
||||
_add_active_child(storage, state="closed")
|
||||
_add_active_child(storage, state="error")
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
def test_non_idle_state_no_enqueue(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
for state in (
|
||||
WorkstreamState.RUNNING,
|
||||
WorkstreamState.THINKING,
|
||||
WorkstreamState.ATTENTION,
|
||||
WorkstreamState.ERROR,
|
||||
):
|
||||
mgr.fire_state(ws.id, state)
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
|
||||
class TestKindFilter:
|
||||
def test_interactive_workstream_skipped(self):
|
||||
mgr = _FakeManager()
|
||||
storage = _FakeStorage()
|
||||
_add_active_child(storage)
|
||||
ws = _FakeWorkstream(kind=WorkstreamKind.INTERACTIVE)
|
||||
mgr.add_ws(ws)
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
# Observer ignored the non-coord workstream entirely.
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
# Storage was NOT queried — kind check happens before list_workstreams.
|
||||
assert storage.list_calls == []
|
||||
|
||||
|
||||
class TestWaitForWorkstreamSkip:
|
||||
def test_skips_when_last_assistant_used_wait(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "kick off"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call-1",
|
||||
"function": {"name": "wait_for_workstream", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
# Don't pile on — model is already using the right tool.
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
def test_fires_when_last_assistant_used_different_tool(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call-1", "function": {"name": "spawn_workstream", "arguments": "{}"}}
|
||||
],
|
||||
},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 1
|
||||
|
||||
|
||||
class TestHardCap:
|
||||
def test_hard_cap_blocks_after_n_fires(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
|
||||
# Bypass cooldown for this test: each call burns a per-type slot
|
||||
# in ``_metacog_state`` so we need to clear it between fires.
|
||||
for _ in range(3):
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
|
||||
# Cap = 3 fires. Even with cooldown bypassed, the 4th doesn't fire.
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
# We enqueued 3 entries total; cap blocked the 4th.
|
||||
snap = ws.session._nudge_queue.pending("any")
|
||||
assert len(snap) == 3
|
||||
|
||||
def test_cap_resets_when_state_leaves_idle_without_wake(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
|
||||
# Burn the cap.
|
||||
for _ in range(3):
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 3
|
||||
|
||||
# Drain the queue (simulate the watcher delivering them).
|
||||
ws.session._nudge_queue.drain({"any"})
|
||||
|
||||
# Real (non-wake) leave-IDLE: tag is empty. Cap resets.
|
||||
ws.session._wake_source_tag = ""
|
||||
mgr.fire_state(ws.id, WorkstreamState.RUNNING)
|
||||
|
||||
# New IDLE — cap is fresh, fires again.
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 1
|
||||
|
||||
def test_cap_does_not_reset_during_wake_driven_exit(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
|
||||
# Burn the cap.
|
||||
for _ in range(3):
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
ws.session._nudge_queue.drain({"any"})
|
||||
|
||||
# Wake-driven leave-IDLE: tag is set during the wake send.
|
||||
ws.session._wake_source_tag = "system_nudge"
|
||||
mgr.fire_state(ws.id, WorkstreamState.RUNNING)
|
||||
ws.session._wake_source_tag = "" # tag cleared at end of wake send
|
||||
|
||||
# Cap should NOT have reset — re-IDLE shouldn't fire.
|
||||
ws.session._metacog_state.clear()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 0
|
||||
|
||||
|
||||
class TestCooldown:
|
||||
def test_cooldown_blocks_within_window(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 1
|
||||
|
||||
# Drain so the queue isn't the gate.
|
||||
ws.session._nudge_queue.drain({"any"})
|
||||
|
||||
# Second fire within the cooldown window → should_nudge returns False.
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 0
|
||||
|
||||
|
||||
class TestStorageFailure:
|
||||
def test_storage_exception_is_swallowed(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
storage.list_raises = True
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
# Must not raise / propagate.
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
|
||||
class TestValidUntilPredicate:
|
||||
def test_predicate_drops_when_children_finish_before_drain(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage, ws_id="child-a", state="running")
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 1
|
||||
|
||||
# Children now complete (storage shows none active).
|
||||
storage.children.clear()
|
||||
|
||||
# Drain at the user seam — predicate re-queries, finds 0 active,
|
||||
# drops the entry without delivering.
|
||||
from turnstone.core.nudge_queue import USER_DRAIN
|
||||
|
||||
delivered = ws.session._nudge_queue.drain(USER_DRAIN)
|
||||
assert delivered == []
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
def test_predicate_delivers_when_children_still_active(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage, ws_id="child-a", state="running")
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
|
||||
# Children still active → predicate returns True → entry delivers.
|
||||
from turnstone.core.nudge_queue import USER_DRAIN
|
||||
|
||||
delivered = ws.session._nudge_queue.drain(USER_DRAIN)
|
||||
assert len(delivered) == 1
|
||||
assert delivered[0][0] == "idle_children"
|
||||
|
||||
def test_predicate_drops_on_storage_failure(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
|
||||
# Storage failure at drain time. Predicate treats raises as
|
||||
# "no longer valid" (drop) — see NudgeQueue.drain's predicate
|
||||
# exception handling.
|
||||
storage.count_raises = True
|
||||
|
||||
from turnstone.core.nudge_queue import USER_DRAIN
|
||||
|
||||
delivered = ws.session._nudge_queue.drain(USER_DRAIN)
|
||||
assert delivered == []
|
||||
|
||||
|
||||
class TestLifecycle:
|
||||
def test_start_idempotent(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
observer.start() # no-op
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
# Double-subscribe would have produced 2 entries.
|
||||
assert len(ws.session._nudge_queue.pending("any")) == 1
|
||||
|
||||
def test_shutdown_unsubscribes(self, coord_setup):
|
||||
mgr, storage, ws = coord_setup
|
||||
_add_active_child(storage)
|
||||
# ≥2 messages so should_nudge's message_count > 1 gate clears.
|
||||
ws.session.messages = [
|
||||
{"role": "user", "content": "go"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
]
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
observer.shutdown()
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
def test_shutdown_idempotent(self, coord_setup):
|
||||
mgr, _storage, _ws = coord_setup
|
||||
observer = CoordinatorIdleObserver(mgr, _storage)
|
||||
observer.start()
|
||||
observer.shutdown()
|
||||
observer.shutdown() # no error
|
||||
+206
-16
@@ -79,9 +79,10 @@ def test_coordinator_js_exposes_inline_approval_helpers():
|
||||
assert "function submitChildApproval" in body or "submitChildApproval(" in body
|
||||
# The shared approve POST helper (parameterized for child ws_ids)
|
||||
assert "function approveWorkstream" in body or "approveWorkstream(" in body
|
||||
# The urgent live-bulk fetch option that fires on activity_state
|
||||
# transitions in/out of "approval"
|
||||
assert "{ urgent: true }" in body or "urgent: true" in body
|
||||
# The 409 stale-call_id retry path uses invalidateLiveBadge +
|
||||
# scheduleLiveFetch (Stage 3 cleanup removed the urgent flag —
|
||||
# cache invalidation makes the TTL gate fall through naturally).
|
||||
assert "invalidateLiveBadge(targetWsId)" in body
|
||||
# Server-side payload field — drift here means the JS reads stale keys
|
||||
assert "pending_approval_detail" in body
|
||||
# Reconnect parity (chunk 4): the SSE re-open handler must drop
|
||||
@@ -90,10 +91,10 @@ def test_coordinator_js_exposes_inline_approval_helpers():
|
||||
# can't render zombie approve/deny buttons on a row whose
|
||||
# approval was resolved during the gap. The implementation
|
||||
# iterates the cache and deletes only !permanent entries —
|
||||
# asserting the literal Map iteration form keeps a refactor
|
||||
# back to liveBadgeCache.clear() (which would re-pay 403s on
|
||||
# every reconnect for denied ids) from sneaking in.
|
||||
assert "liveBadgeCache.delete" in body
|
||||
# asserting the literal helper call keeps a refactor back to
|
||||
# _liveBadgeCacheClear() (which would re-pay 403s on every
|
||||
# reconnect for denied ids) from sneaking in.
|
||||
assert "_liveBadgeCacheDelete" in body
|
||||
# Edge-case matrix sentinel labels — POLICY-BLOCKED renders when
|
||||
# an item has error set + needs_approval=False (server-side
|
||||
# tool policy already blocked the call); "(judge unavailable)"
|
||||
@@ -113,15 +114,19 @@ def test_coordinator_js_exposes_inline_approval_helpers():
|
||||
# coord-self ws_id (the coord lives on the console process).
|
||||
# Children live on cluster nodes and 404 without the prefix.
|
||||
assert "/v1/api/route/workstreams/" in body
|
||||
# Late-judge polling — the LLM judge runs async on the child
|
||||
# node and never pushes a signal that reaches the coord, so
|
||||
# the row's pending_approval_detail with judge_pending=true
|
||||
# would freeze on heuristic verdicts forever without this
|
||||
# poll loop. The poller is GLOBAL (not per-row) so off-screen
|
||||
# rows still refresh — a per-row poller's scheduleLiveFetch
|
||||
# call short-circuits on non-visible rows, leaving them stuck.
|
||||
assert "_maybeStartJudgePoll" in body
|
||||
assert "_judgePollTick" in body
|
||||
# Late-arriving LLM judge verdicts — Stage 3 Step 5 promoted
|
||||
# ``intent_verdict`` and ``approval_resolved`` to first-class
|
||||
# cluster-bus event types, so the coord adapter dispatches them
|
||||
# as ``child_ws_intent_verdict`` / ``child_ws_approval_resolved``
|
||||
# on the parent's SSE stream. The browser handlers write
|
||||
# directly to liveBadgeCache (bypassing scheduleLiveFetch's
|
||||
# visibility gate cleanly) so off-screen rows pick up verdicts
|
||||
# without polling. Replaced the old ``_judgePollTick`` 90-second
|
||||
# global poll loop and its visibility-gate-bypass workaround.
|
||||
assert "handleChildIntentVerdict" in body
|
||||
assert "handleChildApprovalResolved" in body
|
||||
assert "child_ws_intent_verdict" in body
|
||||
assert "child_ws_approval_resolved" in body
|
||||
# Reload parity for the coord-self approval gate: init() must
|
||||
# consume the authoritative GET /workstreams snapshot's
|
||||
# pending_approval_detail so a freshly opened tab can render
|
||||
@@ -152,3 +157,188 @@ def test_coordinator_js_exposes_inline_approval_helpers():
|
||||
# any prior denial. bug-1 / bug-3 from the second /review pass.
|
||||
assert "Denied by user" in body
|
||||
assert "callOutcomes" in body
|
||||
# User-message attachment pills — both live send (coordSend) and
|
||||
# history replay route through appendUserMessageWithAttachments.
|
||||
# Renaming or dropping the helper would silently regress the
|
||||
# attachment affordance to the pre-fix plain-text bubble, which
|
||||
# would only surface in manual testing of an attached-file flow.
|
||||
# The CSS class is the visual anchor (coordinator.css) — keeping
|
||||
# both literals in the smoke layer covers JS↔CSS drift in either
|
||||
# direction.
|
||||
assert "function appendUserMessageWithAttachments" in body
|
||||
assert "msg-user-attach" in body
|
||||
# PR #487 — whitespace-only assistant content (Qwen3 with vLLM
|
||||
# ``--reasoning-parser`` strips ``<think>…</think>`` and emits only
|
||||
# ``"\n\n"`` as content before a tool call) must be skipped on
|
||||
# history replay or the empty ``.msg.assistant`` card surfaces as
|
||||
# a phantom row. The literal substring ``content.trim()`` is the
|
||||
# single-line guard the rendering branch uses; a refactor that
|
||||
# drops the trim() (e.g. simplifies to ``if (!content)``) silently
|
||||
# regresses the phantom-card fix on the multi-node coord path.
|
||||
# Mirrors ``test_app_js.py``'s same-shape pin on ``app.js``.
|
||||
assert "content.trim()" in body
|
||||
# PR #487 — coord history replay must render the assistant content
|
||||
# card BEFORE the tool batch, not after, so DOM order matches the
|
||||
# chronological order the model emitted (text → dispatch → results).
|
||||
# Pre-fix the tool_calls branch sat at the role-agnostic top of the
|
||||
# loop and rendered ahead of the assistant text that announced the
|
||||
# batch, putting parallel fan-outs visually above their narrating
|
||||
# message. The fix hoisted the synthesis into ``renderAssistantToolBatch``
|
||||
# called from inside the assistant branch AFTER the content card —
|
||||
# asserting the helper name lets a refactor that re-inlines or
|
||||
# renames it surface here instead of via manual reload testing.
|
||||
assert "function renderAssistantToolBatch" in body
|
||||
assert "renderAssistantToolBatch(m)" in body
|
||||
|
||||
|
||||
def test_coordinator_js_handle_child_state_no_longer_reads_sse_pending_approval_detail():
|
||||
"""Stage 3 cleanup — ``pending_approval_detail`` is no longer
|
||||
piggybacked on child_ws_state events. Approval items now arrive
|
||||
via bulk fetch on the activity_state="approval" transition;
|
||||
verdicts via the explicit ``child_ws_intent_verdict`` event class;
|
||||
resolution via ``child_ws_approval_resolved``. A refactor that
|
||||
re-introduces the piggyback would silently re-open the
|
||||
duplicate-path race the dedicated event classes were added to
|
||||
eliminate.
|
||||
|
||||
Structural assertions (regex against multi-line source) — symbol-
|
||||
presence alone wouldn't catch a guard that keeps the names but
|
||||
inverts the comparison or drops the ``prev.live`` check. This
|
||||
codebase has no JS test framework, so locking the guard's shape
|
||||
here is the next-best thing to a behavioral test."""
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
coord_js = Path(__file__).resolve().parent.parent / (
|
||||
"turnstone/console/static/coordinator/coordinator.js"
|
||||
)
|
||||
body = coord_js.read_text(encoding="utf-8")
|
||||
|
||||
# The piggyback read is gone from handleChildState. (The string
|
||||
# may still appear elsewhere — e.g. handleChildIntentVerdict
|
||||
# reading from cache, or comments — but never as ``ev.pending_approval_detail``.)
|
||||
assert "ev.pending_approval_detail" not in body
|
||||
# The pre-fix urgent-fetch on activity_state transitions is gone.
|
||||
assert "enteredApproval" not in body
|
||||
assert "leftApproval" not in body
|
||||
# ``pendingApproval`` flag derivation must check BOTH state and
|
||||
# activity_state. The worker thread can fire the state transition
|
||||
# to "attention" before approve_tools updates activity_state, so
|
||||
# checking only activity_state misses children that legitimately
|
||||
# need approval. Pin the disjunction so the regression doesn't
|
||||
# silently re-introduce.
|
||||
assert re.search(
|
||||
r'existing\.state\s*===\s*"attention"\s*\|\|\s*'
|
||||
r'existing\.activity_state\s*===\s*"approval"',
|
||||
body,
|
||||
), (
|
||||
"handleChildState must derive pendingApproval from "
|
||||
"(state==='attention' || activity_state==='approval')"
|
||||
)
|
||||
|
||||
# SSE-authoritative window constant is defined and used.
|
||||
assert re.search(r"\bconst\s+SSE_AUTHORITATIVE_MS\s*=\s*\d+", body), (
|
||||
"SSE_AUTHORITATIVE_MS constant must be defined as a numeric literal"
|
||||
)
|
||||
|
||||
# SSE writers tag entries with sseUpdatedAt: Date.now() so the
|
||||
# merge guard in flushLiveFetches preserves them against stale
|
||||
# bulk-fetch responses. handleChildState only stamps when it
|
||||
# AUTHORITATIVELY clears the detail (off-approval transition);
|
||||
# writers that stamp unconditionally are intent_verdict (verdict
|
||||
# stamp), approval_resolved (clear), and the optimistic-clear
|
||||
# path in submitChildApproval. Pinning the literal Date.now()
|
||||
# call keeps a refactor that drops the SSE-source tag entirely
|
||||
# from sneaking in.
|
||||
assert re.search(
|
||||
r"sseUpdatedAt:\s*Date\.now\(\)",
|
||||
body,
|
||||
), "Critical SSE writers must stamp sseUpdatedAt: Date.now()"
|
||||
|
||||
# flushLiveFetches' merge guard structure: SSE-set pending_approval
|
||||
# / _detail wins over a stale bulk-poll snapshot when (live) AND
|
||||
# (prev exists) AND (prev.sseUpdatedAt set) AND (within window)
|
||||
# AND (prev.live exists). Inverting the comparison or dropping
|
||||
# any of these guards reopens the clobber bug.
|
||||
merge_guard = re.search(
|
||||
r"if\s*\(\s*live\s*&&\s*prev\s*&&\s*prev\.sseUpdatedAt\s*&&\s*"
|
||||
r"now\s*-\s*prev\.sseUpdatedAt\s*<\s*SSE_AUTHORITATIVE_MS\s*&&\s*"
|
||||
r"prev\.live\s*\)",
|
||||
body,
|
||||
)
|
||||
assert merge_guard is not None, (
|
||||
"flushLiveFetches merge guard must be the conjunction "
|
||||
"(live && prev && prev.sseUpdatedAt && now - prev.sseUpdatedAt < "
|
||||
"SSE_AUTHORITATIVE_MS && prev.live). An inverted comparison or "
|
||||
"missing prev.live check would let a stale bulk-poll clobber a "
|
||||
"fresh SSE-set approval."
|
||||
)
|
||||
|
||||
# The merge body must preserve BOTH pending_approval and
|
||||
# pending_approval_detail from prev — preserving only one would
|
||||
# render a row with a phantom badge but no buttons (or vice versa).
|
||||
merge_body = re.search(
|
||||
r"mergedLive\s*=\s*Object\.assign\(\s*\{\}\s*,\s*live\s*,\s*\{"
|
||||
r"[^}]*pending_approval:\s*prev\.live\.pending_approval[^}]*"
|
||||
r"pending_approval_detail:\s*prev\.live\.pending_approval_detail",
|
||||
body,
|
||||
)
|
||||
assert merge_body is not None, (
|
||||
"Merge body must preserve both pending_approval AND "
|
||||
"pending_approval_detail from prev.live — preserving only one "
|
||||
"creates a half-rendered approval row."
|
||||
)
|
||||
|
||||
# flushLiveFetches must forward sseUpdatedAt onto the new cache
|
||||
# entry so the SSE-source tag survives the bulk-poll write back —
|
||||
# without this, every bulk-poll resets the window and the next
|
||||
# late-arriving poll silently clobbers.
|
||||
assert re.search(
|
||||
r"sseUpdatedAt:\s*prev\s*\?\s*prev\.sseUpdatedAt",
|
||||
body,
|
||||
), (
|
||||
"flushLiveFetches must forward prev.sseUpdatedAt onto the new "
|
||||
"cache entry (preserving the SSE-source window across bulk-poll "
|
||||
"cycles) — without this, the second bulk-poll after an SSE "
|
||||
"transition silently clobbers."
|
||||
)
|
||||
|
||||
|
||||
def test_coord_history_renders_user_interjection_advisory_after_tool_block():
|
||||
"""Queued user messages spliced into the last tool-result envelope
|
||||
of a batch (Seam 1) persist on the tool DB row as a wrapped
|
||||
``<tool_output>`` envelope. ``decorate_history_messages`` extracts
|
||||
the advisory back out and the wire layer projects it onto
|
||||
``m.advisories``; the coord history loop must invoke the shared
|
||||
``replayAdvisoriesAfterTool`` helper (defined in
|
||||
``shared/utils.js``) so each ``user_interjection`` renders through
|
||||
``appendUserMessageWithAttachments`` and the bubble looks identical
|
||||
to a Seam 2/3 user row.
|
||||
|
||||
This test pins the call site so a refactor that drops the helper
|
||||
invocation regresses the queued-during-batch replay shape silently.
|
||||
Mirrors ``test_app_js.py``'s same-shape pin on interactive's
|
||||
``replayHistory``."""
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
coord_js = Path(__file__).resolve().parent.parent / (
|
||||
"turnstone/console/static/coordinator/coordinator.js"
|
||||
)
|
||||
body = coord_js.read_text(encoding="utf-8")
|
||||
|
||||
assert "replayAdvisoriesAfterTool(m.advisories" in body, (
|
||||
"Coord history loop must invoke replayAdvisoriesAfterTool with "
|
||||
"m.advisories so queued messages spliced into the tool envelope "
|
||||
"render as user bubbles after the tool block."
|
||||
)
|
||||
# The renderer callback routes through appendUserMessageWithAttachments
|
||||
# so the bubble matches a normal user-row replay.
|
||||
assert re.search(
|
||||
r"appendUserMessageWithAttachments\(\s*text",
|
||||
body,
|
||||
), (
|
||||
"Coord history loop's renderer callback must route the extracted "
|
||||
"advisory text through appendUserMessageWithAttachments so the "
|
||||
"rendered bubble matches a normal user-row replay."
|
||||
)
|
||||
|
||||
@@ -215,7 +215,12 @@ def test_spawn_exec_does_not_surface_misleading_status_field(coord_session):
|
||||
summary tempted callers to write ``if result["status"] == "idle"``
|
||||
which silently never matched. The summary now omits the field
|
||||
entirely; lifecycle state lives on the workstream row and is read
|
||||
via inspect_workstream."""
|
||||
via inspect_workstream.
|
||||
|
||||
Also asserts the return key is ``child_ws_id`` (not ``ws_id``) so
|
||||
the coordinator LLM doesn't recency-bias toward feeding the spawn
|
||||
output back into another ``spawn_workstream(ws_id=...)`` call.
|
||||
"""
|
||||
sess, coord, _ui = coord_session
|
||||
coord.spawn.return_value = {
|
||||
"ws_id": "child-7",
|
||||
@@ -227,8 +232,9 @@ def test_spawn_exec_does_not_surface_misleading_status_field(coord_session):
|
||||
_call_id, output = sess._exec_spawn_workstream(item)
|
||||
body = json.loads(output)
|
||||
assert "status" not in body
|
||||
assert "ws_id" not in body
|
||||
# The substantive fields are still here.
|
||||
assert body["ws_id"] == "child-7"
|
||||
assert body["child_ws_id"] == "child-7"
|
||||
assert body["node_id"] == "node-1"
|
||||
|
||||
|
||||
@@ -248,6 +254,10 @@ def test_spawn_batch_exec_does_not_surface_misleading_status_field(coord_session
|
||||
body = json.loads(output)
|
||||
assert "0" in body["results"]
|
||||
assert "status" not in body["results"]["0"]
|
||||
# Per-result entries surface ``child_ws_id``, not ``ws_id`` — same
|
||||
# recency-bias rationale as the spawn_workstream test above.
|
||||
assert body["results"]["0"]["child_ws_id"] == "c-x"
|
||||
assert "ws_id" not in body["results"]["0"]
|
||||
|
||||
|
||||
def test_spawn_exec_surfaces_client_error(coord_session):
|
||||
@@ -260,6 +270,21 @@ def test_spawn_exec_surfaces_client_error(coord_session):
|
||||
assert ui.tool_results[-1][3] is True # is_error
|
||||
|
||||
|
||||
def test_spawn_exec_treats_missing_ws_id_on_success_path_as_error(coord_session):
|
||||
"""A malformed upstream response (200-success-shape with no
|
||||
``ws_id``) used to emit ``{"child_ws_id": null}`` to the LLM,
|
||||
which then chased a null id through follow-up tools. Now matches
|
||||
the matching guard in ``_exec_spawn_batch``: surface as a tool
|
||||
error so the model retries instead of acting on garbage."""
|
||||
sess, coord, ui = coord_session
|
||||
# No ``error`` field, but ``ws_id`` is missing — the silent-null path.
|
||||
coord.spawn.return_value = {"name": "c", "node_id": "node-1", "status": 200}
|
||||
item = sess._prepare_tool(_tc("spawn_workstream", {"initial_message": "hi"}))
|
||||
_call_id, output = sess._exec_spawn_workstream(item)
|
||||
assert "no ws_id" in output
|
||||
assert ui.tool_results[-1][3] is True # is_error
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# inspect_workstream
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1410,9 +1435,13 @@ def test_spawn_batch_exec_serialises_spawns_and_returns_results(coord_session):
|
||||
assert body["denied"] == []
|
||||
# Keyed by input index (stringified).
|
||||
assert set(body["results"].keys()) == {"0", "1", "2"}
|
||||
assert body["results"]["0"]["ws_id"] == "child-0"
|
||||
assert body["results"]["0"]["child_ws_id"] == "child-0"
|
||||
assert body["results"]["1"]["node_id"] == "n-1"
|
||||
assert body["results"]["2"]["ws_id"] == "child-2"
|
||||
assert body["results"]["2"]["child_ws_id"] == "child-2"
|
||||
# Confirm we don't leak the old ``ws_id`` key alongside the new
|
||||
# ``child_ws_id`` — see test_spawn_exec_does_not_surface_misleading_status_field
|
||||
# for the rationale on the rename.
|
||||
assert "ws_id" not in body["results"]["0"]
|
||||
|
||||
|
||||
def test_spawn_batch_exec_surfaces_per_item_errors_in_denied(coord_session):
|
||||
|
||||
@@ -15,6 +15,16 @@ class TestIsSecret:
|
||||
assert _is_secret("TURNSTONE_JWT_SECRET") is True
|
||||
assert _is_secret("AWS_SECRET_ACCESS_KEY") is True
|
||||
|
||||
def test_tool_config_paths_scrubbed(self):
|
||||
"""Tool-config env vars whose target files load executable
|
||||
directives must be scrubbed even though they don't match a
|
||||
secret-suffix pattern. Defence-in-depth alongside on-CLI
|
||||
``--no-config`` for ripgrep and friends."""
|
||||
assert _is_secret("RIPGREP_CONFIG_PATH") is True
|
||||
assert _is_secret("GIT_CONFIG") is True
|
||||
assert _is_secret("GIT_CONFIG_GLOBAL") is True
|
||||
assert _is_secret("GIT_CONFIG_SYSTEM") is True
|
||||
|
||||
def test_suffix_matching(self):
|
||||
assert _is_secret("MY_CUSTOM_API_KEY") is True
|
||||
assert _is_secret("DB_PASSWORD") is True
|
||||
|
||||
@@ -8,6 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -317,6 +318,13 @@ class TestPromptTemplateCRUD:
|
||||
def test_get_prompt_template_nonexistent(self, db):
|
||||
assert db.get_prompt_template("missing") is None
|
||||
|
||||
def test_create_prompt_template_duplicate_id_raises_conflict(self, db):
|
||||
from turnstone.core.storage._protocol import StorageConflictError
|
||||
|
||||
db.create_prompt_template("dup", "first", "general", "A")
|
||||
with pytest.raises(StorageConflictError, match="prompt_template conflict"):
|
||||
db.create_prompt_template("dup", "second", "general", "B")
|
||||
|
||||
def test_list_prompt_templates_ordered_by_name(self, db):
|
||||
db.create_prompt_template("t2", "beta", "general", "B")
|
||||
db.create_prompt_template("t1", "alpha", "general", "A")
|
||||
|
||||
@@ -0,0 +1,677 @@
|
||||
"""Unit tests for ``turnstone.core.history_decoration``.
|
||||
|
||||
The decoration helpers are shared between two surfaces — interactive's
|
||||
SSE replay (``_build_history``) and the lifted ``/history`` REST
|
||||
endpoint (``make_history_handler``, used by both interactive and
|
||||
coord). Pinning the wire shape here lets a future schema/projection
|
||||
change land in one file rather than spread across the two surfaces.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from turnstone.core.history_decoration import (
|
||||
build_output_assessment_payload,
|
||||
build_verdict_payload,
|
||||
decorate_history_messages,
|
||||
decorate_tool_call,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildVerdictPayload:
|
||||
"""The wire-shape projection that's the single source of truth for
|
||||
what intent_verdict fields ship to the client."""
|
||||
|
||||
def test_skips_unflagged_baseline(self) -> None:
|
||||
"""``risk_level`` "none" is the unflagged-tool baseline; the
|
||||
client filters those anyway, so projecting None at the wire
|
||||
layer keeps the payload tight on long workstreams."""
|
||||
row = {"risk_level": "none", "recommendation": "approve", "tier": "heuristic"}
|
||||
assert build_verdict_payload(row) is None
|
||||
|
||||
def test_drops_call_id_and_func_name(self) -> None:
|
||||
"""The client already has these on ``tc.id`` / ``tc.name``;
|
||||
re-shipping them per-tool_call would balloon long replays."""
|
||||
row = {
|
||||
"call_id": "call_abc",
|
||||
"func_name": "bash",
|
||||
"risk_level": "medium",
|
||||
"recommendation": "review",
|
||||
"confidence": 0.8,
|
||||
"intent_summary": "summary",
|
||||
"tier": "heuristic",
|
||||
}
|
||||
out = build_verdict_payload(row)
|
||||
assert out is not None
|
||||
assert "call_id" not in out
|
||||
assert "func_name" not in out
|
||||
# Sanity — the kept fields are the ones renderVerdictBadge reads.
|
||||
assert out["risk_level"] == "medium"
|
||||
assert out["recommendation"] == "review"
|
||||
assert out["confidence"] == 0.8
|
||||
assert out["intent_summary"] == "summary"
|
||||
assert out["tier"] == "heuristic"
|
||||
|
||||
def test_includes_reasoning_for_either_tier_when_present(self) -> None:
|
||||
"""Heuristic verdicts in this project emit structured
|
||||
rationales (one per matched pattern) — e.g.
|
||||
``policy.py`` writes a reasoning string per heuristic hit.
|
||||
Ship the field for either tier when it has content; only
|
||||
omit when the row didn't write one."""
|
||||
for tier in ("heuristic", "llm"):
|
||||
row = {
|
||||
"risk_level": "high",
|
||||
"tier": tier,
|
||||
"reasoning": "The command exfiltrates ~/.ssh/id_rsa over an external connection.",
|
||||
}
|
||||
out = build_verdict_payload(row)
|
||||
assert out is not None
|
||||
assert "id_rsa" in out["reasoning"]
|
||||
|
||||
def test_omits_reasoning_when_empty(self) -> None:
|
||||
"""An absent / empty reasoning string shouldn't ship as
|
||||
``reasoning: ""`` — the rationale ``<details>`` block on the
|
||||
client renders an empty disclosure when the field is present
|
||||
but empty."""
|
||||
row = {"risk_level": "high", "tier": "heuristic", "reasoning": ""}
|
||||
out = build_verdict_payload(row)
|
||||
assert out is not None
|
||||
assert "reasoning" not in out
|
||||
|
||||
def test_includes_judge_model_when_present(self) -> None:
|
||||
"""``judge_model`` rides through so the batch tier badge can
|
||||
render ``⚖ llm:claude-haiku-4`` on history-only replays
|
||||
rather than the bare ``⚖ llm`` label."""
|
||||
row = {"risk_level": "high", "tier": "llm", "judge_model": "claude-haiku-4"}
|
||||
out = build_verdict_payload(row)
|
||||
assert out is not None
|
||||
assert out["judge_model"] == "claude-haiku-4"
|
||||
|
||||
def test_omits_judge_model_when_empty(self) -> None:
|
||||
row = {"risk_level": "medium", "tier": "heuristic", "judge_model": ""}
|
||||
out = build_verdict_payload(row)
|
||||
assert out is not None
|
||||
assert "judge_model" not in out
|
||||
|
||||
|
||||
class TestBuildOutputAssessmentPayload:
|
||||
"""Output-guard wire shape — flags decoded from JSON string at
|
||||
this layer so the client never has to parse twice."""
|
||||
|
||||
def test_skips_unflagged_baseline(self) -> None:
|
||||
row = {"risk_level": "none", "flags": "[]"}
|
||||
assert build_output_assessment_payload(row) is None
|
||||
|
||||
def test_decodes_flags_from_json(self) -> None:
|
||||
row = {"risk_level": "high", "flags": '["api_key","email"]', "redacted": 1}
|
||||
out = build_output_assessment_payload(row)
|
||||
assert out is not None
|
||||
assert out["flags"] == ["api_key", "email"]
|
||||
assert out["redacted"] is True
|
||||
assert out["risk_level"] == "high"
|
||||
|
||||
def test_handles_malformed_flags_json(self) -> None:
|
||||
"""Bad JSON in ``flags`` must not block the rest of the
|
||||
assessment from rendering — degrade to empty list."""
|
||||
row = {"risk_level": "medium", "flags": "not-json", "redacted": 0}
|
||||
out = build_output_assessment_payload(row)
|
||||
assert out is not None
|
||||
assert out["flags"] == []
|
||||
assert out["redacted"] is False
|
||||
|
||||
|
||||
class TestDecorateToolCall:
|
||||
"""In-place mutation of either OpenAI-format or flattened tool_call
|
||||
entries — both shapes carry ``id`` at the top level."""
|
||||
|
||||
def test_attaches_verdict_when_present(self) -> None:
|
||||
tc: dict[str, object] = {"id": "call_1", "function": {"name": "bash", "arguments": "{}"}}
|
||||
verdicts = {
|
||||
"call_1": {
|
||||
"risk_level": "medium",
|
||||
"recommendation": "review",
|
||||
"confidence": 0.7,
|
||||
"intent_summary": "summary",
|
||||
"tier": "heuristic",
|
||||
}
|
||||
}
|
||||
decorate_tool_call(tc, verdicts, {})
|
||||
assert "verdict" in tc
|
||||
assert tc["verdict"]["risk_level"] == "medium" # type: ignore[index]
|
||||
|
||||
def test_skips_when_no_call_id_match(self) -> None:
|
||||
tc: dict[str, object] = {"id": "call_other", "name": "bash"}
|
||||
verdicts = {
|
||||
"call_1": {"risk_level": "medium", "tier": "heuristic"},
|
||||
}
|
||||
decorate_tool_call(tc, verdicts, {})
|
||||
assert "verdict" not in tc
|
||||
|
||||
def test_skips_unflagged_verdict(self) -> None:
|
||||
"""``build_verdict_payload`` returns None for unflagged rows;
|
||||
decorate_tool_call must not stamp ``verdict`` in that case."""
|
||||
tc: dict[str, object] = {"id": "call_1", "name": "bash"}
|
||||
verdicts = {"call_1": {"risk_level": "none", "tier": "heuristic"}}
|
||||
decorate_tool_call(tc, verdicts, {})
|
||||
assert "verdict" not in tc
|
||||
|
||||
def test_handles_empty_id(self) -> None:
|
||||
"""A tool_call with no id can't be paired against the lookup
|
||||
table — must not raise (or stamp the wrong row's verdict)."""
|
||||
tc: dict[str, object] = {"id": "", "name": "bash"}
|
||||
verdicts = {"call_1": {"risk_level": "high", "tier": "heuristic"}}
|
||||
decorate_tool_call(tc, verdicts, {})
|
||||
assert "verdict" not in tc
|
||||
|
||||
|
||||
class TestDecorateHistoryMessages:
|
||||
"""End-to-end mutation of a /history-shaped message list — covers
|
||||
the full transform applied by ``make_history_handler``."""
|
||||
|
||||
def test_decorates_tool_calls_with_verdict_and_assessment(self) -> None:
|
||||
verdicts = {
|
||||
"call_a": {
|
||||
"risk_level": "high",
|
||||
"recommendation": "deny",
|
||||
"confidence": 0.95,
|
||||
"intent_summary": "exfil",
|
||||
"tier": "llm",
|
||||
"reasoning": "ssh key access",
|
||||
}
|
||||
}
|
||||
assessments = {
|
||||
"call_a": {"risk_level": "high", "flags": '["secret"]', "redacted": 1},
|
||||
}
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "running",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_a",
|
||||
"function": {"name": "bash", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": "long output"},
|
||||
{"role": "tool", "tool_call_id": "call_b", "content": "short"},
|
||||
]
|
||||
decorate_history_messages(messages, verdicts, assessments)
|
||||
# Assistant tool_calls got both decorations.
|
||||
tc = messages[1]["tool_calls"][0] # type: ignore[index]
|
||||
assert tc["verdict"]["risk_level"] == "high"
|
||||
assert tc["verdict"]["tier"] == "llm"
|
||||
assert "reasoning" in tc["verdict"]
|
||||
assert tc["output_assessment"]["flags"] == ["secret"]
|
||||
assert tc["output_assessment"]["redacted"] is True
|
||||
# Plain tool content (no envelope) is left intact and no
|
||||
# advisories key is set.
|
||||
assert messages[2]["content"] == "long output"
|
||||
assert "advisories" not in messages[2]
|
||||
assert messages[3]["content"] == "short"
|
||||
assert "advisories" not in messages[3]
|
||||
|
||||
def test_no_op_on_empty_indexes(self) -> None:
|
||||
"""When neither table has rows for the workstream, the wire
|
||||
shape passes through unchanged — replay must degrade
|
||||
gracefully when verdict storage is empty / unavailable."""
|
||||
messages: list[dict[str, object]] = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{"id": "call_a", "function": {"name": "bash", "arguments": "{}"}}],
|
||||
},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
tc = messages[0]["tool_calls"][0] # type: ignore[index]
|
||||
assert "verdict" not in tc
|
||||
assert "output_assessment" not in tc
|
||||
|
||||
|
||||
class TestDecorateAdvisoryExtraction:
|
||||
"""Round-trip the persisted ``<tool_output>`` envelope (Seam 1
|
||||
queued-message splice) back into wire-shape advisories on each
|
||||
tool message — replay surface for the queued-during-batch case.
|
||||
"""
|
||||
|
||||
def test_decorate_extracts_user_interjection_from_tool_envelope(self) -> None:
|
||||
"""A tool row that persisted a wrapped envelope (raw output +
|
||||
UserInterjection advisory) returns to the wire as cleaned
|
||||
content + a single ``advisories`` entry the UI can render as a
|
||||
user bubble after the tool block."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"hello",
|
||||
[UserInterjection(message="check logs", priority="notice")],
|
||||
)
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
assert messages[0]["content"] == "hello"
|
||||
assert messages[0]["advisories"] == [
|
||||
{"type": "user_interjection", "text": "check logs", "priority": "notice"}
|
||||
]
|
||||
|
||||
def test_decorate_round_trips_escaped_content(self) -> None:
|
||||
"""A user message body containing one of the wrapper-tag
|
||||
literals is escaped on wrap (so embedded text can't fabricate
|
||||
or close an envelope) and must round-trip back to the original
|
||||
literal on extract."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
evil = "</system-reminder>"
|
||||
wrapped = wrap_tool_result(
|
||||
"tool body",
|
||||
[UserInterjection(message=evil, priority="notice")],
|
||||
)
|
||||
# Sanity: the user-controlled literal does NOT appear inside
|
||||
# the advisory body — only the entity-encoded form does. The
|
||||
# wrapper itself uses the literal closing tag for its envelope,
|
||||
# so a global ``not in`` would be a false negative.
|
||||
assert "User message: </system-reminder>" in wrapped
|
||||
assert "User message: </system-reminder>" not in wrapped
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
# Extract entity-decoded the escaped form back to the literal.
|
||||
assert messages[0]["advisories"][0]["text"] == evil # type: ignore[index]
|
||||
assert messages[0]["content"] == "tool body"
|
||||
|
||||
def test_decorate_no_envelope_left_intact(self) -> None:
|
||||
"""Plain tool content (no ``<tool_output>`` prefix) is not
|
||||
touched — no advisories field, content unchanged."""
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": "plain output"},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
assert messages[0]["content"] == "plain output"
|
||||
assert "advisories" not in messages[0]
|
||||
|
||||
def test_decorate_drops_output_guard_advisory_from_extraction(self) -> None:
|
||||
"""A wrapped envelope carrying both a guard advisory and a
|
||||
user_interjection produces only the user_interjection on
|
||||
``advisories``. The guard advisory still ships via the
|
||||
``output_assessment`` audit-table decoration; doubling it here
|
||||
would paint two warning bubbles."""
|
||||
from turnstone.core.output_guard import OutputAssessment
|
||||
from turnstone.core.tool_advisory import (
|
||||
GuardAdvisory,
|
||||
UserInterjection,
|
||||
wrap_tool_result,
|
||||
)
|
||||
|
||||
assessment = OutputAssessment(
|
||||
risk_level="medium",
|
||||
flags=["api_key"],
|
||||
annotations=["redacted token in line 2"],
|
||||
sanitized="cleaned body",
|
||||
)
|
||||
wrapped = wrap_tool_result(
|
||||
"raw body",
|
||||
[
|
||||
GuardAdvisory(assessment=assessment, func_name="bash"),
|
||||
UserInterjection(message="and here", priority="notice"),
|
||||
],
|
||||
)
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
adv = messages[0]["advisories"]
|
||||
assert len(adv) == 1 # type: ignore[arg-type]
|
||||
assert adv[0]["type"] == "user_interjection" # type: ignore[index]
|
||||
|
||||
def test_decorate_handles_important_priority(self) -> None:
|
||||
"""The MUST-address preamble round-trips to ``priority=important``."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"out",
|
||||
[UserInterjection(message="urgent", priority="important")],
|
||||
)
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
adv = messages[0]["advisories"][0] # type: ignore[index]
|
||||
assert adv["priority"] == "important"
|
||||
assert adv["text"] == "urgent"
|
||||
|
||||
def test_decorate_suppresses_empty_advisory_body(self) -> None:
|
||||
"""``queue_message`` doesn't reject empty / whitespace-only
|
||||
text, so an advisory with an empty body can round-trip through
|
||||
``wrap_tool_result``. ``_classify_advisory`` must filter those
|
||||
out so replay doesn't paint a featureless empty user bubble.
|
||||
|
||||
Removing the ``if not body.strip(): return None`` guard in
|
||||
``_classify_advisory`` breaks this test."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"tool body",
|
||||
[UserInterjection(message="", priority="notice")],
|
||||
)
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
# Envelope is still stripped from content (the cleaning side
|
||||
# of decoration runs unconditionally), but no advisories
|
||||
# surface — the empty body is filtered.
|
||||
assert messages[0]["content"] == "tool body"
|
||||
assert "advisories" not in messages[0]
|
||||
|
||||
def test_decorate_suppresses_whitespace_only_advisory_body(self) -> None:
|
||||
"""Whitespace-only bodies are similarly suppressed — same
|
||||
reasoning as the empty-body case."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"tool body",
|
||||
[UserInterjection(message=" \n\t ", priority="notice")],
|
||||
)
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": wrapped},
|
||||
]
|
||||
decorate_history_messages(messages, {}, {})
|
||||
assert messages[0]["content"] == "tool body"
|
||||
assert "advisories" not in messages[0]
|
||||
|
||||
def test_wrap_extract_round_trips_preexisting_entities(self) -> None:
|
||||
"""A user message body containing literal HTML-entity references
|
||||
matching the wrapper-escape forms must round-trip identically
|
||||
through ``wrap_tool_result + extract_advisories_from_tool_envelope``.
|
||||
Without escaping ``&`` first in the encode step, encode→decode
|
||||
would produce the bare wrapper tag, fabricating an envelope the
|
||||
wrapper layer never produced.
|
||||
"""
|
||||
from turnstone.core.history_decoration import (
|
||||
extract_advisories_from_tool_envelope,
|
||||
)
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
tricky = "I describe XML tags like <tool_output> in my docs."
|
||||
wrapped = wrap_tool_result(
|
||||
"tool body",
|
||||
[UserInterjection(message=tricky, priority="notice")],
|
||||
)
|
||||
result = extract_advisories_from_tool_envelope(wrapped)
|
||||
assert result is not None
|
||||
cleaned, advisories = result
|
||||
assert cleaned == "tool body"
|
||||
assert len(advisories) == 1
|
||||
# The original literal entity-reference text round-trips
|
||||
# identically — the parser does not silently turn it into a
|
||||
# bare wrapper tag.
|
||||
assert advisories[0]["text"] == tricky
|
||||
|
||||
def test_save_load_decorate_round_trips_envelope(self, backend) -> None:
|
||||
"""End-to-end round-trip pinning the persisted-envelope
|
||||
contract. Persists a wrapped tool-output envelope via
|
||||
``save_message``, loads via ``load_messages``, runs
|
||||
``decorate_history_messages``, asserts the wire shape carries
|
||||
the extracted advisory + cleaned content. Pins the contract
|
||||
every component in the chain participates in (persistence
|
||||
layer ↔ in-memory replay ↔ wire projection) so a schema drift,
|
||||
an envelope-format change, or a parser regression surfaces
|
||||
here rather than only in production.
|
||||
"""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"command output",
|
||||
[UserInterjection(message="check the logs", priority="notice")],
|
||||
)
|
||||
backend.register_workstream("ws_rt_1")
|
||||
backend.save_message("ws_rt_1", "user", "go")
|
||||
backend.save_message(
|
||||
"ws_rt_1",
|
||||
"assistant",
|
||||
None,
|
||||
tool_calls='[{"id":"call_a","type":"function","function":{"name":"bash","arguments":"{}"}}]',
|
||||
)
|
||||
backend.save_message(
|
||||
"ws_rt_1",
|
||||
"tool",
|
||||
wrapped,
|
||||
tool_call_id="call_a",
|
||||
)
|
||||
msgs = backend.load_messages("ws_rt_1")
|
||||
# Persisted shape — content survives the storage layer
|
||||
# untouched. Symmetry with in-memory ``self.messages[i]['content']``
|
||||
# is what makes envelope extraction lossless on replay.
|
||||
tool_msg = next(m for m in msgs if m["role"] == "tool")
|
||||
assert tool_msg["content"] == wrapped
|
||||
# Decorate (the /history shared transform) — extracts the
|
||||
# advisory and strips the envelope.
|
||||
decorate_history_messages(msgs, {}, {})
|
||||
tool_msg = next(m for m in msgs if m["role"] == "tool")
|
||||
assert tool_msg["content"] == "command output"
|
||||
assert tool_msg["advisories"] == [
|
||||
{"type": "user_interjection", "text": "check the logs", "priority": "notice"}
|
||||
]
|
||||
|
||||
|
||||
class TestExtractReasoningForHistory:
|
||||
"""``extract_reasoning_for_history`` — Phase 1 surfaces stored
|
||||
Anthropic thinking blocks on assistant messages and strips
|
||||
``_provider_content`` from the wire payload.
|
||||
|
||||
Drives through the real ``AnthropicProvider.extract_reasoning_text``
|
||||
(no mock-of-extractor) — the helper test and the provider unit
|
||||
test (``tests/test_provider_anthropic_reasoning.py``) together
|
||||
catch a regression at either layer distinctly.
|
||||
"""
|
||||
|
||||
def _anthropic_thinking_msg(self, text: str = "let me think") -> dict[str, object]:
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": "Final answer.",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": text, "signature": "sig"},
|
||||
{"type": "text", "text": "Final answer."},
|
||||
],
|
||||
}
|
||||
|
||||
def test_extract_thinking_surfaces_reasoning_field(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [self._anthropic_thinking_msg("let me think")]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "let me think"
|
||||
|
||||
def test_strips_provider_content_after_extraction(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [self._anthropic_thinking_msg("anything")]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_strips_provider_content_when_flag_false(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [self._anthropic_thinking_msg("anything")]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=False)
|
||||
# Strip is unconditional; reasoning is the conditional bit.
|
||||
assert "_provider_content" not in messages[0]
|
||||
assert "reasoning" not in messages[0]
|
||||
|
||||
def test_first_block_thinking_dispatches_to_anthropic(self) -> None:
|
||||
# Even when text and tool_use blocks follow, the first-block-type
|
||||
# discriminator routes thinking-prefixed payloads correctly.
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"_provider_content": [
|
||||
{"type": "thinking", "thinking": "first", "signature": "s"},
|
||||
{"type": "text", "text": "spoken"},
|
||||
{"type": "tool_use", "id": "t1", "name": "f", "input": {}},
|
||||
],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "first"
|
||||
|
||||
def test_first_block_reasoning_dispatches_to_openai_responses(self) -> None:
|
||||
# Phase 3: dispatcher routes type=="reasoning" to the
|
||||
# OpenAI Responses extractor, which now returns the
|
||||
# summary[*].text concatenation. Pre-Phase-3 this asserted
|
||||
# "" (the stub); the assertion was tightened once the wire
|
||||
# path landed.
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"_provider_content": [
|
||||
{"type": "reasoning", "summary": [{"type": "summary_text", "text": "s"}]}
|
||||
],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "s"
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_unknown_first_block_type_no_op(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"_provider_content": [{"type": "text", "text": "no reasoning here"}],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "reasoning" not in messages[0]
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_skips_messages_without_provider_content(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [{"role": "assistant", "content": "plain"}]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "reasoning" not in messages[0]
|
||||
assert messages[0]["content"] == "plain"
|
||||
|
||||
def test_user_and_tool_messages_untouched(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "out"},
|
||||
self._anthropic_thinking_msg("only this one"),
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "reasoning" not in messages[0]
|
||||
assert "reasoning" not in messages[1]
|
||||
assert messages[2]["reasoning"] == "only this one"
|
||||
|
||||
def test_empty_provider_content_no_extraction(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [{"role": "assistant", "content": "x", "_provider_content": []}]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "reasoning" not in messages[0]
|
||||
# Empty-list provider_content is still stripped from the wire.
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_first_block_not_a_dict_skipped(self) -> None:
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages: list[dict[str, object]] = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"_provider_content": ["bogus"],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert "reasoning" not in messages[0]
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_first_block_reasoning_text_dispatches_to_openai_chat(self) -> None:
|
||||
# Phase 3 path 3: synthetic ``reasoning_text`` blocks (stamped
|
||||
# by ChatSession._maybe_synth_reasoning_block for vLLM /
|
||||
# llama.cpp / Gemini-compat conversations) dispatch to
|
||||
# OpenAIChatCompletionsProvider.extract_reasoning_text.
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "answer",
|
||||
"_provider_content": [
|
||||
{"type": "reasoning_text", "text": "synth thought", "source": "vllm"},
|
||||
],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "synth thought"
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_dispatcher_scans_past_unrecognized_first_blocks(self) -> None:
|
||||
# Regression for Copilot finding: dispatcher used to inspect
|
||||
# only provider_content[0]['type']. OpenAI Responses captures
|
||||
# EVERY output_item.done event into provider_blocks (not just
|
||||
# reasoning), so a hypothetical [message, reasoning, ...]
|
||||
# ordering would have silently dropped the reasoning. Now
|
||||
# walks the list for the first recognised reasoning-bearing
|
||||
# type and dispatches the whole list to that provider.
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "answer",
|
||||
"_provider_content": [
|
||||
# First block is a non-reasoning OpenAI Responses item.
|
||||
{"type": "message", "role": "assistant", "content": "answer"},
|
||||
# Reasoning sits later in the list.
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "r_1",
|
||||
"summary": [{"type": "summary_text", "text": "deferred"}],
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "deferred"
|
||||
assert "_provider_content" not in messages[0]
|
||||
|
||||
def test_first_block_redacted_thinking_dispatches_to_anthropic(self) -> None:
|
||||
# Anthropic's extended-thinking API documents that
|
||||
# ``redacted_thinking`` blocks (sealed by the safety system)
|
||||
# can appear before, after, or interleaved with regular
|
||||
# ``thinking`` blocks. When the redacted block lands first,
|
||||
# the dispatcher must still route to AnthropicProvider so the
|
||||
# surrounding real thinking text surfaces — without this the
|
||||
# reasoning bubble silently disappears on history rehydration.
|
||||
# Pinned by registering "redacted_thinking" as a second key
|
||||
# in _BLOCK_TYPE_PROVIDER_FACTORY pointing at the Anthropic
|
||||
# factory; Anthropic's extractor's type=="thinking" filter
|
||||
# already correctly skips the redacted block.
|
||||
from turnstone.core.history_decoration import extract_reasoning_for_history
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "answer",
|
||||
"_provider_content": [
|
||||
{"type": "redacted_thinking", "data": "sealed-blob"},
|
||||
{"type": "thinking", "thinking": "real thought", "signature": "s"},
|
||||
{"type": "text", "text": "answer"},
|
||||
],
|
||||
}
|
||||
]
|
||||
extract_reasoning_for_history(messages, surface_persisted_reasoning_flag=True)
|
||||
assert messages[0]["reasoning"] == "real thought"
|
||||
assert "_provider_content" not in messages[0]
|
||||
@@ -0,0 +1,391 @@
|
||||
"""Boundary-crossing integration test for the wake trigger pipeline.
|
||||
|
||||
Drives a *real* :class:`SessionManager` + a *real* :class:`ChatSession`
|
||||
+ a *real* :class:`IdleNudgeWatcher` end-to-end. The only stub is the
|
||||
LLM provider (patched ``_create_stream_with_retry``); every other layer
|
||||
is production code:
|
||||
|
||||
* ``SessionManager.set_state`` snapshotting + iterating subscribers
|
||||
* ``IdleNudgeWatcher._on_state`` peeking the queue
|
||||
* ``session_worker.send`` atomic-spawn + daemon thread
|
||||
* ``ChatSession.deliver_wake_nudge_from_queue`` opening / closing
|
||||
``_wake_source_tag``
|
||||
* ``ChatSession.send`` chat loop short-circuiting metacog detection
|
||||
* ``_append_user_turn`` stamping ``_source = "system_nudge"``
|
||||
* ``_attach_pending_user_reminders`` draining ``USER_DRAIN``
|
||||
* ``_apply_reminders_for_provider`` splicing the rendered envelope
|
||||
onto empty content
|
||||
|
||||
Per ``feedback_tests_through_boundaries.md``: direct injection tests
|
||||
that bypass these boundaries silently mask wiring bugs. This test is
|
||||
the structural integration gate.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.test_session_manager import FakeStorage
|
||||
from turnstone.core.idle_nudge_watcher import IdleNudgeWatcher
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.session_manager import SessionManager
|
||||
from turnstone.core.workstream import Workstream, WorkstreamKind, WorkstreamState
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Minimal fake adapter / UI for this integration test. Storage reuses
|
||||
# the canonical FakeStorage from test_session_manager.py to avoid the
|
||||
# drift risk of a parallel fake.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeUI:
|
||||
"""Minimal UI surface for ChatSession + SessionManager.cleanup_ui."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.events: list[tuple[str, Any]] = []
|
||||
|
||||
def _unblock(self) -> None: # SessionManager.close calls this
|
||||
pass
|
||||
|
||||
def broadcast_ws_closed(self) -> None:
|
||||
pass
|
||||
|
||||
# ChatSession callbacks (no-op for this test)
|
||||
def on_turn_start(self) -> None:
|
||||
pass
|
||||
|
||||
def on_turn_committed(self) -> None:
|
||||
pass
|
||||
|
||||
def on_thinking_start(self) -> None:
|
||||
pass
|
||||
|
||||
def on_thinking_end(self) -> None:
|
||||
pass
|
||||
|
||||
def on_state_change(self, state: str) -> None:
|
||||
self.events.append(("state", state))
|
||||
|
||||
def on_user_reminder(self, reminders: Any, source: str | None = None) -> None:
|
||||
self.events.append(("user_reminder", reminders, source))
|
||||
|
||||
def on_error(self, message: str) -> None:
|
||||
pass
|
||||
|
||||
def on_rename(self, name: str) -> None:
|
||||
pass
|
||||
|
||||
def on_output_warning(self, call_id: Any, assessment: Any) -> None:
|
||||
pass
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
# Catch-all for any UI hook not enumerated above so the chat
|
||||
# loop's ``self.ui.<something>()`` call doesn't blow up.
|
||||
return MagicMock()
|
||||
|
||||
|
||||
class _BuildRealSessionAdapter:
|
||||
"""Adapter that returns a real :class:`ChatSession` instead of a stub.
|
||||
|
||||
Tracks emit_* events the integration test asserts on. Mirrors the
|
||||
``SessionKindAdapter`` + ``SessionEventEmitter`` Protocol surface
|
||||
that production ``WebUI`` / coord adapters expose.
|
||||
"""
|
||||
|
||||
def __init__(self, kind: WorkstreamKind = WorkstreamKind.INTERACTIVE) -> None:
|
||||
self.kind = kind
|
||||
self.events: list[str] = []
|
||||
self.cleaned_up: list[str] = []
|
||||
|
||||
def emit_created(self, ws: Workstream) -> None:
|
||||
self.events.append(f"created:{ws.id}")
|
||||
|
||||
def emit_rehydrated(self, ws: Workstream) -> None:
|
||||
self.events.append(f"rehydrated:{ws.id}")
|
||||
|
||||
def emit_state(self, ws: Workstream, state: WorkstreamState) -> None:
|
||||
self.events.append(f"state:{ws.id}:{state.value}")
|
||||
|
||||
def emit_closed(self, ws_id: str, *, reason: str = "closed", name: str = "") -> None:
|
||||
self.events.append(f"closed:{ws_id}")
|
||||
|
||||
def cleanup_ui(self, ws: Workstream) -> None:
|
||||
# Real production cleanup_ui calls ws.session.cancel() + close().
|
||||
# We don't need that here — the test exits cleanly via pytest
|
||||
# teardown without exercising the cleanup path. Just record
|
||||
# the call for any test that wants to assert on it.
|
||||
self.cleaned_up.append(ws.id)
|
||||
|
||||
def build_ui(self, ws: Workstream) -> Any:
|
||||
return _FakeUI()
|
||||
|
||||
def build_session(
|
||||
self,
|
||||
ws: Workstream,
|
||||
*,
|
||||
skill: Any = None,
|
||||
model: Any = None,
|
||||
client_type: Any = None,
|
||||
**extra: Any,
|
||||
) -> Any:
|
||||
# Mirror SessionManager.create's keyword set so config-threading
|
||||
# bugs surface here rather than being silently swallowed by
|
||||
# **kwargs. ``model`` flows to the real ChatSession; the rest
|
||||
# are accepted but not used by this test.
|
||||
client = MagicMock()
|
||||
return ChatSession(
|
||||
client=client,
|
||||
model=str(model) if model else "test-model",
|
||||
ui=ws.ui,
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
tool_timeout=30,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def real_mgr() -> tuple[SessionManager, _BuildRealSessionAdapter]:
|
||||
"""Real SessionManager wired to an adapter that builds real ChatSessions.
|
||||
|
||||
No StateWriter is wired so ``set_state`` writes directly to storage
|
||||
on the calling thread (we want subscriber dispatch to fire in the
|
||||
same thread the test invokes ``set_state`` on).
|
||||
"""
|
||||
adapter = _BuildRealSessionAdapter()
|
||||
storage = FakeStorage()
|
||||
mgr = SessionManager(
|
||||
adapter,
|
||||
storage=storage,
|
||||
max_active=5,
|
||||
event_emitter=adapter,
|
||||
)
|
||||
return mgr, adapter
|
||||
|
||||
|
||||
def _wait_for_worker_done(ws: Workstream, timeout: float = 5.0) -> None:
|
||||
"""Poll ``ws._worker_running`` until it clears or timeout elapses."""
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
with ws._lock:
|
||||
if not ws._worker_running:
|
||||
return
|
||||
time.sleep(0.01)
|
||||
raise AssertionError(f"worker thread for ws={ws.id[:8]} didn't exit within {timeout}s")
|
||||
|
||||
|
||||
def test_idle_event_through_real_session_manager_drives_wake_send(real_mgr, tmp_db):
|
||||
"""The full wake pipeline, no direct-injection shortcuts.
|
||||
|
||||
Boundary path under test:
|
||||
enqueue → mgr.set_state(IDLE)
|
||||
→ SessionManager._state_subscribers iteration (real)
|
||||
→ IdleNudgeWatcher._on_state (real)
|
||||
→ session_worker.send (real)
|
||||
→ real daemon thread
|
||||
→ ChatSession.deliver_wake_nudge_from_queue (real)
|
||||
→ ChatSession.send("") (real, with patched LLM stream)
|
||||
→ _append_user_turn stamps ``_source``
|
||||
→ _attach_pending_user_reminders drains ``{"user","any"}``
|
||||
→ _apply_reminders_for_provider splices envelope onto empty content
|
||||
"""
|
||||
mgr, _adapter = real_mgr
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
|
||||
try:
|
||||
ws = mgr.create(user_id="u1", name="wake-int", skill=None)
|
||||
assert ws.session is not None
|
||||
# Patch the LLM-facing surface so send() runs the chat loop end-to-end
|
||||
# without any real provider. We patch on the just-built ChatSession;
|
||||
# the patches are reverted by the `with` block.
|
||||
with (
|
||||
patch.object(ws.session, "_create_stream_with_retry", return_value=iter([])),
|
||||
patch.object(
|
||||
ws.session,
|
||||
"_stream_response",
|
||||
return_value={"role": "assistant", "content": "ok"},
|
||||
),
|
||||
patch.object(ws.session, "_update_token_table"),
|
||||
patch.object(ws.session, "_print_status_line"),
|
||||
patch.object(ws.session, "_visible_memory_count", return_value=0),
|
||||
patch("turnstone.core.session.save_message"),
|
||||
):
|
||||
# Suppress the auto-title side-thread; orthogonal to wake.
|
||||
ws.session._title_generated = True
|
||||
|
||||
# Enqueue an any-channel nudge — the future ``idle_children`` shape.
|
||||
ws.session._nudge_queue.enqueue("idle_children", "your kids", "any")
|
||||
assert len(ws.session._nudge_queue) == 1
|
||||
|
||||
# Trigger IDLE. This runs subscriber dispatch synchronously on
|
||||
# the calling thread → IdleNudgeWatcher._on_state → session_worker.send
|
||||
# → spawn daemon thread → deliver_wake_nudge_from_queue.
|
||||
mgr.set_state(ws.id, WorkstreamState.IDLE)
|
||||
|
||||
# Wait for the daemon thread to clear ``_worker_running`` so the
|
||||
# post-conditions are stable.
|
||||
_wait_for_worker_done(ws)
|
||||
|
||||
# Queue fully drained by the wake.
|
||||
assert len(ws.session._nudge_queue) == 0
|
||||
|
||||
# The synthesized empty user message landed in history with the
|
||||
# ``_source`` audit tag and the reminder side-channel populated.
|
||||
user_msgs = [m for m in ws.session.messages if m.get("role") == "user"]
|
||||
assert user_msgs, "expected a synthesized user message from the wake"
|
||||
wake_msg = user_msgs[-1]
|
||||
assert wake_msg["content"] == ""
|
||||
assert wake_msg.get("_source") == "system_nudge"
|
||||
assert wake_msg.get("_reminders") == [{"type": "idle_children", "text": "your kids"}]
|
||||
|
||||
# The wake-source tag is reset post-send so subsequent activity
|
||||
# behaves normally.
|
||||
assert ws.session._wake_source_tag == ""
|
||||
finally:
|
||||
watcher.shutdown()
|
||||
|
||||
|
||||
def test_idle_event_with_empty_queue_does_not_dispatch_wake(real_mgr, tmp_db):
|
||||
"""Non-empty queue is the gate. An IDLE event on a workstream with
|
||||
nothing queued must NOT call ``session_worker.send``.
|
||||
|
||||
Patches the dispatch primitive directly rather than racing a
|
||||
``time.sleep`` against an erroneous spawn — the question is
|
||||
whether the watcher's gate fired, which is a deterministic
|
||||
decision the patch captures.
|
||||
"""
|
||||
mgr, _adapter = real_mgr
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
|
||||
try:
|
||||
ws = mgr.create(user_id="u1", name="empty-int", skill=None)
|
||||
# No enqueue.
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.set_state(ws.id, WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 0, "wake must not dispatch for an empty queue"
|
||||
finally:
|
||||
watcher.shutdown()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def coord_mgr() -> tuple[SessionManager, _BuildRealSessionAdapter, FakeStorage]:
|
||||
"""Real coord-side SessionManager with the adapter's kind set to
|
||||
COORDINATOR. Same shape as ``real_mgr`` but for the coord half of
|
||||
the lifespan. No StateWriter wired so subscriber dispatch fires
|
||||
synchronously on the test thread.
|
||||
"""
|
||||
adapter = _BuildRealSessionAdapter(kind=WorkstreamKind.COORDINATOR)
|
||||
storage = FakeStorage()
|
||||
mgr = SessionManager(
|
||||
adapter,
|
||||
storage=storage,
|
||||
max_active=5,
|
||||
event_emitter=adapter,
|
||||
)
|
||||
return mgr, adapter, storage
|
||||
|
||||
|
||||
def test_coord_idle_with_active_children_emits_envelope_via_real_managers(coord_mgr, tmp_db):
|
||||
"""Full coord-path integration test (matches design doc §7.4).
|
||||
|
||||
Drives the production install order — ``CoordinatorIdleObserver``
|
||||
registered FIRST, then ``IdleNudgeWatcher`` — and asserts the
|
||||
full chain: observer enqueues on IDLE → watcher peeks → wake
|
||||
spawns a worker → ``deliver_wake_nudge_from_queue`` drains and
|
||||
runs the synthetic empty-user turn → reminder envelope reaches
|
||||
the synthesized user message via the side-channel.
|
||||
|
||||
The boundary-crossing path tested here mirrors what
|
||||
``console/server.py``'s lifespan does at production startup; if
|
||||
the install order is ever reversed, this test fails.
|
||||
"""
|
||||
from turnstone.console.coordinator_idle_observer import CoordinatorIdleObserver
|
||||
from turnstone.core.workstream import WorkstreamKind as _Kind
|
||||
|
||||
mgr, adapter, storage = coord_mgr
|
||||
# Observer FIRST, then watcher. Same order as
|
||||
# ``console/server.py:4435-4443`` — production correctness depends
|
||||
# on subscribers firing in registration order on the same IDLE.
|
||||
observer = CoordinatorIdleObserver(mgr, storage)
|
||||
observer.start()
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
|
||||
try:
|
||||
coord = mgr.create(user_id="u1", name="parent-coord", skill=None)
|
||||
assert coord.session is not None
|
||||
|
||||
# Two interactive children of the coord, both running. Use
|
||||
# the storage's register_workstream API so the rows match
|
||||
# production shape (the observer queries via list_workstreams).
|
||||
storage.register_workstream(
|
||||
"child-a",
|
||||
user_id="u1",
|
||||
name="research-pricing",
|
||||
kind=_Kind.INTERACTIVE,
|
||||
parent_ws_id=coord.id,
|
||||
state="running",
|
||||
)
|
||||
storage.register_workstream(
|
||||
"child-b",
|
||||
user_id="u1",
|
||||
name="draft-rfc",
|
||||
kind=_Kind.INTERACTIVE,
|
||||
parent_ws_id=coord.id,
|
||||
state="thinking",
|
||||
)
|
||||
|
||||
# Pretend the coord has already had a real conversation so
|
||||
# ``should_nudge``'s message_count > 1 gate passes.
|
||||
coord.session.messages.append({"role": "user", "content": "spawn 2"})
|
||||
coord.session.messages.append({"role": "assistant", "content": "ok"})
|
||||
|
||||
with (
|
||||
patch.object(coord.session, "_create_stream_with_retry", return_value=iter([])),
|
||||
patch.object(
|
||||
coord.session,
|
||||
"_stream_response",
|
||||
return_value={"role": "assistant", "content": "ack"},
|
||||
),
|
||||
patch.object(coord.session, "_full_messages", return_value=[]),
|
||||
patch.object(coord.session, "_update_token_table"),
|
||||
patch.object(coord.session, "_print_status_line"),
|
||||
patch.object(coord.session, "_visible_memory_count", return_value=0),
|
||||
patch("turnstone.core.session.save_message"),
|
||||
):
|
||||
coord.session._title_generated = True
|
||||
mgr.set_state(coord.id, WorkstreamState.IDLE)
|
||||
_wait_for_worker_done(coord)
|
||||
|
||||
# Queue drained — the wake delivered the observer's enqueue.
|
||||
assert len(coord.session._nudge_queue) == 0
|
||||
# The synthetic empty-user turn landed with a reminder containing
|
||||
# both children.
|
||||
user_msgs = [m for m in coord.session.messages if m.get("role") == "user"]
|
||||
# Two real msgs (user + assistant context above) plus the wake.
|
||||
wake_msg = user_msgs[-1]
|
||||
assert wake_msg["content"] == ""
|
||||
assert wake_msg.get("_source") == "system_nudge"
|
||||
reminders = wake_msg.get("_reminders") or []
|
||||
assert len(reminders) == 1
|
||||
assert reminders[0]["type"] == "idle_children"
|
||||
text = reminders[0]["text"]
|
||||
assert "research-pricing" in text
|
||||
assert "draft-rfc" in text
|
||||
assert "child-a" in text
|
||||
assert "child-b" in text
|
||||
assert "wait_for_workstream" in text
|
||||
finally:
|
||||
watcher.shutdown()
|
||||
observer.shutdown()
|
||||
@@ -0,0 +1,165 @@
|
||||
"""Unit tests for :class:`IdleNudgeWatcher`.
|
||||
|
||||
Drives a fake :class:`SessionManager` that mimics the real one's
|
||||
``subscribe_to_state`` / ``get`` contract. The watcher itself
|
||||
dispatches via ``turnstone.core.session_worker.send``; we patch that
|
||||
module-level function to capture calls without spawning real threads.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import threading
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.idle_nudge_watcher import IdleNudgeWatcher
|
||||
from turnstone.core.nudge_queue import NudgeQueue
|
||||
from turnstone.core.workstream import WorkstreamState
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self) -> None:
|
||||
self._nudge_queue = NudgeQueue()
|
||||
self.deliver_wake_nudge_from_queue_called = 0
|
||||
|
||||
def deliver_wake_nudge_from_queue(self) -> None:
|
||||
self.deliver_wake_nudge_from_queue_called += 1
|
||||
|
||||
|
||||
class _FakeWorkstream:
|
||||
def __init__(self, ws_id: str = "ws-test") -> None:
|
||||
self.id = ws_id
|
||||
self.session: _FakeSession | None = _FakeSession()
|
||||
self._lock = threading.Lock()
|
||||
self._worker_running = False
|
||||
self._closed = False
|
||||
self.worker_thread: Any = None
|
||||
|
||||
|
||||
class _FakeManager:
|
||||
"""Mimics SessionManager's subscribe-to-state surface without a DB."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._workstreams: dict[str, _FakeWorkstream] = {}
|
||||
self._subscribers: list[Any] = []
|
||||
self._subscribers_lock = threading.Lock()
|
||||
|
||||
def add_ws(self, ws: _FakeWorkstream) -> None:
|
||||
self._workstreams[ws.id] = ws
|
||||
|
||||
def get(self, ws_id: str) -> _FakeWorkstream | None:
|
||||
return self._workstreams.get(ws_id)
|
||||
|
||||
def subscribe_to_state(self, callback: Any) -> None:
|
||||
with self._subscribers_lock:
|
||||
self._subscribers.append(callback)
|
||||
|
||||
def unsubscribe_from_state(self, callback: Any) -> None:
|
||||
with self._subscribers_lock, contextlib.suppress(ValueError):
|
||||
self._subscribers.remove(callback)
|
||||
|
||||
def fire_state(self, ws_id: str, state: WorkstreamState) -> None:
|
||||
"""Mirror SessionManager.set_state's subscriber-fan-out behaviour."""
|
||||
with self._subscribers_lock:
|
||||
subs = list(self._subscribers)
|
||||
for cb in subs:
|
||||
# Match contextlib.suppress(Exception) in real SessionManager.
|
||||
with contextlib.suppress(Exception):
|
||||
cb(ws_id, state)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_mgr_and_ws() -> tuple[_FakeManager, _FakeWorkstream]:
|
||||
mgr = _FakeManager()
|
||||
ws = _FakeWorkstream()
|
||||
mgr.add_ws(ws)
|
||||
return mgr, ws
|
||||
|
||||
|
||||
class TestIdleNudgeWatcher:
|
||||
def test_idle_event_with_empty_queue_no_op(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 0
|
||||
|
||||
def test_idle_event_with_pending_nudge_dispatches(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
ws.session._nudge_queue.enqueue("idle_children", "your kids", "any")
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 1
|
||||
kwargs = mock_send.call_args.kwargs
|
||||
# `enqueue=lambda: None` — verify by calling and checking no-op.
|
||||
assert kwargs["enqueue"]() is None
|
||||
# `run` should call deliver_wake_nudge_from_queue when invoked.
|
||||
kwargs["run"]()
|
||||
assert ws.session.deliver_wake_nudge_from_queue_called == 1
|
||||
assert kwargs["thread_name"].startswith("wake-nudge-")
|
||||
|
||||
def test_non_idle_state_ignored(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
ws.session._nudge_queue.enqueue("foo", "bar", "any")
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
for state in (
|
||||
WorkstreamState.RUNNING,
|
||||
WorkstreamState.THINKING,
|
||||
WorkstreamState.ATTENTION,
|
||||
WorkstreamState.ERROR,
|
||||
):
|
||||
mgr.fire_state(ws.id, state)
|
||||
assert mock_send.call_count == 0
|
||||
|
||||
def test_unknown_ws_ignored(self, fake_mgr_and_ws):
|
||||
mgr, _ws = fake_mgr_and_ws
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state("ghost", WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 0
|
||||
|
||||
def test_session_none_ignored(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
ws.session = None # workstream loaded but session not yet built
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 0
|
||||
|
||||
def test_start_is_idempotent(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
watcher.start() # no-op
|
||||
ws.session._nudge_queue.enqueue("foo", "bar", "any")
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
# Only one subscriber was registered despite the double-start.
|
||||
assert mock_send.call_count == 1
|
||||
|
||||
def test_shutdown_unsubscribes(self, fake_mgr_and_ws):
|
||||
mgr, ws = fake_mgr_and_ws
|
||||
ws.session._nudge_queue.enqueue("foo", "bar", "any")
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
watcher.shutdown()
|
||||
with patch("turnstone.core.session_worker.send") as mock_send:
|
||||
mgr.fire_state(ws.id, WorkstreamState.IDLE)
|
||||
assert mock_send.call_count == 0
|
||||
|
||||
def test_shutdown_is_idempotent(self, fake_mgr_and_ws):
|
||||
mgr, _ws = fake_mgr_and_ws
|
||||
watcher = IdleNudgeWatcher(mgr)
|
||||
watcher.start()
|
||||
watcher.shutdown()
|
||||
watcher.shutdown() # no error
|
||||
@@ -777,6 +777,58 @@ class TestModelAliasResolution:
|
||||
assert judge._client_factory_args["api_key"] == "alias-key"
|
||||
assert judge._client_factory_args["provider_name"] == "openai"
|
||||
|
||||
def test_unknown_alias_inherits_session_model(self):
|
||||
"""``judge.model`` is alias-only. A value that doesn't resolve
|
||||
through the registry inherits the session model (same path as
|
||||
an empty config.model) rather than getting pinned onto the
|
||||
session provider as a raw model id — that legacy behavior
|
||||
silently broke whenever the session provider didn't speak the
|
||||
configured model id (Anthropic session, ``judge.model =
|
||||
"gpt-5-mini"`` → every verdict came back as ``llm_fallback``)."""
|
||||
session_provider = _make_mock_provider()
|
||||
session_provider.provider_name = "anthropic"
|
||||
session_client = MagicMock()
|
||||
session_client.base_url = "https://session.example/v1"
|
||||
session_client.api_key = "session-key"
|
||||
|
||||
registry = MagicMock()
|
||||
registry.has_alias.return_value = False # judge.model isn't an alias
|
||||
|
||||
config = JudgeConfig(enabled=True, model="gpt-5-mini")
|
||||
judge = IntentJudge(
|
||||
config=config,
|
||||
session_provider=session_provider,
|
||||
session_client=session_client,
|
||||
session_model="session-default-model",
|
||||
context_window=100_000,
|
||||
model_registry=registry,
|
||||
)
|
||||
|
||||
assert judge._provider is session_provider
|
||||
assert judge._model == "session-default-model"
|
||||
# Context window mirrors the session, not the (uncalled) caps lookup.
|
||||
assert judge._judge_context_window == 100_000
|
||||
|
||||
def test_empty_model_inherits_session_model(self):
|
||||
"""Empty ``config.model`` is the documented self-consistency path."""
|
||||
session_provider = _make_mock_provider()
|
||||
session_provider.provider_name = "openai"
|
||||
session_client = MagicMock()
|
||||
session_client.base_url = "https://session.example/v1"
|
||||
session_client.api_key = "session-key"
|
||||
|
||||
config = JudgeConfig(enabled=True, model="")
|
||||
judge = IntentJudge(
|
||||
config=config,
|
||||
session_provider=session_provider,
|
||||
session_client=session_client,
|
||||
session_model="session-default-model",
|
||||
context_window=100_000,
|
||||
)
|
||||
|
||||
assert judge._provider is session_provider
|
||||
assert judge._model == "session-default-model"
|
||||
|
||||
def test_coordinator_tool_call_returns_llm_verdict_not_fallback(self):
|
||||
"""Happy-path regression for coordinator tool calls: with a properly
|
||||
resolved provider, the verdict tier must be ``llm`` — the
|
||||
|
||||
+122
-1
@@ -51,7 +51,11 @@ class TestIntentVerdictCRUD:
|
||||
assert v["tier"] == "heuristic"
|
||||
assert v["judge_model"] == ""
|
||||
assert v["latency_ms"] == 2
|
||||
assert v["user_decision"] == ""
|
||||
# ``user_decision`` defaults to ``"pending"`` (not the empty
|
||||
# string) so an audit reader can distinguish in-flight rows
|
||||
# from pre-convention legacy rows that carry the column's
|
||||
# server_default of ``""``.
|
||||
assert v["user_decision"] == "pending"
|
||||
assert "created" in v
|
||||
|
||||
def test_get_nonexistent(self, db):
|
||||
@@ -114,6 +118,123 @@ class TestIntentVerdictCRUD:
|
||||
assert ok is False
|
||||
|
||||
|
||||
class TestIntentVerdictUpsert:
|
||||
"""``upsert_intent_verdict`` — the LLM-tier-aware persistence path.
|
||||
|
||||
Backs the heuristic → llm_fallback "upgrade in place" pattern.
|
||||
The async judge's fallback verdicts deliberately reuse the
|
||||
heuristic ``verdict_id``; a plain INSERT would collide on the
|
||||
PK and the upgrade would be lost to a silently-swallowed
|
||||
exception (Postgres logged ``intent_verdicts_pkey`` violations
|
||||
for every fallback delivery on stable/1.5 smoke tests).
|
||||
"""
|
||||
|
||||
def test_upsert_on_fresh_id_inserts(self, db):
|
||||
"""No conflict — behaves like a regular INSERT."""
|
||||
db.upsert_intent_verdict(**_make_verdict_kwargs())
|
||||
v = db.get_intent_verdict("v_001")
|
||||
assert v is not None
|
||||
assert v["tier"] == "heuristic"
|
||||
assert v["user_decision"] == "pending"
|
||||
|
||||
def test_upsert_on_conflict_upgrades_tier_reasoning_judge_model(self, db):
|
||||
"""On PK conflict: tier, reasoning, judge_model update — every
|
||||
other field is preserved. Mirrors what the judge emits when
|
||||
promoting heuristic → llm_fallback."""
|
||||
db.upsert_intent_verdict(
|
||||
**_make_verdict_kwargs(
|
||||
tier="heuristic",
|
||||
reasoning="initial heuristic reasoning",
|
||||
judge_model="",
|
||||
)
|
||||
)
|
||||
db.upsert_intent_verdict(
|
||||
**_make_verdict_kwargs(
|
||||
tier="llm_fallback",
|
||||
reasoning="initial heuristic reasoning (LLM judge did not return a verdict)",
|
||||
judge_model="gpt-5-judge",
|
||||
)
|
||||
)
|
||||
v = db.get_intent_verdict("v_001")
|
||||
assert v is not None
|
||||
# The three fields that should change.
|
||||
assert v["tier"] == "llm_fallback"
|
||||
assert "LLM judge did not return" in v["reasoning"]
|
||||
assert v["judge_model"] == "gpt-5-judge"
|
||||
|
||||
def test_upsert_on_conflict_preserves_user_decision(self, db):
|
||||
"""LOAD-BEARING: a manually-resolved approval (user_decision=
|
||||
``"approved"``) or auto-approve-stamped row (user_decision=
|
||||
``"policy"``/``"blanket"``/etc.) must NOT be clobbered back to
|
||||
``"pending"`` when the late LLM-fallback verdict lands.
|
||||
``IntentVerdict.to_dict()`` doesn't project user_decision, so
|
||||
the upsert's defaulted ``"pending"`` would silently overwrite
|
||||
the real value if user_decision were in the on-conflict
|
||||
SET clause."""
|
||||
db.upsert_intent_verdict(**_make_verdict_kwargs())
|
||||
ok = db.update_intent_verdict("v_001", user_decision="approved")
|
||||
assert ok is True
|
||||
# Simulate the late LLM-fallback delivery — same verdict_id,
|
||||
# default user_decision (the IntentVerdict.to_dict() shape).
|
||||
db.upsert_intent_verdict(
|
||||
**_make_verdict_kwargs(
|
||||
tier="llm_fallback",
|
||||
reasoning="extended (LLM judge did not return a verdict)",
|
||||
judge_model="gpt-5-judge",
|
||||
)
|
||||
)
|
||||
v = db.get_intent_verdict("v_001")
|
||||
assert v is not None
|
||||
assert v["user_decision"] == "approved" # NOT clobbered to "pending"
|
||||
assert v["tier"] == "llm_fallback" # but the upgrade did land
|
||||
|
||||
def test_upsert_on_conflict_preserves_identity_and_carried_fields(self, db):
|
||||
"""Identity columns (ws_id, call_id, func_name, func_args) and
|
||||
carried-verbatim columns (intent_summary, risk_level,
|
||||
confidence, recommendation, evidence, latency_ms) are
|
||||
excluded from the on-conflict SET — verify they aren't
|
||||
changed even when the second upsert passes different values
|
||||
(defensive against a future judge bug that ships divergent
|
||||
carried fields)."""
|
||||
db.upsert_intent_verdict(**_make_verdict_kwargs())
|
||||
db.upsert_intent_verdict(
|
||||
**_make_verdict_kwargs(
|
||||
# Same verdict_id (conflict trigger), divergent everything else.
|
||||
ws_id="ws-different",
|
||||
call_id="tc_different",
|
||||
func_name="bash_v2",
|
||||
func_args='{"command":"rm -rf /"}',
|
||||
intent_summary="totally different summary",
|
||||
risk_level="critical",
|
||||
confidence=0.0,
|
||||
recommendation="deny",
|
||||
evidence='["dangerous"]',
|
||||
latency_ms=99999,
|
||||
# The three fields that DO update.
|
||||
tier="llm_fallback",
|
||||
reasoning="upgraded reasoning",
|
||||
judge_model="judge-v2",
|
||||
)
|
||||
)
|
||||
v = db.get_intent_verdict("v_001")
|
||||
assert v is not None
|
||||
# All preserved from the first upsert (identity + carried).
|
||||
assert v["ws_id"] == "ws-abc"
|
||||
assert v["call_id"] == "tc_001"
|
||||
assert v["func_name"] == "bash"
|
||||
assert v["func_args"] == '{"command":"echo hello"}'
|
||||
assert v["intent_summary"] == "Echo a greeting to stdout"
|
||||
assert v["risk_level"] == "low"
|
||||
assert v["confidence"] == 0.85
|
||||
assert v["recommendation"] == "approve"
|
||||
assert v["evidence"] == '["The command only prints text."]'
|
||||
assert v["latency_ms"] == 2
|
||||
# Only the three updated.
|
||||
assert v["tier"] == "llm_fallback"
|
||||
assert v["reasoning"] == "upgraded reasoning"
|
||||
assert v["judge_model"] == "judge-v2"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bulk insert
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -420,8 +420,9 @@ class TestSkillCatalogDisclosure:
|
||||
session.system_messages = []
|
||||
session._agent_system_messages = []
|
||||
session.reasoning_effort = "medium"
|
||||
session._pending_tool_advisories = []
|
||||
session._pending_user_advisories = []
|
||||
from turnstone.core.nudge_queue import NudgeQueue
|
||||
|
||||
session._nudge_queue = NudgeQueue()
|
||||
session._tool_search = None
|
||||
session._mcp_client = None
|
||||
session._notify_on_complete = "{}"
|
||||
|
||||
+1049
-2
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,211 @@
|
||||
"""Integration tests for the Phase 9 admin bulk-revoke endpoint.
|
||||
|
||||
POST /v1/api/admin/mcp-servers/{name}/bulk-revoke clears every user's
|
||||
OAuth token for a server (admin-side counterpart to the per-user
|
||||
DELETE /v1/api/mcp/oauth/connections/{server_name} that shipped in
|
||||
Phase 8).
|
||||
|
||||
Coverage:
|
||||
- requires ``admin.mcp`` permission (401/403 without).
|
||||
- 404 when the named server is missing.
|
||||
- 400 when the server's ``auth_type`` is not ``oauth_user``.
|
||||
- 200 + ``rows_deleted`` + ``consented_users_before`` on success.
|
||||
- Audit row written with
|
||||
``upstream_revoke_outcome="bulk_admin_no_upstream"``.
|
||||
- Token rows are gone from ``mcp_user_tokens`` post-call.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.routing import Mount, Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from turnstone.console.server import admin_mcp_bulk_revoke
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
|
||||
class _InjectAdminMcp(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id="admin-user",
|
||||
scopes=frozenset({"approve"}),
|
||||
token_source="config",
|
||||
permissions=frozenset({"read", "write", "approve", "admin.mcp"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
class _InjectNoAdminMcp(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id="regular-user",
|
||||
scopes=frozenset({"approve"}),
|
||||
token_source="jwt",
|
||||
permissions=frozenset({"read", "write", "approve"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _build_app(storage: SQLiteBackend, *, with_admin_mcp: bool = True) -> Starlette:
|
||||
mw = _InjectAdminMcp if with_admin_mcp else _InjectNoAdminMcp
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route(
|
||||
"/api/admin/mcp-servers/{name}/bulk-revoke",
|
||||
admin_mcp_bulk_revoke,
|
||||
methods=["POST"],
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
middleware=[Middleware(mw)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
return app
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _seed_oauth_server(
|
||||
backend: SQLiteBackend,
|
||||
*,
|
||||
name: str = "srv-oauth",
|
||||
server_id: str = "srv-oauth-id",
|
||||
) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url="https://example.com/mcp",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
|
||||
|
||||
def _seed_static_server(
|
||||
backend: SQLiteBackend,
|
||||
*,
|
||||
name: str = "srv-static",
|
||||
server_id: str = "srv-static-id",
|
||||
) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url="https://example.com/mcp",
|
||||
auth_type="static",
|
||||
)
|
||||
|
||||
|
||||
def _seed_user_tokens(backend: SQLiteBackend, server_name: str, users: int) -> None:
|
||||
for i in range(users):
|
||||
backend.create_mcp_user_token(
|
||||
f"user-{i}",
|
||||
server_name,
|
||||
access_token_ct=b"ct",
|
||||
refresh_token_ct=None,
|
||||
expires_at=None,
|
||||
scopes=None,
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://example.com/mcp",
|
||||
)
|
||||
|
||||
|
||||
def test_requires_admin_mcp_permission(storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
client = TestClient(_build_app(storage, with_admin_mcp=False))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/srv-oauth/bulk-revoke")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
def test_404_on_missing_server(storage: SQLiteBackend) -> None:
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/never-existed/bulk-revoke")
|
||||
assert resp.status_code == 404
|
||||
assert resp.json() == {"error": "No such server"}
|
||||
|
||||
|
||||
def test_400_on_static_server(storage: SQLiteBackend) -> None:
|
||||
_seed_static_server(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/srv-static/bulk-revoke")
|
||||
assert resp.status_code == 400
|
||||
body = resp.json()
|
||||
assert "oauth_user" in body["error"]
|
||||
|
||||
|
||||
def test_400_on_invalid_server_name(storage: SQLiteBackend) -> None:
|
||||
# double-underscore is reserved for the prefixed-tool-name encoding.
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/bad__name/bulk-revoke")
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
def test_200_on_success_with_no_consented_users(storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/srv-oauth/bulk-revoke")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["rows_deleted"] == 0
|
||||
assert body["consented_users_before"] == 0
|
||||
|
||||
|
||||
def test_200_clears_all_user_tokens(storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_user_tokens(storage, "srv-oauth", users=3)
|
||||
# Token for another server must survive the bulk-revoke.
|
||||
_seed_oauth_server(storage, name="srv-other", server_id="srv-other-id")
|
||||
_seed_user_tokens(storage, "srv-other", users=2)
|
||||
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/srv-oauth/bulk-revoke")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["rows_deleted"] == 3
|
||||
assert body["consented_users_before"] == 3
|
||||
|
||||
# Target server's tokens are gone; bystander's tokens survive.
|
||||
assert storage.count_mcp_consented_users_by_server("srv-oauth") == 0
|
||||
assert storage.count_mcp_consented_users_by_server("srv-other") == 2
|
||||
|
||||
|
||||
def test_audits_with_bulk_admin_no_upstream(storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_user_tokens(storage, "srv-oauth", users=2)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.post("/v1/api/admin/mcp-servers/srv-oauth/bulk-revoke")
|
||||
assert resp.status_code == 200
|
||||
|
||||
# Pull the most-recent audit row for the bulk_revoked action and
|
||||
# verify it carries the deferral marker.
|
||||
events = storage.list_audit_events(limit=10)
|
||||
bulk_rows = [e for e in events if e.get("action") == "mcp_server.oauth.bulk_revoked"]
|
||||
assert len(bulk_rows) == 1
|
||||
detail = bulk_rows[0].get("detail")
|
||||
if isinstance(detail, str):
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(detail)
|
||||
assert detail.get("upstream_revoke_outcome") == "bulk_admin_no_upstream"
|
||||
assert detail.get("rows_deleted") == 2
|
||||
assert detail.get("consented_users_before") == 2
|
||||
assert detail.get("name") == "srv-oauth"
|
||||
+931
-148
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,117 @@
|
||||
"""Structural gate against the Phase 7b sibling-bug pattern.
|
||||
|
||||
Phase 7b's bug-1 was a single ``f"MCP X error: {e}"`` site dropping a
|
||||
structured-error JSON. Phase 8 introduces the ``consent_url`` field on
|
||||
the same JSON envelope: every ``_structured_error(...)`` invocation
|
||||
that emits ``mcp_consent_required`` or ``mcp_insufficient_scope`` MUST
|
||||
also pass a ``consent_url=`` kwarg, otherwise the dashboard renderer
|
||||
can't surface a re-consent button.
|
||||
|
||||
This test is purely structural — it scans the source of
|
||||
:mod:`turnstone.core.mcp_client` and asserts every consent-required /
|
||||
insufficient-scope ``_structured_error`` call carries
|
||||
``consent_url=``. It catches future regressions where a new exec path
|
||||
adds a fourth call site and forgets the kwarg.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import turnstone.core.mcp_client as _mcp_client_module
|
||||
|
||||
_USER_ACTIONABLE_CODES = ("mcp_consent_required", "mcp_insufficient_scope")
|
||||
|
||||
|
||||
def _read_source() -> str:
|
||||
path = Path(_mcp_client_module.__file__)
|
||||
return path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def _find_structured_error_blocks(source: str) -> list[tuple[int, str]]:
|
||||
"""Return ``(line_no, block)`` pairs for every ``_structured_error(...)``.
|
||||
|
||||
Each block is the call's argument list expanded across however many
|
||||
lines the formatter chose. Uses a paren-counting walk so multi-line
|
||||
kwargs and nested expressions are captured correctly.
|
||||
"""
|
||||
blocks: list[tuple[int, str]] = []
|
||||
needle = "_structured_error("
|
||||
idx = 0
|
||||
while True:
|
||||
loc = source.find(needle, idx)
|
||||
if loc < 0:
|
||||
break
|
||||
# Skip the function definition itself.
|
||||
if source[loc - 4 : loc] == "def ":
|
||||
idx = loc + len(needle)
|
||||
continue
|
||||
line_no = source.count("\n", 0, loc) + 1
|
||||
depth = 1
|
||||
end = loc + len(needle)
|
||||
while end < len(source) and depth > 0:
|
||||
ch = source[end]
|
||||
if ch == "(":
|
||||
depth += 1
|
||||
elif ch == ")":
|
||||
depth -= 1
|
||||
end += 1
|
||||
blocks.append((line_no, source[loc:end]))
|
||||
idx = end
|
||||
return blocks
|
||||
|
||||
|
||||
def test_every_user_actionable_structured_error_passes_consent_url() -> None:
|
||||
source = _read_source()
|
||||
blocks = _find_structured_error_blocks(source)
|
||||
user_actionable_blocks = [
|
||||
(ln, blk)
|
||||
for ln, blk in blocks
|
||||
if any(f'code="{code}"' in blk for code in _USER_ACTIONABLE_CODES)
|
||||
]
|
||||
|
||||
# Sanity check: ensure we actually scanned the file the audit cares
|
||||
# about (a stale path or import would otherwise silently pass with
|
||||
# zero matches).
|
||||
assert user_actionable_blocks, (
|
||||
"No mcp_consent_required / mcp_insufficient_scope _structured_error "
|
||||
"call sites found — has the audit been pointed at the wrong file?"
|
||||
)
|
||||
|
||||
missing: list[tuple[int, str]] = []
|
||||
for ln, blk in user_actionable_blocks:
|
||||
if "consent_url=" not in blk:
|
||||
# Strip whitespace and truncate so the failure message is
|
||||
# readable in CI.
|
||||
collapsed = re.sub(r"\s+", " ", blk).strip()
|
||||
missing.append((ln, collapsed[:200]))
|
||||
|
||||
assert not missing, (
|
||||
"Sibling-bug regression: the following consent-required / "
|
||||
"insufficient-scope _structured_error sites are missing the "
|
||||
"consent_url= kwarg.\n" + "\n".join(f" line {ln}: {snippet}" for ln, snippet in missing)
|
||||
)
|
||||
|
||||
|
||||
def test_audit_finds_all_known_user_actionable_sites() -> None:
|
||||
"""Lock the count so accidental deletions are caught.
|
||||
|
||||
There are 13 user-actionable ``_structured_error`` call sites today
|
||||
(4 each in the tool / resource / prompt token-classify branches +
|
||||
3 in the post-retry-failed branches + 1 in ``_handle_auth_403``'s
|
||||
insufficient-scope branch). If a new exec path is added the count
|
||||
can rise; if a branch is removed the count can fall — both are
|
||||
fine, but require an intentional bump of this number to confirm
|
||||
the change went through review.
|
||||
"""
|
||||
source = _read_source()
|
||||
blocks = _find_structured_error_blocks(source)
|
||||
user_actionable_count = sum(
|
||||
1 for _, blk in blocks if any(f'code="{code}"' in blk for code in _USER_ACTIONABLE_CODES)
|
||||
)
|
||||
assert user_actionable_count == 13, (
|
||||
f"Expected 13 user-actionable _structured_error sites, got "
|
||||
f"{user_actionable_count}. If this is intentional, bump the "
|
||||
f"expected count and document why in the commit message."
|
||||
)
|
||||
@@ -0,0 +1,235 @@
|
||||
"""Tests for ``turnstone.core.mcp_crypto`` cipher + config loading.
|
||||
|
||||
Covers token-at-rest encryption for OAuth-MCP.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
|
||||
import pytest
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from turnstone.core.mcp_crypto import (
|
||||
MCPTokenCipher,
|
||||
MCPTokenCipherConfig,
|
||||
MCPTokenDecryptError,
|
||||
MCPTokenKeyConfigError,
|
||||
_key_fingerprint,
|
||||
_validate_key,
|
||||
load_mcp_token_cipher_config,
|
||||
)
|
||||
|
||||
|
||||
def _new_raw_key() -> bytes:
|
||||
"""Return a fresh 32-byte Fernet key as raw bytes (post-base64-decode)."""
|
||||
return base64.urlsafe_b64decode(Fernet.generate_key())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cipher round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCipherRoundTrip:
|
||||
def test_round_trip_single_key(self) -> None:
|
||||
cipher = MCPTokenCipher(MCPTokenCipherConfig(keys=(_new_raw_key(),)))
|
||||
plaintext = b"access_token_12345"
|
||||
ct = cipher.encrypt(plaintext)
|
||||
assert ct != plaintext
|
||||
assert cipher.decrypt(ct) == plaintext
|
||||
|
||||
def test_round_trip_unicode_token(self) -> None:
|
||||
cipher = MCPTokenCipher(MCPTokenCipherConfig(keys=(_new_raw_key(),)))
|
||||
# Tokens may legitimately carry UTF-8 bytes (e.g. JWT with
|
||||
# non-ASCII claim values). Round-trip a multi-byte sequence.
|
||||
plaintext = "tok_é中💯".encode()
|
||||
ct = cipher.encrypt(plaintext)
|
||||
assert cipher.decrypt(ct) == plaintext
|
||||
|
||||
def test_wrong_key_raises_decrypt_error(self) -> None:
|
||||
cipher_a = MCPTokenCipher(MCPTokenCipherConfig(keys=(_new_raw_key(),)))
|
||||
cipher_b = MCPTokenCipher(MCPTokenCipherConfig(keys=(_new_raw_key(),)))
|
||||
ct = cipher_a.encrypt(b"secret")
|
||||
with pytest.raises(MCPTokenDecryptError) as exc_info:
|
||||
cipher_b.decrypt(ct)
|
||||
# Audit-trail correlation: error must carry the fingerprints of
|
||||
# the keys actually attempted, not a placeholder.
|
||||
assert exc_info.value.key_fingerprints_attempted
|
||||
assert exc_info.value.key_fingerprints_attempted == cipher_b.key_fingerprints
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Rotation (MultiFernet behavior)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRotation:
|
||||
def test_rotation_forward(self) -> None:
|
||||
"""Encrypt with a new-only cipher, decrypt with a [v2, v1] cluster.
|
||||
|
||||
Mirrors the operational situation where a node already has the
|
||||
rotated key list installed and a peer just wrote a row under v2.
|
||||
"""
|
||||
v1 = _new_raw_key()
|
||||
v2 = _new_raw_key()
|
||||
new_only = MCPTokenCipher(MCPTokenCipherConfig(keys=(v2,)))
|
||||
cluster = MCPTokenCipher(MCPTokenCipherConfig(keys=(v2, v1)))
|
||||
ct = new_only.encrypt(b"hello")
|
||||
assert cluster.decrypt(ct) == b"hello"
|
||||
|
||||
def test_rotation_backward_keeps_old_decryptable(self) -> None:
|
||||
"""A row written under the OLD key (v1) must still decrypt after
|
||||
rotation places v2 first and keeps v1 as fallback."""
|
||||
v1 = _new_raw_key()
|
||||
v2 = _new_raw_key()
|
||||
old_only = MCPTokenCipher(MCPTokenCipherConfig(keys=(v1,)))
|
||||
rotated = MCPTokenCipher(MCPTokenCipherConfig(keys=(v2, v1)))
|
||||
ct = old_only.encrypt(b"legacy")
|
||||
assert rotated.decrypt(ct) == b"legacy"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config loader
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _patch_load_config(monkeypatch: pytest.MonkeyPatch, payload: dict) -> None:
|
||||
"""Override ``turnstone.core.config.load_config`` to return ``payload``
|
||||
when the ``"security"`` section is requested."""
|
||||
|
||||
def fake(section: str | None = None) -> dict:
|
||||
if section == "security":
|
||||
return payload
|
||||
return {}
|
||||
|
||||
import turnstone.core.config as cfg_mod
|
||||
|
||||
monkeypatch.setattr(cfg_mod, "load_config", fake)
|
||||
|
||||
|
||||
class TestLoadConfig:
|
||||
def test_load_singular_key(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
key = Fernet.generate_key().decode()
|
||||
_patch_load_config(monkeypatch, {"mcp_token_encryption_key": key})
|
||||
cfg = load_mcp_token_cipher_config()
|
||||
assert cfg is not None
|
||||
assert len(cfg.keys) == 1
|
||||
|
||||
def test_load_plural_overrides_singular(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
plural = [Fernet.generate_key().decode(), Fernet.generate_key().decode()]
|
||||
_patch_load_config(
|
||||
monkeypatch,
|
||||
{
|
||||
"mcp_token_encryption_keys": plural,
|
||||
"mcp_token_encryption_key": Fernet.generate_key().decode(),
|
||||
},
|
||||
)
|
||||
cfg = load_mcp_token_cipher_config()
|
||||
assert cfg is not None
|
||||
assert len(cfg.keys) == 2 # plural wins, singular ignored
|
||||
|
||||
def test_load_returns_none_when_absent(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_patch_load_config(monkeypatch, {})
|
||||
assert load_mcp_token_cipher_config() is None
|
||||
|
||||
def test_load_empty_plural_falls_through_to_singular(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Operator wrote ``mcp_token_encryption_keys = []`` AND set a
|
||||
singular value: empty plural is treated as absent."""
|
||||
key = Fernet.generate_key().decode()
|
||||
_patch_load_config(
|
||||
monkeypatch,
|
||||
{"mcp_token_encryption_keys": [], "mcp_token_encryption_key": key},
|
||||
)
|
||||
cfg = load_mcp_token_cipher_config()
|
||||
assert cfg is not None
|
||||
assert len(cfg.keys) == 1
|
||||
|
||||
def test_load_invalid_base64_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_patch_load_config(monkeypatch, {"mcp_token_encryption_key": "###not-base64###"})
|
||||
with pytest.raises(MCPTokenKeyConfigError) as exc_info:
|
||||
load_mcp_token_cipher_config()
|
||||
# Operator-facing hint is part of every error message.
|
||||
assert "regenerate with:" in str(exc_info.value)
|
||||
|
||||
def test_load_wrong_length_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
# 24 raw bytes → 32 base64 chars; not 32 raw bytes after decode.
|
||||
short_key = base64.urlsafe_b64encode(b"\x00" * 24).decode()
|
||||
_patch_load_config(monkeypatch, {"mcp_token_encryption_key": short_key})
|
||||
with pytest.raises(MCPTokenKeyConfigError) as exc_info:
|
||||
load_mcp_token_cipher_config()
|
||||
assert "32 bytes" in str(exc_info.value)
|
||||
|
||||
def test_load_non_list_plural_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_patch_load_config(monkeypatch, {"mcp_token_encryption_keys": "single-string-not-list"})
|
||||
with pytest.raises(MCPTokenKeyConfigError) as exc_info:
|
||||
load_mcp_token_cipher_config()
|
||||
assert "list" in str(exc_info.value).lower()
|
||||
|
||||
def test_load_non_string_in_plural_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_patch_load_config(monkeypatch, {"mcp_token_encryption_keys": [12345]})
|
||||
with pytest.raises(MCPTokenKeyConfigError):
|
||||
load_mcp_token_cipher_config()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fingerprint stability
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFingerprint:
|
||||
def test_key_fingerprint_stable_and_short(self) -> None:
|
||||
key = _new_raw_key()
|
||||
fp1 = _key_fingerprint(key)
|
||||
fp2 = _key_fingerprint(key)
|
||||
assert fp1 == fp2
|
||||
# 8 bytes -> 16 hex characters.
|
||||
assert len(fp1) == 16
|
||||
assert all(c in "0123456789abcdef" for c in fp1)
|
||||
|
||||
def test_different_keys_have_different_fingerprints(self) -> None:
|
||||
fp1 = _key_fingerprint(_new_raw_key())
|
||||
fp2 = _key_fingerprint(_new_raw_key())
|
||||
assert fp1 != fp2
|
||||
|
||||
def test_cipher_fingerprints_match_keys(self) -> None:
|
||||
v1 = _new_raw_key()
|
||||
v2 = _new_raw_key()
|
||||
cipher = MCPTokenCipher(MCPTokenCipherConfig(keys=(v1, v2)))
|
||||
assert cipher.key_fingerprints == (
|
||||
_key_fingerprint(v1),
|
||||
_key_fingerprint(v2),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct ``_validate_key`` — exercises edge cases not reachable via loader
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestValidateKey:
|
||||
def test_empty_string_rejected(self) -> None:
|
||||
with pytest.raises(MCPTokenKeyConfigError):
|
||||
_validate_key("", label="x")
|
||||
|
||||
def test_whitespace_only_rejected(self) -> None:
|
||||
with pytest.raises(MCPTokenKeyConfigError):
|
||||
_validate_key(" ", label="x")
|
||||
|
||||
def test_label_propagated_in_error(self) -> None:
|
||||
with pytest.raises(MCPTokenKeyConfigError) as exc_info:
|
||||
_validate_key("###", label="my_label_42")
|
||||
assert "my_label_42" in str(exc_info.value)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCPTokenCipher constructor guard
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCipherConstructorGuard:
|
||||
def test_empty_keys_rejected(self) -> None:
|
||||
with pytest.raises(MCPTokenKeyConfigError):
|
||||
MCPTokenCipher(MCPTokenCipherConfig(keys=()))
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from tests.conftest import _seed_static_state
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -105,14 +106,18 @@ class TestRemoveServerSync:
|
||||
"""remove_server_sync cleans up all per-server state dicts."""
|
||||
mgr = MCPClientManager({"test": {"command": "echo"}})
|
||||
# Simulate state as if the server was connected
|
||||
mgr._per_server_tools["test"] = [_fake_openai_tool()]
|
||||
mgr._per_server_resources["test"] = [_fake_resource_dict()]
|
||||
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
|
||||
mgr._supports_list_changed["test"] = True
|
||||
mgr._supports_resources["test"] = True
|
||||
mgr._supports_resource_list_changed["test"] = True
|
||||
mgr._supports_prompts["test"] = True
|
||||
mgr._supports_prompt_list_changed["test"] = True
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"test",
|
||||
tools=[_fake_openai_tool()],
|
||||
resources=[_fake_resource_dict()],
|
||||
prompts=[_fake_prompt_dict()],
|
||||
supports_list_changed=True,
|
||||
supports_resources=True,
|
||||
supports_resource_list_changed=True,
|
||||
supports_prompts=True,
|
||||
supports_prompt_list_changed=True,
|
||||
)
|
||||
mgr._rebuild_tools()
|
||||
mgr._rebuild_resources()
|
||||
mgr._rebuild_prompts()
|
||||
@@ -127,14 +132,7 @@ class TestRemoveServerSync:
|
||||
assert len(mgr.get_tools()) == 0
|
||||
assert mgr.resource_count == 0
|
||||
assert mgr.prompt_count == 0
|
||||
assert "test" not in mgr._per_server_tools
|
||||
assert "test" not in mgr._per_server_resources
|
||||
assert "test" not in mgr._per_server_prompts
|
||||
assert "test" not in mgr._supports_list_changed
|
||||
assert "test" not in mgr._supports_resources
|
||||
assert "test" not in mgr._supports_resource_list_changed
|
||||
assert "test" not in mgr._supports_prompts
|
||||
assert "test" not in mgr._supports_prompt_list_changed
|
||||
assert "test" not in mgr._static_servers
|
||||
|
||||
def test_removes_config_to_prevent_reconnect(self) -> None:
|
||||
"""remove_server_sync removes from _server_configs to prevent reconnect."""
|
||||
@@ -146,8 +144,8 @@ class TestRemoveServerSync:
|
||||
def test_preserves_other_servers(self) -> None:
|
||||
"""Removing one server does not affect another server's state."""
|
||||
mgr = MCPClientManager({"srv_a": {}, "srv_b": {}})
|
||||
mgr._per_server_tools["srv_a"] = [_fake_openai_tool("mcp__srv_a__foo")]
|
||||
mgr._per_server_tools["srv_b"] = [_fake_openai_tool("mcp__srv_b__bar")]
|
||||
_seed_static_state(mgr, "srv_a", tools=[_fake_openai_tool("mcp__srv_a__foo")])
|
||||
_seed_static_state(mgr, "srv_b", tools=[_fake_openai_tool("mcp__srv_b__bar")])
|
||||
mgr._rebuild_tools()
|
||||
|
||||
assert len(mgr.get_tools()) == 2
|
||||
@@ -179,13 +177,17 @@ class TestGetServerStatus:
|
||||
"""Status of a connected server reports correct tool/resource/prompt counts."""
|
||||
mgr = MCPClientManager({"test": {}})
|
||||
# Simulate connected state
|
||||
mgr._sessions["test"] = object() # any truthy value
|
||||
mgr._per_server_tools["test"] = [
|
||||
_fake_openai_tool("mcp__test__a"),
|
||||
_fake_openai_tool("mcp__test__b"),
|
||||
]
|
||||
mgr._per_server_resources["test"] = [_fake_resource_dict()]
|
||||
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"test",
|
||||
session=object(), # any truthy value
|
||||
tools=[
|
||||
_fake_openai_tool("mcp__test__a"),
|
||||
_fake_openai_tool("mcp__test__b"),
|
||||
],
|
||||
resources=[_fake_resource_dict()],
|
||||
prompts=[_fake_prompt_dict()],
|
||||
)
|
||||
|
||||
status = mgr.get_server_status("test")
|
||||
assert status["connected"] is True
|
||||
@@ -225,8 +227,7 @@ class TestGetAllServerStatus:
|
||||
def test_mixed_connected_and_disconnected(self) -> None:
|
||||
"""Status correctly reflects a mix of connected and disconnected servers."""
|
||||
mgr = MCPClientManager({"up": {}, "down": {}})
|
||||
mgr._sessions["up"] = object()
|
||||
mgr._per_server_tools["up"] = [_fake_openai_tool("mcp__up__x")]
|
||||
_seed_static_state(mgr, "up", session=object(), tools=[_fake_openai_tool("mcp__up__x")])
|
||||
|
||||
statuses = mgr.get_all_server_status()
|
||||
assert statuses["up"]["connected"] is True
|
||||
|
||||
@@ -0,0 +1,242 @@
|
||||
"""Unit tests for ``turnstone.core.mcp_http_parsers``.
|
||||
|
||||
The parser replaces the prior hand-rolled scanners that used
|
||||
``header.lower().find("scope")`` to locate parameter names — that approach
|
||||
misparsed ``scope`` embedded inside other tokens (``xscope``) or inside
|
||||
quoted-string values of preceding params. Each adversarial case below
|
||||
asserts the new tokenizer respects RFC 7235 ``challenge → auth-param``
|
||||
boundaries; the docstrings document the equivalent input that broke the
|
||||
naive parser. Negative-test verification: temporarily reverting
|
||||
``parse_www_authenticate_scope`` to delegate to ``header.lower().find("scope")``
|
||||
makes ``test_scope_inside_realm_value`` and ``test_scope_inside_xscope`` fail.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.mcp_http_parsers import (
|
||||
parse_www_authenticate_bearer,
|
||||
parse_www_authenticate_error,
|
||||
parse_www_authenticate_scope,
|
||||
)
|
||||
|
||||
|
||||
class TestParseScope:
|
||||
def test_basic_scope(self) -> None:
|
||||
header = 'Bearer error="insufficient_scope", scope="files:read mail:send"'
|
||||
assert parse_www_authenticate_scope(header) == ("files:read", "mail:send")
|
||||
|
||||
def test_no_scope_param(self) -> None:
|
||||
assert parse_www_authenticate_scope('Bearer error="invalid_token"') == ()
|
||||
|
||||
def test_unterminated_quoted_string_returns_empty(self) -> None:
|
||||
assert parse_www_authenticate_scope('Bearer scope="files:read') == ()
|
||||
|
||||
def test_escaped_chars_in_value_drops_invalid_scope_token(self) -> None:
|
||||
# RFC 7230 §3.2.6 backslash escapes decode the literal scope to
|
||||
# ``files:read "weird"``. RFC 6749 §3.3 ``scope-token`` forbids
|
||||
# ``"``, so ``"weird"`` is dropped and only ``files:read``
|
||||
# survives the post-split validation.
|
||||
header = r'Bearer scope="files:read \"weird\""'
|
||||
assert parse_www_authenticate_scope(header) == ("files:read",)
|
||||
|
||||
def test_empty_string(self) -> None:
|
||||
assert parse_www_authenticate_scope("") == ()
|
||||
|
||||
def test_unquoted_scope_value(self) -> None:
|
||||
# Unquoted single token.
|
||||
assert parse_www_authenticate_scope("Bearer scope=files:read") == ("files:read",)
|
||||
|
||||
# --- the four headline misparse cases ---
|
||||
|
||||
def test_scope_inside_xscope(self) -> None:
|
||||
"""``Bearer xscope="value"`` must NOT be read as ``scope``.
|
||||
|
||||
The naive ``find("scope")`` matched at position 7 inside
|
||||
``xscope`` and returned ``("value",)``.
|
||||
"""
|
||||
assert parse_www_authenticate_scope('Bearer xscope="value"') == ()
|
||||
|
||||
def test_scope_inside_realm_value(self) -> None:
|
||||
"""``Bearer realm="my scope=fake", scope="real"`` must return ``("real",)``.
|
||||
|
||||
The naive parser found ``scope=`` inside the quoted ``realm``
|
||||
value first and returned ``("fake",)``.
|
||||
"""
|
||||
header = 'Bearer realm="my scope=fake", scope="real"'
|
||||
assert parse_www_authenticate_scope(header) == ("real",)
|
||||
|
||||
def test_scope_inside_quoted_realm_with_escaped_quotes(self) -> None:
|
||||
"""``Bearer realm="foo scope=\\"admin:write\\" bar"`` returns ``()``.
|
||||
|
||||
The inner ``scope=`` is wholly inside the quoted-string value of
|
||||
``realm`` — there is no top-level ``scope`` auth-param, so the
|
||||
result is empty.
|
||||
"""
|
||||
header = r'Bearer realm="foo scope=\"admin:write\" bar"'
|
||||
assert parse_www_authenticate_scope(header) == ()
|
||||
|
||||
def test_scope_token_validation_drops_control_bytes(self) -> None:
|
||||
"""Tokens containing CR / LF / tab / DEL / quote are dropped.
|
||||
|
||||
RFC 6749 §3.3 restricts ``scope-token`` to visible ASCII
|
||||
excluding ``"`` and ``\\``. The splitter applies that
|
||||
validation so a malicious AS cannot smuggle CRLF (or the like)
|
||||
through a future log / notification path that prints the scope
|
||||
list verbatim. ``"a\\rb"`` and ``"\\nc"`` fail validation;
|
||||
``"d"`` survives. The legitimate space separator splits ``d``
|
||||
into its own token.
|
||||
"""
|
||||
# Build via concatenation so the assertion stays intelligible.
|
||||
header = 'Bearer scope="a\rb \nc d"'
|
||||
assert parse_www_authenticate_scope(header) == ("d",)
|
||||
|
||||
|
||||
class TestParseError:
|
||||
def test_basic_quoted_error(self) -> None:
|
||||
assert (
|
||||
parse_www_authenticate_error('Bearer error="insufficient_scope"')
|
||||
== "insufficient_scope"
|
||||
)
|
||||
|
||||
def test_other_quoted_error_tokens(self) -> None:
|
||||
assert parse_www_authenticate_error('Bearer error="invalid_token"') == "invalid_token"
|
||||
assert parse_www_authenticate_error('Bearer error="invalid_request"') == "invalid_request"
|
||||
|
||||
def test_no_error_param(self) -> None:
|
||||
assert parse_www_authenticate_error("Bearer realm=foo") is None
|
||||
|
||||
def test_error_description_does_not_match_error(self) -> None:
|
||||
"""``error_description`` is its own auth-param key, not ``error``.
|
||||
|
||||
The tokenizer reads ``_`` as part of the token (RFC 7230 ``tchar``),
|
||||
so ``error_description`` becomes one key, ``error`` another.
|
||||
"""
|
||||
assert parse_www_authenticate_error('Bearer error_description="bad"') is None
|
||||
|
||||
def test_unquoted_error(self) -> None:
|
||||
# Some ASes don't quote the error token.
|
||||
assert (
|
||||
parse_www_authenticate_error("Bearer error=insufficient_scope") == "insufficient_scope"
|
||||
)
|
||||
|
||||
def test_empty_string(self) -> None:
|
||||
assert parse_www_authenticate_error("") is None
|
||||
|
||||
def test_error_inside_realm_value(self) -> None:
|
||||
"""``Bearer realm="my error=fake", error="real"`` must return ``"real"``.
|
||||
|
||||
Naive parser grabbed ``fake`` from inside the ``realm`` quoted
|
||||
value.
|
||||
"""
|
||||
header = 'Bearer realm="my error=fake", error="real"'
|
||||
assert parse_www_authenticate_error(header) == "real"
|
||||
|
||||
|
||||
class TestBearerDict:
|
||||
def test_returns_lowercased_keys(self) -> None:
|
||||
header = 'Bearer Realm="x", Error="y", Scope="a b"'
|
||||
params = parse_www_authenticate_bearer(header)
|
||||
assert params == {"realm": "x", "error": "y", "scope": "a b"}
|
||||
|
||||
def test_non_bearer_scheme_returns_empty(self) -> None:
|
||||
assert parse_www_authenticate_bearer('Basic realm="x"') == {}
|
||||
|
||||
def test_no_scheme(self) -> None:
|
||||
assert parse_www_authenticate_bearer('realm="x"') == {}
|
||||
|
||||
def test_bearer_only_no_params(self) -> None:
|
||||
assert parse_www_authenticate_bearer("Bearer ") == {}
|
||||
|
||||
def test_bearer_with_no_space_returns_empty(self) -> None:
|
||||
# ``BearerToken`` is not a Bearer challenge (no separator).
|
||||
assert parse_www_authenticate_bearer("BearerToken") == {}
|
||||
|
||||
def test_first_value_wins_on_duplicate(self) -> None:
|
||||
# If a malformed AS sends two ``scope=`` params we keep the first.
|
||||
# The earlier ``find()``-based scanner would have returned the
|
||||
# last; either choice is legal for malformed input but we need
|
||||
# to be consistent.
|
||||
header = 'Bearer scope="first", scope="second"'
|
||||
assert parse_www_authenticate_bearer(header) == {"scope": "first"}
|
||||
|
||||
def test_trailing_comma(self) -> None:
|
||||
header = 'Bearer error="x",'
|
||||
assert parse_www_authenticate_bearer(header) == {"error": "x"}
|
||||
|
||||
def test_multiple_commas(self) -> None:
|
||||
header = 'Bearer ,, error="x",,, scope="y"'
|
||||
assert parse_www_authenticate_bearer(header) == {"error": "x", "scope": "y"}
|
||||
|
||||
def test_embedded_escaped_quote(self) -> None:
|
||||
header = r'Bearer realm="he said \"hi\""'
|
||||
assert parse_www_authenticate_bearer(header) == {"realm": 'he said "hi"'}
|
||||
|
||||
def test_param_without_value_skipped(self) -> None:
|
||||
header = 'Bearer realm, error="x"'
|
||||
# ``realm`` without ``=`` is dropped; ``error`` survives.
|
||||
assert parse_www_authenticate_bearer(header) == {"error": "x"}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"header,expected",
|
||||
[
|
||||
("", {}),
|
||||
("Bearer", {}),
|
||||
('Bearer realm=""', {"realm": ""}),
|
||||
('Bearer realm="", scope=""', {"realm": "", "scope": ""}),
|
||||
],
|
||||
)
|
||||
def test_edge_cases(self, header: str, expected: dict[str, str]) -> None:
|
||||
assert parse_www_authenticate_bearer(header) == expected
|
||||
|
||||
|
||||
class TestPathologicalInput:
|
||||
def test_oversized_pathological_input_rejected_under_50ms(self) -> None:
|
||||
"""Headers longer than the defensive cap return ``{}`` immediately.
|
||||
|
||||
The cap is set to 4096 bytes — real ASes emit a few hundred bytes
|
||||
at most. This guards both ``parse_www_authenticate_bearer``
|
||||
callers against pathological input from a misbehaving server.
|
||||
The previous ``header.lower().find("scope", i)`` loop was
|
||||
O(N**2) — a 100 KB header with no ``=`` took ~330 ms because
|
||||
each ``find`` rescanned the entire suffix. The single-pass
|
||||
tokenizer (capped at 4 KB) reduces this to a one-shot length
|
||||
check that returns ``{}`` in microseconds, so the budget is
|
||||
generous regardless of which side of the cap was hit.
|
||||
"""
|
||||
big = "Bearer scope=" + "a" * 10_000
|
||||
start = time.perf_counter()
|
||||
result = parse_www_authenticate_scope(big)
|
||||
elapsed = time.perf_counter() - start
|
||||
assert result == ()
|
||||
assert elapsed < 0.05, f"oversized-header reject took {elapsed * 1000:.1f}ms"
|
||||
|
||||
def test_within_cap_long_header_under_50ms(self) -> None:
|
||||
"""A 4 KB header with thousands of ``find`` candidates still parses fast.
|
||||
|
||||
Stays under the cap so the tokenizer actually runs end to end —
|
||||
the goal is to prove the inner loop is O(N), not just that the
|
||||
cap rejects oversized input.
|
||||
"""
|
||||
# Pack the header right up to the cap with non-matching
|
||||
# auth-params, then put the real ``scope`` at the end.
|
||||
filler_parts = []
|
||||
size = len("Bearer ")
|
||||
i = 0
|
||||
while size < 3900:
|
||||
part = f'xscope{i}="ignore", '
|
||||
if size + len(part) > 3900:
|
||||
break
|
||||
filler_parts.append(part)
|
||||
size += len(part)
|
||||
i += 1
|
||||
header = "Bearer " + "".join(filler_parts) + 'scope="real"'
|
||||
assert len(header) <= 4096
|
||||
start = time.perf_counter()
|
||||
result = parse_www_authenticate_scope(header)
|
||||
elapsed = time.perf_counter() - start
|
||||
assert result == ("real",)
|
||||
assert elapsed < 0.05, f"4kb tokenize took {elapsed * 1000:.1f}ms"
|
||||
@@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.conftest import _seed_static_state
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
@@ -109,13 +110,15 @@ class TestFullLifecycleResourcesPrompts:
|
||||
|
||||
def test_rebuild_resources_produces_merged_state(self, mgr: MCPClientManager) -> None:
|
||||
"""_rebuild_resources merges per-server resources into a unified list."""
|
||||
mgr._per_server_resources["alpha"] = [
|
||||
_make_resource("file:///a.txt", "a", "alpha"),
|
||||
_make_resource("file:///b.txt", "b", "alpha"),
|
||||
]
|
||||
mgr._per_server_resources["beta"] = [
|
||||
_make_resource("file:///c.txt", "c", "beta"),
|
||||
]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
resources=[
|
||||
_make_resource("file:///a.txt", "a", "alpha"),
|
||||
_make_resource("file:///b.txt", "b", "alpha"),
|
||||
],
|
||||
)
|
||||
_seed_static_state(mgr, "beta", resources=[_make_resource("file:///c.txt", "c", "beta")])
|
||||
|
||||
mgr._rebuild_resources()
|
||||
|
||||
@@ -130,13 +133,19 @@ class TestFullLifecycleResourcesPrompts:
|
||||
|
||||
def test_rebuild_prompts_produces_merged_state(self, mgr: MCPClientManager) -> None:
|
||||
"""_rebuild_prompts merges per-server prompts into a unified list."""
|
||||
mgr._per_server_prompts["alpha"] = [
|
||||
_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello"),
|
||||
]
|
||||
mgr._per_server_prompts["beta"] = [
|
||||
_make_prompt("mcp__beta__summarize", "summarize", "beta", "Summarize text"),
|
||||
_make_prompt("mcp__beta__translate", "translate", "beta", "Translate text"),
|
||||
]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello")],
|
||||
)
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"beta",
|
||||
prompts=[
|
||||
_make_prompt("mcp__beta__summarize", "summarize", "beta", "Summarize text"),
|
||||
_make_prompt("mcp__beta__translate", "translate", "beta", "Translate text"),
|
||||
],
|
||||
)
|
||||
|
||||
mgr._rebuild_prompts()
|
||||
|
||||
@@ -164,10 +173,12 @@ class TestFullLifecycleResourcesPrompts:
|
||||
try:
|
||||
# Populate session and resource map
|
||||
session = _make_mock_session()
|
||||
mgr._sessions["alpha"] = session
|
||||
mgr._per_server_resources["alpha"] = [
|
||||
_make_resource("file:///readme.md", "readme", "alpha"),
|
||||
]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
session=session,
|
||||
resources=[_make_resource("file:///readme.md", "readme", "alpha")],
|
||||
)
|
||||
mgr._rebuild_resources()
|
||||
|
||||
result = mgr.read_resource_sync("file:///readme.md", timeout=5)
|
||||
@@ -194,18 +205,22 @@ class TestFullLifecycleResourcesPrompts:
|
||||
|
||||
try:
|
||||
session = _make_mock_session()
|
||||
mgr._sessions["alpha"] = session
|
||||
# Register a template resource (no concrete resources)
|
||||
mgr._per_server_resources["alpha"] = [
|
||||
{
|
||||
"uri": "db://tables/{table}/rows/{id}",
|
||||
"name": "row",
|
||||
"description": "Fetch a row",
|
||||
"mimeType": "application/json",
|
||||
"server": "alpha",
|
||||
"template": True,
|
||||
},
|
||||
]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
session=session,
|
||||
resources=[
|
||||
{
|
||||
"uri": "db://tables/{table}/rows/{id}",
|
||||
"name": "row",
|
||||
"description": "Fetch a row",
|
||||
"mimeType": "application/json",
|
||||
"server": "alpha",
|
||||
"template": True,
|
||||
},
|
||||
],
|
||||
)
|
||||
mgr._rebuild_resources()
|
||||
|
||||
# Template should not be in _resource_map
|
||||
@@ -230,10 +245,12 @@ class TestFullLifecycleResourcesPrompts:
|
||||
|
||||
try:
|
||||
session = _make_mock_session()
|
||||
mgr._sessions["alpha"] = session
|
||||
mgr._per_server_prompts["alpha"] = [
|
||||
_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello"),
|
||||
]
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
session=session,
|
||||
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello")],
|
||||
)
|
||||
mgr._rebuild_prompts()
|
||||
|
||||
messages = mgr.get_prompt_sync(
|
||||
@@ -314,40 +331,42 @@ class TestFullLifecycleResourcesPrompts:
|
||||
def test_shutdown_clears_all_state(self, mgr: MCPClientManager) -> None:
|
||||
"""shutdown() clears sessions, tools, resources, prompts, and listeners."""
|
||||
# Populate state
|
||||
mgr._sessions["alpha"] = MagicMock()
|
||||
mgr._per_server_tools["alpha"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "mcp__alpha__search",
|
||||
"description": "Search",
|
||||
"parameters": {},
|
||||
_seed_static_state(
|
||||
mgr,
|
||||
"alpha",
|
||||
session=MagicMock(),
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "mcp__alpha__search",
|
||||
"description": "Search",
|
||||
"parameters": {},
|
||||
},
|
||||
}
|
||||
],
|
||||
resources=[
|
||||
_make_resource("file:///a.txt", "a", "alpha"),
|
||||
{
|
||||
"uri": "db://tables/{table}",
|
||||
"name": "table",
|
||||
"description": "",
|
||||
"mimeType": "",
|
||||
"server": "alpha",
|
||||
"template": True,
|
||||
},
|
||||
}
|
||||
]
|
||||
],
|
||||
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha")],
|
||||
)
|
||||
mgr._rebuild_tools()
|
||||
mgr._per_server_resources["alpha"] = [
|
||||
_make_resource("file:///a.txt", "a", "alpha"),
|
||||
{
|
||||
"uri": "db://tables/{table}",
|
||||
"name": "table",
|
||||
"description": "",
|
||||
"mimeType": "",
|
||||
"server": "alpha",
|
||||
"template": True,
|
||||
},
|
||||
]
|
||||
mgr._rebuild_resources()
|
||||
mgr._per_server_prompts["alpha"] = [
|
||||
_make_prompt("mcp__alpha__greet", "greet", "alpha"),
|
||||
]
|
||||
mgr._rebuild_prompts()
|
||||
mgr._listeners.append(lambda: None)
|
||||
mgr._resource_listeners.append(lambda: None)
|
||||
mgr._prompt_listeners.append(lambda: None)
|
||||
|
||||
# Verify populated
|
||||
assert len(mgr._sessions) == 1
|
||||
assert len(mgr._static_servers) == 1
|
||||
assert len(mgr._tools) == 1
|
||||
assert len(mgr._resources) == 2 # 1 concrete + 1 template
|
||||
assert len(mgr._template_prefixes) == 1
|
||||
@@ -355,7 +374,7 @@ class TestFullLifecycleResourcesPrompts:
|
||||
|
||||
mgr.shutdown()
|
||||
|
||||
assert len(mgr._sessions) == 0
|
||||
assert len(mgr._static_servers) == 0
|
||||
assert len(mgr._tools) == 0
|
||||
assert len(mgr._tool_map) == 0
|
||||
assert len(mgr._resources) == 0
|
||||
@@ -376,19 +395,15 @@ class TestFullLifecycleResourcesPrompts:
|
||||
mgr.add_resource_listener(lambda: resource_fired.append(1))
|
||||
mgr.add_prompt_listener(lambda: prompt_fired.append(1))
|
||||
|
||||
mgr._per_server_tools["alpha"] = []
|
||||
_seed_static_state(mgr, "alpha", tools=[])
|
||||
mgr._rebuild_tools()
|
||||
assert len(tool_fired) == 1
|
||||
|
||||
mgr._per_server_resources["alpha"] = [
|
||||
_make_resource("file:///x.txt", "x", "alpha"),
|
||||
]
|
||||
_seed_static_state(mgr, "alpha", resources=[_make_resource("file:///x.txt", "x", "alpha")])
|
||||
mgr._rebuild_resources()
|
||||
assert len(resource_fired) == 1
|
||||
|
||||
mgr._per_server_prompts["alpha"] = [
|
||||
_make_prompt("mcp__alpha__p1", "p1", "alpha"),
|
||||
]
|
||||
_seed_static_state(mgr, "alpha", prompts=[_make_prompt("mcp__alpha__p1", "p1", "alpha")])
|
||||
mgr._rebuild_prompts()
|
||||
assert len(prompt_fired) == 1
|
||||
|
||||
|
||||
@@ -0,0 +1,780 @@
|
||||
"""Integration tests for the MCP OAuth ``/connections`` endpoints.
|
||||
|
||||
Covers the list and revoke handlers that surface user-owned MCP server
|
||||
consents to the settings UI:
|
||||
|
||||
* ``GET /v1/api/mcp/oauth/connections`` — non-secret projection only.
|
||||
* ``DELETE /v1/api/mcp/oauth/connections/{server_name}`` — best-effort
|
||||
upstream revoke (RFC 7009) followed by the authoritative local
|
||||
delete; cross-user attempts return 404 with the exact same body
|
||||
shape as a never-existed row to avoid leaking tenant existence.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.routing import Mount, Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.mcp_oauth import (
|
||||
handle_mcp_oauth_list_connections,
|
||||
handle_mcp_oauth_revoke_connection,
|
||||
)
|
||||
from turnstone.core.oidc import OIDCConfig
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures + helpers (mirror tests/test_mcp_oauth_handlers.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _InjectAuthMiddleware(BaseHTTPMiddleware):
|
||||
"""Stamp a fixed authenticated user on every request."""
|
||||
|
||||
def __init__(self, app: Any, user_id: str = "user-1") -> None:
|
||||
super().__init__(app)
|
||||
self._user_id = user_id
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id=self._user_id,
|
||||
scopes=frozenset({"write"}),
|
||||
token_source="config",
|
||||
permissions=frozenset({"read", "write"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
class _NoAuthMiddleware(BaseHTTPMiddleware):
|
||||
"""Leave ``request.state.auth_result`` unset so handlers see anon."""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
async def _list_handler(request: Request) -> Response:
|
||||
return await handle_mcp_oauth_list_connections(request)
|
||||
|
||||
|
||||
async def _revoke_handler(request: Request) -> Response:
|
||||
return await handle_mcp_oauth_revoke_connection(request)
|
||||
|
||||
|
||||
def _build_app(
|
||||
*,
|
||||
storage: SQLiteBackend,
|
||||
http_client: httpx.AsyncClient | MagicMock,
|
||||
token_store: MCPTokenStore | None,
|
||||
user_id: str = "user-1",
|
||||
mcp_client: Any = None,
|
||||
authenticated: bool = True,
|
||||
) -> Starlette:
|
||||
middleware: list[Middleware]
|
||||
if authenticated:
|
||||
middleware = [Middleware(_InjectAuthMiddleware, user_id=user_id)]
|
||||
else:
|
||||
middleware = [Middleware(_NoAuthMiddleware)]
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route("/api/mcp/oauth/connections", _list_handler),
|
||||
Route(
|
||||
"/api/mcp/oauth/connections/{server_name}",
|
||||
_revoke_handler,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
middleware=middleware,
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
app.state.mcp_token_store = token_store
|
||||
app.state.mcp_oauth_http_client = http_client
|
||||
app.state.mcp_oauth_refresh_locks = {}
|
||||
app.state.mcp_oauth_dcr_locks = {}
|
||||
app.state.mcp_oauth_metadata_cache = {}
|
||||
app.state.mcp_oauth_last_cleanup_monotonic = 0.0
|
||||
app.state.oidc_config = OIDCConfig(enabled=False, redirect_base="https://testserver")
|
||||
if mcp_client is not None:
|
||||
app.state.mcp_client = mcp_client
|
||||
return app
|
||||
|
||||
|
||||
def _make_token_store(backend: SQLiteBackend) -> MCPTokenStore:
|
||||
return MCPTokenStore(backend, make_mcp_token_cipher(), node_id="test")
|
||||
|
||||
|
||||
def _seed_oauth_user_server(
|
||||
backend: SQLiteBackend,
|
||||
*,
|
||||
name: str = "srv-oauth",
|
||||
server_id: str = "srv-id-1",
|
||||
cached_issuer: str | None = "https://as.example.com",
|
||||
) -> str:
|
||||
backend.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="client-abc",
|
||||
oauth_scopes="openid profile",
|
||||
oauth_audience="https://mcp.example.com",
|
||||
oauth_authorization_server_url=None,
|
||||
)
|
||||
if cached_issuer is not None:
|
||||
backend.update_mcp_server(server_id, oauth_as_issuer_cached=cached_issuer)
|
||||
return server_id
|
||||
|
||||
|
||||
def _seed_user_token(
|
||||
token_store: MCPTokenStore,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "srv-oauth",
|
||||
refresh_token: str | None = "refresh-secret",
|
||||
) -> None:
|
||||
token_store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token="access-secret",
|
||||
refresh_token=refresh_token,
|
||||
expires_at="2099-12-31T00:00:00",
|
||||
scopes="openid profile",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
|
||||
def _good_as_metadata_doc(
|
||||
*, revocation_endpoint: str | None = "https://as.example.com/revoke"
|
||||
) -> dict[str, Any]:
|
||||
doc: dict[str, Any] = {
|
||||
"issuer": "https://as.example.com",
|
||||
"authorization_endpoint": "https://as.example.com/authorize",
|
||||
"token_endpoint": "https://as.example.com/token",
|
||||
"registration_endpoint": "https://as.example.com/register",
|
||||
"jwks_uri": "https://as.example.com/jwks",
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["none", "client_secret_basic"],
|
||||
}
|
||||
if revocation_endpoint is not None:
|
||||
doc["revocation_endpoint"] = revocation_endpoint
|
||||
return doc
|
||||
|
||||
|
||||
def _mk_response(
|
||||
status_code: int = 200,
|
||||
json_body: Any = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> MagicMock:
|
||||
import json as _json
|
||||
|
||||
resp = MagicMock(spec=httpx.Response)
|
||||
resp.status_code = status_code
|
||||
resp.headers = headers or {}
|
||||
body_str = _json.dumps(json_body) if json_body is not None else ""
|
||||
resp.content = body_str.encode("utf-8")
|
||||
if json_body is not None:
|
||||
resp.json.return_value = json_body
|
||||
else:
|
||||
resp.json.side_effect = ValueError("no body")
|
||||
resp.text = body_str
|
||||
return resp
|
||||
|
||||
|
||||
def _public_addr_patch():
|
||||
return patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("93.184.216.34", 0))])
|
||||
|
||||
|
||||
def _drain_revoke_upstream_tasks(client: TestClient, timeout: float = 2.0) -> None:
|
||||
"""Block until all in-flight upstream-revoke tasks complete.
|
||||
|
||||
Phase 8 perf-1 made the RFC 7009 AS round-trip a fire-and-forget
|
||||
task so the user-visible 204 isn't gated on the AS. The tasks were
|
||||
scheduled on the TestClient's portal loop; we re-enter that loop
|
||||
via :attr:`TestClient.portal` to await them. Tests that assert
|
||||
against the upstream POST must call this helper before the
|
||||
assertion.
|
||||
"""
|
||||
from turnstone.core.mcp_oauth import _revoke_upstream_tasks
|
||||
|
||||
portal = getattr(client, "portal", None)
|
||||
if portal is None:
|
||||
return
|
||||
|
||||
async def _drain() -> None:
|
||||
pending = list(_revoke_upstream_tasks)
|
||||
if pending:
|
||||
async with asyncio.timeout(timeout):
|
||||
await asyncio.gather(*pending, return_exceptions=True)
|
||||
|
||||
portal.call(_drain)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
backend = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
backend.create_user("user-1", "user1", "User One", "hash")
|
||||
backend.create_user("user-2", "user2", "User Two", "hash")
|
||||
return backend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def http_client_mock() -> MagicMock:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock()
|
||||
client.post = AsyncMock()
|
||||
return client
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /connections
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestListConnections:
|
||||
def test_list_connections_unauthenticated_401(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
token_store = _make_token_store(storage)
|
||||
app = _build_app(
|
||||
storage=storage,
|
||||
http_client=http_client_mock,
|
||||
token_store=token_store,
|
||||
authenticated=False,
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
assert resp.status_code == 401
|
||||
assert resp.json() == {"error": "Authentication required"}
|
||||
|
||||
def test_list_connections_no_token_store_503(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=None)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
assert resp.status_code == 503
|
||||
|
||||
def test_list_connections_empty_user_returns_empty_list(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
token_store = _make_token_store(storage)
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"connections": []}
|
||||
|
||||
def test_list_connections_returns_users_consents(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage, name="srv-a", server_id="srv-id-a")
|
||||
_seed_oauth_user_server(storage, name="srv-b", server_id="srv-id-b")
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, server_name="srv-a")
|
||||
_seed_user_token(token_store, server_name="srv-b")
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert "connections" in body
|
||||
servers = sorted(row["server_name"] for row in body["connections"])
|
||||
assert servers == ["srv-a", "srv-b"]
|
||||
|
||||
def test_list_connections_isolates_by_user(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, user_id="user-1", server_name="srv-oauth")
|
||||
_seed_user_token(token_store, user_id="user-2", server_name="srv-oauth")
|
||||
|
||||
# User-1 sees only user-1's row.
|
||||
app = _build_app(
|
||||
storage=storage, http_client=http_client_mock, token_store=token_store, user_id="user-1"
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
rows = resp.json()["connections"]
|
||||
assert all(row["user_id"] == "user-1" for row in rows)
|
||||
assert len(rows) == 1
|
||||
|
||||
# User-2 sees only user-2's row.
|
||||
app2 = _build_app(
|
||||
storage=storage, http_client=http_client_mock, token_store=token_store, user_id="user-2"
|
||||
)
|
||||
client2 = TestClient(app2, raise_server_exceptions=False)
|
||||
resp2 = client2.get("/v1/api/mcp/oauth/connections")
|
||||
rows2 = resp2.json()["connections"]
|
||||
assert all(row["user_id"] == "user-2" for row in rows2)
|
||||
assert len(rows2) == 1
|
||||
|
||||
def test_list_connections_does_not_leak_secret_fields(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.get("/v1/api/mcp/oauth/connections")
|
||||
rows = resp.json()["connections"]
|
||||
assert rows
|
||||
for row in rows:
|
||||
for forbidden in (
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"access_token_ct",
|
||||
"refresh_token_ct",
|
||||
):
|
||||
assert forbidden not in row, f"secret field {forbidden!r} leaked in {row!r}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DELETE /connections/{server_name}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRevokeConnection:
|
||||
def test_revoke_connection_unauthenticated_401(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store)
|
||||
app = _build_app(
|
||||
storage=storage,
|
||||
http_client=http_client_mock,
|
||||
token_store=token_store,
|
||||
authenticated=False,
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_revoke_connection_missing_row_404(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
token_store = _make_token_store(storage)
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-nonexistent")
|
||||
assert resp.status_code == 404
|
||||
assert resp.json() == {"error": "No such connection"}
|
||||
|
||||
def test_revoke_connection_local_delete_succeeds_204(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
# No refresh token → upstream revoke is skipped entirely.
|
||||
_seed_user_token(token_store, refresh_token=None)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 204
|
||||
# Local row is gone.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
# Upstream not contacted.
|
||||
http_client_mock.post.assert_not_called()
|
||||
|
||||
def test_revoke_connection_with_revocation_endpoint_calls_upstream(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token="refresh-secret")
|
||||
|
||||
http_client_mock.get.return_value = _mk_response(200, _good_as_metadata_doc())
|
||||
http_client_mock.post.return_value = _mk_response(200)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
# ``with TestClient(...)`` keeps a persistent portal so the
|
||||
# fire-and-forget upstream-revoke task isn't cancelled when
|
||||
# the request handler returns. See ``_drain_revoke_upstream_tasks``.
|
||||
# The SSRF-validator's ``socket.getaddrinfo`` patch must wrap
|
||||
# the drain too — the discovery call now runs on the background
|
||||
# task and resolves the AS hostname after the request returns.
|
||||
with (
|
||||
TestClient(app, raise_server_exceptions=False) as client,
|
||||
_public_addr_patch(),
|
||||
):
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
|
||||
assert resp.status_code == 204
|
||||
# Local row is gone.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
# The upstream RFC 7009 POST is fire-and-forget post-Phase-8 perf-1
|
||||
# so the test must drain the in-flight task set before asserting.
|
||||
_drain_revoke_upstream_tasks(client)
|
||||
# Upstream POSTed to revocation_endpoint with refresh-token grant.
|
||||
assert http_client_mock.post.await_count == 1
|
||||
call = http_client_mock.post.await_args
|
||||
assert call.args[0] == "https://as.example.com/revoke"
|
||||
data = call.kwargs.get("data") or {}
|
||||
assert data.get("token") == "refresh-secret"
|
||||
assert data.get("token_type_hint") == "refresh_token"
|
||||
assert data.get("client_id") == "client-abc"
|
||||
|
||||
def test_revoke_connection_without_revocation_endpoint_skips_upstream(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token="refresh-secret")
|
||||
|
||||
http_client_mock.get.return_value = _mk_response(
|
||||
200, _good_as_metadata_doc(revocation_endpoint=None)
|
||||
)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
# ``with TestClient(...)`` keeps the portal alive for the
|
||||
# background task drain.
|
||||
with (
|
||||
TestClient(app, raise_server_exceptions=False) as client,
|
||||
_public_addr_patch(),
|
||||
):
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
|
||||
assert resp.status_code == 204
|
||||
# Local row gone, upstream POST never made.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
# Drain the fire-and-forget discovery task before asserting on
|
||||
# the AS POST — the task runs ``discover_authorization_server``
|
||||
# but does NOT proceed to POST because revocation_endpoint is
|
||||
# absent.
|
||||
_drain_revoke_upstream_tasks(client)
|
||||
http_client_mock.post.assert_not_called()
|
||||
|
||||
def test_revoke_connection_upstream_failure_still_204(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token="refresh-secret")
|
||||
|
||||
http_client_mock.get.return_value = _mk_response(200, _good_as_metadata_doc())
|
||||
# AS returns 500 — local delete must still succeed.
|
||||
http_client_mock.post.return_value = _mk_response(500)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
with _public_addr_patch():
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
|
||||
assert resp.status_code == 204
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
|
||||
def test_revoke_connection_audit_event_emitted_with_user_revoked_reason(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token=None)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 204
|
||||
|
||||
# Audit row was written via the storage API (tests don't poke at
|
||||
# the SQLite schema directly — the table name is an internal
|
||||
# detail).
|
||||
events = storage.list_audit_events(action="mcp_server.oauth.token_revoked")
|
||||
assert len(events) == 1
|
||||
ev = events[0]
|
||||
assert ev["user_id"] == "user-1"
|
||||
# resource_id is the immutable server_id PK, not the name.
|
||||
assert ev["resource_id"] == "srv-id-1"
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(ev["detail"]) if isinstance(ev["detail"], str) else ev["detail"]
|
||||
assert detail["reason"] == "user_revoked"
|
||||
assert detail["upstream_revoke_outcome"] == "no_refresh_token"
|
||||
assert detail["server_name"] == "srv-oauth"
|
||||
|
||||
def test_revoke_connection_cross_user_attempt_404(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
# Owned by user-2, not user-1.
|
||||
_seed_user_token(token_store, user_id="user-2", server_name="srv-oauth")
|
||||
|
||||
app = _build_app(
|
||||
storage=storage, http_client=http_client_mock, token_store=token_store, user_id="user-1"
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
# Cross-user attempt MUST surface as a generic 404, byte-identical
|
||||
# body to the never-existed case (no tenant existence leak).
|
||||
assert resp.status_code == 404
|
||||
assert resp.json() == {"error": "No such connection"}
|
||||
# Drain pending tasks defensively, then confirm the upstream
|
||||
# endpoint was NEVER contacted on the 404-cross-user path. A
|
||||
# bug that scheduled the AS round-trip before the cross-user
|
||||
# check would leak existence via the AS-side 200/4xx response.
|
||||
_drain_revoke_upstream_tasks(client)
|
||||
http_client_mock.post.assert_not_called()
|
||||
# User-2's row is untouched.
|
||||
assert token_store.get_user_token("user-2", "srv-oauth") is not None
|
||||
|
||||
def test_revoke_connection_evicts_pool_session(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token=None)
|
||||
|
||||
mcp_client_mock = MagicMock()
|
||||
# ``evict_user_session`` is the public sync surface on
|
||||
# MCPClientManager; mirror its signature here so the handler's
|
||||
# ``hasattr`` gate triggers.
|
||||
mcp_client_mock.evict_user_session = MagicMock(return_value=None)
|
||||
|
||||
app = _build_app(
|
||||
storage=storage,
|
||||
http_client=http_client_mock,
|
||||
token_store=token_store,
|
||||
mcp_client=mcp_client_mock,
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 204
|
||||
|
||||
mcp_client_mock.evict_user_session.assert_called_once_with("user-1", "srv-oauth")
|
||||
|
||||
def test_revoke_connection_pool_eviction_failure_does_not_block_204(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token=None)
|
||||
|
||||
mcp_client_mock = MagicMock()
|
||||
mcp_client_mock.evict_user_session = MagicMock(side_effect=RuntimeError("loop closed"))
|
||||
|
||||
app = _build_app(
|
||||
storage=storage,
|
||||
http_client=http_client_mock,
|
||||
token_store=token_store,
|
||||
mcp_client=mcp_client_mock,
|
||||
)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 204
|
||||
# Local delete still happened.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
|
||||
def test_revoke_connection_204_not_gated_on_slow_upstream(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
"""The user-visible 204 must return promptly even when the
|
||||
upstream AS round-trip is slow / hanging. Pre-perf-1 the
|
||||
handler awaited ``revoke_token_at_as`` synchronously, so a
|
||||
stuck AS could block the user's revoke confirmation. The
|
||||
fire-and-forget refactor moves the call onto a background task
|
||||
so the 204 returns in well under 1s regardless of AS latency.
|
||||
Bound is conservative for CI runner jitter.
|
||||
"""
|
||||
import time
|
||||
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token="refresh-secret")
|
||||
|
||||
http_client_mock.get.return_value = _mk_response(200, _good_as_metadata_doc())
|
||||
|
||||
async def _slow_post(*_args: Any, **_kwargs: Any) -> Any:
|
||||
# Simulate a slow / unreachable AS — must NOT gate the
|
||||
# user-visible 204 on this round-trip.
|
||||
await asyncio.sleep(5.0)
|
||||
return _mk_response(200)
|
||||
|
||||
http_client_mock.post = AsyncMock(side_effect=_slow_post)
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
with _public_addr_patch():
|
||||
start = time.monotonic()
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
assert resp.status_code == 204
|
||||
# 1s ceiling — the 204 must return on the local-delete path
|
||||
# without waiting on the AS POST (which sleeps 5s above). Bound
|
||||
# is intentionally generous for CI runner jitter; the actual
|
||||
# path is on the order of milliseconds.
|
||||
assert elapsed < 1.0, (
|
||||
f"204 returned in {elapsed:.3f}s — should be <1s; the "
|
||||
"fire-and-forget upstream revoke isn't decoupled from the "
|
||||
"response."
|
||||
)
|
||||
# The local row IS gone — the authoritative delete ran before
|
||||
# the 204 returned, even though the AS round-trip is still
|
||||
# in flight.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
# Cancel any in-flight tasks so the test client can exit cleanly.
|
||||
from turnstone.core.mcp_oauth import _revoke_upstream_tasks
|
||||
|
||||
portal = getattr(client, "portal", None)
|
||||
if portal is not None:
|
||||
for task in list(_revoke_upstream_tasks):
|
||||
portal.call(task.cancel)
|
||||
|
||||
def test_revoke_connection_sheds_upstream_when_task_set_full(
|
||||
self, storage: SQLiteBackend, http_client_mock: MagicMock
|
||||
) -> None:
|
||||
"""Round-2 q-2 regression: the soft cap on ``_revoke_upstream_tasks``
|
||||
is the only protection against unbounded background-task pile-up
|
||||
under a coordinated mass-revoke. When the set is full, the local
|
||||
delete still runs but no upstream task is scheduled; the audit
|
||||
detail records ``upstream_revoke_outcome="shed_by_cap"`` and
|
||||
the AS endpoint is never contacted.
|
||||
"""
|
||||
from turnstone.core.mcp_oauth import (
|
||||
_REVOKE_UPSTREAM_TASKS_MAX,
|
||||
_revoke_upstream_tasks,
|
||||
)
|
||||
|
||||
_seed_oauth_user_server(storage)
|
||||
token_store = _make_token_store(storage)
|
||||
_seed_user_token(token_store, refresh_token="refresh-secret")
|
||||
|
||||
app = _build_app(storage=storage, http_client=http_client_mock, token_store=token_store)
|
||||
sentinel_event_holder: dict[str, asyncio.Event] = {}
|
||||
|
||||
# Use ``with TestClient(...)`` so the portal stays alive — we
|
||||
# need to schedule sentinel tasks on the portal's loop and the
|
||||
# tasks must outlive the request to actually fill the set.
|
||||
with (
|
||||
TestClient(app, raise_server_exceptions=False) as client,
|
||||
_public_addr_patch(),
|
||||
):
|
||||
portal = client.portal
|
||||
assert portal is not None
|
||||
|
||||
async def _create_sentinel_event() -> asyncio.Event:
|
||||
event = asyncio.Event()
|
||||
sentinel_event_holder["event"] = event
|
||||
return event
|
||||
|
||||
sentinel_event = portal.call(_create_sentinel_event)
|
||||
|
||||
async def _wait_on_event() -> None:
|
||||
await sentinel_event.wait()
|
||||
|
||||
async def _fill_task_set() -> list[asyncio.Task[None]]:
|
||||
tasks: list[asyncio.Task[None]] = []
|
||||
for _ in range(_REVOKE_UPSTREAM_TASKS_MAX):
|
||||
t = asyncio.create_task(_wait_on_event())
|
||||
_revoke_upstream_tasks.add(t)
|
||||
tasks.append(t)
|
||||
return tasks
|
||||
|
||||
sentinels = portal.call(_fill_task_set)
|
||||
assert len(_revoke_upstream_tasks) >= _REVOKE_UPSTREAM_TASKS_MAX
|
||||
|
||||
try:
|
||||
resp = client.delete("/v1/api/mcp/oauth/connections/srv-oauth")
|
||||
assert resp.status_code == 204
|
||||
# Local row is still gone — authoritative delete ran.
|
||||
assert token_store.get_user_token("user-1", "srv-oauth") is None
|
||||
# AS endpoint MUST NOT have been contacted.
|
||||
http_client_mock.post.assert_not_called()
|
||||
# Audit detail records the categorical shed outcome.
|
||||
events = storage.list_audit_events(action="mcp_server.oauth.token_revoked")
|
||||
assert len(events) == 1
|
||||
detail = events[0]["detail"]
|
||||
if isinstance(detail, str):
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(detail)
|
||||
assert detail["upstream_revoke_outcome"] == "shed_by_cap"
|
||||
finally:
|
||||
# Release sentinels so the portal can shut down cleanly.
|
||||
async def _release() -> None:
|
||||
sentinel_event.set()
|
||||
for t in sentinels:
|
||||
t.cancel()
|
||||
await asyncio.gather(*sentinels, return_exceptions=True)
|
||||
|
||||
portal.call(_release)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# evict_user_session helper sanity checks
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEvictUserSession:
|
||||
def test_evict_user_session_no_loop_is_silent_noop(self) -> None:
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
|
||||
mgr = MCPClientManager.__new__(MCPClientManager)
|
||||
mgr._loop = None # type: ignore[attr-defined]
|
||||
# Must not raise.
|
||||
mgr.evict_user_session("user-1", "srv-oauth")
|
||||
|
||||
def test_evict_user_session_dispatches_to_loop(self) -> None:
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
|
||||
mgr = MCPClientManager.__new__(MCPClientManager)
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
mgr._loop = loop # type: ignore[attr-defined]
|
||||
mgr._user_pool_entries = {} # type: ignore[attr-defined]
|
||||
mgr._last_pool_notification_refresh = {} # type: ignore[attr-defined]
|
||||
evicted: list[tuple[str, str]] = []
|
||||
|
||||
def _fake_evict(key: tuple[str, str]) -> None:
|
||||
evicted.append(key)
|
||||
|
||||
mgr._evict_session = _fake_evict # type: ignore[method-assign]
|
||||
|
||||
# Run the dispatch on a separate thread so the loop can drain.
|
||||
import threading
|
||||
|
||||
done = threading.Event()
|
||||
|
||||
def _run_loop() -> None:
|
||||
loop.call_later(0.05, loop.stop)
|
||||
loop.run_forever()
|
||||
done.set()
|
||||
|
||||
t = threading.Thread(target=_run_loop, daemon=True)
|
||||
t.start()
|
||||
mgr.evict_user_session("user-1", "srv-oauth")
|
||||
done.wait(timeout=1.0)
|
||||
|
||||
assert evicted == [("user-1", "srv-oauth")]
|
||||
finally:
|
||||
if not loop.is_closed():
|
||||
loop.close()
|
||||
@@ -0,0 +1,626 @@
|
||||
"""Discovery tests for the per-(user, server) MCP OAuth flow.
|
||||
|
||||
Covers PRM (RFC 9728) and AS metadata (RFC 8414) discovery, including:
|
||||
- override URL takes precedence
|
||||
- PRM happy path: server URL -> .well-known/oauth-protected-resource
|
||||
-> ``authorization_servers[0]``
|
||||
- PRM 401 + ``WWW-Authenticate: Bearer resource_metadata="..."`` follows
|
||||
the URL.
|
||||
- AS metadata without S256 -> :class:`MCPOAuthDiscoveryError`.
|
||||
- SSRF rejection on AS issuer URL.
|
||||
- In-memory cache hit/miss + persistent cache write to
|
||||
``mcp_servers.oauth_as_issuer_cached``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from turnstone.core.mcp_oauth import (
|
||||
ASMetadata,
|
||||
MCPOAuthDiscoveryError,
|
||||
_parse_prm_url_from_www_authenticate,
|
||||
discover_authorization_server,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mk_response(
|
||||
status_code: int = 200,
|
||||
json_body: Any = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> MagicMock:
|
||||
"""Build a MagicMock that quacks like ``httpx.Response``."""
|
||||
resp = MagicMock(spec=httpx.Response)
|
||||
resp.status_code = status_code
|
||||
resp.headers = headers or {}
|
||||
resp.content = (str(json_body) if json_body is not None else "").encode("utf-8")
|
||||
if json_body is not None:
|
||||
resp.json.return_value = json_body
|
||||
else:
|
||||
resp.json.side_effect = ValueError("no body")
|
||||
resp.text = str(json_body) if json_body is not None else ""
|
||||
return resp
|
||||
|
||||
|
||||
def _good_as_metadata_doc() -> dict[str, Any]:
|
||||
return {
|
||||
"issuer": "https://as.example.com",
|
||||
"authorization_endpoint": "https://as.example.com/authorize",
|
||||
"token_endpoint": "https://as.example.com/token",
|
||||
"jwks_uri": "https://as.example.com/jwks",
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["none", "client_secret_basic"],
|
||||
"registration_endpoint": "https://as.example.com/register",
|
||||
}
|
||||
|
||||
|
||||
def _public_addr_patch():
|
||||
return patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("93.184.216.34", 0))])
|
||||
|
||||
|
||||
def _mk_storage_mock(server_id: str = "srv-id") -> MagicMock:
|
||||
storage = MagicMock()
|
||||
storage.update_mcp_server.return_value = True
|
||||
return storage
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PRM parser
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParsePRMUrl:
|
||||
def test_extracts_resource_metadata_url(self) -> None:
|
||||
header = (
|
||||
'Bearer error="invalid_token", '
|
||||
'resource_metadata="https://srv.example.com/.well-known/oauth-protected-resource"'
|
||||
)
|
||||
url = _parse_prm_url_from_www_authenticate(header)
|
||||
assert url == "https://srv.example.com/.well-known/oauth-protected-resource"
|
||||
|
||||
def test_returns_none_when_absent(self) -> None:
|
||||
assert _parse_prm_url_from_www_authenticate('Bearer realm="x"') is None
|
||||
|
||||
def test_handles_empty_header(self) -> None:
|
||||
assert _parse_prm_url_from_www_authenticate("") is None
|
||||
|
||||
def test_handles_escaped_quote_in_value(self) -> None:
|
||||
"""RFC 7230 quoted-string allows ``\\"`` — naive ``[^"]+`` truncates.
|
||||
|
||||
A malicious or buggy resource server could send an embedded
|
||||
escaped quote; the parser must yield the unescaped value, not
|
||||
the prefix up to the escaped quote.
|
||||
"""
|
||||
header = 'Bearer resource_metadata="https://srv.example.com/with\\"quote"'
|
||||
url = _parse_prm_url_from_www_authenticate(header)
|
||||
assert url == 'https://srv.example.com/with"quote'
|
||||
|
||||
def test_handles_escaped_backslash(self) -> None:
|
||||
header = 'Bearer resource_metadata="https://srv.example.com/back\\\\slash"'
|
||||
url = _parse_prm_url_from_www_authenticate(header)
|
||||
assert url == "https://srv.example.com/back\\slash"
|
||||
|
||||
def test_unterminated_quoted_string_returns_none(self) -> None:
|
||||
# Closing quote missing — naive regex would still match, but
|
||||
# the proper parser should reject malformed input.
|
||||
header = 'Bearer resource_metadata="https://srv.example.com/no-close'
|
||||
assert _parse_prm_url_from_www_authenticate(header) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# discover_authorization_server happy paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDiscoveryOverride:
|
||||
def test_override_url_skips_prm(self) -> None:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, _good_as_metadata_doc()))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert isinstance(meta, ASMetadata)
|
||||
assert meta.token_endpoint == "https://as.example.com/token"
|
||||
# Only the AS metadata URL was hit, not PRM.
|
||||
called_urls = [c.args[0] for c in client.get.call_args_list]
|
||||
assert all("oauth-authorization-server" in u for u in called_urls)
|
||||
|
||||
|
||||
class TestDiscoveryPRM:
|
||||
def test_prm_happy_path(self) -> None:
|
||||
async def _get(url, *args, **kwargs):
|
||||
if url.endswith("/oauth-protected-resource"):
|
||||
return _mk_response(
|
||||
200,
|
||||
{
|
||||
"resource": "https://mcp.example.com",
|
||||
"authorization_servers": ["https://as.example.com"],
|
||||
},
|
||||
)
|
||||
if url.endswith("/oauth-authorization-server"):
|
||||
return _mk_response(200, _good_as_metadata_doc())
|
||||
raise AssertionError(f"unexpected URL: {url}")
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(side_effect=_get)
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert meta.issuer == "https://as.example.com"
|
||||
|
||||
def test_prm_401_follows_www_authenticate(self) -> None:
|
||||
async def _get(url, *args, **kwargs):
|
||||
if url == "https://mcp.example.com/.well-known/oauth-protected-resource":
|
||||
return _mk_response(
|
||||
401,
|
||||
headers={
|
||||
"www-authenticate": (
|
||||
'Bearer error="invalid_token", '
|
||||
"resource_metadata="
|
||||
'"https://meta.example.com/prm"'
|
||||
)
|
||||
},
|
||||
json_body=None,
|
||||
)
|
||||
if url == "https://meta.example.com/prm":
|
||||
return _mk_response(
|
||||
200,
|
||||
{
|
||||
"authorization_servers": ["https://as.example.com"],
|
||||
},
|
||||
)
|
||||
if url.endswith("/oauth-authorization-server"):
|
||||
return _mk_response(200, _good_as_metadata_doc())
|
||||
raise AssertionError(f"unexpected URL: {url}")
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(side_effect=_get)
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert meta.token_endpoint == "https://as.example.com/token"
|
||||
|
||||
def test_prm_401_without_resource_metadata_raises(self) -> None:
|
||||
async def _get(url, *args, **kwargs):
|
||||
return _mk_response(401, headers={"www-authenticate": "Basic realm=x"})
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(side_effect=_get)
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError, match="resource_metadata"):
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AS metadata validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestASMetadataValidation:
|
||||
def test_no_s256_raises(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
doc["code_challenge_methods_supported"] = ["plain"]
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError, match="S256"):
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_missing_endpoints_raises(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
del doc["token_endpoint"]
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError, match="missing required"):
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_third_party_endpoint_rejected(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
doc["token_endpoint"] = "https://attacker.example.com/token"
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError, match="token_endpoint"):
|
||||
asyncio.run(_run())
|
||||
|
||||
def test_ssrf_on_override_rejected(self) -> None:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock()
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
# Resolve to private 10.x — SSRF guard fires before any HTTP call.
|
||||
with patch(
|
||||
"socket.getaddrinfo",
|
||||
return_value=[(2, 1, 6, "", ("10.0.0.1", 0))],
|
||||
):
|
||||
await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://internal.corp.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError):
|
||||
asyncio.run(_run())
|
||||
client.get.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Caching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMetadataCache:
|
||||
def test_cache_miss_then_hit(self) -> None:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, _good_as_metadata_doc()))
|
||||
storage = _mk_storage_mock()
|
||||
cache: dict[str, tuple[ASMetadata, float]] = {}
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
first = await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
metadata_cache=cache,
|
||||
)
|
||||
second = await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer="https://as.example.com",
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
metadata_cache=cache,
|
||||
)
|
||||
return first, second
|
||||
|
||||
first, second = asyncio.run(_run())
|
||||
assert first.token_endpoint == second.token_endpoint
|
||||
# First call hit AS metadata; second call hit the cache.
|
||||
assert client.get.call_count == 1
|
||||
|
||||
def test_cache_expiry_refetches(self) -> None:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, _good_as_metadata_doc()))
|
||||
storage = _mk_storage_mock()
|
||||
# Pre-populate cache with a very stale entry.
|
||||
stale_meta = ASMetadata(
|
||||
issuer="https://as.example.com",
|
||||
authorization_endpoint="https://as.example.com/authorize",
|
||||
token_endpoint="https://as.example.com/token",
|
||||
registration_endpoint=None,
|
||||
revocation_endpoint=None,
|
||||
jwks_uri=None,
|
||||
code_challenge_methods_supported=("S256",),
|
||||
token_endpoint_auth_methods_supported=(),
|
||||
)
|
||||
cache = {"https://as.example.com": (stale_meta, time.monotonic() - 10**6)}
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
metadata_cache=cache,
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
# Stale entry was bypassed -> we hit the network.
|
||||
assert client.get.call_count == 1
|
||||
assert meta.token_endpoint == "https://as.example.com/token"
|
||||
|
||||
def test_persistent_cache_write_on_first_resolution(self) -> None:
|
||||
async def _get(url, *args, **kwargs):
|
||||
if url.endswith("/oauth-protected-resource"):
|
||||
return _mk_response(200, {"authorization_servers": ["https://as.example.com"]})
|
||||
return _mk_response(200, _good_as_metadata_doc())
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(side_effect=_get)
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
asyncio.run(_run())
|
||||
# update_mcp_server was called once with the cached issuer.
|
||||
storage.update_mcp_server.assert_called_once_with(
|
||||
"srv-id", oauth_as_issuer_cached="https://as.example.com"
|
||||
)
|
||||
|
||||
def test_persistent_cache_skip_when_already_cached(self) -> None:
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, _good_as_metadata_doc()))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run():
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer="https://as.example.com",
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
asyncio.run(_run())
|
||||
storage.update_mcp_server.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sec-3 — cached_issuer re-validated on read
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCachedIssuerSSRFRevalidation:
|
||||
"""A cached issuer URL must still pass SSRF validation on every read.
|
||||
|
||||
Defense-in-depth: an admin who points ``oauth_as_issuer_cached`` at a
|
||||
private address (or a hostname that has rebound to one) should not
|
||||
bypass the guard just because the value was already in the row.
|
||||
"""
|
||||
|
||||
def test_cached_issuer_rejected_clears_row_and_falls_through_to_prm(self) -> None:
|
||||
async def _get(url: str, *args: Any, **kwargs: Any) -> MagicMock:
|
||||
if url.endswith("/oauth-protected-resource"):
|
||||
return _mk_response(200, {"authorization_servers": ["https://as.example.com"]})
|
||||
if url.endswith("/oauth-authorization-server"):
|
||||
return _mk_response(200, _good_as_metadata_doc())
|
||||
raise AssertionError(f"unexpected URL: {url}")
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(side_effect=_get)
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
# cached_issuer points at a private host. SSRF guard fires on
|
||||
# the cached value first, the row is cleared, and PRM
|
||||
# discovery runs as a fallback.
|
||||
async def _run() -> Any:
|
||||
with patch(
|
||||
"socket.getaddrinfo",
|
||||
# Private resolution for "internal.corp", public for everything else.
|
||||
side_effect=lambda host, *a, **kw: [
|
||||
(2, 1, 6, "", ("10.0.0.1" if "internal" in host else "93.184.216.34", 0))
|
||||
],
|
||||
):
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url=None,
|
||||
cached_issuer="https://internal.corp.example.com",
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert meta.token_endpoint == "https://as.example.com/token"
|
||||
# The bad cached_issuer was cleared from the row.
|
||||
clear_calls = [
|
||||
c
|
||||
for c in storage.update_mcp_server.call_args_list
|
||||
if c.kwargs.get("oauth_as_issuer_cached") is None
|
||||
]
|
||||
assert clear_calls, "cached_issuer should have been cleared"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# revocation_endpoint parsing (RFC 8414)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestASMetadataRevocationEndpoint:
|
||||
def test_as_metadata_parses_revocation_endpoint(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
doc["revocation_endpoint"] = "https://as.example.com/revoke"
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run() -> ASMetadata:
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert meta.revocation_endpoint == "https://as.example.com/revoke"
|
||||
|
||||
def test_as_metadata_revocation_endpoint_absent(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
doc.pop("revocation_endpoint", None)
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run() -> ASMetadata:
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
meta = asyncio.run(_run())
|
||||
assert meta.revocation_endpoint is None
|
||||
|
||||
def test_as_metadata_revocation_endpoint_rejected_when_cross_origin(self) -> None:
|
||||
doc = _good_as_metadata_doc()
|
||||
doc["revocation_endpoint"] = "https://attacker.example.com/revoke"
|
||||
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.get = AsyncMock(return_value=_mk_response(200, doc))
|
||||
storage = _mk_storage_mock()
|
||||
|
||||
async def _run() -> ASMetadata:
|
||||
with _public_addr_patch():
|
||||
return await discover_authorization_server(
|
||||
server_name="srv-x",
|
||||
server_url="https://mcp.example.com/sse",
|
||||
override_url="https://as.example.com",
|
||||
cached_issuer=None,
|
||||
http_client=client,
|
||||
storage=storage,
|
||||
server_id="srv-id",
|
||||
trusted_hosts=frozenset(),
|
||||
)
|
||||
|
||||
with pytest.raises(MCPOAuthDiscoveryError, match="revocation_endpoint"):
|
||||
asyncio.run(_run())
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,133 @@
|
||||
"""Smoke tests for the new OAuth-MCP storage tables.
|
||||
|
||||
Phase 2 only adds the schema — token CRUD lands in Phase 3 and pending-
|
||||
state CRUD in Phase 4. These tests verify the tables exist after
|
||||
``init_storage`` and accept the documented row shape via raw SQL.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import mcp_oauth_pending, mcp_user_tokens
|
||||
|
||||
|
||||
class TestMcpUserTokensTable:
|
||||
def test_table_exists_and_accepts_row(self, backend) -> None:
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.insert(mcp_user_tokens),
|
||||
{
|
||||
"user_id": "u1",
|
||||
"server_name": "srv-a",
|
||||
"access_token_ct": b"\x00ciphertext-a",
|
||||
"refresh_token_ct": b"\x00ciphertext-r",
|
||||
"expires_at": "2026-05-04T12:00:00",
|
||||
"scopes": "openid profile",
|
||||
"as_issuer": "https://auth.example.com",
|
||||
"audience": "https://mcp.example.com",
|
||||
"created": "2026-05-04T11:00:00",
|
||||
"last_refreshed": None,
|
||||
},
|
||||
)
|
||||
conn.commit()
|
||||
row = conn.execute(
|
||||
sa.select(mcp_user_tokens).where(
|
||||
(mcp_user_tokens.c.user_id == "u1") & (mcp_user_tokens.c.server_name == "srv-a")
|
||||
)
|
||||
).one()
|
||||
assert row.access_token_ct == b"\x00ciphertext-a"
|
||||
assert row.refresh_token_ct == b"\x00ciphertext-r"
|
||||
assert row.scopes == "openid profile"
|
||||
assert row.audience == "https://mcp.example.com"
|
||||
|
||||
def test_composite_pk_distinguishes_user_server(self, backend) -> None:
|
||||
"""Same user, different server => two rows; same (user, server) => conflict."""
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.insert(mcp_user_tokens),
|
||||
[
|
||||
{
|
||||
"user_id": "u1",
|
||||
"server_name": "srv-a",
|
||||
"access_token_ct": b"a",
|
||||
"refresh_token_ct": None,
|
||||
"expires_at": None,
|
||||
"scopes": None,
|
||||
"as_issuer": "https://auth.example.com",
|
||||
"audience": "https://a.example.com",
|
||||
"created": "2026-05-04T11:00:00",
|
||||
"last_refreshed": None,
|
||||
},
|
||||
{
|
||||
"user_id": "u1",
|
||||
"server_name": "srv-b",
|
||||
"access_token_ct": b"b",
|
||||
"refresh_token_ct": None,
|
||||
"expires_at": None,
|
||||
"scopes": None,
|
||||
"as_issuer": "https://auth.example.com",
|
||||
"audience": "https://b.example.com",
|
||||
"created": "2026-05-04T11:00:00",
|
||||
"last_refreshed": None,
|
||||
},
|
||||
],
|
||||
)
|
||||
conn.commit()
|
||||
count = conn.execute(sa.select(sa.func.count()).select_from(mcp_user_tokens)).scalar()
|
||||
assert count == 2
|
||||
|
||||
|
||||
class TestMcpOauthPendingTable:
|
||||
def test_table_exists_and_accepts_row(self, backend) -> None:
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.insert(mcp_oauth_pending),
|
||||
{
|
||||
"state": "rand-state-xyz",
|
||||
"user_id": "u1",
|
||||
"server_name": "srv-a",
|
||||
"code_verifier": "verifier-blob",
|
||||
"return_url": "/admin/mcp-servers",
|
||||
"created_at": "2026-05-04T11:00:00",
|
||||
},
|
||||
)
|
||||
conn.commit()
|
||||
row = conn.execute(
|
||||
sa.select(mcp_oauth_pending).where(mcp_oauth_pending.c.state == "rand-state-xyz")
|
||||
).one()
|
||||
assert row.user_id == "u1"
|
||||
assert row.server_name == "srv-a"
|
||||
assert row.return_url == "/admin/mcp-servers"
|
||||
|
||||
def test_state_pk_unique(self, backend) -> None:
|
||||
"""A second insert with the same state value raises IntegrityError."""
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.insert(mcp_oauth_pending),
|
||||
{
|
||||
"state": "dup-state",
|
||||
"user_id": "u1",
|
||||
"server_name": "srv-a",
|
||||
"code_verifier": "v",
|
||||
"return_url": "/x",
|
||||
"created_at": "2026-05-04T11:00:00",
|
||||
},
|
||||
)
|
||||
conn.commit()
|
||||
import pytest
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
with pytest.raises(IntegrityError), backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.insert(mcp_oauth_pending),
|
||||
{
|
||||
"state": "dup-state",
|
||||
"user_id": "u2",
|
||||
"server_name": "srv-b",
|
||||
"code_verifier": "v",
|
||||
"return_url": "/y",
|
||||
"created_at": "2026-05-04T11:01:00",
|
||||
},
|
||||
)
|
||||
conn.commit()
|
||||
@@ -0,0 +1,53 @@
|
||||
"""PKCE pair-generation tests for the MCP OAuth flow.
|
||||
|
||||
Verifies the contract documented in RFC 7636 §4.1 and §4.2:
|
||||
|
||||
- ``code_verifier`` is a high-entropy 43..128 character urlsafe-base64 string.
|
||||
- ``code_challenge`` is the BASE64URL-NO-PADDING encoding of
|
||||
``SHA256(verifier)``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import string
|
||||
|
||||
from turnstone.core.mcp_oauth import generate_pkce_pair
|
||||
|
||||
_URLSAFE_CHARS = set(string.ascii_letters + string.digits + "-_")
|
||||
|
||||
|
||||
class TestGeneratePkcePair:
|
||||
def test_returns_tuple_of_strings(self) -> None:
|
||||
verifier, challenge = generate_pkce_pair()
|
||||
assert isinstance(verifier, str)
|
||||
assert isinstance(challenge, str)
|
||||
|
||||
def test_verifier_length_in_rfc_range(self) -> None:
|
||||
for _ in range(20):
|
||||
verifier, _ = generate_pkce_pair()
|
||||
assert 43 <= len(verifier) <= 128
|
||||
|
||||
def test_verifier_is_urlsafe(self) -> None:
|
||||
for _ in range(20):
|
||||
verifier, _ = generate_pkce_pair()
|
||||
assert all(ch in _URLSAFE_CHARS for ch in verifier)
|
||||
|
||||
def test_challenge_matches_sha256_of_verifier(self) -> None:
|
||||
for _ in range(20):
|
||||
verifier, challenge = generate_pkce_pair()
|
||||
digest = hashlib.sha256(verifier.encode("ascii")).digest()
|
||||
expected = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
|
||||
assert challenge == expected
|
||||
|
||||
def test_challenge_has_no_padding(self) -> None:
|
||||
for _ in range(20):
|
||||
_, challenge = generate_pkce_pair()
|
||||
assert "=" not in challenge
|
||||
|
||||
def test_pairs_are_unique(self) -> None:
|
||||
pairs = {generate_pkce_pair() for _ in range(50)}
|
||||
# 50 random draws shouldn't collide; if they do we have a much
|
||||
# bigger problem than this assertion.
|
||||
assert len(pairs) == 50
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,399 @@
|
||||
"""Tests for :func:`turnstone.core.mcp_oauth.revoke_token_at_as`.
|
||||
|
||||
The helper is best-effort RFC 7009 token revocation. It must:
|
||||
- skip cleanly when the AS metadata doesn't advertise a revocation endpoint
|
||||
- POST the form body when one is present (with optional client_secret)
|
||||
- never raise on non-2xx, network errors, or timeouts — caller doesn't
|
||||
want try/except in cleanup paths
|
||||
- never use ``exc_info=True`` — chained ``__context__`` may carry an
|
||||
``httpx.Request`` whose ``Authorization`` header holds a bearer; the
|
||||
bearer-leak invariant requires structured fields with type names only
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
from turnstone.core.mcp_oauth import (
|
||||
ASMetadata,
|
||||
MCPOAuthDiscoveryError,
|
||||
_attempt_upstream_revoke,
|
||||
revoke_token_at_as,
|
||||
)
|
||||
|
||||
|
||||
def _make_as_metadata(
|
||||
*,
|
||||
revocation_endpoint: str | None = "https://as.example.com/revoke",
|
||||
) -> ASMetadata:
|
||||
return ASMetadata(
|
||||
issuer="https://as.example.com",
|
||||
authorization_endpoint="https://as.example.com/authorize",
|
||||
token_endpoint="https://as.example.com/token",
|
||||
registration_endpoint=None,
|
||||
revocation_endpoint=revocation_endpoint,
|
||||
jwks_uri=None,
|
||||
code_challenge_methods_supported=("S256",),
|
||||
token_endpoint_auth_methods_supported=("client_secret_basic",),
|
||||
)
|
||||
|
||||
|
||||
def _mk_response(status_code: int) -> MagicMock:
|
||||
resp = MagicMock(spec=httpx.Response)
|
||||
resp.status_code = status_code
|
||||
return resp
|
||||
|
||||
|
||||
class TestRevocationUnsupported:
|
||||
def test_revoke_token_skipped_when_revocation_endpoint_none(self) -> None:
|
||||
as_meta = _make_as_metadata(revocation_endpoint=None)
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock()
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
client.post.assert_not_called()
|
||||
info_events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_unsupported" in info_events
|
||||
|
||||
|
||||
class TestRevocationSuccess:
|
||||
def test_revoke_token_succeeds_on_200(self) -> None:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(return_value=_mk_response(200))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret="s-secret",
|
||||
)
|
||||
)
|
||||
|
||||
# POST shape — URL + form body keys.
|
||||
client.post.assert_awaited_once()
|
||||
call_args = client.post.call_args
|
||||
assert call_args.args[0] == "https://as.example.com/revoke"
|
||||
body = call_args.kwargs["data"]
|
||||
assert body == {
|
||||
"token": "r-secret",
|
||||
"token_type_hint": "refresh_token",
|
||||
"client_id": "client-1",
|
||||
"client_secret": "s-secret",
|
||||
}
|
||||
info_events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_succeeded" in info_events
|
||||
|
||||
def test_revoke_token_omits_client_secret_when_none(self) -> None:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(return_value=_mk_response(200))
|
||||
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
body = client.post.call_args.kwargs["data"]
|
||||
assert "client_secret" not in body
|
||||
assert body["token"] == "r-secret"
|
||||
assert body["token_type_hint"] == "refresh_token"
|
||||
assert body["client_id"] == "client-1"
|
||||
|
||||
def test_revoke_token_succeeds_on_204(self) -> None:
|
||||
# RFC 7009 says the AS MAY return any 2xx; treat the whole range
|
||||
# as success.
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(return_value=_mk_response(204))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
info_events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_succeeded" in info_events
|
||||
|
||||
|
||||
class TestRevocationFailureLogged:
|
||||
def _run_and_capture(self, status: int) -> list[Any]:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(return_value=_mk_response(status))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
return mock_log.info.call_args_list
|
||||
|
||||
def test_revoke_token_logs_on_400_does_not_raise(self) -> None:
|
||||
calls = self._run_and_capture(400)
|
||||
events = [c.args[0] for c in calls]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
# Must include status field.
|
||||
failed_call = next(c for c in calls if c.args[0] == "mcp_server.oauth.revocation_failed")
|
||||
assert failed_call.kwargs.get("status") == 400
|
||||
|
||||
def test_revoke_token_logs_on_401_does_not_raise(self) -> None:
|
||||
calls = self._run_and_capture(401)
|
||||
events = [c.args[0] for c in calls]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
failed_call = next(c for c in calls if c.args[0] == "mcp_server.oauth.revocation_failed")
|
||||
assert failed_call.kwargs.get("status") == 401
|
||||
|
||||
def test_revoke_token_logs_on_403_does_not_raise(self) -> None:
|
||||
calls = self._run_and_capture(403)
|
||||
events = [c.args[0] for c in calls]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
|
||||
def test_revoke_token_logs_on_5xx_does_not_raise(self) -> None:
|
||||
calls = self._run_and_capture(500)
|
||||
events = [c.args[0] for c in calls]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
failed_call = next(c for c in calls if c.args[0] == "mcp_server.oauth.revocation_failed")
|
||||
assert failed_call.kwargs.get("status") == 500
|
||||
|
||||
|
||||
class TestRevocationExceptionPaths:
|
||||
def test_revoke_token_handles_network_error(self) -> None:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(side_effect=httpx.ConnectError("boom"))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
failed_call = next(
|
||||
c
|
||||
for c in mock_log.info.call_args_list
|
||||
if c.args[0] == "mcp_server.oauth.revocation_failed"
|
||||
)
|
||||
assert failed_call.kwargs.get("error") == "ConnectError"
|
||||
|
||||
def test_revoke_token_handles_httpx_timeout(self) -> None:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(side_effect=httpx.TimeoutException("slow"))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
failed_call = next(
|
||||
c
|
||||
for c in mock_log.info.call_args_list
|
||||
if c.args[0] == "mcp_server.oauth.revocation_failed"
|
||||
)
|
||||
assert failed_call.kwargs.get("error") == "TimeoutException"
|
||||
|
||||
def test_revoke_token_handles_asyncio_timeout(self) -> None:
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
|
||||
async def _slow(*_args: Any, **_kwargs: Any) -> Any:
|
||||
await asyncio.sleep(10.0)
|
||||
raise AssertionError("should have timed out")
|
||||
|
||||
client.post = AsyncMock(side_effect=_slow)
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
timeout_seconds=0.05,
|
||||
)
|
||||
)
|
||||
|
||||
events = [c.args[0] for c in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.revocation_failed" in events
|
||||
failed_call = next(
|
||||
c
|
||||
for c in mock_log.info.call_args_list
|
||||
if c.args[0] == "mcp_server.oauth.revocation_failed"
|
||||
)
|
||||
# ``asyncio.timeout`` raises ``TimeoutError`` (Python's builtin)
|
||||
# on cancellation.
|
||||
assert failed_call.kwargs.get("error") == "TimeoutError"
|
||||
|
||||
def test_revoke_token_no_exc_info_in_logs(self) -> None:
|
||||
"""Bearer-leak invariant: the revoke path must NEVER set
|
||||
``exc_info=True``. Chained ``__context__`` may include an
|
||||
``httpx.Request`` whose ``Authorization`` header holds a
|
||||
bearer; the traceback formatter would render it.
|
||||
"""
|
||||
as_meta = _make_as_metadata()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
client.post = AsyncMock(side_effect=httpx.ConnectError("boom"))
|
||||
|
||||
with patch("turnstone.core.mcp_oauth.log") as mock_log:
|
||||
asyncio.run(
|
||||
revoke_token_at_as(
|
||||
as_metadata=as_meta,
|
||||
http_client=client,
|
||||
refresh_token="r-secret",
|
||||
client_id="client-1",
|
||||
client_secret=None,
|
||||
)
|
||||
)
|
||||
|
||||
# No info call may carry exc_info.
|
||||
for call in mock_log.info.call_args_list:
|
||||
assert "exc_info" not in call.kwargs, (
|
||||
f"mcp_server.oauth log info({call.args[0]!r}) used exc_info — "
|
||||
"this violates the bearer-leak invariant"
|
||||
)
|
||||
# Defensively: also check warning + exception levels for the
|
||||
# same call site.
|
||||
for call in mock_log.warning.call_args_list:
|
||||
assert "exc_info" not in call.kwargs
|
||||
mock_log.exception.assert_not_called()
|
||||
|
||||
|
||||
class TestAttemptUpstreamRevokeNeverRaises:
|
||||
"""Round-2 q-3 regression: ``_attempt_upstream_revoke``'s docstring
|
||||
claims ``Never raises``. Background-task semantics make this load-
|
||||
bearing — a propagated exception logs ``Task exception was never
|
||||
retrieved`` because the ``set.discard`` done-callback doesn't read
|
||||
``task.exception()``.
|
||||
|
||||
The wrapper's narrow inner ``except`` clauses (``MCPOAuthDiscoveryError``,
|
||||
``MCPTokenDecryptError``) leave room for any other exception type
|
||||
raised by ``discover_authorization_server`` /
|
||||
``storage.get_mcp_oauth_client_secret_ct`` / ``token_store.cipher.decrypt``
|
||||
to escape. The outer ``try/except Exception`` is what keeps the
|
||||
contract honest. These tests pin that gate.
|
||||
"""
|
||||
|
||||
def _build_args(self) -> dict[str, Any]:
|
||||
token_store = MagicMock()
|
||||
token_store.cipher = MagicMock()
|
||||
token_store.cipher.decrypt.return_value = b"shh"
|
||||
storage = MagicMock()
|
||||
storage.get_mcp_oauth_client_secret_ct.return_value = None
|
||||
return {
|
||||
"http_client": MagicMock(spec=httpx.AsyncClient),
|
||||
"metadata_cache": None,
|
||||
"storage": storage,
|
||||
"token_store": token_store,
|
||||
"server_name": "srv-oauth",
|
||||
"server_row": {
|
||||
"url": "https://mcp.example.com",
|
||||
"oauth_client_id": "client-1",
|
||||
"oauth_authorization_server_url": None,
|
||||
"oauth_as_issuer_cached": None,
|
||||
},
|
||||
"server_id_for_audit": "srv-id-1",
|
||||
"refresh_token": "r-secret",
|
||||
}
|
||||
|
||||
def test_attempt_upstream_revoke_swallows_unexpected_exception(self) -> None:
|
||||
"""A generic exception from a path the inner handlers don't
|
||||
cover MUST be caught at the outer boundary and logged with type
|
||||
name only (no exc_info=True per the bearer-leak invariant).
|
||||
"""
|
||||
args = self._build_args()
|
||||
|
||||
async def _boom(*_a: Any, **_kw: Any) -> Any:
|
||||
raise RuntimeError("network blew up")
|
||||
|
||||
with (
|
||||
patch("turnstone.core.mcp_oauth.discover_authorization_server", side_effect=_boom),
|
||||
patch("turnstone.core.mcp_oauth.log") as mock_log,
|
||||
):
|
||||
# MUST NOT raise.
|
||||
asyncio.run(_attempt_upstream_revoke(**args))
|
||||
|
||||
events = [call.args[0] for call in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.upstream_revoke_failed" in events, (
|
||||
"outer try/except must log mcp_server.oauth.upstream_revoke_failed "
|
||||
"with the exception type name when an unexpected exception escapes "
|
||||
"the narrow inner handlers"
|
||||
)
|
||||
for call in mock_log.info.call_args_list:
|
||||
assert "exc_info" not in call.kwargs, (
|
||||
"outer-block log must not use exc_info=True — chained "
|
||||
"__context__ may carry an httpx.Request bearer"
|
||||
)
|
||||
|
||||
def test_attempt_upstream_revoke_logs_discovery_failure(self) -> None:
|
||||
"""Round-2 bug-1: ``MCPOAuthDiscoveryError`` MUST emit
|
||||
``upstream_revoke_discovery_failed`` so operators have visibility
|
||||
into a silent-discovery-failure path that previously logged
|
||||
nothing while the audit row recorded ``upstream_revoke_outcome=scheduled``.
|
||||
"""
|
||||
args = self._build_args()
|
||||
|
||||
async def _disc_fail(*_a: Any, **_kw: Any) -> Any:
|
||||
raise MCPOAuthDiscoveryError("PRM fetch 503")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_oauth.discover_authorization_server",
|
||||
side_effect=_disc_fail,
|
||||
),
|
||||
patch("turnstone.core.mcp_oauth.log") as mock_log,
|
||||
):
|
||||
asyncio.run(_attempt_upstream_revoke(**args))
|
||||
|
||||
events = [call.args[0] for call in mock_log.info.call_args_list]
|
||||
assert "mcp_server.oauth.upstream_revoke_discovery_failed" in events
|
||||
assert "mcp_server.oauth.upstream_revoke_failed" not in events
|
||||
@@ -0,0 +1,212 @@
|
||||
"""Storage CRUD tests for the per-(user, server) MCP OAuth pending-state table.
|
||||
|
||||
Validates the storage-protocol additions for the per-(user, server)
|
||||
OAuth flow:
|
||||
|
||||
- ``create_mcp_oauth_pending_state``
|
||||
- ``pop_mcp_oauth_pending_state`` (atomic, with TTL)
|
||||
- ``cleanup_expired_mcp_oauth_pending_states``
|
||||
- ``get_mcp_oauth_client_secret_ct``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
class TestCreateAndPop:
|
||||
def test_round_trip(self, backend) -> None:
|
||||
backend.create_mcp_oauth_pending_state(
|
||||
"state-1",
|
||||
"user-a",
|
||||
"srv-x",
|
||||
"verifier-blob",
|
||||
"/admin/mcp-servers",
|
||||
)
|
||||
row = backend.pop_mcp_oauth_pending_state("state-1", max_age_seconds=600)
|
||||
assert row is not None
|
||||
assert row["state"] == "state-1"
|
||||
assert row["user_id"] == "user-a"
|
||||
assert row["server_name"] == "srv-x"
|
||||
assert row["code_verifier"] == "verifier-blob"
|
||||
assert row["return_url"] == "/admin/mcp-servers"
|
||||
|
||||
def test_pop_consumes_row(self, backend) -> None:
|
||||
backend.create_mcp_oauth_pending_state("s2", "u", "s", "v", "/r")
|
||||
first = backend.pop_mcp_oauth_pending_state("s2")
|
||||
assert first is not None
|
||||
# Second pop must miss — row was consumed.
|
||||
second = backend.pop_mcp_oauth_pending_state("s2")
|
||||
assert second is None
|
||||
|
||||
def test_pop_missing_returns_none(self, backend) -> None:
|
||||
assert backend.pop_mcp_oauth_pending_state("never-existed") is None
|
||||
|
||||
|
||||
class TestTTL:
|
||||
def test_pop_rejects_expired_row(self, backend) -> None:
|
||||
backend.create_mcp_oauth_pending_state("old-state", "u", "s", "v", "/r")
|
||||
# Backdate it so it's older than the TTL window.
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE mcp_oauth_pending SET created_at = '2020-01-01T00:00:00' "
|
||||
"WHERE state = 'old-state'"
|
||||
)
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# Default TTL is 600s — the row is decades old.
|
||||
row = backend.pop_mcp_oauth_pending_state("old-state")
|
||||
assert row is None
|
||||
|
||||
# Even though pop returned None, the row must have been wiped — a
|
||||
# second pop with a giant TTL must still see nothing.
|
||||
again = backend.pop_mcp_oauth_pending_state("old-state", max_age_seconds=10**9)
|
||||
assert again is None
|
||||
|
||||
def test_pop_accepts_fresh_row(self, backend) -> None:
|
||||
backend.create_mcp_oauth_pending_state("fresh", "u", "s", "v", "/r")
|
||||
row = backend.pop_mcp_oauth_pending_state("fresh", max_age_seconds=600)
|
||||
assert row is not None
|
||||
assert row["state"] == "fresh"
|
||||
|
||||
|
||||
class TestCleanup:
|
||||
def test_cleanup_deletes_only_expired(self, backend) -> None:
|
||||
backend.create_mcp_oauth_pending_state("old", "u", "s", "v", "/r")
|
||||
backend.create_mcp_oauth_pending_state("new", "u", "s", "v", "/r")
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE mcp_oauth_pending SET created_at = '2020-01-01T00:00:00' "
|
||||
"WHERE state = 'old'"
|
||||
)
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
deleted = backend.cleanup_expired_mcp_oauth_pending_states(max_age_seconds=600)
|
||||
assert deleted == 1
|
||||
# Old gone, new still around.
|
||||
assert backend.pop_mcp_oauth_pending_state("old") is None
|
||||
survivor = backend.pop_mcp_oauth_pending_state("new")
|
||||
assert survivor is not None
|
||||
|
||||
def test_cleanup_no_rows(self, backend) -> None:
|
||||
assert backend.cleanup_expired_mcp_oauth_pending_states() == 0
|
||||
|
||||
|
||||
class TestGetOAuthClientSecretCt:
|
||||
def test_returns_none_when_unset(self, backend) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-id",
|
||||
name="srv-x",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
assert backend.get_mcp_oauth_client_secret_ct("srv-id") is None
|
||||
|
||||
def test_returns_ciphertext_after_set(self, backend) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-id",
|
||||
name="srv-x",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
ct = b"\x00\xff\x42encrypted-blob"
|
||||
ok = backend.set_mcp_oauth_client_secret_ct("srv-id", ct)
|
||||
assert ok is True
|
||||
out = backend.get_mcp_oauth_client_secret_ct("srv-id")
|
||||
assert out == ct
|
||||
|
||||
def test_returns_none_for_missing_server(self, backend) -> None:
|
||||
assert backend.get_mcp_oauth_client_secret_ct("does-not-exist") is None
|
||||
|
||||
|
||||
def _create_user_token_row(
|
||||
backend,
|
||||
*,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
created: str,
|
||||
) -> None:
|
||||
"""Insert a token row + backdate ``created`` so ordering is deterministic.
|
||||
|
||||
The storage helper stamps ``created`` from ``datetime.now(UTC)``; for
|
||||
multi-row ordering tests we backdate via raw SQL so the inserts stay
|
||||
independent of clock resolution.
|
||||
"""
|
||||
backend.create_mcp_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token_ct=b"ct-access",
|
||||
refresh_token_ct=b"ct-refresh",
|
||||
expires_at="2026-05-04T12:00:00",
|
||||
scopes="openid",
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE mcp_user_tokens SET created = :created "
|
||||
"WHERE user_id = :uid AND server_name = :sn"
|
||||
),
|
||||
{"created": created, "uid": user_id, "sn": server_name},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
|
||||
class TestListMCPUserTokenMetadataByUser:
|
||||
def test_list_mcp_user_token_metadata_by_user_empty(self, backend) -> None:
|
||||
assert backend.list_mcp_user_token_metadata_by_user("nobody") == []
|
||||
|
||||
def test_list_mcp_user_token_metadata_by_user_single_server(self, backend) -> None:
|
||||
_create_user_token_row(
|
||||
backend, user_id="u1", server_name="srv-a", created="2026-05-01T00:00:00"
|
||||
)
|
||||
rows = backend.list_mcp_user_token_metadata_by_user("u1")
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["user_id"] == "u1"
|
||||
assert rows[0]["server_name"] == "srv-a"
|
||||
assert rows[0]["as_issuer"] == "https://auth.example.com"
|
||||
assert rows[0]["audience"] == "https://mcp.example.com"
|
||||
assert rows[0]["scopes"] == "openid"
|
||||
# Projection MUST omit ciphertext columns — the SQL no longer
|
||||
# selects them, so the TypedDict has no key.
|
||||
assert "access_token_ct" not in rows[0]
|
||||
assert "refresh_token_ct" not in rows[0]
|
||||
|
||||
def test_list_mcp_user_token_metadata_by_user_multiple_servers(self, backend) -> None:
|
||||
_create_user_token_row(
|
||||
backend, user_id="u1", server_name="srv-c", created="2026-05-03T00:00:00"
|
||||
)
|
||||
_create_user_token_row(
|
||||
backend, user_id="u1", server_name="srv-a", created="2026-05-01T00:00:00"
|
||||
)
|
||||
_create_user_token_row(
|
||||
backend, user_id="u1", server_name="srv-b", created="2026-05-02T00:00:00"
|
||||
)
|
||||
rows = backend.list_mcp_user_token_metadata_by_user("u1")
|
||||
assert [r["server_name"] for r in rows] == ["srv-a", "srv-b", "srv-c"]
|
||||
|
||||
def test_list_mcp_user_token_metadata_by_user_isolates_by_user(self, backend) -> None:
|
||||
_create_user_token_row(
|
||||
backend, user_id="user-a", server_name="srv-a", created="2026-05-01T00:00:00"
|
||||
)
|
||||
_create_user_token_row(
|
||||
backend, user_id="user-a", server_name="srv-b", created="2026-05-02T00:00:00"
|
||||
)
|
||||
_create_user_token_row(
|
||||
backend, user_id="user-b", server_name="srv-a", created="2026-05-03T00:00:00"
|
||||
)
|
||||
rows_a = backend.list_mcp_user_token_metadata_by_user("user-a")
|
||||
assert {r["server_name"] for r in rows_a} == {"srv-a", "srv-b"}
|
||||
assert all(r["user_id"] == "user-a" for r in rows_a)
|
||||
|
||||
rows_b = backend.list_mcp_user_token_metadata_by_user("user-b")
|
||||
assert len(rows_b) == 1
|
||||
assert rows_b[0]["user_id"] == "user-b"
|
||||
assert rows_b[0]["server_name"] == "srv-a"
|
||||
@@ -0,0 +1,304 @@
|
||||
"""Boundary tests for the Phase 9 pending-consent write path.
|
||||
|
||||
Drives ``MCPClientManager._dispatch_pool_sync`` (and the helper it
|
||||
calls, ``_record_pending_consent_best_effort``) and asserts that
|
||||
deferred-consent records reach storage only on non-interactive callers.
|
||||
|
||||
Per ``feedback_tests_through_boundaries.md``, at least one test must
|
||||
drive the real sync dispatcher → real ``_is_structured_error`` →
|
||||
real ``_record_pending_consent_best_effort`` plumb-through; the
|
||||
``_helpers`` unit tests below cover the classifier in isolation, but
|
||||
the end-to-end test is the structural gate that catches
|
||||
plumb-through regressions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import threading
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.mcp_client import (
|
||||
_PENDING_CONSENT_PERSIST_CODES,
|
||||
MCPClientManager,
|
||||
_parse_pending_consent_envelope,
|
||||
)
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.mcp_oauth import TokenLookupResult
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper-level unit tests (cheap, no event loop)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseEnvelope:
|
||||
def test_consent_required_no_scopes(self) -> None:
|
||||
env = json.dumps({"error": {"code": "mcp_consent_required", "server": "x", "detail": "d"}})
|
||||
assert _parse_pending_consent_envelope(env) == ("mcp_consent_required", None)
|
||||
|
||||
def test_insufficient_scope_with_scopes(self) -> None:
|
||||
env = json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"code": "mcp_insufficient_scope",
|
||||
"server": "x",
|
||||
"detail": "d",
|
||||
"scopes_required": ["read", "write"],
|
||||
}
|
||||
}
|
||||
)
|
||||
assert _parse_pending_consent_envelope(env) == (
|
||||
"mcp_insufficient_scope",
|
||||
["read", "write"],
|
||||
)
|
||||
|
||||
def test_operator_codes_filtered(self) -> None:
|
||||
# Key-unknown / url-insecure / *_forbidden are operator-actionable,
|
||||
# NOT user-consent-shaped. They must not produce pending-consent
|
||||
# rows, regardless of whether the caller is interactive.
|
||||
for code in (
|
||||
"mcp_token_undecryptable_key_unknown",
|
||||
"mcp_oauth_url_insecure",
|
||||
"mcp_tool_call_forbidden",
|
||||
"mcp_resource_read_forbidden",
|
||||
"mcp_prompt_get_forbidden",
|
||||
):
|
||||
env = json.dumps({"error": {"code": code, "server": "x", "detail": "d"}})
|
||||
assert _parse_pending_consent_envelope(env) is None, code
|
||||
|
||||
def test_malformed_json_returns_none(self) -> None:
|
||||
assert _parse_pending_consent_envelope("not json") is None
|
||||
assert _parse_pending_consent_envelope("") is None
|
||||
|
||||
def test_persist_codes_set_is_expected(self) -> None:
|
||||
# Pin the contract — adding a new persistable code here is a
|
||||
# deliberate design decision and should require a test update.
|
||||
assert {
|
||||
"mcp_consent_required",
|
||||
"mcp_insufficient_scope",
|
||||
} == _PENDING_CONSENT_PERSIST_CODES
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# End-to-end plumb-through (drives _dispatch_pool_sync)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _seed_oauth_server(backend: Any, *, name: str = "pool-srv") -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-" + name,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
command="",
|
||||
args="[]",
|
||||
url="https://example.com/mcp",
|
||||
headers="{}",
|
||||
env="{}",
|
||||
auto_approve=False,
|
||||
enabled=True,
|
||||
created_by="admin",
|
||||
)
|
||||
backend.update_mcp_server("srv-" + name, auth_type="oauth_user")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def running_loop_mgr():
|
||||
cfg: dict[str, Any] = {}
|
||||
mgr = MCPClientManager(cfg)
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True, name="phase9-test-loop")
|
||||
thread.start()
|
||||
mgr._loop = loop
|
||||
try:
|
||||
yield mgr, loop, thread
|
||||
finally:
|
||||
|
||||
async def _drain(m: MCPClientManager) -> None:
|
||||
task = m._user_pool_eviction_task
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
m._user_pool_eviction_task = None
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(_drain(mgr), loop).result(timeout=2)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
|
||||
|
||||
def _wire_mgr(mgr: MCPClientManager, backend: Any) -> None:
|
||||
cipher = make_mcp_token_cipher()
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
app_state = SimpleNamespace(
|
||||
auth_storage=backend,
|
||||
mcp_token_store=MCPTokenStore(backend, cipher, node_id="test"),
|
||||
mcp_oauth_http_client=MagicMock(),
|
||||
mcp_oauth_refresh_locks={},
|
||||
mcp_oauth_metadata_cache={},
|
||||
)
|
||||
mgr.set_storage(backend)
|
||||
mgr.set_app_state(app_state)
|
||||
|
||||
|
||||
def test_dispatch_persists_pending_for_non_interactive_caller(
|
||||
running_loop_mgr: Any, backend: Any
|
||||
) -> None:
|
||||
"""Non-interactive caller hits ``mcp_consent_required`` → a
|
||||
``mcp_pending_consent`` row appears for ``(user_id, server_name)``."""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
_seed_oauth_server(backend)
|
||||
_wire_mgr(mgr, backend)
|
||||
|
||||
async def _missing_token(**kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_missing_token,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo",
|
||||
{"payload": "hi"},
|
||||
user_id="user-a",
|
||||
timeout=10,
|
||||
is_interactive_for_consent=False,
|
||||
)
|
||||
|
||||
# Structured error envelope surfaces as RuntimeError to the caller.
|
||||
payload = json.loads(str(exc_info.value)).get("error", {})
|
||||
assert payload.get("code") == "mcp_consent_required"
|
||||
|
||||
# Persistent row written for the dashboard badge.
|
||||
rows = backend.list_mcp_pending_consent_by_user("user-a")
|
||||
assert len(rows) == 1
|
||||
r = rows[0]
|
||||
assert r["user_id"] == "user-a"
|
||||
assert r["server_name"] == "pool-srv"
|
||||
assert r["error_code"] == "mcp_consent_required"
|
||||
assert r["occurrence_count"] == 1
|
||||
|
||||
|
||||
def test_dispatch_does_not_persist_for_interactive_caller(
|
||||
running_loop_mgr: Any, backend: Any
|
||||
) -> None:
|
||||
"""Interactive caller hits the same error path → NO row written.
|
||||
|
||||
Interactive (WEB / CLI) sessions surface the consent prompt in-flight
|
||||
via the Phase 8 SSE renderer; persisting would just produce
|
||||
immediately-stale dashboard badges.
|
||||
"""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
_seed_oauth_server(backend)
|
||||
_wire_mgr(mgr, backend)
|
||||
|
||||
async def _missing_token(**kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_missing_token,
|
||||
),
|
||||
pytest.raises(RuntimeError),
|
||||
):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo",
|
||||
{"payload": "hi"},
|
||||
user_id="user-a",
|
||||
timeout=10,
|
||||
is_interactive_for_consent=True,
|
||||
)
|
||||
|
||||
assert backend.list_mcp_pending_consent_by_user("user-a") == []
|
||||
|
||||
|
||||
def test_dispatch_returns_envelope_unchanged_on_storage_failure(
|
||||
running_loop_mgr: Any, backend: Any
|
||||
) -> None:
|
||||
"""When ``upsert_mcp_pending_consent`` raises, the agent-observable
|
||||
contract is unchanged: the structured-error ``RuntimeError`` still
|
||||
surfaces with the original ``mcp_consent_required`` code. The doc-
|
||||
string promises best-effort persistence; this test pins that
|
||||
promise so a regression that propagates the storage exception would
|
||||
fail visibly.
|
||||
"""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
_seed_oauth_server(backend)
|
||||
_wire_mgr(mgr, backend)
|
||||
|
||||
async def _missing_token(**kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
original_upsert = backend.upsert_mcp_pending_consent
|
||||
|
||||
def _raise(*_a: Any, **_kw: Any) -> None:
|
||||
raise RuntimeError("storage offline")
|
||||
|
||||
backend.upsert_mcp_pending_consent = _raise # type: ignore[method-assign]
|
||||
try:
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_missing_token,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo",
|
||||
{"payload": "hi"},
|
||||
user_id="user-a",
|
||||
timeout=10,
|
||||
is_interactive_for_consent=False,
|
||||
)
|
||||
finally:
|
||||
backend.upsert_mcp_pending_consent = original_upsert # type: ignore[method-assign]
|
||||
|
||||
payload = json.loads(str(exc_info.value)).get("error", {})
|
||||
assert payload.get("code") == "mcp_consent_required"
|
||||
|
||||
|
||||
def test_dispatch_does_not_persist_for_operator_actionable_code(
|
||||
running_loop_mgr: Any, backend: Any
|
||||
) -> None:
|
||||
"""Decrypt-failure → operator-actionable; even non-interactive callers
|
||||
must NOT produce a user-facing pending-consent record (the user can't
|
||||
resolve this by re-consenting).
|
||||
"""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
_seed_oauth_server(backend)
|
||||
_wire_mgr(mgr, backend)
|
||||
|
||||
async def _decrypt_failure(**kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="decrypt_failure")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_decrypt_failure,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo",
|
||||
{"payload": "hi"},
|
||||
user_id="user-a",
|
||||
timeout=10,
|
||||
is_interactive_for_consent=False,
|
||||
)
|
||||
|
||||
payload = json.loads(str(exc_info.value)).get("error", {})
|
||||
assert payload.get("code") == "mcp_token_undecryptable_key_unknown"
|
||||
# The operator-actionable code does NOT produce a pending-consent row.
|
||||
assert backend.list_mcp_pending_consent_by_user("user-a") == []
|
||||
@@ -0,0 +1,259 @@
|
||||
"""HTTP tests for the Phase 9 pending-consent endpoints.
|
||||
|
||||
Covers:
|
||||
- ``GET /v1/api/mcp/oauth/pending`` (install gate + read path)
|
||||
- ``DELETE /v1/api/mcp/oauth/pending/{server_name}`` (single clear)
|
||||
- ``DELETE /v1/api/mcp/oauth/pending`` (bulk clear)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.routing import Mount, Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from turnstone.core.auth import AuthResult
|
||||
from turnstone.core.mcp_oauth import (
|
||||
handle_mcp_oauth_clear_all_pending,
|
||||
handle_mcp_oauth_clear_pending,
|
||||
handle_mcp_oauth_list_pending,
|
||||
)
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
|
||||
class _InjectAuthMiddleware(BaseHTTPMiddleware):
|
||||
"""Stamp a fixed authenticated user on every request."""
|
||||
|
||||
def __init__(self, app: Any, user_id: str = "user-1") -> None:
|
||||
super().__init__(app)
|
||||
self._user_id = user_id
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id=self._user_id,
|
||||
scopes=frozenset({"write"}),
|
||||
token_source="config",
|
||||
permissions=frozenset({"read", "write"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _build_app(storage: SQLiteBackend, *, user_id: str = "user-1") -> Starlette:
|
||||
class _Mw(_InjectAuthMiddleware):
|
||||
def __init__(self, app: Any) -> None:
|
||||
super().__init__(app, user_id=user_id)
|
||||
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route("/api/mcp/oauth/pending", handle_mcp_oauth_list_pending),
|
||||
Route(
|
||||
"/api/mcp/oauth/pending",
|
||||
handle_mcp_oauth_clear_all_pending,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
Route(
|
||||
"/api/mcp/oauth/pending/{server_name}",
|
||||
handle_mcp_oauth_clear_pending,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
middleware=[Middleware(_Mw)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
return app
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
backend = SQLiteBackend(str(tmp_path / "test.db"))
|
||||
backend.create_user("user-1", "user1", "User One", "hash")
|
||||
backend.create_user("user-2", "user2", "User Two", "hash")
|
||||
return backend
|
||||
|
||||
|
||||
def _seed_oauth_server(backend: SQLiteBackend, *, name: str = "srv-x") -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-id-" + name,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url="https://example.com/mcp",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
|
||||
|
||||
def _seed_pending(
|
||||
backend: SQLiteBackend,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "srv-x",
|
||||
error_code: str = "mcp_consent_required",
|
||||
now_iso: str = "2026-05-11T12:00:00",
|
||||
) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
error_code=error_code,
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso=now_iso,
|
||||
)
|
||||
|
||||
|
||||
class TestListPending:
|
||||
def test_install_gate_short_circuits_on_no_oauth_servers(self, storage: SQLiteBackend) -> None:
|
||||
# Seed a pending row but NO oauth_user MCP server — the gate
|
||||
# must short-circuit to {pending: 0} regardless.
|
||||
_seed_pending(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.get("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"pending": 0, "servers": []}
|
||||
|
||||
def test_lists_pending_records_for_authenticated_user(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_pending(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.get("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["pending"] == 1
|
||||
assert len(body["servers"]) == 1
|
||||
assert body["servers"][0]["server_name"] == "srv-x"
|
||||
assert body["servers"][0]["error_code"] == "mcp_consent_required"
|
||||
|
||||
def test_does_not_leak_cross_user_records(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_pending(storage, user_id="user-2")
|
||||
client = TestClient(_build_app(storage, user_id="user-1"))
|
||||
resp = client.get("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"pending": 0, "servers": []}
|
||||
|
||||
|
||||
class TestClearPending:
|
||||
def test_delete_single(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_pending(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending/srv-x")
|
||||
assert resp.status_code == 204
|
||||
assert storage.list_mcp_pending_consent_by_user("user-1") == []
|
||||
|
||||
def test_delete_missing_still_returns_204(self, storage: SQLiteBackend) -> None:
|
||||
# Idempotent — must not leak cross-user existence info via 404.
|
||||
_seed_oauth_server(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending/never-existed")
|
||||
assert resp.status_code == 204
|
||||
|
||||
def test_delete_does_not_touch_cross_user_rows(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_pending(storage, user_id="user-1")
|
||||
_seed_pending(storage, user_id="user-2")
|
||||
client = TestClient(_build_app(storage, user_id="user-1"))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending/srv-x")
|
||||
assert resp.status_code == 204
|
||||
# User-2's row survives.
|
||||
assert len(storage.list_mcp_pending_consent_by_user("user-2")) == 1
|
||||
|
||||
|
||||
class TestAuditTrail:
|
||||
def test_single_dismiss_audits(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_pending(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending/srv-x")
|
||||
assert resp.status_code == 204
|
||||
|
||||
events = storage.list_audit_events(limit=10)
|
||||
rows = [
|
||||
e for e in events if e.get("action") == "mcp_server.oauth.pending_consent_dismissed"
|
||||
]
|
||||
assert len(rows) == 1
|
||||
detail = rows[0].get("detail")
|
||||
if isinstance(detail, str):
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(detail)
|
||||
assert detail.get("mode") == "single"
|
||||
assert detail.get("cleared") == 1
|
||||
|
||||
def test_single_dismiss_audits_even_when_no_row_existed(self, storage: SQLiteBackend) -> None:
|
||||
# Cross-tenant non-observability requires a 204 in the never-existed
|
||||
# case — the audit row distinguishes a real dismiss from a stuffed
|
||||
# attempt by recording ``cleared=0``.
|
||||
_seed_oauth_server(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending/never-existed")
|
||||
assert resp.status_code == 204
|
||||
|
||||
events = storage.list_audit_events(limit=10)
|
||||
rows = [
|
||||
e for e in events if e.get("action") == "mcp_server.oauth.pending_consent_dismissed"
|
||||
]
|
||||
assert len(rows) == 1
|
||||
detail = rows[0].get("detail")
|
||||
if isinstance(detail, str):
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(detail)
|
||||
assert detail.get("mode") == "single"
|
||||
assert detail.get("cleared") == 0
|
||||
|
||||
def test_bulk_dismiss_audits(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_oauth_server(storage, name="srv-y")
|
||||
_seed_pending(storage, server_name="srv-x")
|
||||
_seed_pending(storage, server_name="srv-y")
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"cleared": 2}
|
||||
|
||||
events = storage.list_audit_events(limit=10)
|
||||
rows = [
|
||||
e for e in events if e.get("action") == "mcp_server.oauth.pending_consent_dismissed"
|
||||
]
|
||||
assert len(rows) == 1
|
||||
detail = rows[0].get("detail")
|
||||
if isinstance(detail, str):
|
||||
import json as _json
|
||||
|
||||
detail = _json.loads(detail)
|
||||
assert detail.get("mode") == "bulk"
|
||||
assert detail.get("cleared") == 2
|
||||
|
||||
|
||||
class TestClearAllPending:
|
||||
def test_bulk_clear(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
_seed_oauth_server(storage, name="srv-y")
|
||||
_seed_pending(storage, server_name="srv-x")
|
||||
_seed_pending(storage, server_name="srv-y")
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"cleared": 2}
|
||||
assert storage.list_mcp_pending_consent_by_user("user-1") == []
|
||||
|
||||
def test_bulk_clear_zero_when_empty(self, storage: SQLiteBackend) -> None:
|
||||
_seed_oauth_server(storage)
|
||||
client = TestClient(_build_app(storage))
|
||||
resp = client.delete("/v1/api/mcp/oauth/pending")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"cleared": 0}
|
||||
@@ -0,0 +1,263 @@
|
||||
"""Storage CRUD tests for the Phase 9 ``mcp_pending_consent`` table.
|
||||
|
||||
Validates protocol additions backing the dashboard pending-consent badge:
|
||||
|
||||
- ``upsert_mcp_pending_consent`` — insert + on-conflict refresh
|
||||
- ``list_mcp_pending_consent_by_user`` — read path
|
||||
- ``delete_mcp_pending_consent`` — single-row clear
|
||||
- ``delete_all_mcp_pending_consent_by_user`` — bulk clear
|
||||
- ``count_mcp_consented_users_by_server`` — admin status pill
|
||||
- ``any_oauth_user_mcp_servers`` — install-level gate
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def _iso(ts: str = "2026-05-11T12:00:00") -> str:
|
||||
return ts
|
||||
|
||||
|
||||
class TestUpsertAndList:
|
||||
def test_insert_round_trip(self, backend) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-x",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required="read write",
|
||||
last_ws_id="ws-1",
|
||||
last_tool_call_id="tool-1",
|
||||
now_iso=_iso(),
|
||||
)
|
||||
rows = backend.list_mcp_pending_consent_by_user("user-a")
|
||||
assert len(rows) == 1
|
||||
r = rows[0]
|
||||
assert r["user_id"] == "user-a"
|
||||
assert r["server_name"] == "srv-x"
|
||||
assert r["error_code"] == "mcp_consent_required"
|
||||
assert r["scopes_required"] == "read write"
|
||||
assert r["last_ws_id"] == "ws-1"
|
||||
assert r["last_tool_call_id"] == "tool-1"
|
||||
assert r["occurrence_count"] == 1
|
||||
assert r["first_seen_at"] == r["last_seen_at"]
|
||||
|
||||
def test_upsert_bumps_count_and_refreshes_recency(self, backend) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-x",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso="2026-05-11T12:00:00",
|
||||
)
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-x",
|
||||
error_code="mcp_insufficient_scope",
|
||||
scopes_required="read",
|
||||
last_ws_id="ws-2",
|
||||
last_tool_call_id="tool-2",
|
||||
now_iso="2026-05-11T13:00:00",
|
||||
)
|
||||
rows = backend.list_mcp_pending_consent_by_user("user-a")
|
||||
assert len(rows) == 1
|
||||
r = rows[0]
|
||||
# Recency fields refreshed to the second call's values; count bumped.
|
||||
assert r["occurrence_count"] == 2
|
||||
assert r["error_code"] == "mcp_insufficient_scope"
|
||||
assert r["scopes_required"] == "read"
|
||||
assert r["last_ws_id"] == "ws-2"
|
||||
assert r["last_tool_call_id"] == "tool-2"
|
||||
assert r["last_seen_at"] == "2026-05-11T13:00:00"
|
||||
# first_seen_at preserved — that's the load-bearing audit value.
|
||||
assert r["first_seen_at"] == "2026-05-11T12:00:00"
|
||||
|
||||
def test_list_orders_by_last_seen_desc(self, backend) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-old",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso="2026-05-11T10:00:00",
|
||||
)
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-new",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso="2026-05-11T11:00:00",
|
||||
)
|
||||
rows = backend.list_mcp_pending_consent_by_user("user-a")
|
||||
assert [r["server_name"] for r in rows] == ["srv-new", "srv-old"]
|
||||
|
||||
def test_per_user_isolation(self, backend) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso=_iso(),
|
||||
)
|
||||
assert backend.list_mcp_pending_consent_by_user("user-b") == []
|
||||
|
||||
|
||||
class TestDelete:
|
||||
def test_delete_single(self, backend) -> None:
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name="srv-x",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso=_iso(),
|
||||
)
|
||||
assert backend.delete_mcp_pending_consent("user-a", "srv-x") is True
|
||||
assert backend.list_mcp_pending_consent_by_user("user-a") == []
|
||||
# Second delete returns False (no row).
|
||||
assert backend.delete_mcp_pending_consent("user-a", "srv-x") is False
|
||||
|
||||
def test_delete_missing_returns_false(self, backend) -> None:
|
||||
assert backend.delete_mcp_pending_consent("never", "missing") is False
|
||||
|
||||
def test_delete_all_by_user(self, backend) -> None:
|
||||
for name in ("srv-a", "srv-b", "srv-c"):
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-a",
|
||||
server_name=name,
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso=_iso(),
|
||||
)
|
||||
# Cross-user row that must NOT be touched.
|
||||
backend.upsert_mcp_pending_consent(
|
||||
user_id="user-b",
|
||||
server_name="srv-z",
|
||||
error_code="mcp_consent_required",
|
||||
scopes_required=None,
|
||||
last_ws_id=None,
|
||||
last_tool_call_id=None,
|
||||
now_iso=_iso(),
|
||||
)
|
||||
assert backend.delete_all_mcp_pending_consent_by_user("user-a") == 3
|
||||
assert backend.list_mcp_pending_consent_by_user("user-a") == []
|
||||
assert len(backend.list_mcp_pending_consent_by_user("user-b")) == 1
|
||||
|
||||
|
||||
class TestCountConsentedUsersByServer:
|
||||
def _seed_server(self, backend, name: str = "srv-x") -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-id-" + name,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
command="",
|
||||
args="[]",
|
||||
url="https://example.com/mcp",
|
||||
headers="{}",
|
||||
env="{}",
|
||||
auto_approve=False,
|
||||
enabled=True,
|
||||
created_by="admin",
|
||||
)
|
||||
backend.update_mcp_server("srv-id-" + name, auth_type="oauth_user")
|
||||
|
||||
def test_counts_distinct_non_expired_users(self, backend) -> None:
|
||||
self._seed_server(backend)
|
||||
future = "2099-01-01T00:00:00"
|
||||
backend.create_mcp_user_token(
|
||||
"alice",
|
||||
"srv-x",
|
||||
access_token_ct=b"ct",
|
||||
refresh_token_ct=None,
|
||||
expires_at=future,
|
||||
scopes=None,
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://example.com/mcp",
|
||||
)
|
||||
backend.create_mcp_user_token(
|
||||
"bob",
|
||||
"srv-x",
|
||||
access_token_ct=b"ct",
|
||||
refresh_token_ct=None,
|
||||
expires_at=None, # null treated as non-expired
|
||||
scopes=None,
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://example.com/mcp",
|
||||
)
|
||||
# Different server — must not count.
|
||||
self._seed_server(backend, name="srv-y")
|
||||
backend.create_mcp_user_token(
|
||||
"carol",
|
||||
"srv-y",
|
||||
access_token_ct=b"ct",
|
||||
refresh_token_ct=None,
|
||||
expires_at=future,
|
||||
scopes=None,
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://example.com/mcp",
|
||||
)
|
||||
assert backend.count_mcp_consented_users_by_server("srv-x") == 2
|
||||
assert backend.count_mcp_consented_users_by_server("srv-y") == 1
|
||||
|
||||
def test_excludes_expired(self, backend) -> None:
|
||||
self._seed_server(backend)
|
||||
backend.create_mcp_user_token(
|
||||
"alice",
|
||||
"srv-x",
|
||||
access_token_ct=b"ct",
|
||||
refresh_token_ct=None,
|
||||
expires_at="2020-01-01T00:00:00", # well in the past
|
||||
scopes=None,
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://example.com/mcp",
|
||||
)
|
||||
assert backend.count_mcp_consented_users_by_server("srv-x") == 0
|
||||
|
||||
def test_zero_when_no_rows(self, backend) -> None:
|
||||
assert backend.count_mcp_consented_users_by_server("missing") == 0
|
||||
|
||||
|
||||
class TestInstallGate:
|
||||
def test_any_oauth_user_returns_false_on_empty(self, backend) -> None:
|
||||
assert backend.any_oauth_user_mcp_servers() is False
|
||||
|
||||
def test_any_oauth_user_ignores_static_rows(self, backend) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-1",
|
||||
name="static-only",
|
||||
transport="streamable-http",
|
||||
command="",
|
||||
args="[]",
|
||||
url="https://example.com",
|
||||
headers='{"Authorization": "Bearer x"}',
|
||||
env="{}",
|
||||
auto_approve=False,
|
||||
enabled=True,
|
||||
created_by="admin",
|
||||
)
|
||||
assert backend.any_oauth_user_mcp_servers() is False
|
||||
|
||||
def test_any_oauth_user_returns_true_when_one_exists(self, backend) -> None:
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-2",
|
||||
name="oauth-srv",
|
||||
transport="streamable-http",
|
||||
command="",
|
||||
args="[]",
|
||||
url="https://example.com",
|
||||
headers="{}",
|
||||
env="{}",
|
||||
auto_approve=False,
|
||||
enabled=True,
|
||||
created_by="admin",
|
||||
)
|
||||
backend.update_mcp_server("srv-2", auth_type="oauth_user")
|
||||
assert backend.any_oauth_user_mcp_servers() is True
|
||||
@@ -0,0 +1,912 @@
|
||||
"""Phase 6 integration tests — real-transport drives 401/403 through the SDK.
|
||||
|
||||
These are the structural exit criterion for Phase 6. They MUST drive
|
||||
through the real ``streamablehttp_client``, the real httpx response-hook
|
||||
path, and a REAL upstream MCP server (a ``FastMCP`` in-process subprocess
|
||||
with a starlette middleware that programmatically returns 401/403 with
|
||||
crafted ``WWW-Authenticate`` headers).
|
||||
|
||||
Direct ``httpx.HTTPStatusError`` injection is FORBIDDEN here — Phase 5
|
||||
bug-1 was masked precisely by that pattern (the production code path
|
||||
was structurally unreachable, but the unit-test injection bypassed the
|
||||
SDK's swallow). The integration tests gate that the production path
|
||||
actually receives the carrier signal end-to-end.
|
||||
|
||||
The fixture upstream is built in-thread (uvicorn on its own asyncio
|
||||
loop in a background thread) — same pattern as
|
||||
``tests/spike_sdk_concurrency.py``. Per the orchestrator's startup-cost
|
||||
note, measured at ~0.05s per fixture spin-up locally; well under the
|
||||
2s threshold for default-collection inclusion.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import uvicorn
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.mcp_oauth import TokenLookupResult
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
# Quiet noisy logs during tests.
|
||||
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("mcp").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixture upstream — programmable BehaviorMiddleware
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BehaviorMiddleware(BaseHTTPMiddleware):
|
||||
"""Inspects per-request behaviour state and returns 401/403 on demand.
|
||||
|
||||
The behaviour is steered by a mutable ``behaviour`` dict on the
|
||||
middleware instance; tests mutate it via the fixture handle.
|
||||
Records every request's Authorization header for assertion.
|
||||
|
||||
Behaviour semantics:
|
||||
* ``"once_401"``: return 401 once, then 200 thereafter.
|
||||
* ``"always_401"``: always return 401.
|
||||
* ``"once_403_insufficient"``: return 403 with insufficient_scope once.
|
||||
* ``"once_403_generic"``: return 403 without error param once.
|
||||
* ``"once_multi_www_auth_403"``: return 403 with TWO
|
||||
``WWW-Authenticate`` headers — first ``Bearer`` challenge
|
||||
carries the SAFE scopes, second carries INJECTED scopes. The
|
||||
dispatcher must report only the first.
|
||||
* ``"never"`` (default): pass through to the real handler.
|
||||
|
||||
``www_authenticate`` overrides the default header crafted per shape.
|
||||
"""
|
||||
|
||||
def __init__(self, app: Any, behaviour: dict[str, Any]) -> None:
|
||||
super().__init__(app)
|
||||
self._behaviour = behaviour
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable[..., Any]) -> Response:
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
|
||||
# Record the Authorization header for assertion. POST is the
|
||||
# tools/call request the dispatcher sends.
|
||||
if request.method == "POST" and "/mcp" in str(request.url):
|
||||
self._behaviour.setdefault("post_auth_headers", []).append(
|
||||
request.headers.get("authorization")
|
||||
)
|
||||
|
||||
mode = self._behaviour.get("mode", "never")
|
||||
if mode == "once_401":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "always_401":
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_insufficient":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate",
|
||||
'Bearer error="insufficient_scope", scope="files:write mail:send"',
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_generic" and not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get("www_authenticate", "Bearer realm=mcp")
|
||||
},
|
||||
)
|
||||
elif mode == "once_multi_www_auth_403" and not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
# Two ``WWW-Authenticate: Bearer ...`` challenges. The
|
||||
# first carries ``error=insufficient_scope`` but NO
|
||||
# ``scope=`` parameter; the second carries the INJECTED
|
||||
# scopes the dispatcher must NOT report. The first
|
||||
# challenge intentionally lacks ``scope`` because
|
||||
# ``parse_www_authenticate_bearer`` uses ``setdefault`` —
|
||||
# if the first challenge HAD a scope, ``setdefault`` would
|
||||
# already win on first-occurrence. The vector this test
|
||||
# guards is the case where a defended absence becomes a
|
||||
# silent presence: a hook regression to ``get(...)`` joins
|
||||
# repeated headers with ``, `` and the parser then folds
|
||||
# the second challenge's scope into the first challenge's
|
||||
# params dict because there is no first-occurrence to
|
||||
# protect.
|
||||
response = StarletteResponse("forbidden", status_code=403)
|
||||
response.headers.append(
|
||||
"www-authenticate",
|
||||
'Bearer realm="legit", error="insufficient_scope"',
|
||||
)
|
||||
response.headers.append(
|
||||
"www-authenticate",
|
||||
'Bearer error="insufficient_scope", scope="org:admin db:write"',
|
||||
)
|
||||
return response
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _find_free_port() -> int:
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def _build_server(port: int, behaviour: dict[str, Any]) -> uvicorn.Server:
|
||||
mcp = FastMCP(name="phase6-target", streamable_http_path="/mcp")
|
||||
|
||||
@mcp.tool()
|
||||
async def echo(payload: str = "default") -> str:
|
||||
return f"echoed:{payload}"
|
||||
|
||||
app = mcp.streamable_http_app()
|
||||
app.add_middleware(BehaviorMiddleware, behaviour=behaviour)
|
||||
config = uvicorn.Config(app, host="127.0.0.1", port=port, log_level="warning", access_log=False)
|
||||
return uvicorn.Server(config)
|
||||
|
||||
|
||||
def _wait_ready(port: int, timeout: float = 5.0) -> None:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=0.5):
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(0.05)
|
||||
raise TimeoutError(f"upstream at 127.0.0.1:{port} not ready after {timeout}s")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def upstream():
|
||||
"""Boot a FastMCP fixture upstream in a background thread.
|
||||
|
||||
Yields ``(url, behaviour)`` where ``behaviour`` is a mutable dict
|
||||
the test mutates to steer the middleware (set ``mode`` to one of
|
||||
the BehaviorMiddleware shapes).
|
||||
"""
|
||||
port = _find_free_port()
|
||||
behaviour: dict[str, Any] = {}
|
||||
server = _build_server(port, behaviour)
|
||||
|
||||
def _run() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.serve())
|
||||
|
||||
t = threading.Thread(target=_run, daemon=True, name="phase6-upstream")
|
||||
t.start()
|
||||
try:
|
||||
_wait_ready(port)
|
||||
yield f"http://127.0.0.1:{port}/mcp", behaviour
|
||||
finally:
|
||||
server.should_exit = True
|
||||
t.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _seed_oauth_server(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
name: str = "pool-srv",
|
||||
server_id: str = "srv-pool",
|
||||
url: str = "https://mcp.example.com/sse",
|
||||
) -> None:
|
||||
storage.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url=url,
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="client-abc",
|
||||
oauth_scopes="openid",
|
||||
oauth_audience=url,
|
||||
)
|
||||
|
||||
|
||||
def _seed_user_token(
|
||||
storage: SQLiteBackend,
|
||||
cipher: Any,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "pool-srv",
|
||||
expires_in_seconds: int = 3600,
|
||||
access_token: str = "access-aaa",
|
||||
refresh_token: str | None = "refresh-rrr",
|
||||
) -> None:
|
||||
expires_at = (datetime.now(UTC) + timedelta(seconds=expires_in_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
)
|
||||
store = MCPTokenStore(storage, cipher, node_id="test")
|
||||
store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token=access_token,
|
||||
refresh_token=refresh_token,
|
||||
expires_at=expires_at,
|
||||
scopes="openid",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
|
||||
def _make_app_state(storage: SQLiteBackend, *, cipher: Any) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
auth_storage=storage,
|
||||
mcp_token_store=MCPTokenStore(storage, cipher, node_id="test"),
|
||||
mcp_oauth_http_client=MagicMock(),
|
||||
mcp_oauth_refresh_locks={},
|
||||
mcp_oauth_metadata_cache={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def running_loop_mgr():
|
||||
cfg: dict[str, Any] = {}
|
||||
mgr = MCPClientManager(cfg)
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True, name="mcp-pool-test-loop")
|
||||
thread.start()
|
||||
mgr._loop = loop
|
||||
try:
|
||||
yield mgr, loop, thread
|
||||
finally:
|
||||
|
||||
async def _drain(m: MCPClientManager) -> None:
|
||||
task = m._user_pool_eviction_task
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
m._user_pool_eviction_task = None
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(_drain(mgr), loop).result(timeout=2)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 21: 401 → refresh-and-retry → success
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_401_refresh_and_retry_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Real upstream returns 401 once with ``WWW-Authenticate: Bearer
|
||||
error="invalid_token"``, then 200. Dispatcher carrier captures the
|
||||
401, ``force_refresh=True`` mints a new bearer (stubbed), retry
|
||||
succeeds. Hard invariant 3: breaker counter remains 0.
|
||||
|
||||
Drives through the REAL ``streamablehttp_client`` and a REAL
|
||||
upstream subprocess (no ``httpx.HTTPStatusError`` injection). This
|
||||
is the structural exit gate for Phase 6 — the equivalent unit
|
||||
tests CANNOT prove the production wiring works because the SDK
|
||||
swallows the underlying exception.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
# Override URL to point at the local upstream (loopback http:// is
|
||||
# exempt from the URL-validator).
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "hi"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
assert "echoed:hi" in result
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
# Server saw at least 2 POSTs to /mcp (initial + retry).
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) >= 2, f"expected >=2 POSTs; got {len(post_headers)}"
|
||||
# Retry carries a different bearer than the initial.
|
||||
initial = post_headers[0]
|
||||
retry = post_headers[1]
|
||||
assert initial != retry, (
|
||||
"retry attached the same bearer as the initial; the dispatcher "
|
||||
"did not pick up the refreshed token."
|
||||
)
|
||||
# Pool entry has a session after the successful retry.
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 22: 401 + refresh failure → mcp_consent_required
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_401_with_refresh_failure_emits_consent_required(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="refresh_failed")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync("mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15)
|
||||
|
||||
# Structured-error envelopes flow back via ``RuntimeError(json_str)``
|
||||
# so the session-layer ``except Exception`` handler routes the
|
||||
# consent card uniformly across tool / resource / prompt dispatchers.
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert payload["error"]["server"] == "pool-srv"
|
||||
# Phase 8 — consent_url surfaces a /start URL the dashboard can open
|
||||
# in a popup. URL-encoded server name; no scopes baked in (the AS
|
||||
# picks up the configured scopes server-side at /start).
|
||||
assert payload["error"]["consent_url"] == "/v1/api/mcp/oauth/start?server=pool-srv"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 23: 403 + insufficient_scope → mcp_insufficient_scope with parsed scopes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_403_insufficient_scope_emits_structured_error(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_insufficient"
|
||||
behaviour["www_authenticate"] = (
|
||||
'Bearer error="insufficient_scope", scope="files:write mail:send"'
|
||||
)
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync("mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_insufficient_scope"
|
||||
assert payload["error"]["scopes_required"] == ["files:write", "mail:send"]
|
||||
# Phase 8 — consent_url carries the step-up scopes URL-encoded so the
|
||||
# dashboard can union them with the configured set at /start.
|
||||
assert payload["error"]["consent_url"] == (
|
||||
"/v1/api/mcp/oauth/start?server=pool-srv&scopes=files%3Awrite%20mail%3Asend"
|
||||
)
|
||||
# No retry — exactly ONE POST attempted before the structured error.
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) == 1, (
|
||||
f"403 must NOT trigger a retry; observed {len(post_headers)} POSTs"
|
||||
)
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 24: 403 without insufficient_scope → generic forbidden
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_403_no_insufficient_scope_emits_generic_forbidden(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_generic"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync("mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_tool_call_forbidden"
|
||||
assert "scopes_required" not in payload["error"]
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) == 1, (
|
||||
f"403 must NOT trigger a retry; observed {len(post_headers)} POSTs"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# sec-1: multi-WWW-Authenticate header injection — only the FIRST
|
||||
# Bearer challenge feeds the structured-error / audit emission.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_403_multi_www_authenticate_drops_injected_scopes(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Upstream returns a 403 with TWO ``WWW-Authenticate: Bearer ...``
|
||||
challenges. The first carries ``error=insufficient_scope`` but NO
|
||||
``scope=`` parameter; the second carries INJECTED scopes
|
||||
(``["org:admin", "db:write"]``). The dispatcher must report
|
||||
``scopes_required == []`` — derived from the first challenge alone
|
||||
— never the second challenge's injected scopes.
|
||||
|
||||
Two layers of defence cooperate (either alone neutralises the
|
||||
vector; both run together so a regression in one cannot silently
|
||||
re-open it):
|
||||
|
||||
1. ``_make_capturing_http_factory._hook`` reads
|
||||
``response.headers.get_list("www-authenticate")[0]`` rather than
|
||||
``response.headers.get(...)`` — the latter joins repeated
|
||||
headers with ``", "`` which the RFC 7235 tokenizer would
|
||||
otherwise consume as a continuation of the first challenge.
|
||||
2. ``parse_www_authenticate_bearer`` stops at the first ``Bearer``
|
||||
challenge boundary even if the input was already joined, so a
|
||||
hook regression to ``get(...)`` would NOT re-open the vector.
|
||||
|
||||
The first challenge intentionally lacks ``scope=`` — the parser
|
||||
uses ``setdefault`` so a first-occurrence ``scope`` would already
|
||||
win and mask a single-layer regression. The undefended-absence
|
||||
case is what proves both layers actually do their job.
|
||||
|
||||
Negative-test (CRITICAL — Phase 5 lesson): verified by reverting
|
||||
the hook to ``response.headers.get("www-authenticate")`` AND
|
||||
removing the ``_looks_like_bearer_challenge_start`` guard in
|
||||
``parse_www_authenticate_bearer``. The test then fails because
|
||||
``scopes_required`` becomes ``["org:admin", "db:write"]`` — the
|
||||
injected scopes from the second challenge silently fold into the
|
||||
first challenge's params dict via httpx's comma-joined header
|
||||
value (the absence of a first-occurrence scope means nothing
|
||||
blocks the fold).
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_multi_www_auth_403"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync("mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_insufficient_scope", (
|
||||
f"expected mcp_insufficient_scope; got {payload!r}"
|
||||
)
|
||||
# ``scopes_required`` derives from the FIRST challenge alone, which
|
||||
# carries no ``scope=`` parameter. The injected second challenge
|
||||
# MUST NOT appear here.
|
||||
assert payload["error"]["scopes_required"] == [], (
|
||||
"Multi-header injection slipped through: dispatcher reported "
|
||||
"scopes from the SECOND Bearer challenge. Got "
|
||||
f"{payload['error']['scopes_required']!r}; expected []."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 25: 401 retry ceiling — never recurse
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_401_retry_ceiling(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Upstream always returns 401; refresh stub keeps minting tokens.
|
||||
After exactly ONE retry, dispatcher emits ``mcp_consent_required``.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
refresh_count = 0
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
nonlocal refresh_count
|
||||
if kwargs.get("force_refresh"):
|
||||
refresh_count += 1
|
||||
return TokenLookupResult(kind="token", token=f"refreshed-{refresh_count}")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.call_tool_sync("mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
# Exactly ONE refresh round-trip.
|
||||
assert refresh_count == 1, f"expected exactly 1 refresh round-trip; got {refresh_count}"
|
||||
# Server saw EXACTLY 2 POSTs (initial + 1 retry).
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) == 2, (
|
||||
f"expected exactly 2 POSTs (initial + 1 retry); got {len(post_headers)}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 26: breaker unaffected by repeated auth failures (slow — 50 cycles)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_breaker_unaffected_by_auth_failures(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""50 sequential dispatches all hit 401 with refresh-failed → 50
|
||||
cycles of ``mcp_consent_required``. ``_consecutive_failures`` MUST
|
||||
stay at 0 throughout (hard invariant 3 verified end-to-end).
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="refresh_failed")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
for _ in range(50):
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 27: static path unaffected by Phase 6 changes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_static_path_unaffected(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Static-path connect against an unauthed upstream succeeds without
|
||||
going through the capturing factory. This is the integration-level
|
||||
mirror of ``test_reconnect_preserves_static_state_identity``.
|
||||
|
||||
Drives the static path against the same fixture upstream (with
|
||||
``behaviour={}`` so middleware passes through) — confirms the
|
||||
static path's session lifecycle is byte-identical even when the
|
||||
pool path's auth introspection is wired up.
|
||||
"""
|
||||
url, _behaviour = upstream
|
||||
# No mode → middleware passes through to FastMCP.
|
||||
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
|
||||
# Manually configure mgr with a static-path server pointing at the
|
||||
# fixture upstream. Use _connect_one (not the pool path).
|
||||
cfg = {"type": "streamable-http", "url": url}
|
||||
|
||||
async def _connect_static() -> None:
|
||||
await mgr._connect_one("static-srv", cfg)
|
||||
|
||||
fut = asyncio.run_coroutine_threadsafe(_connect_static(), loop)
|
||||
fut.result(timeout=15)
|
||||
|
||||
state_before = mgr._static_servers.get("static-srv")
|
||||
assert state_before is not None
|
||||
assert state_before.session is not None
|
||||
# Snapshot identity.
|
||||
state_id_before = id(state_before)
|
||||
session_before = state_before.session
|
||||
|
||||
# Reconnect — the canonical regression check is that the
|
||||
# StaticServerState object identity is preserved.
|
||||
fut = asyncio.run_coroutine_threadsafe(_connect_static(), loop)
|
||||
fut.result(timeout=15)
|
||||
|
||||
state_after = mgr._static_servers.get("static-srv")
|
||||
assert state_after is not None
|
||||
assert id(state_after) == state_id_before, (
|
||||
"Static path StaticServerState identity changed across reconnect; "
|
||||
"hard invariant 1 violated."
|
||||
)
|
||||
assert state_after.session is not None
|
||||
assert state_after.session is not session_before, (
|
||||
"Reconnect did not actually replace the session"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 27b: static dispatch unaffected by Phase 8 consent_url kwarg
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_static_dispatch_unaffected_by_consent_url_kwarg(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Static-auth tool dispatch must be byte-identical post-Phase 8.
|
||||
|
||||
The Phase 8 changes only ADD a ``consent_url`` kwarg to
|
||||
``_structured_error`` invocations on the pool path. Static dispatch
|
||||
must not pick up the field — there's no consent flow for
|
||||
``auth_type='none'`` / ``'static'`` servers, and exposing one would
|
||||
confuse the dashboard renderer. Asserts a successful tool result is
|
||||
a plain string with no JSON envelope and no ``consent_url`` substring.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "never" # passthrough — succeeds
|
||||
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cfg = {"type": "streamable-http", "url": url}
|
||||
|
||||
async def _connect_static() -> None:
|
||||
await mgr._connect_one("static-srv", cfg)
|
||||
|
||||
fut = asyncio.run_coroutine_threadsafe(_connect_static(), loop)
|
||||
fut.result(timeout=15)
|
||||
|
||||
# Drive call_tool_sync without a user_id — the static path is taken.
|
||||
result = mgr.call_tool_sync("mcp__static-srv__echo", {"payload": "static-x"}, timeout=15)
|
||||
|
||||
# Static path returns the FastMCP fixture's echo string.
|
||||
assert "echoed:static-x" in result
|
||||
# No JSON envelope leaked through; specifically no consent_url field.
|
||||
assert "consent_url" not in result, (
|
||||
f"Static-auth tool dispatch surfaced a consent_url; result: {result!r}"
|
||||
)
|
||||
# Defensive: result is not a JSON-encoded structured error.
|
||||
try:
|
||||
parsed = json.loads(result)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
parsed = None
|
||||
if isinstance(parsed, dict):
|
||||
assert "error" not in parsed, (
|
||||
f"Static-auth dispatch returned a structured-error envelope; got {parsed!r}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 28: pool reuse — 401 on a SECOND dispatch (carrier owned by entry)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_integration_pool_reuse_401_refresh_and_retry_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Reused pool sessions still capture 401 correctly.
|
||||
|
||||
Dispatch 1 hits a passthrough upstream (200) and populates
|
||||
``entry.session``. Dispatch 2 reuses that session — no fresh
|
||||
connect, so a per-dispatch ``_AuthCapture`` would never reach
|
||||
the httpx response hook (the hook closes over the carrier passed
|
||||
at first connect, which lives on the entry). A correctly-wired
|
||||
entry-owned carrier is the only shape that lets dispatch 2's 401
|
||||
surface to the dispatcher.
|
||||
|
||||
Two independent production bugs gate this test passing; both must
|
||||
hold for reused-session 401 recovery to work end-to-end:
|
||||
|
||||
1. The carrier must live on the pool entry (not per-dispatch) so
|
||||
the response hook bound at first connect writes to the same
|
||||
object the dispatcher reads across reuse. Verified by reverting
|
||||
``PoolEntryState.auth_capture`` to a per-dispatch
|
||||
``_AuthCapture()`` allocation: the carrier-fired event never
|
||||
reaches the dispatcher and the test times out.
|
||||
|
||||
2. The dispatcher must race ``call_tool`` against the carrier's
|
||||
fired event. The SDK's ``_receive_loop`` runs in BaseSession's
|
||||
TaskGroup nested inside ``streamablehttp_client``'s TaskGroup;
|
||||
when an upstream 4xx fires, the outer TaskGroup cancels
|
||||
``_receive_loop`` mid-finally before it can deliver
|
||||
``CONNECTION_CLOSED`` to the response stream's waiting
|
||||
receiver. anyio's ``send_nowait`` skips waiters with pending
|
||||
cancellation — but our dispatch task (created via
|
||||
``run_coroutine_threadsafe`` for the reused-session case) has
|
||||
NO pending cancellation, so the send delivers but the receiver
|
||||
never wakes (the waiter's Event is set on stale state). Result:
|
||||
a forever-hung ``response_stream_reader.receive()``. Verified
|
||||
by reverting the ``asyncio.wait({call_task, fired_task})``
|
||||
race in ``_dispatch_pool_with_entry`` to a bare ``await
|
||||
session.call_tool(...)``: the test times out.
|
||||
|
||||
This test is the structural gate against the per-dispatch carrier
|
||||
pattern: it looks right in code review and passes single-dispatch
|
||||
integration tests, but breaks silently on session reuse — and the
|
||||
SDK-level hang the carrier fix exposes silently strands the
|
||||
dispatcher even when the carrier is correct.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "never" # passthrough — dispatch 1 succeeds
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
# Dispatch 1: passthrough success. Establishes the pooled session.
|
||||
result1 = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "first"}, user_id="user-1", timeout=15
|
||||
)
|
||||
assert "echoed:first" in result1
|
||||
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
session_after_first = entry.session
|
||||
assert session_after_first is not None, (
|
||||
"test setup: dispatch 1 did not populate entry.session; "
|
||||
"subsequent dispatch will not exercise the reuse path"
|
||||
)
|
||||
|
||||
# Reconfigure upstream to 401 once on the next call. Reset the
|
||||
# auth-headers log so we can count dispatch-2's POSTs cleanly.
|
||||
behaviour["post_auth_headers"] = []
|
||||
behaviour["mode"] = "once_401"
|
||||
behaviour["_fired"] = False
|
||||
|
||||
# Dispatch 2: same (user, server). The hook from dispatch 1's
|
||||
# connect is still bound to entry.auth_capture. The 401 fires;
|
||||
# the dispatcher's auth_401 path triggers refresh-and-retry.
|
||||
result2 = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "second"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
assert "echoed:second" in result2, (
|
||||
f"reused-session 401 retry did not succeed. result: {result2!r}. "
|
||||
"If this is JSON with mcp_consent_required, the dispatcher "
|
||||
"fell through to consent_required emission; if a generic "
|
||||
"tool error, the carrier was empty (auth branch unreachable)."
|
||||
)
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0, (
|
||||
"auth failures must not trip the per-server breaker"
|
||||
)
|
||||
|
||||
# Dispatch 2 produces multiple POSTs: the original 401 with the
|
||||
# rejected bearer, then the retry's full connect handshake
|
||||
# (initialize + notifications/initialized + tools/list) followed by
|
||||
# the actual tools/call — all under the refreshed bearer. The retry
|
||||
# reconnects because the auth_401 handler evicted the broken session.
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) >= 2, (
|
||||
f"expected >=2 POSTs after dispatch 2 (401 + retry); "
|
||||
f"got {len(post_headers)}: {post_headers}"
|
||||
)
|
||||
# First POST is the original bearer that got 401'd.
|
||||
assert post_headers[0] == "Bearer access-aaa", (
|
||||
f"first POST was {post_headers[0]!r}; expected the original bearer"
|
||||
)
|
||||
# Every subsequent POST carries the refreshed bearer (the retry
|
||||
# ran with force_refresh=True and reconnected with the new token).
|
||||
refreshed = post_headers[1:]
|
||||
assert all(h == "Bearer refreshed-bearer" for h in refreshed), (
|
||||
f"retry POSTs carried unexpected bearer(s); observed: {post_headers}"
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,720 @@
|
||||
"""Phase 7b integration tests — real-transport prompt get 401/403/etc.
|
||||
|
||||
Mirror of :mod:`tests.test_mcp_pool_auth_resource_integration` for the
|
||||
prompt path (RFC §3.3). Drives through the real ``streamablehttp_client``,
|
||||
real httpx response-hook plumbing, and a real upstream subprocess
|
||||
(``FastMCP`` with a programmable ``BehaviorMiddleware``). Direct
|
||||
``httpx.HTTPStatusError`` injection is forbidden (invariant 14).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import uvicorn
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.mcp_oauth import TokenLookupResult
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("mcp").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
class BehaviorMiddleware(BaseHTTPMiddleware):
|
||||
"""Programmable upstream behaviour — see
|
||||
:mod:`tests.test_mcp_pool_auth_integration` for the semantics. This
|
||||
copy serves the prompt integration tests.
|
||||
"""
|
||||
|
||||
def __init__(self, app: Any, behaviour: dict[str, Any]) -> None:
|
||||
super().__init__(app)
|
||||
self._behaviour = behaviour
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable[..., Any]) -> Response:
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
|
||||
if request.method == "POST" and "/mcp" in str(request.url):
|
||||
self._behaviour.setdefault("post_auth_headers", []).append(
|
||||
request.headers.get("authorization")
|
||||
)
|
||||
|
||||
mode = self._behaviour.get("mode", "never")
|
||||
if mode == "once_401":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "always_401":
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_insufficient":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate",
|
||||
'Bearer error="insufficient_scope", scope="prompts:read"',
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_generic" and not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get("www_authenticate", "Bearer realm=mcp")
|
||||
},
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _find_free_port() -> int:
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def _build_server(port: int, behaviour: dict[str, Any]) -> uvicorn.Server:
|
||||
mcp = FastMCP(name="phase7b-prompt-target", streamable_http_path="/mcp")
|
||||
|
||||
@mcp.prompt()
|
||||
def greet(who: str = "world") -> str:
|
||||
return f"Hello, {who}!"
|
||||
|
||||
@mcp.prompt()
|
||||
def summarize(topic: str = "today") -> str:
|
||||
return f"Please summarize {topic}."
|
||||
|
||||
app = mcp.streamable_http_app()
|
||||
app.add_middleware(BehaviorMiddleware, behaviour=behaviour)
|
||||
config = uvicorn.Config(app, host="127.0.0.1", port=port, log_level="warning", access_log=False)
|
||||
return uvicorn.Server(config)
|
||||
|
||||
|
||||
def _wait_ready(port: int, timeout: float = 5.0) -> None:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=0.5):
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(0.05)
|
||||
raise TimeoutError(f"upstream at 127.0.0.1:{port} not ready after {timeout}s")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def upstream():
|
||||
port = _find_free_port()
|
||||
behaviour: dict[str, Any] = {}
|
||||
server = _build_server(port, behaviour)
|
||||
|
||||
def _run() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.serve())
|
||||
|
||||
t = threading.Thread(target=_run, daemon=True, name="phase7b-prompt-upstream")
|
||||
t.start()
|
||||
try:
|
||||
_wait_ready(port)
|
||||
yield f"http://127.0.0.1:{port}/mcp", behaviour
|
||||
finally:
|
||||
server.should_exit = True
|
||||
t.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _seed_oauth_server(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
name: str = "pool-srv",
|
||||
server_id: str = "srv-pool",
|
||||
url: str = "https://mcp.example.com/sse",
|
||||
) -> None:
|
||||
storage.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url=url,
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="client-abc",
|
||||
oauth_scopes="openid",
|
||||
oauth_audience=url,
|
||||
)
|
||||
|
||||
|
||||
def _seed_user_token(
|
||||
storage: SQLiteBackend,
|
||||
cipher: Any,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "pool-srv",
|
||||
expires_in_seconds: int = 3600,
|
||||
access_token: str = "access-aaa",
|
||||
refresh_token: str | None = "refresh-rrr",
|
||||
) -> None:
|
||||
expires_at = (datetime.now(UTC) + timedelta(seconds=expires_in_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
)
|
||||
store = MCPTokenStore(storage, cipher, node_id="test")
|
||||
store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token=access_token,
|
||||
refresh_token=refresh_token,
|
||||
expires_at=expires_at,
|
||||
scopes="openid",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
|
||||
def _make_app_state(storage: SQLiteBackend, *, cipher: Any) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
auth_storage=storage,
|
||||
mcp_token_store=MCPTokenStore(storage, cipher, node_id="test"),
|
||||
mcp_oauth_http_client=MagicMock(),
|
||||
mcp_oauth_refresh_locks={},
|
||||
mcp_oauth_metadata_cache={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def running_loop_mgr():
|
||||
cfg: dict[str, Any] = {}
|
||||
mgr = MCPClientManager(cfg)
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True, name="mcp-pool-test-loop")
|
||||
thread.start()
|
||||
mgr._loop = loop
|
||||
try:
|
||||
yield mgr, loop, thread
|
||||
finally:
|
||||
|
||||
async def _drain(m: MCPClientManager) -> None:
|
||||
task = m._user_pool_eviction_task
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
m._user_pool_eviction_task = None
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(_drain(mgr), loop).result(timeout=2)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
|
||||
|
||||
def _seed_pool_prompt_map(
|
||||
mgr: MCPClientManager,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
prefixed_name: str,
|
||||
original_name: str,
|
||||
) -> None:
|
||||
"""Pre-seed ``_user_prompt_map`` so ``_resolve_pool_target_prompt``
|
||||
finds the prefixed name. Production wires this through
|
||||
``_connect_one_pool``; the integration tests seed it directly so the
|
||||
test focuses on the dispatch behaviour after resolution succeeds.
|
||||
"""
|
||||
|
||||
async def _seed() -> None:
|
||||
entry = await mgr._ensure_pool_entry((user_id, server_name))
|
||||
entry.prompts = [
|
||||
{
|
||||
"name": prefixed_name,
|
||||
"original_name": original_name,
|
||||
"server": server_name,
|
||||
"description": "",
|
||||
"arguments": [],
|
||||
}
|
||||
]
|
||||
mgr._rebuild_user_prompt_map(user_id)
|
||||
|
||||
assert mgr._loop is not None
|
||||
asyncio.run_coroutine_threadsafe(_seed(), mgr._loop).result(timeout=5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-1: 401 → refresh → retry → success (prompt path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_401_refresh_and_retry_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Real upstream returns 401 once, then 200. Carrier captures 401,
|
||||
force_refresh=True mints a new bearer, retry returns the prompt
|
||||
messages. Hard invariant 3: breaker counter remains 0.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
messages = mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "everyone"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
assert isinstance(messages, list)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["role"] == "user"
|
||||
assert "everyone" in messages[0]["content"]
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) >= 2, f"expected >=2 POSTs; got {len(post_headers)}"
|
||||
assert post_headers[0] != post_headers[1], (
|
||||
"retry attached the same bearer as the initial; the dispatcher "
|
||||
"did not pick up the refreshed token."
|
||||
)
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-2: persistent 401 → mcp_consent_required (prompt path) → RuntimeError
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_persistent_401_emits_consent_required(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert payload["error"]["server"] == "pool-srv"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-3: 403 + insufficient_scope → mcp_insufficient_scope (prompt path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_403_insufficient_scope_emits_structured_error(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_insufficient"
|
||||
behaviour["www_authenticate"] = 'Bearer error="insufficient_scope", scope="prompts:read"'
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_insufficient_scope"
|
||||
assert payload["error"]["scopes_required"] == ["prompts:read"]
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) == 1, (
|
||||
f"403 must NOT trigger a retry; observed {len(post_headers)} POSTs"
|
||||
)
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-3b: 403 generic → mcp_prompt_get_forbidden
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_403_generic_forbidden(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_generic"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
# Per the kind="prompt" wiring of `_handle_auth_403`, the
|
||||
# operation-specific code surfaces here rather than the tool path's
|
||||
# generic mcp_tool_call_forbidden.
|
||||
assert payload["error"]["code"] == "mcp_prompt_get_forbidden"
|
||||
assert "scopes_required" not in payload["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-6: breaker isolation — auth failures NEVER trip the breaker
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_breaker_unaffected_by_auth_failures(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Repeated 401 + refresh-failed cycles leave breaker at 0
|
||||
(hard invariant 3 verified end-to-end for the prompt path)."""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="refresh_failed")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
# Re-seed each iteration: symmetric eviction (Phase 7b) clears
|
||||
# ``_user_prompt_map`` on auth failure so the next dispatch's
|
||||
# resolver would miss without a fresh seed. Production reconnect
|
||||
# repopulates this; the test simulates that out-of-band.
|
||||
for _ in range(10):
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Negative tests — token lookup edge cases (prompt path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_missing_token_emits_consent_required(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, _behaviour = upstream
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
|
||||
|
||||
def test_prompt_get_decrypt_failure_emits_token_undecryptable(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, _behaviour = upstream
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="decrypt_failure")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_token_undecryptable_key_unknown"
|
||||
|
||||
|
||||
def test_prompt_get_http_url_emits_url_insecure(
|
||||
running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""An ``http://`` (non-loopback) oauth_user URL must surface
|
||||
``mcp_oauth_url_insecure`` BEFORE the bearer is attached.
|
||||
"""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url="http://example.com/mcp")
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_prompt_map(mgr, "user-1", "pool-srv", "mcp__pool-srv__greet", "greet")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as excinfo,
|
||||
):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
|
||||
payload = json.loads(str(excinfo.value))
|
||||
assert payload["error"]["code"] == "mcp_oauth_url_insecure"
|
||||
|
||||
|
||||
def test_prompt_get_unknown_name_raises_value_error(
|
||||
running_loop_mgr: Any,
|
||||
) -> None:
|
||||
"""When the prefixed name doesn't resolve to either pool or static,
|
||||
the static-path code raises ``ValueError``. Per-user-first
|
||||
resolution (scope decision 0.1) means user_id-bearing callers still
|
||||
hit this path when their pool catalog doesn't carry the name."""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
with pytest.raises(ValueError, match="Unknown MCP prompt"):
|
||||
mgr.get_prompt_sync(
|
||||
"mcp__nonexistent__missing",
|
||||
None,
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-PR-E2E: real discovery + dispatch in same connect (no _seed_pool_prompt_map)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prompt_get_e2e_discovery_then_dispatch_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Drive REAL discovery + dispatch end-to-end through the pool path.
|
||||
|
||||
Mirror of the tool path's
|
||||
``test_integration_pool_reuse_401_refresh_and_retry_succeeds``: skips
|
||||
the ``_seed_pool_prompt_map`` shortcut and lets ``_connect_one_pool``
|
||||
populate ``_user_prompt_map`` from the real ``prompts/list``
|
||||
upstream response. Verifies that the entry's discovered prompts
|
||||
match what the FastMCP fixture advertises AND that
|
||||
``_user_prompt_map[user_id]`` is populated with the prefixed name
|
||||
after dispatch — proving the discovery path actually fired.
|
||||
|
||||
This is the structural gate against a regression where prompt
|
||||
dispatch silently bypasses discovery (e.g., a mis-wired resolver
|
||||
that finds the (server, original) via prefix-parsing alone never
|
||||
populates the per-user catalog).
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "never" # passthrough — discovery + dispatch both succeed
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
# NB: no `_seed_pool_prompt_map` — the resolver finds (server, original)
|
||||
# via the `mcp__{server}__{prompt}` prefix and hands off to
|
||||
# ``_dispatch_pool_prompt_sync``, which lazy-connects via
|
||||
# ``_connect_one_pool``. The connect runs the real ``prompts/list``
|
||||
# against the FastMCP fixture and populates the per-user catalog.
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
messages = mgr.get_prompt_sync(
|
||||
"mcp__pool-srv__greet",
|
||||
{"who": "world"},
|
||||
user_id="user-1",
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
assert isinstance(messages, list)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["role"] == "user"
|
||||
assert "world" in messages[0]["content"]
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
# Discovery populated the entry's prompts with both fixtures
|
||||
# (``greet`` and ``summarize``) — proves real ``prompts/list``
|
||||
# ran during the connect, not just the targeted ``prompts/get``.
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is not None
|
||||
assert entry.prompts is not None
|
||||
discovered_names = {p["name"] for p in entry.prompts}
|
||||
assert "mcp__pool-srv__greet" in discovered_names
|
||||
assert "mcp__pool-srv__summarize" in discovered_names
|
||||
|
||||
# ``_rebuild_user_prompt_map`` ran during the connect, populating the
|
||||
# per-user catalog. This is the signal that discovery wired into the
|
||||
# routing tables — without it, a follow-up ``get_prompt_sync`` would
|
||||
# need to re-resolve via prefix parsing every time.
|
||||
user_prompt_map = mgr._user_prompt_map.get("user-1") or {}
|
||||
assert "mcp__pool-srv__greet" in user_prompt_map
|
||||
assert "mcp__pool-srv__summarize" in user_prompt_map
|
||||
@@ -0,0 +1,690 @@
|
||||
"""Phase 7b integration tests — real-transport resource read 401/403/etc.
|
||||
|
||||
Mirror of :mod:`tests.test_mcp_pool_auth_integration` for the resource
|
||||
path (RFC §3.2). Drives through the real ``streamablehttp_client``,
|
||||
real httpx response-hook plumbing, and a real upstream subprocess
|
||||
(``FastMCP`` with a programmable ``BehaviorMiddleware``). Direct
|
||||
``httpx.HTTPStatusError`` injection is forbidden (invariant 14).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import uvicorn
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.mcp_client import MCPClientManager
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.mcp_oauth import TokenLookupResult
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("mcp").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
class BehaviorMiddleware(BaseHTTPMiddleware):
|
||||
"""Programmable upstream behaviour — see
|
||||
:mod:`tests.test_mcp_pool_auth_integration` for the semantics. This
|
||||
copy serves the resource integration tests.
|
||||
"""
|
||||
|
||||
def __init__(self, app: Any, behaviour: dict[str, Any]) -> None:
|
||||
super().__init__(app)
|
||||
self._behaviour = behaviour
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable[..., Any]) -> Response:
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
|
||||
if request.method == "POST" and "/mcp" in str(request.url):
|
||||
self._behaviour.setdefault("post_auth_headers", []).append(
|
||||
request.headers.get("authorization")
|
||||
)
|
||||
|
||||
mode = self._behaviour.get("mode", "never")
|
||||
if mode == "once_401":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "always_401":
|
||||
return StarletteResponse(
|
||||
"unauthorized",
|
||||
status_code=401,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate", 'Bearer error="invalid_token"'
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_insufficient":
|
||||
if not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get(
|
||||
"www_authenticate",
|
||||
'Bearer error="insufficient_scope", scope="files:read"',
|
||||
)
|
||||
},
|
||||
)
|
||||
elif mode == "once_403_generic" and not self._behaviour.get("_fired"):
|
||||
self._behaviour["_fired"] = True
|
||||
return StarletteResponse(
|
||||
"forbidden",
|
||||
status_code=403,
|
||||
headers={
|
||||
"www-authenticate": self._behaviour.get("www_authenticate", "Bearer realm=mcp")
|
||||
},
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _find_free_port() -> int:
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def _build_server(port: int, behaviour: dict[str, Any]) -> uvicorn.Server:
|
||||
mcp = FastMCP(name="phase7b-resource-target", streamable_http_path="/mcp")
|
||||
|
||||
@mcp.resource("res://hello")
|
||||
def hello() -> str:
|
||||
return "world"
|
||||
|
||||
@mcp.resource("res://json/data")
|
||||
def jdata() -> str:
|
||||
return '{"k": 1}'
|
||||
|
||||
# Echo tool exists so the e2e test can trigger ``_connect_one_pool``
|
||||
# (and the full tool + resource + prompt discovery) via prefix-parsed
|
||||
# ``call_tool_sync`` BEFORE the resource read. The other tests in this
|
||||
# module use ``_seed_pool_resource_map`` and never invoke tools, so
|
||||
# adding the tool is invisible to them.
|
||||
@mcp.tool()
|
||||
async def echo(payload: str = "default") -> str:
|
||||
return f"echoed:{payload}"
|
||||
|
||||
app = mcp.streamable_http_app()
|
||||
app.add_middleware(BehaviorMiddleware, behaviour=behaviour)
|
||||
config = uvicorn.Config(app, host="127.0.0.1", port=port, log_level="warning", access_log=False)
|
||||
return uvicorn.Server(config)
|
||||
|
||||
|
||||
def _wait_ready(port: int, timeout: float = 5.0) -> None:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=0.5):
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(0.05)
|
||||
raise TimeoutError(f"upstream at 127.0.0.1:{port} not ready after {timeout}s")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def upstream():
|
||||
port = _find_free_port()
|
||||
behaviour: dict[str, Any] = {}
|
||||
server = _build_server(port, behaviour)
|
||||
|
||||
def _run() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.serve())
|
||||
|
||||
t = threading.Thread(target=_run, daemon=True, name="phase7b-resource-upstream")
|
||||
t.start()
|
||||
try:
|
||||
_wait_ready(port)
|
||||
yield f"http://127.0.0.1:{port}/mcp", behaviour
|
||||
finally:
|
||||
server.should_exit = True
|
||||
t.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _seed_oauth_server(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
name: str = "pool-srv",
|
||||
server_id: str = "srv-pool",
|
||||
url: str = "https://mcp.example.com/sse",
|
||||
) -> None:
|
||||
storage.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url=url,
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="client-abc",
|
||||
oauth_scopes="openid",
|
||||
oauth_audience=url,
|
||||
)
|
||||
|
||||
|
||||
def _seed_user_token(
|
||||
storage: SQLiteBackend,
|
||||
cipher: Any,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "pool-srv",
|
||||
expires_in_seconds: int = 3600,
|
||||
access_token: str = "access-aaa",
|
||||
refresh_token: str | None = "refresh-rrr",
|
||||
) -> None:
|
||||
expires_at = (datetime.now(UTC) + timedelta(seconds=expires_in_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
)
|
||||
store = MCPTokenStore(storage, cipher, node_id="test")
|
||||
store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token=access_token,
|
||||
refresh_token=refresh_token,
|
||||
expires_at=expires_at,
|
||||
scopes="openid",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
|
||||
def _make_app_state(storage: SQLiteBackend, *, cipher: Any) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
auth_storage=storage,
|
||||
mcp_token_store=MCPTokenStore(storage, cipher, node_id="test"),
|
||||
mcp_oauth_http_client=MagicMock(),
|
||||
mcp_oauth_refresh_locks={},
|
||||
mcp_oauth_metadata_cache={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def running_loop_mgr():
|
||||
cfg: dict[str, Any] = {}
|
||||
mgr = MCPClientManager(cfg)
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True, name="mcp-pool-test-loop")
|
||||
thread.start()
|
||||
mgr._loop = loop
|
||||
try:
|
||||
yield mgr, loop, thread
|
||||
finally:
|
||||
|
||||
async def _drain(m: MCPClientManager) -> None:
|
||||
task = m._user_pool_eviction_task
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
m._user_pool_eviction_task = None
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(_drain(mgr), loop).result(timeout=2)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
|
||||
|
||||
def _seed_pool_resource_map(
|
||||
mgr: MCPClientManager, user_id: str, server_name: str, uri: str
|
||||
) -> None:
|
||||
"""Pre-seed ``_user_resource_map`` so ``_resolve_pool_target_resource``
|
||||
finds the URI. Production wires this through ``_connect_one_pool``;
|
||||
the integration tests seed it directly so the test focuses on the
|
||||
dispatch behaviour after resolution succeeds.
|
||||
"""
|
||||
|
||||
async def _seed() -> None:
|
||||
entry = await mgr._ensure_pool_entry((user_id, server_name))
|
||||
entry.resources = [
|
||||
{
|
||||
"uri": uri,
|
||||
"name": "",
|
||||
"description": "",
|
||||
"mimeType": "",
|
||||
"server": server_name,
|
||||
}
|
||||
]
|
||||
mgr._rebuild_user_resource_map(user_id)
|
||||
|
||||
assert mgr._loop is not None
|
||||
asyncio.run_coroutine_threadsafe(_seed(), mgr._loop).result(timeout=5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-1: 401 → refresh → retry → success (resource path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_401_refresh_and_retry_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Real upstream returns 401 once, then 200. Carrier captures 401,
|
||||
force_refresh=True mints a new bearer, retry returns the resource.
|
||||
Hard invariant 3: breaker counter remains 0.
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
result = mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
|
||||
assert result == "world"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) >= 2, f"expected >=2 POSTs; got {len(post_headers)}"
|
||||
assert post_headers[0] != post_headers[1], (
|
||||
"retry attached the same bearer as the initial; the dispatcher "
|
||||
"did not pick up the refreshed token."
|
||||
)
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-2: persistent 401 → mcp_consent_required (resource path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_persistent_401_emits_consent_required(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="token", token="refreshed-bearer")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert payload["error"]["server"] == "pool-srv"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-3: 403 + insufficient_scope → mcp_insufficient_scope (resource path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_403_insufficient_scope_emits_structured_error(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_insufficient"
|
||||
behaviour["www_authenticate"] = 'Bearer error="insufficient_scope", scope="files:read"'
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_insufficient_scope"
|
||||
assert payload["error"]["scopes_required"] == ["files:read"]
|
||||
post_headers = behaviour.get("post_auth_headers", [])
|
||||
assert len(post_headers) == 1, (
|
||||
f"403 must NOT trigger a retry; observed {len(post_headers)} POSTs"
|
||||
)
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-3b: 403 generic → mcp_resource_read_forbidden
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_403_generic_forbidden(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "once_403_generic"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
# Per the kind="resource" wiring of `_handle_auth_403`, the
|
||||
# operation-specific code surfaces here rather than the tool path's
|
||||
# generic mcp_tool_call_forbidden.
|
||||
assert payload["error"]["code"] == "mcp_resource_read_forbidden"
|
||||
assert "scopes_required" not in payload["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-6: breaker isolation — auth failures NEVER trip the breaker
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_breaker_unaffected_by_auth_failures(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Repeated 401 + refresh-failed cycles leave breaker at 0
|
||||
(hard invariant 3 verified end-to-end for the resource path)."""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "always_401"
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**kwargs: Any) -> TokenLookupResult:
|
||||
if kwargs.get("force_refresh"):
|
||||
return TokenLookupResult(kind="refresh_failed")
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
# Re-seed each iteration: symmetric eviction (Phase 7b) clears
|
||||
# ``_user_resource_map`` on auth failure so the next dispatch's
|
||||
# resolver would miss without a fresh seed. Production reconnect
|
||||
# repopulates this; the test simulates that out-of-band.
|
||||
for _ in range(10):
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Negative tests — token lookup edge cases (resource path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_missing_token_emits_consent_required(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, _behaviour = upstream
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=10)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
|
||||
|
||||
def test_resource_read_decrypt_failure_emits_token_undecryptable(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
url, _behaviour = upstream
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="decrypt_failure")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=10)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_token_undecryptable_key_unknown"
|
||||
|
||||
|
||||
def test_resource_read_http_url_emits_url_insecure(
|
||||
running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""An ``http://`` (non-loopback) oauth_user URL must surface
|
||||
``mcp_oauth_url_insecure`` BEFORE the bearer is attached.
|
||||
"""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url="http://example.com/mcp")
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
_seed_pool_resource_map(mgr, "user-1", "pool-srv", "res://hello")
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
),
|
||||
pytest.raises(RuntimeError) as exc_info,
|
||||
):
|
||||
mgr.read_resource_sync("res://hello", user_id="user-1", timeout=5)
|
||||
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_oauth_url_insecure"
|
||||
|
||||
|
||||
def test_resource_read_unknown_uri_raises_value_error(
|
||||
running_loop_mgr: Any,
|
||||
) -> None:
|
||||
"""When the URI doesn't resolve to either pool or static, the
|
||||
static-path code raises ``ValueError``. Per-user-first resolution
|
||||
(scope decision 0.1) means user_id-bearing callers still hit this
|
||||
path when their pool catalog doesn't carry the URI."""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
with pytest.raises(ValueError, match="Unknown MCP resource"):
|
||||
mgr.read_resource_sync("res://nonexistent", user_id="user-1", timeout=5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I-RP-E2E: real discovery + dispatch in same connect (no _seed_pool_resource_map)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resource_read_e2e_discovery_then_dispatch_succeeds(
|
||||
upstream: Any, running_loop_mgr: Any, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Drive REAL discovery + dispatch end-to-end through the pool path.
|
||||
|
||||
Mirror of the tool path's
|
||||
``test_integration_pool_reuse_401_refresh_and_retry_succeeds``: skips
|
||||
the ``_seed_pool_resource_map`` shortcut and lets ``_connect_one_pool``
|
||||
populate ``_user_resource_map`` from the real ``resources/list``
|
||||
upstream response. Verifies that the entry's discovered resources
|
||||
match what the FastMCP fixture advertises AND that
|
||||
``_user_resource_map[user_id]`` is populated with the URI(s) after
|
||||
discovery — proving the discovery path actually fired.
|
||||
|
||||
Resource URIs do NOT carry a server-name prefix (unlike tools and
|
||||
prompts), so the resource resolver cannot derive (server, uri) by
|
||||
parsing alone. The test triggers the connect via a prefix-parsed
|
||||
``call_tool_sync`` first (which runs the full
|
||||
tools+resources+prompts discovery against the FastMCP fixture),
|
||||
then drives ``read_resource_sync`` against a URI that the
|
||||
upstream advertised — proving that real discovery wired the URI
|
||||
into the per-user catalog.
|
||||
|
||||
Structural gate against a regression where resource discovery is
|
||||
silently skipped (e.g., a capability-gating bug that drops the
|
||||
``resources/list`` call but keeps the connect succeeding).
|
||||
"""
|
||||
url, behaviour = upstream
|
||||
behaviour["mode"] = "never" # passthrough — discovery + dispatch both succeed
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url=url)
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(_make_app_state(storage, cipher=cipher))
|
||||
# NB: no `_seed_pool_resource_map` — the connect runs the real
|
||||
# ``resources/list`` against the FastMCP fixture and populates the
|
||||
# per-user catalog. The tool call below triggers that connect because
|
||||
# ``_resolve_pool_target`` derives (server, original) from the
|
||||
# ``mcp__pool-srv__echo`` prefix and lazy-connects via
|
||||
# ``_connect_one_pool``.
|
||||
|
||||
async def _fake_classified(**_kwargs: Any) -> TokenLookupResult:
|
||||
return TokenLookupResult(kind="token", token="access-aaa")
|
||||
|
||||
with patch(
|
||||
"turnstone.core.mcp_client.get_user_access_token_classified",
|
||||
side_effect=_fake_classified,
|
||||
):
|
||||
# Step 1: trigger the connect via prefix-parsed tool dispatch.
|
||||
# Discovery (tools + resources + prompts) populates the per-user
|
||||
# catalogs.
|
||||
tool_result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "ignite"}, user_id="user-1", timeout=15
|
||||
)
|
||||
assert "echoed:ignite" in tool_result
|
||||
|
||||
# Step 2: now that discovery has populated ``_user_resource_map``,
|
||||
# the resource resolver finds ``res://hello`` and dispatches the
|
||||
# read on the SAME pool entry / session.
|
||||
result = mgr.read_resource_sync("res://hello", user_id="user-1", timeout=15)
|
||||
|
||||
assert result == "world"
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
|
||||
# Discovery populated the entry's resources with both fixtures
|
||||
# (``res://hello`` and ``res://json/data``) — proves real
|
||||
# ``resources/list`` ran during the connect, not just the targeted
|
||||
# ``resources/read``.
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is not None
|
||||
assert entry.resources is not None
|
||||
discovered_uris = {r["uri"] for r in entry.resources if not r.get("template")}
|
||||
assert "res://hello" in discovered_uris
|
||||
assert "res://json/data" in discovered_uris
|
||||
|
||||
# ``_rebuild_user_resource_map`` ran during the connect, populating
|
||||
# the per-user catalog. This is the signal that discovery wired into
|
||||
# the routing tables — without it, ``read_resource_sync`` would have
|
||||
# raised ValueError because the resolver had no entry for the URI.
|
||||
user_resource_map = mgr._user_resource_map.get("user-1") or {}
|
||||
assert "res://hello" in user_resource_map
|
||||
assert "res://json/data" in user_resource_map
|
||||
@@ -0,0 +1,286 @@
|
||||
"""Tests for ``MCPTokenStore`` ciphertext-aware CRUD.
|
||||
|
||||
Phase 3 of the OAuth-MCP RFC: validates the encrypt/decrypt boundary
|
||||
between :class:`MCPTokenStore` and the storage protocol's ciphertext-only
|
||||
columns. Exercises the row-not-deleted-on-decrypt-failure invariant.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
|
||||
import pytest
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from turnstone.core.mcp_crypto import (
|
||||
MCPTokenCipher,
|
||||
MCPTokenCipherConfig,
|
||||
MCPTokenDecryptError,
|
||||
MCPTokenStore,
|
||||
)
|
||||
|
||||
|
||||
def _make_cipher() -> MCPTokenCipher:
|
||||
raw = base64.urlsafe_b64decode(Fernet.generate_key())
|
||||
return MCPTokenCipher(MCPTokenCipherConfig(keys=(raw,)))
|
||||
|
||||
|
||||
def _make_store(backend, *, audit: bool = False) -> tuple[MCPTokenStore, MCPTokenCipher]:
|
||||
cipher = _make_cipher()
|
||||
store = MCPTokenStore(
|
||||
backend,
|
||||
cipher,
|
||||
node_id="test-node",
|
||||
audit_storage=backend if audit else None,
|
||||
)
|
||||
return store, cipher
|
||||
|
||||
|
||||
def _seed_server(backend, *, server_id: str = "srv-id-1", name: str = "srv-a") -> str:
|
||||
backend.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
return server_id
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# User-token CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestUserTokenCRUD:
|
||||
def test_create_and_get_round_trip(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
store.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="access-aaa",
|
||||
refresh_token="refresh-bbb",
|
||||
expires_at="2026-05-04T12:00:00",
|
||||
scopes="openid profile",
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
plain = store.get_user_token("u1", "srv-a")
|
||||
assert plain is not None
|
||||
assert plain["user_id"] == "u1"
|
||||
assert plain["server_name"] == "srv-a"
|
||||
assert plain["access_token"] == "access-aaa"
|
||||
assert plain["refresh_token"] == "refresh-bbb"
|
||||
assert plain["scopes"] == "openid profile"
|
||||
assert plain["audience"] == "https://mcp.example.com"
|
||||
|
||||
def test_create_with_no_refresh_token(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
store.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="access-only",
|
||||
refresh_token=None,
|
||||
expires_at=None,
|
||||
scopes=None,
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
plain = store.get_user_token("u1", "srv-a")
|
||||
assert plain is not None
|
||||
assert plain["access_token"] == "access-only"
|
||||
assert plain["refresh_token"] is None
|
||||
|
||||
def test_get_missing_returns_none(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
assert store.get_user_token("nobody", "srv-a") is None
|
||||
|
||||
def test_update_after_refresh(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
store.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="old-access",
|
||||
refresh_token="old-refresh",
|
||||
expires_at="2026-05-04T12:00:00",
|
||||
scopes="openid",
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
ok = store.update_user_token_after_refresh(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="new-access",
|
||||
refresh_token="new-refresh",
|
||||
expires_at="2026-05-04T13:00:00",
|
||||
)
|
||||
assert ok is True
|
||||
plain = store.get_user_token("u1", "srv-a")
|
||||
assert plain is not None
|
||||
assert plain["access_token"] == "new-access"
|
||||
assert plain["refresh_token"] == "new-refresh"
|
||||
assert plain["expires_at"] == "2026-05-04T13:00:00"
|
||||
# Preserved columns:
|
||||
assert plain["scopes"] == "openid"
|
||||
assert plain["as_issuer"] == "https://auth.example.com"
|
||||
# last_refreshed got stamped:
|
||||
assert plain["last_refreshed"] is not None
|
||||
|
||||
def test_update_after_refresh_missing_row_returns_false(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
ok = store.update_user_token_after_refresh(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="x",
|
||||
refresh_token=None,
|
||||
expires_at=None,
|
||||
)
|
||||
assert ok is False
|
||||
|
||||
def test_delete(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
store.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="a",
|
||||
refresh_token=None,
|
||||
expires_at=None,
|
||||
scopes=None,
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
assert store.delete_user_token("u1", "srv-a") is True
|
||||
assert store.get_user_token("u1", "srv-a") is None
|
||||
# Idempotent: deleting again returns False.
|
||||
assert store.delete_user_token("u1", "srv-a") is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client-secret writer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestClientSecretWriter:
|
||||
def test_set_oauth_client_secret_round_trip(self, backend) -> None:
|
||||
store, cipher = _make_store(backend)
|
||||
server_id = _seed_server(backend)
|
||||
ok = store.set_oauth_client_secret(server_id, "plaintext-secret")
|
||||
assert ok is True
|
||||
# Read raw via get_mcp_server: ciphertext != plaintext, decrypts back.
|
||||
raw = backend.get_mcp_server(server_id)
|
||||
assert raw is not None
|
||||
ct = raw["oauth_client_secret_ct"]
|
||||
assert isinstance(ct, (bytes, bytearray, memoryview))
|
||||
ct_bytes = bytes(ct)
|
||||
assert ct_bytes != b"plaintext-secret"
|
||||
assert cipher.decrypt(ct_bytes) == b"plaintext-secret"
|
||||
|
||||
def test_set_oauth_client_secret_clear_with_none(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
server_id = _seed_server(backend)
|
||||
store.set_oauth_client_secret(server_id, "x")
|
||||
assert store.set_oauth_client_secret(server_id, None) is True
|
||||
raw = backend.get_mcp_server(server_id)
|
||||
assert raw is not None
|
||||
assert raw["oauth_client_secret_ct"] is None
|
||||
|
||||
def test_set_oauth_client_secret_missing_server_returns_false(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
ok = store.set_oauth_client_secret("does-not-exist", "x")
|
||||
assert ok is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Decrypt failure: row preservation invariant
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDecryptFailureInvariant:
|
||||
def test_get_user_token_with_wrong_key_raises_decrypt_error(self, backend) -> None:
|
||||
"""CRITICAL: when no installed key can decrypt a stored row,
|
||||
``get_user_token`` MUST NOT auto-delete the row. The row is
|
||||
still valid; this node just doesn't have the right key.
|
||||
"""
|
||||
# Write under cipher A.
|
||||
store_a, _cipher_a = _make_store(backend)
|
||||
store_a.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="secret-access",
|
||||
refresh_token="secret-refresh",
|
||||
expires_at="2026-05-04T12:00:00",
|
||||
scopes="openid",
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
raw_before = backend.get_mcp_user_token("u1", "srv-a")
|
||||
assert raw_before is not None
|
||||
ct_before = bytes(raw_before["access_token_ct"])
|
||||
|
||||
# Read under cipher B (different key).
|
||||
store_b, cipher_b = _make_store(backend)
|
||||
with pytest.raises(MCPTokenDecryptError) as exc_info:
|
||||
store_b.get_user_token("u1", "srv-a")
|
||||
# The exception carries the keys we tried — useful for audit.
|
||||
assert exc_info.value.key_fingerprints_attempted == cipher_b.key_fingerprints
|
||||
|
||||
# Row MUST still exist with ciphertext intact.
|
||||
raw_after = backend.get_mcp_user_token("u1", "srv-a")
|
||||
assert raw_after is not None
|
||||
assert bytes(raw_after["access_token_ct"]) == ct_before
|
||||
|
||||
def test_decrypt_failure_emits_audit_when_configured(self, backend) -> None:
|
||||
"""When ``audit_storage`` is set, decrypt failures emit a
|
||||
``mcp_server.oauth.token_decrypt_failure`` audit event."""
|
||||
store_a, _ = _make_store(backend)
|
||||
store_a.create_user_token(
|
||||
"u1",
|
||||
"srv-a",
|
||||
access_token="x",
|
||||
refresh_token=None,
|
||||
expires_at=None,
|
||||
scopes=None,
|
||||
as_issuer="https://a",
|
||||
audience="https://m",
|
||||
)
|
||||
|
||||
store_b, cipher_b = _make_store(backend, audit=True)
|
||||
with pytest.raises(MCPTokenDecryptError):
|
||||
store_b.get_user_token("u1", "srv-a")
|
||||
|
||||
events = backend.list_audit_events(limit=10)
|
||||
actions = {ev.get("action") for ev in events}
|
||||
assert "mcp_server.oauth.token_decrypt_failure" in actions
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client-secret reader — q-9
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestClientSecretReader:
|
||||
def test_get_oauth_client_secret_returns_none_when_row_absent(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
assert store.get_oauth_client_secret("does-not-exist") is None
|
||||
|
||||
def test_get_oauth_client_secret_returns_none_when_column_null(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
server_id = _seed_server(backend)
|
||||
# No set_oauth_client_secret call — column stays NULL.
|
||||
assert store.get_oauth_client_secret(server_id) is None
|
||||
|
||||
def test_get_oauth_client_secret_round_trip(self, backend) -> None:
|
||||
store, _ = _make_store(backend)
|
||||
server_id = _seed_server(backend)
|
||||
store.set_oauth_client_secret(server_id, "shhh-its-secret")
|
||||
assert store.get_oauth_client_secret(server_id) == "shhh-its-secret"
|
||||
|
||||
def test_get_oauth_client_secret_raises_on_key_mismatch(self, backend) -> None:
|
||||
store_a, _ = _make_store(backend)
|
||||
server_id = _seed_server(backend)
|
||||
store_a.set_oauth_client_secret(server_id, "secret-under-key-a")
|
||||
|
||||
# Cipher B has a different key — decrypt fails loudly.
|
||||
store_b, _ = _make_store(backend)
|
||||
with pytest.raises(MCPTokenDecryptError):
|
||||
store_b.get_oauth_client_secret(server_id)
|
||||
@@ -0,0 +1,107 @@
|
||||
"""Tests for ``MCPTokenStore.list_user_token_metadata``.
|
||||
|
||||
Validates the non-secret projection used by the settings UI: ciphertext
|
||||
columns are stripped, ordering is preserved, and the empty case returns
|
||||
``[]``. Decrypt is intentionally skipped — the list view must never need
|
||||
the access/refresh secrets.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
|
||||
import sqlalchemy as sa
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from turnstone.core.mcp_crypto import (
|
||||
MCPTokenCipher,
|
||||
MCPTokenCipherConfig,
|
||||
MCPTokenStore,
|
||||
)
|
||||
|
||||
|
||||
def _make_cipher() -> MCPTokenCipher:
|
||||
raw = base64.urlsafe_b64decode(Fernet.generate_key())
|
||||
return MCPTokenCipher(MCPTokenCipherConfig(keys=(raw,)))
|
||||
|
||||
|
||||
def _make_store(backend) -> MCPTokenStore:
|
||||
return MCPTokenStore(backend, _make_cipher(), node_id="test-node")
|
||||
|
||||
|
||||
def _seed_token(
|
||||
store: MCPTokenStore,
|
||||
backend,
|
||||
*,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
created: str,
|
||||
) -> None:
|
||||
"""Create a token via the store and backdate ``created`` for ordering."""
|
||||
store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token="access-secret",
|
||||
refresh_token="refresh-secret",
|
||||
expires_at="2026-05-04T12:00:00",
|
||||
scopes="openid profile",
|
||||
as_issuer="https://auth.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE mcp_user_tokens SET created = :created "
|
||||
"WHERE user_id = :uid AND server_name = :sn"
|
||||
),
|
||||
{"created": created, "uid": user_id, "sn": server_name},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
|
||||
class TestListUserTokenMetadata:
|
||||
def test_list_user_token_metadata_returns_non_secret_fields_only(self, backend) -> None:
|
||||
store = _make_store(backend)
|
||||
_seed_token(
|
||||
store, backend, user_id="u1", server_name="srv-a", created="2026-05-01T00:00:00"
|
||||
)
|
||||
rows = store.list_user_token_metadata("u1")
|
||||
assert len(rows) == 1
|
||||
meta = rows[0]
|
||||
# Secrets MUST be absent.
|
||||
assert "access_token" not in meta
|
||||
assert "refresh_token" not in meta
|
||||
assert "access_token_ct" not in meta
|
||||
assert "refresh_token_ct" not in meta
|
||||
# Non-secret columns surface verbatim.
|
||||
assert meta["user_id"] == "u1"
|
||||
assert meta["server_name"] == "srv-a"
|
||||
assert meta["scopes"] == "openid profile"
|
||||
assert meta["as_issuer"] == "https://auth.example.com"
|
||||
assert meta["audience"] == "https://mcp.example.com"
|
||||
assert meta["expires_at"] == "2026-05-04T12:00:00"
|
||||
assert meta["created"] == "2026-05-01T00:00:00"
|
||||
assert meta["last_refreshed"] is None
|
||||
|
||||
def test_list_user_token_metadata_empty(self, backend) -> None:
|
||||
store = _make_store(backend)
|
||||
assert store.list_user_token_metadata("nobody") == []
|
||||
|
||||
def test_list_user_token_metadata_preserves_creation_order(self, backend) -> None:
|
||||
store = _make_store(backend)
|
||||
_seed_token(
|
||||
store, backend, user_id="u1", server_name="srv-c", created="2026-05-03T00:00:00"
|
||||
)
|
||||
_seed_token(
|
||||
store, backend, user_id="u1", server_name="srv-a", created="2026-05-01T00:00:00"
|
||||
)
|
||||
_seed_token(
|
||||
store, backend, user_id="u1", server_name="srv-b", created="2026-05-02T00:00:00"
|
||||
)
|
||||
rows = store.list_user_token_metadata("u1")
|
||||
assert [r["server_name"] for r in rows] == ["srv-a", "srv-b", "srv-c"]
|
||||
assert [r["created"] for r in rows] == [
|
||||
"2026-05-01T00:00:00",
|
||||
"2026-05-02T00:00:00",
|
||||
"2026-05-03T00:00:00",
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,972 @@
|
||||
"""Tests for the per-(user, server) MCP session pool.
|
||||
|
||||
Covers Phase 5 of the OAuth-MCP rollout: pool data structures,
|
||||
``_ensure_pool_entry`` lazy allocation, ``_connect_one_pool`` plumbing,
|
||||
the dispatch state machine in ``_dispatch_pool``, idle / LRU eviction,
|
||||
failure classification, and ``user_id`` thread-through.
|
||||
|
||||
The static path (``auth_type ∈ {none, static}``) MUST stay
|
||||
byte-identical — see ``test_mcp_client.py``'s
|
||||
``test_reconnect_preserves_static_state_identity``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
from contextlib import AsyncExitStack
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.conftest import make_mcp_token_cipher
|
||||
from turnstone.core.mcp_client import MCPClientManager, PoolEntryState
|
||||
from turnstone.core.mcp_crypto import MCPTokenStore
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures and helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
"""A fresh SQLite backend per test (not the shared singleton)."""
|
||||
return SQLiteBackend(str(tmp_path / "test.db"))
|
||||
|
||||
|
||||
def _seed_oauth_server(
|
||||
storage: SQLiteBackend,
|
||||
*,
|
||||
name: str = "pool-srv",
|
||||
server_id: str = "srv-pool",
|
||||
url: str = "https://mcp.example.com/sse",
|
||||
) -> None:
|
||||
storage.create_mcp_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport="streamable-http",
|
||||
url=url,
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="client-abc",
|
||||
oauth_scopes="openid",
|
||||
oauth_audience=url,
|
||||
)
|
||||
|
||||
|
||||
def _seed_user_token(
|
||||
storage: SQLiteBackend,
|
||||
cipher: Any,
|
||||
*,
|
||||
user_id: str = "user-1",
|
||||
server_name: str = "pool-srv",
|
||||
expires_in_seconds: int = 3600,
|
||||
access_token: str = "access-aaa",
|
||||
) -> None:
|
||||
expires_at = (datetime.now(UTC) + timedelta(seconds=expires_in_seconds)).strftime(
|
||||
"%Y-%m-%dT%H:%M:%S"
|
||||
)
|
||||
store = MCPTokenStore(storage, cipher, node_id="test")
|
||||
store.create_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token=access_token,
|
||||
refresh_token="refresh-rrr",
|
||||
expires_at=expires_at,
|
||||
scopes="openid",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
|
||||
def _make_app_state(storage: SQLiteBackend, *, cipher: Any) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
auth_storage=storage,
|
||||
mcp_token_store=MCPTokenStore(storage, cipher, node_id="test"),
|
||||
mcp_oauth_http_client=MagicMock(),
|
||||
mcp_oauth_refresh_locks={},
|
||||
mcp_oauth_metadata_cache={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def running_loop_mgr():
|
||||
"""Background-loop fixture matching the static-path test convention.
|
||||
|
||||
Tests that need a wired-up app_state assign it via ``mgr.set_app_state``.
|
||||
"""
|
||||
cfg: dict[str, Any] = {}
|
||||
mgr = MCPClientManager(cfg)
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True, name="mcp-pool-test-loop")
|
||||
thread.start()
|
||||
mgr._loop = loop
|
||||
try:
|
||||
yield mgr, loop, thread
|
||||
finally:
|
||||
# Drain the eviction task before stopping the loop so its log/stream
|
||||
# handlers don't fire after pytest has torn its handlers down. Mirrors
|
||||
# the production ``shutdown()`` shape.
|
||||
async def _drain(m: MCPClientManager) -> None:
|
||||
task = m._user_pool_eviction_task
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
m._user_pool_eviction_task = None
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(_drain(mgr), loop).result(timeout=2)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
|
||||
|
||||
def _run_on_loop(loop: asyncio.AbstractEventLoop, coro: Any) -> Any:
|
||||
"""Submit *coro* to *loop*, wait for the result with a 5s timeout."""
|
||||
fut = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||
return fut.result(timeout=5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pool data structures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPoolDataStructures:
|
||||
"""``_user_pool_entries``, ``_user_pool_locks``, eviction-task state."""
|
||||
|
||||
def test_pool_state_starts_empty(self) -> None:
|
||||
mgr = MCPClientManager({})
|
||||
assert mgr._user_pool_entries == {}
|
||||
assert mgr._user_pool_last_used == {}
|
||||
assert mgr._user_pool_locks == {}
|
||||
assert mgr._user_pool_eviction_task is None
|
||||
|
||||
def test_set_app_state_persists(self) -> None:
|
||||
mgr = MCPClientManager({})
|
||||
sentinel = SimpleNamespace(token_store=object())
|
||||
mgr.set_app_state(sentinel)
|
||||
assert mgr._app_state is sentinel
|
||||
|
||||
def test_ensure_pool_entry_allocates_lock_on_loop(self, running_loop_mgr) -> None:
|
||||
"""``asyncio.Lock`` MUST be created on the mcp-loop (RFC §2.0 #2)."""
|
||||
mgr, loop, _thread = running_loop_mgr
|
||||
key = ("user-A", "pool-srv")
|
||||
entry = _run_on_loop(loop, mgr._ensure_pool_entry(key))
|
||||
assert isinstance(entry, PoolEntryState)
|
||||
assert entry.key == key
|
||||
assert isinstance(entry.open_lock, asyncio.Lock)
|
||||
# Calling again returns the same entry / lock object.
|
||||
entry2 = _run_on_loop(loop, mgr._ensure_pool_entry(key))
|
||||
assert entry2 is entry
|
||||
assert entry2.open_lock is entry.open_lock
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lazy connect (`_connect_one_pool`)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _AsyncCM:
|
||||
"""Awaitable async context manager that returns ``value`` from __aenter__."""
|
||||
|
||||
def __init__(self, value: Any) -> None:
|
||||
self._value = value
|
||||
|
||||
async def __aenter__(self) -> Any:
|
||||
return self._value
|
||||
|
||||
async def __aexit__(self, *exc: Any) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
class TestLazyConnect:
|
||||
def test_connect_pool_injects_authorization_header(self, running_loop_mgr) -> None:
|
||||
from unittest.mock import patch
|
||||
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
|
||||
observed_kwargs: dict[str, Any] = {}
|
||||
|
||||
async def _probe(*_args: Any, **_kwargs: Any) -> None:
|
||||
return None
|
||||
|
||||
fake_session = MagicMock()
|
||||
fake_session.initialize = AsyncMock(return_value=None)
|
||||
# Phase 7b: ``_connect_one_pool`` discovers tools, resources,
|
||||
# and prompts after ``initialize()`` returns (resources/prompts
|
||||
# capability-gated). The capability stub returns a tools-only
|
||||
# advertisement so the test can keep its narrow focus on the
|
||||
# bearer-injection contract; resources/prompts paths are
|
||||
# exercised by the real-transport tests in
|
||||
# ``tests/test_mcp_user_catalog.py``.
|
||||
fake_caps = MagicMock()
|
||||
fake_caps.resources = None
|
||||
fake_caps.prompts = None
|
||||
fake_session.get_server_capabilities = MagicMock(return_value=fake_caps)
|
||||
fake_session.list_tools = AsyncMock(return_value=MagicMock(tools=[]))
|
||||
|
||||
def _stream_factory(*, url: str, headers: dict[str, str]) -> _AsyncCM:
|
||||
observed_kwargs["url"] = url
|
||||
observed_kwargs["headers"] = dict(headers)
|
||||
return _AsyncCM((AsyncMock(), AsyncMock(), lambda: None))
|
||||
|
||||
with (
|
||||
patch("turnstone.core.mcp_client.streamablehttp_client", side_effect=_stream_factory),
|
||||
patch.object(mgr, "_tcp_probe", side_effect=_probe),
|
||||
patch("turnstone.core.mcp_client.ClientSession", return_value=_AsyncCM(fake_session)),
|
||||
):
|
||||
cfg = {
|
||||
"type": "streamable-http",
|
||||
"url": "https://mcp.example.com/sse",
|
||||
"headers": {},
|
||||
}
|
||||
entry = _run_on_loop(
|
||||
loop,
|
||||
mgr._connect_one_pool(("user-1", "pool-srv"), cfg, "access-aaa"),
|
||||
)
|
||||
|
||||
assert entry.session is fake_session
|
||||
assert observed_kwargs["headers"]["Authorization"] == "Bearer access-aaa"
|
||||
|
||||
def test_connect_pool_rejects_non_http_transport(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cfg = {"type": "stdio", "command": "echo"}
|
||||
with pytest.raises(RuntimeError, match="streamable-http"):
|
||||
_run_on_loop(
|
||||
loop,
|
||||
mgr._connect_one_pool(("user-1", "pool-srv"), cfg, "access-aaa"),
|
||||
)
|
||||
|
||||
def test_pool_path_does_not_touch_static_servers(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
# Pre-seed a static-path entry so accidental writes are observable.
|
||||
from turnstone.core.mcp_client import StaticServerState
|
||||
|
||||
sentinel = StaticServerState(name="static-srv", session=MagicMock())
|
||||
mgr._static_servers["static-srv"] = sentinel
|
||||
|
||||
async def _seed_pool() -> None:
|
||||
entry = await mgr._ensure_pool_entry(("user-1", "pool-srv"))
|
||||
entry.session = MagicMock()
|
||||
entry.last_used = time.monotonic()
|
||||
|
||||
_run_on_loop(loop, _seed_pool())
|
||||
# Pool side has its own state; the static dict is untouched.
|
||||
assert mgr._static_servers["static-srv"] is sentinel
|
||||
assert mgr._user_pool_entries[("user-1", "pool-srv")].session is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Eviction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEviction:
|
||||
def test_idle_eviction_closes_stale_entries(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
mgr._user_pool_idle_ttl_s = 0.0 # everything is stale
|
||||
|
||||
async def _seed() -> list[PoolEntryState]:
|
||||
entries = []
|
||||
for i in range(3):
|
||||
entry = await mgr._ensure_pool_entry((f"u{i}", "pool-srv"))
|
||||
entry.session = MagicMock()
|
||||
entries.append(entry)
|
||||
return entries
|
||||
|
||||
_run_on_loop(loop, _seed())
|
||||
|
||||
async def _evict() -> None:
|
||||
await mgr._evict_idle_pool_entries()
|
||||
|
||||
_run_on_loop(loop, _evict())
|
||||
assert mgr._user_pool_entries == {}
|
||||
|
||||
def test_eviction_skips_locked_entries(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
mgr._user_pool_idle_ttl_s = 0.0
|
||||
|
||||
async def _seed_and_lock() -> tuple[asyncio.Lock, asyncio.Event]:
|
||||
entry = await mgr._ensure_pool_entry(("u-busy", "pool-srv"))
|
||||
entry.session = MagicMock()
|
||||
held = asyncio.Event()
|
||||
|
||||
async def _hold() -> None:
|
||||
async with entry.open_lock:
|
||||
held.set()
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
asyncio.create_task(_hold())
|
||||
await held.wait()
|
||||
return entry.open_lock, held
|
||||
|
||||
_run_on_loop(loop, _seed_and_lock())
|
||||
|
||||
async def _evict() -> None:
|
||||
await mgr._evict_idle_pool_entries()
|
||||
|
||||
_run_on_loop(loop, _evict())
|
||||
# Entry survives because eviction skipped the locked key.
|
||||
assert ("u-busy", "pool-srv") in mgr._user_pool_entries
|
||||
|
||||
def test_lru_cap_evicts_oldest(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
mgr._user_pool_idle_ttl_s = 999_999.0 # TTL effectively disabled
|
||||
mgr._user_pool_lru_max = 2
|
||||
|
||||
async def _seed() -> None:
|
||||
base = time.monotonic()
|
||||
for i in range(5):
|
||||
key = (f"u{i}", "pool-srv")
|
||||
entry = await mgr._ensure_pool_entry(key)
|
||||
entry.session = MagicMock()
|
||||
# Recent timestamps so TTL doesn't fire — only LRU should.
|
||||
entry.last_used = base + i
|
||||
mgr._user_pool_last_used[key] = base + i
|
||||
|
||||
_run_on_loop(loop, _seed())
|
||||
|
||||
async def _evict() -> None:
|
||||
await mgr._evict_idle_pool_entries()
|
||||
|
||||
_run_on_loop(loop, _evict())
|
||||
assert len(mgr._user_pool_entries) <= 2
|
||||
# The two newest survive (u3, u4).
|
||||
assert ("u4", "pool-srv") in mgr._user_pool_entries
|
||||
assert ("u3", "pool-srv") in mgr._user_pool_entries
|
||||
|
||||
def test_eviction_resilient_to_close_errors(self, running_loop_mgr) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
mgr._user_pool_idle_ttl_s = 0.0
|
||||
|
||||
broken_stack = MagicMock(spec=AsyncExitStack)
|
||||
broken_stack.aclose = AsyncMock(side_effect=RuntimeError("close failed"))
|
||||
|
||||
async def _seed() -> None:
|
||||
for i in range(2):
|
||||
entry = await mgr._ensure_pool_entry((f"u{i}", "pool-srv"))
|
||||
entry.session = MagicMock()
|
||||
entry.stack = broken_stack
|
||||
|
||||
_run_on_loop(loop, _seed())
|
||||
|
||||
async def _evict() -> None:
|
||||
await mgr._evict_idle_pool_entries()
|
||||
|
||||
# Eviction must not raise even if close fails.
|
||||
_run_on_loop(loop, _evict())
|
||||
# All entries removed from the dict regardless.
|
||||
assert mgr._user_pool_entries == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dispatch state machine
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDispatchStateMachine:
|
||||
"""One row per state in the §1.5 / RFC §6 state machine."""
|
||||
|
||||
def _wire_pool(
|
||||
self, mgr: MCPClientManager, storage: SQLiteBackend, cipher: Any
|
||||
) -> SimpleNamespace:
|
||||
mgr.set_storage(storage)
|
||||
state = _make_app_state(storage, cipher=cipher)
|
||||
mgr.set_app_state(state)
|
||||
return state
|
||||
|
||||
def test_no_token_emits_consent_required(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
self._wire_pool(mgr, storage, cipher)
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert payload["error"]["server"] == "pool-srv"
|
||||
|
||||
def test_decrypt_failure_does_not_emit_consent(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
from turnstone.core.mcp_crypto import MCPTokenDecryptError
|
||||
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
_seed_user_token(storage, cipher)
|
||||
state = self._wire_pool(mgr, storage, cipher)
|
||||
|
||||
def _raise(*args, **kwargs):
|
||||
raise MCPTokenDecryptError(
|
||||
"no installed key can decrypt",
|
||||
key_fingerprints_attempted=("aabbccdd",),
|
||||
)
|
||||
|
||||
state.mcp_token_store.get_user_token = _raise
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_token_undecryptable_key_unknown"
|
||||
# Operator fingerprints stay server-side (audit log + structured log);
|
||||
# the agent-facing payload must NOT carry them onward to the LLM
|
||||
# provider.
|
||||
assert "key_fingerprints_attempted" not in payload["error"]
|
||||
|
||||
def test_refresh_failure_emits_consent(self, running_loop_mgr, storage: SQLiteBackend) -> None:
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
# Seed an expired token with no refresh — the classified getter
|
||||
# treats this as "refresh_failed" (deletes the row, returns the
|
||||
# tagged result).
|
||||
_seed_user_token(storage, cipher, expires_in_seconds=-1000)
|
||||
state = self._wire_pool(mgr, storage, cipher)
|
||||
# Drop the refresh token to force the no-refresh-token branch.
|
||||
state.mcp_token_store.delete_user_token("user-1", "pool-srv")
|
||||
state.mcp_token_store.create_user_token(
|
||||
"user-1",
|
||||
"pool-srv",
|
||||
access_token="access-aaa",
|
||||
refresh_token=None,
|
||||
expires_at=(datetime.now(UTC) - timedelta(seconds=1000)).strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
scopes="openid",
|
||||
as_issuer="https://as.example.com",
|
||||
audience="https://mcp.example.com",
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
|
||||
def test_token_present_dispatches_to_session(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
_seed_user_token(storage, cipher, expires_in_seconds=3600)
|
||||
self._wire_pool(mgr, storage, cipher)
|
||||
|
||||
# Pre-seed a connected pool entry so dispatch never touches the
|
||||
# SDK or the network.
|
||||
fake_session = MagicMock()
|
||||
|
||||
async def _call_tool(name, args):
|
||||
content = MagicMock()
|
||||
content.text = "tool-result"
|
||||
res = MagicMock()
|
||||
res.content = [content]
|
||||
res.isError = False
|
||||
return res
|
||||
|
||||
fake_session.call_tool = _call_tool
|
||||
|
||||
async def _seed_entry() -> None:
|
||||
entry = await mgr._ensure_pool_entry(("user-1", "pool-srv"))
|
||||
entry.session = fake_session
|
||||
|
||||
_run_on_loop(loop, _seed_entry())
|
||||
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{"q": "hi"},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
assert result == "tool-result"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Failure classification
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestClassifyFailure:
|
||||
def test_transport_failure_classified_as_transport(self) -> None:
|
||||
mgr = MCPClientManager({})
|
||||
for exc in (
|
||||
BrokenPipeError(),
|
||||
ConnectionResetError(),
|
||||
EOFError(),
|
||||
TimeoutError("net"),
|
||||
):
|
||||
assert mgr._classify_failure(exc) == "transport"
|
||||
|
||||
def test_protocol_error_classified_as_protocol(self) -> None:
|
||||
from mcp import McpError
|
||||
from mcp.types import ErrorData
|
||||
|
||||
mgr = MCPClientManager({})
|
||||
err = McpError(ErrorData(code=-32600, message="bad request"))
|
||||
assert mgr._classify_failure(err) == "protocol"
|
||||
|
||||
def test_other_classified_as_other(self) -> None:
|
||||
mgr = MCPClientManager({})
|
||||
assert mgr._classify_failure(ValueError("nope")) == "other"
|
||||
|
||||
def test_http_401_classified_as_auth_401(self) -> None:
|
||||
"""Defense-in-depth: ``HTTPStatusError`` classification still works
|
||||
even though Phase 6 normally consults the carrier instead.
|
||||
|
||||
Phase 6 split ``"auth"`` into ``"auth_401"`` / ``"auth_403"``
|
||||
so the dispatcher can refresh-and-retry only on 401.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
mgr = MCPClientManager({})
|
||||
req = httpx.Request("POST", "https://mcp.example.com/sse")
|
||||
resp = httpx.Response(401, request=req)
|
||||
exc = httpx.HTTPStatusError("unauthorized", request=req, response=resp)
|
||||
assert mgr._classify_failure(exc) == "auth_401"
|
||||
|
||||
def test_http_403_classified_as_auth_403(self) -> None:
|
||||
import httpx
|
||||
|
||||
mgr = MCPClientManager({})
|
||||
req = httpx.Request("POST", "https://mcp.example.com/sse")
|
||||
resp = httpx.Response(403, request=req)
|
||||
exc = httpx.HTTPStatusError("forbidden", request=req, response=resp)
|
||||
assert mgr._classify_failure(exc) == "auth_403"
|
||||
|
||||
def test_http_500_not_classified_as_auth(self) -> None:
|
||||
import httpx
|
||||
|
||||
mgr = MCPClientManager({})
|
||||
req = httpx.Request("POST", "https://mcp.example.com/sse")
|
||||
resp = httpx.Response(500, request=req)
|
||||
exc = httpx.HTTPStatusError("server", request=req, response=resp)
|
||||
# 5xx is not auth — falls through to "other".
|
||||
assert mgr._classify_failure(exc) == "other"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Wired-failure paths in _dispatch_pool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDispatchFailureWiring:
|
||||
"""``_classify_failure`` is consulted in production, not just tests."""
|
||||
|
||||
def _wire_pool(
|
||||
self, mgr: MCPClientManager, storage: SQLiteBackend, cipher: Any
|
||||
) -> SimpleNamespace:
|
||||
mgr.set_storage(storage)
|
||||
state = _make_app_state(storage, cipher=cipher)
|
||||
mgr.set_app_state(state)
|
||||
return state
|
||||
|
||||
def _seed_connected_session(
|
||||
self, mgr: MCPClientManager, loop: asyncio.AbstractEventLoop, exc: BaseException
|
||||
) -> None:
|
||||
async def _seed() -> None:
|
||||
entry = await mgr._ensure_pool_entry(("user-1", "pool-srv"))
|
||||
sess = MagicMock()
|
||||
|
||||
async def _raise(*_args: Any, **_kwargs: Any) -> Any:
|
||||
raise exc
|
||||
|
||||
sess.call_tool = _raise
|
||||
entry.session = sess
|
||||
|
||||
_run_on_loop(loop, _seed())
|
||||
|
||||
def test_dispatch_pool_transport_failure_trips_breaker(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
_seed_user_token(storage, cipher)
|
||||
self._wire_pool(mgr, storage, cipher)
|
||||
|
||||
self._seed_connected_session(mgr, loop, BrokenPipeError("dead"))
|
||||
|
||||
with pytest.raises(BrokenPipeError):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
# Transport failure ticks the breaker.
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HTTPS enforcement (sec-1)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHttpsEnforcement:
|
||||
def test_pool_rejects_http_url_for_oauth_user(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url="http://insecure.example.com/sse")
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
state = _make_app_state(storage, cipher=cipher)
|
||||
mgr.set_app_state(state)
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
payload = json.loads(str(exc_info.value))
|
||||
assert payload["error"]["code"] == "mcp_oauth_url_insecure"
|
||||
assert payload["error"]["server"] == "pool-srv"
|
||||
|
||||
def test_pool_accepts_loopback_http(self, running_loop_mgr, storage: SQLiteBackend) -> None:
|
||||
"""``http://127.0.0.1`` and ``http://localhost`` should not be blocked."""
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv", url="http://127.0.0.1:8000/sse")
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
state = _make_app_state(storage, cipher=cipher)
|
||||
mgr.set_app_state(state)
|
||||
|
||||
# Pre-seed a connected pool entry so dispatch succeeds without
|
||||
# touching the network.
|
||||
fake_session = MagicMock()
|
||||
|
||||
async def _call_tool(name, args):
|
||||
content = MagicMock()
|
||||
content.text = "ok"
|
||||
res = MagicMock()
|
||||
res.content = [content]
|
||||
res.isError = False
|
||||
return res
|
||||
|
||||
fake_session.call_tool = _call_tool
|
||||
|
||||
async def _seed_entry() -> None:
|
||||
entry = await mgr._ensure_pool_entry(("user-1", "pool-srv"))
|
||||
entry.session = fake_session
|
||||
|
||||
_run_on_loop(loop, _seed_entry())
|
||||
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
# Loopback URL not rejected — dispatch reaches the (fake) session.
|
||||
assert result == "ok"
|
||||
|
||||
def test_validate_oauth_user_url_helper(self) -> None:
|
||||
from turnstone.core.mcp_client import _validate_oauth_user_url
|
||||
|
||||
# Acceptable: https + the exact loopback hostnames.
|
||||
_validate_oauth_user_url("https://mcp.example.com/sse")
|
||||
_validate_oauth_user_url("http://localhost/sse")
|
||||
_validate_oauth_user_url("http://127.0.0.1:9000/sse")
|
||||
_validate_oauth_user_url("http://[::1]/sse")
|
||||
|
||||
# Rejected: non-https + non-loopback. The ``*.localhost`` suffix
|
||||
# bypass is intentionally NOT honored (RFC 6761 localhost-zone
|
||||
# resolution is configuration-dependent — custom resolvers,
|
||||
# /etc/hosts, Docker overlays may map ``foo.localhost`` to
|
||||
# non-loopback IPs).
|
||||
for bad in (
|
||||
"http://mcp.example.com/sse",
|
||||
"http://app.localhost/sse",
|
||||
"ws://mcp.example.com/sse",
|
||||
"ftp://mcp.example.com/sse",
|
||||
"//mcp.example.com/sse",
|
||||
):
|
||||
with pytest.raises(ValueError, match="https://"):
|
||||
_validate_oauth_user_url(bad)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_pool_target parser (q-9)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolvePoolTarget:
|
||||
def _make_mgr_with_oauth_server(
|
||||
self, storage: SQLiteBackend, *, name: str = "pool-srv"
|
||||
) -> MCPClientManager:
|
||||
_seed_oauth_server(storage, name=name)
|
||||
mgr = MCPClientManager({})
|
||||
mgr.set_storage(storage)
|
||||
return mgr
|
||||
|
||||
def test_malformed_prefix(self, storage: SQLiteBackend) -> None:
|
||||
mgr = self._make_mgr_with_oauth_server(storage)
|
||||
# Wrong prefix.
|
||||
assert mgr._resolve_pool_target("xyz__pool-srv__t", None, None) is None
|
||||
|
||||
def test_too_few_separators(self, storage: SQLiteBackend) -> None:
|
||||
mgr = self._make_mgr_with_oauth_server(storage)
|
||||
# mcp__server with no original_name segment.
|
||||
assert mgr._resolve_pool_target("mcp__pool-srv", None, None) is None
|
||||
|
||||
def test_empty_server_segment(self, storage: SQLiteBackend) -> None:
|
||||
mgr = self._make_mgr_with_oauth_server(storage)
|
||||
# mcp____tool — server segment is empty.
|
||||
assert mgr._resolve_pool_target("mcp____tool", None, None) is None
|
||||
|
||||
def test_original_with_double_underscore_round_trips(self, storage: SQLiteBackend) -> None:
|
||||
mgr = self._make_mgr_with_oauth_server(storage)
|
||||
target = mgr._resolve_pool_target("mcp__pool-srv__do__thing", None, None)
|
||||
assert target is not None
|
||||
assert target[0] == "pool-srv"
|
||||
# Original-name keeps its embedded ``__``.
|
||||
assert target[1] == "do__thing"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LRU + lock interlock (q-7)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLruInterlock:
|
||||
def test_lru_cap_skips_locked_oldest(self, running_loop_mgr) -> None:
|
||||
"""LRU eviction must skip a locked entry the same way TTL does."""
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
mgr._user_pool_idle_ttl_s = 999_999.0 # disable TTL
|
||||
mgr._user_pool_lru_max = 2
|
||||
|
||||
async def _seed_and_lock_oldest() -> tuple[asyncio.Lock, asyncio.Event]:
|
||||
base = time.monotonic()
|
||||
for i in range(3):
|
||||
key = (f"u{i}", "pool-srv")
|
||||
entry = await mgr._ensure_pool_entry(key)
|
||||
entry.session = MagicMock()
|
||||
# Older index ⇒ older timestamp.
|
||||
entry.last_used = base + i
|
||||
mgr._user_pool_last_used[key] = base + i
|
||||
# Lock the oldest (u0) so eviction must skip it and pick a younger one.
|
||||
oldest = mgr._user_pool_entries[("u0", "pool-srv")]
|
||||
held = asyncio.Event()
|
||||
|
||||
async def _hold() -> None:
|
||||
async with oldest.open_lock:
|
||||
held.set()
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
asyncio.create_task(_hold())
|
||||
await held.wait()
|
||||
return oldest.open_lock, held
|
||||
|
||||
_run_on_loop(loop, _seed_and_lock_oldest())
|
||||
|
||||
async def _evict() -> None:
|
||||
await mgr._evict_idle_pool_entries()
|
||||
|
||||
_run_on_loop(loop, _evict())
|
||||
# Locked u0 must survive.
|
||||
assert ("u0", "pool-srv") in mgr._user_pool_entries
|
||||
# The oldest unlocked entry (u1) was evicted to bring count down to cap.
|
||||
assert ("u1", "pool-srv") not in mgr._user_pool_entries
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Concurrent dispatch on shared session (M4 / perf-1)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConcurrentDispatch:
|
||||
def test_pool_concurrent_dispatch_to_same_user_server_is_serialized(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Phase 6: two tool calls on the SAME (user, server) MUST serialize
|
||||
on ``open_lock`` so the auth-introspection carrier never crosses
|
||||
between concurrent dispatches.
|
||||
|
||||
Phase 5 perf-1 released ``open_lock`` before ``call_tool`` so two
|
||||
concurrent same-key calls multiplexed on a shared
|
||||
``ClientSession``. Phase 6 reverts that for the auth-aware path
|
||||
because the per-dispatch ``_AuthCapture`` is keyed off the
|
||||
``httpx.AsyncClient`` event hook — releasing the lock would let
|
||||
a concurrent dispatch overwrite the carrier mid-flight,
|
||||
attributing one caller's 401 to another (a security bug).
|
||||
|
||||
Verified by reverting ``_dispatch_pool_with_entry`` to the
|
||||
Phase 5 shape (release ``open_lock`` before ``call_tool`` —
|
||||
i.e. move the ``in_flight += 1`` / ``call_tool`` / decrement
|
||||
block out of the ``async with`` body) and confirming this test
|
||||
observes ``max_concurrency == 2``.
|
||||
"""
|
||||
mgr, loop, _ = running_loop_mgr
|
||||
cipher = make_mcp_token_cipher()
|
||||
_seed_oauth_server(storage, name="pool-srv")
|
||||
_seed_user_token(storage, cipher)
|
||||
mgr.set_storage(storage)
|
||||
state = _make_app_state(storage, cipher=cipher)
|
||||
mgr.set_app_state(state)
|
||||
|
||||
observed_max_concurrency = 0
|
||||
in_flight = 0
|
||||
in_flight_lock = threading.Lock()
|
||||
|
||||
async def _call_tool(name, args):
|
||||
nonlocal observed_max_concurrency, in_flight
|
||||
with in_flight_lock:
|
||||
in_flight += 1
|
||||
observed_max_concurrency = max(observed_max_concurrency, in_flight)
|
||||
try:
|
||||
# Hold a moment so concurrent calls would overlap if
|
||||
# they weren't serialized on ``open_lock``.
|
||||
await asyncio.sleep(0.1)
|
||||
content = MagicMock()
|
||||
content.text = "ok"
|
||||
res = MagicMock()
|
||||
res.content = [content]
|
||||
res.isError = False
|
||||
return res
|
||||
finally:
|
||||
with in_flight_lock:
|
||||
in_flight -= 1
|
||||
|
||||
fake_session = MagicMock()
|
||||
fake_session.call_tool = _call_tool
|
||||
|
||||
async def _seed_entry() -> None:
|
||||
entry = await mgr._ensure_pool_entry(("user-1", "pool-srv"))
|
||||
entry.session = fake_session
|
||||
|
||||
_run_on_loop(loop, _seed_entry())
|
||||
|
||||
results: list[str] = []
|
||||
errors: list[Exception] = []
|
||||
|
||||
def _dispatch() -> None:
|
||||
try:
|
||||
results.append(
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
)
|
||||
except Exception as exc: # pragma: no cover — diagnostic only
|
||||
errors.append(exc)
|
||||
|
||||
t1 = threading.Thread(target=_dispatch)
|
||||
t2 = threading.Thread(target=_dispatch)
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join(timeout=5)
|
||||
t2.join(timeout=5)
|
||||
|
||||
assert errors == []
|
||||
assert results == ["ok", "ok"]
|
||||
# ``open_lock`` held across ``call_tool`` — the second dispatch
|
||||
# waits for the first to release before entering call_tool.
|
||||
assert observed_max_concurrency == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# user_id thread-through (signature)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestUserIdThreadThrough:
|
||||
def test_default_user_id_takes_static_path(self, running_loop_mgr) -> None:
|
||||
"""``user_id=None`` must leave the static-path call byte-identical."""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
# Static-path tool registered the standard way.
|
||||
mgr._tool_map["mcp__static__t"] = ("static-srv", "t")
|
||||
from turnstone.core.mcp_client import StaticServerState
|
||||
|
||||
fake_session = MagicMock()
|
||||
|
||||
async def _call_tool(name, args):
|
||||
content = MagicMock()
|
||||
content.text = "static-output"
|
||||
res = MagicMock()
|
||||
res.content = [content]
|
||||
res.isError = False
|
||||
return res
|
||||
|
||||
fake_session.call_tool = _call_tool
|
||||
mgr._static_servers["static-srv"] = StaticServerState(
|
||||
name="static-srv", session=fake_session
|
||||
)
|
||||
|
||||
# No user_id, no app_state — pool branch is skipped entirely.
|
||||
result = mgr.call_tool_sync("mcp__static__t", {"q": "hi"}, user_id=None, timeout=5)
|
||||
assert result == "static-output"
|
||||
|
||||
def test_user_id_with_static_path_does_not_use_pool(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Caller passes user_id but the resolved server is static — pool
|
||||
branch must not run because ``_lookup_server_row`` reports
|
||||
``auth_type != 'oauth_user'``."""
|
||||
mgr, _loop, _ = running_loop_mgr
|
||||
storage.create_mcp_server(
|
||||
server_id="srv-static",
|
||||
name="static-srv",
|
||||
transport="stdio",
|
||||
url="",
|
||||
command="echo",
|
||||
auth_type="static",
|
||||
)
|
||||
mgr.set_storage(storage)
|
||||
mgr.set_app_state(SimpleNamespace())
|
||||
|
||||
mgr._tool_map["mcp__static-srv__t"] = ("static-srv", "t")
|
||||
from turnstone.core.mcp_client import StaticServerState
|
||||
|
||||
fake_session = MagicMock()
|
||||
|
||||
async def _call_tool(name, args):
|
||||
content = MagicMock()
|
||||
content.text = "static-output"
|
||||
res = MagicMock()
|
||||
res.content = [content]
|
||||
res.isError = False
|
||||
return res
|
||||
|
||||
fake_session.call_tool = _call_tool
|
||||
mgr._static_servers["static-srv"] = StaticServerState(
|
||||
name="static-srv", session=fake_session
|
||||
)
|
||||
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__static-srv__t",
|
||||
{"q": "hi"},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
assert result == "static-output"
|
||||
# No pool entries were created.
|
||||
assert mgr._user_pool_entries == {}
|
||||
@@ -1,6 +1,9 @@
|
||||
"""Tests for turnstone.core.memory_relevance — scoring, formatting, context extraction."""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from turnstone.core.memory_relevance import (
|
||||
MemoryConfig,
|
||||
build_memory_context,
|
||||
extract_recent_context,
|
||||
score_memories,
|
||||
@@ -192,3 +195,265 @@ class TestExtractRecentContext:
|
||||
|
||||
def test_empty_messages(self):
|
||||
assert extract_recent_context([]) == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Composition candidate-selection (_init_system_messages)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_mem(name: str, content: str = "", memory_id: str | None = None) -> dict[str, str]:
|
||||
return {
|
||||
"name": name,
|
||||
"memory_id": memory_id or f"mid_{name}",
|
||||
"type": "project",
|
||||
"scope": "global",
|
||||
"scope_id": "",
|
||||
"description": "",
|
||||
"content": content or name,
|
||||
"updated": "2024-01-01T00:00:00",
|
||||
}
|
||||
|
||||
|
||||
def _make_session(fetch_limit: int = 5, relevance_k: int = 3, **kwargs: object):
|
||||
"""Composition tests need a real ChatSession (constructor calls
|
||||
``_init_system_messages`` once, unpatched, before the test gets a chance
|
||||
to install patches). ``tmp_db`` initializes the storage singleton that
|
||||
constructor needs; tests then patch the visibility helpers and call
|
||||
``_init_system_messages`` a second time to exercise the new logic.
|
||||
"""
|
||||
from tests._helpers import make_chat_session
|
||||
|
||||
return make_chat_session(
|
||||
memory_config=MemoryConfig(fetch_limit=fetch_limit, relevance_k=relevance_k),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class TestCompositionCandidateSelection:
|
||||
"""Verify the query-aware candidate set in _init_system_messages."""
|
||||
|
||||
def test_recency_ceiling_regression(self, tmp_db):
|
||||
"""Old relevant memory not in recency top-N still injected via search path."""
|
||||
session = _make_session(fetch_limit=5, relevance_k=3)
|
||||
session.messages = [{"role": "user", "content": "postgres database configuration"}]
|
||||
|
||||
old_mem = _make_mem(
|
||||
"ancient_db_config",
|
||||
content="postgres database configuration connection host port",
|
||||
memory_id="m_old",
|
||||
)
|
||||
# Recency top-5 do not include old_mem
|
||||
recent = [_make_mem(f"recent_{i}", memory_id=f"mr{i}") for i in range(5)]
|
||||
|
||||
with (
|
||||
patch.object(session, "_search_visible_memories", return_value=[old_mem]),
|
||||
patch.object(session, "_list_visible_memories", return_value=recent),
|
||||
):
|
||||
session._init_system_messages()
|
||||
|
||||
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
|
||||
# With the fix, old_mem enters the candidate pool via search and wins BM25
|
||||
assert "ancient_db_config" in joined
|
||||
|
||||
def test_empty_query_falls_back_to_recency(self, tmp_db):
|
||||
"""No user messages → empty context → recency path, search never called."""
|
||||
session = _make_session()
|
||||
session.messages = [] # extract_recent_context returns ""
|
||||
|
||||
recency = [_make_mem("note_alpha"), _make_mem("note_beta")]
|
||||
|
||||
with (
|
||||
patch.object(session, "_list_visible_memories", return_value=recency),
|
||||
patch.object(session, "_search_visible_memories") as search_mock,
|
||||
):
|
||||
session._init_system_messages()
|
||||
|
||||
search_mock.assert_not_called()
|
||||
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
|
||||
assert "note_alpha" in joined
|
||||
|
||||
def test_sparse_match_union_fills_candidate_pool(self, tmp_db):
|
||||
"""Search returning < fetch_limit results unions with recency fillers."""
|
||||
session = _make_session(fetch_limit=5, relevance_k=4)
|
||||
session.messages = [{"role": "user", "content": "unique_term xyzzy"}]
|
||||
|
||||
hit_a = _make_mem("hit_alpha", content="unique_term xyzzy alpha", memory_id="m_ha")
|
||||
hit_b = _make_mem("hit_beta", content="unique_term xyzzy beta", memory_id="m_hb")
|
||||
search_hits = [hit_a, hit_b] # 2 < fetch_limit=5 → triggers union
|
||||
|
||||
# Recency overlaps on hit_a/hit_b and adds 3 fillers
|
||||
filler = [_make_mem(f"filler_{i}", memory_id=f"mf{i}") for i in range(3)]
|
||||
recency = [hit_a, hit_b] + filler
|
||||
|
||||
with (
|
||||
patch.object(session, "_search_visible_memories", return_value=search_hits),
|
||||
patch.object(session, "_list_visible_memories", return_value=recency),
|
||||
):
|
||||
session._init_system_messages()
|
||||
|
||||
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
|
||||
# Both hits match "unique_term xyzzy" well → appear after BM25 ranking
|
||||
assert "hit_alpha" in joined
|
||||
assert "hit_beta" in joined
|
||||
|
||||
def test_recency_preserved_when_search_returns_noise_above_relevance_k(self, tmp_db):
|
||||
"""Pool guarantee: recency-50 always reaches BM25, even when search
|
||||
returns enough noise hits to clear ``relevance_k``.
|
||||
|
||||
Closes the narrow regression vs. the original bug — without the
|
||||
``fetch_limit`` threshold, a stopword-dominated cap-search that
|
||||
returned >= relevance_k irrelevant hits would short-circuit and
|
||||
evict the recency-only memory the bug had been surfacing.
|
||||
"""
|
||||
session = _make_session(fetch_limit=10, relevance_k=3)
|
||||
session.messages = [{"role": "user", "content": "configure host"}]
|
||||
|
||||
# Search returns relevance_k=3 noise hits — enough to skip recency
|
||||
# under the OLD threshold, not enough to fill fetch_limit=10.
|
||||
noise = [
|
||||
_make_mem(f"noise_{i}", content="generic content", memory_id=f"mn{i}") for i in range(3)
|
||||
]
|
||||
# The memory the user actually wants — distinctive, in recency,
|
||||
# but its content doesn't share any token with the noise hits.
|
||||
wanted = _make_mem(
|
||||
"host_config_v2",
|
||||
content="host=localhost port=5432 db=production",
|
||||
memory_id="m_wanted",
|
||||
)
|
||||
recency = [wanted] + [_make_mem(f"recent_{i}", memory_id=f"mr{i}") for i in range(5)]
|
||||
|
||||
with (
|
||||
patch.object(session, "_search_visible_memories", return_value=noise),
|
||||
patch.object(session, "_list_visible_memories", return_value=recency),
|
||||
):
|
||||
session._init_system_messages()
|
||||
|
||||
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
|
||||
# ``wanted`` reached BM25 via the union and matched "host" → injected.
|
||||
assert "host_config_v2" in joined
|
||||
|
||||
def test_recency_tail_preserved_when_search_adds_distinct_hits(self, tmp_db):
|
||||
"""SUPERSET invariant: every recency item is in the candidate pool
|
||||
when search adds hits, even if the resulting union exceeds
|
||||
fetch_limit. Truncating the union at fetch_limit (the prior
|
||||
behavior) evicted the recency tail — which is exactly where
|
||||
ancient-but-recently-touched memories live, the recall this PR
|
||||
sets out to improve.
|
||||
"""
|
||||
session = _make_session(fetch_limit=10, relevance_k=3)
|
||||
session.messages = [{"role": "user", "content": "alpha"}]
|
||||
|
||||
# 5 search hits, none of which appear in recency.
|
||||
search_hits = [
|
||||
_make_mem(f"search_{i}", content="alpha", memory_id=f"ms{i}") for i in range(5)
|
||||
]
|
||||
# 10 recency items; without the union uncap, the 5 oldest of these
|
||||
# would be displaced by the 5 search hits.
|
||||
recency = [_make_mem(f"recency_{i}", memory_id=f"mr{i}") for i in range(10)]
|
||||
|
||||
with (
|
||||
patch.object(session, "_search_visible_memories", return_value=search_hits),
|
||||
patch.object(session, "_list_visible_memories", return_value=recency),
|
||||
):
|
||||
candidates, source = session._select_memory_candidates("alpha")
|
||||
|
||||
candidate_ids = {c["memory_id"] for c in candidates}
|
||||
# Pool is search_hits ∪ recency — 15 items, no truncation.
|
||||
assert len(candidates) == 15
|
||||
assert source == "union"
|
||||
# Every recency item present (no tail eviction).
|
||||
for i in range(10):
|
||||
assert f"mr{i}" in candidate_ids, f"recency item {i} evicted"
|
||||
# And every search hit is also in the pool.
|
||||
for i in range(5):
|
||||
assert f"ms{i}" in candidate_ids, f"search hit {i} missing"
|
||||
|
||||
def test_coord_scope_isolated_visibility(self, tmp_db):
|
||||
"""Coord composition queries the coord scope alone, never the
|
||||
global/workstream/user union."""
|
||||
from turnstone.core.workstream import WorkstreamKind
|
||||
|
||||
coord = _make_session(
|
||||
fetch_limit=5,
|
||||
relevance_k=3,
|
||||
ws_id="coord-1",
|
||||
user_id="user-1",
|
||||
kind=WorkstreamKind.COORDINATOR,
|
||||
)
|
||||
scopes = coord._visible_scopes()
|
||||
assert scopes == [("coordinator", "coord-1")]
|
||||
# And: search uses those same scopes (no global/user fan-in)
|
||||
coord.messages = [{"role": "user", "content": "anything"}]
|
||||
with patch(
|
||||
"turnstone.core.session.search_visible_structured_memories",
|
||||
return_value=[],
|
||||
) as search_mock:
|
||||
coord._search_visible_memories("anything", limit=5)
|
||||
search_mock.assert_called_once()
|
||||
# Second positional arg is the scopes list
|
||||
assert search_mock.call_args.args[1] == [("coordinator", "coord-1")]
|
||||
|
||||
|
||||
class TestMemorySearchToolExecution:
|
||||
"""End-to-end test of ``memory(action='search')`` through _exec_memory.
|
||||
|
||||
Drives the actual tool dispatch (not just the storage facade) so the
|
||||
OR-of-terms fix and the coalesced ``memory.search`` log get exercised
|
||||
together.
|
||||
"""
|
||||
|
||||
def test_search_action_returns_or_of_terms_results(self, tmp_db):
|
||||
"""Multi-word query returns rows where ANY term matches — not all."""
|
||||
from turnstone.core.memory import save_structured_memory
|
||||
|
||||
save_structured_memory("postgres_notes", "host=localhost port=5432")
|
||||
save_structured_memory("redis_notes", "host=redis port=6379")
|
||||
save_structured_memory("unrelated", "completely different")
|
||||
|
||||
session = _make_session()
|
||||
item = session._prepare_memory(
|
||||
"call-1",
|
||||
{"action": "search", "query": "postgres no_such_word_a no_such_word_b"},
|
||||
)
|
||||
# Sanity: prepare returned a search-ready dispatch (not an error item)
|
||||
assert item.get("action") == "search"
|
||||
|
||||
call_id, msg = session._exec_memory(item)
|
||||
assert call_id == "call-1"
|
||||
assert "postgres_notes" in msg
|
||||
# Other memories don't match any query term
|
||||
assert "unrelated" not in msg
|
||||
|
||||
|
||||
class TestPerTurnSearchCache:
|
||||
"""The per-turn cache spares redundant SQL across mid-turn rebuilds."""
|
||||
|
||||
def test_repeated_search_in_same_turn_hits_cache(self, tmp_db):
|
||||
from turnstone.core.memory import save_structured_memory
|
||||
|
||||
save_structured_memory("hello_mem", "alpha beta gamma")
|
||||
session = _make_session()
|
||||
with patch(
|
||||
"turnstone.core.session.search_visible_structured_memories",
|
||||
return_value=[],
|
||||
) as backend_mock:
|
||||
session._search_visible_memories("alpha beta", limit=5)
|
||||
session._search_visible_memories("alpha beta", limit=5)
|
||||
session._search_visible_memories("alpha beta", limit=5)
|
||||
# 3 calls but only 1 backend hit — cache absorbed the rest
|
||||
assert backend_mock.call_count == 1
|
||||
|
||||
def test_user_turn_invalidates_cache(self, tmp_db):
|
||||
from turnstone.core.memory import save_structured_memory
|
||||
|
||||
save_structured_memory("hello_mem", "alpha")
|
||||
session = _make_session()
|
||||
with patch(
|
||||
"turnstone.core.session.search_visible_structured_memories",
|
||||
return_value=[],
|
||||
) as backend_mock:
|
||||
session._search_visible_memories("alpha", limit=5)
|
||||
session._invalidate_memory_cache() # simulates new user turn
|
||||
session._search_visible_memories("alpha", limit=5)
|
||||
assert backend_mock.call_count == 2
|
||||
|
||||
@@ -4,12 +4,16 @@ from turnstone.core.metacognition import (
|
||||
NUDGE_COMPLETION,
|
||||
NUDGE_CORRECTION,
|
||||
NUDGE_DENIAL,
|
||||
NUDGE_IDLE_CHILDREN_DISPLAY_CAP,
|
||||
NUDGE_IDLE_CHILDREN_WAIT_CAP,
|
||||
NUDGE_REPEAT,
|
||||
NUDGE_RESUME,
|
||||
NUDGE_START,
|
||||
NUDGE_TOOL_ERROR,
|
||||
RepeatDetector,
|
||||
detect_completion,
|
||||
detect_correction,
|
||||
format_idle_children_nudge,
|
||||
format_nudge,
|
||||
should_nudge,
|
||||
)
|
||||
@@ -308,3 +312,276 @@ class TestRepeatNudge:
|
||||
"""Repeat nudge should fire even with zero memories."""
|
||||
state: dict[str, float] = {}
|
||||
assert should_nudge("repeat", state, message_count=5, memory_count=0) is True
|
||||
|
||||
|
||||
class TestRepeatDetector:
|
||||
"""Repeat-detection streak machine — fires only when the same signature
|
||||
is recorded ``threshold`` times *consecutively* (default 3). Recording
|
||||
any different signature resets the streak, so an interrupted repeat
|
||||
isn't flagged as a stuck loop."""
|
||||
|
||||
def test_below_threshold_does_not_fire(self):
|
||||
det = RepeatDetector()
|
||||
assert det.record("a") is False
|
||||
assert det.record("a") is False # second call still under threshold
|
||||
|
||||
def test_at_threshold_fires(self):
|
||||
det = RepeatDetector()
|
||||
det.record("a")
|
||||
det.record("a")
|
||||
assert det.record("a") is True
|
||||
|
||||
def test_continues_to_fire_past_threshold(self):
|
||||
# Caller is responsible for clearing after a fire — until they do,
|
||||
# subsequent identical calls keep returning True.
|
||||
det = RepeatDetector()
|
||||
det.record("a")
|
||||
det.record("a")
|
||||
assert det.record("a") is True
|
||||
assert det.record("a") is True
|
||||
|
||||
def test_clear_resets_count(self):
|
||||
det = RepeatDetector()
|
||||
det.record("a")
|
||||
det.record("a")
|
||||
det.clear()
|
||||
assert det.record("a") is False # back to 1 after clear
|
||||
|
||||
def test_intervening_sig_resets_streak(self):
|
||||
# The streak is consecutive: recording any other sig mid-streak
|
||||
# discards the in-progress count. An alternating pattern like
|
||||
# [A, A, B, A, A] is two short streaks of 2, not a streak of 4.
|
||||
det = RepeatDetector()
|
||||
det.record("a")
|
||||
det.record("a")
|
||||
assert det.record("b") is False # b at count 1; a's streak is gone
|
||||
assert det.record("a") is False # a starts fresh at 1
|
||||
assert det.record("a") is False # a at 2
|
||||
assert det.record("a") is True # a hits 3 — fresh streak completes
|
||||
|
||||
def test_errored_signature_counts_toward_repeat(self):
|
||||
# Regression: when metacog was split out of the system message,
|
||||
# the error-output skip got reintroduced and stuck-loop detection
|
||||
# silently broke for tools that kept failing. Detector itself is
|
||||
# signature-only — error vs. success is the caller's policy.
|
||||
det = RepeatDetector()
|
||||
# Caller records an errored call's sig the same as a successful one;
|
||||
# the streak is what matters.
|
||||
for _ in range(3):
|
||||
last = det.record("bash:ls /nonexistent")
|
||||
assert last is True
|
||||
|
||||
def test_custom_threshold(self):
|
||||
det = RepeatDetector(threshold=2)
|
||||
assert det.record("a") is False
|
||||
assert det.record("a") is True
|
||||
|
||||
def test_threshold_one_fires_immediately(self):
|
||||
det = RepeatDetector(threshold=1)
|
||||
assert det.record("a") is True
|
||||
|
||||
|
||||
class TestFormatIdleChildrenNudge:
|
||||
"""``format_idle_children_nudge`` renders the wake-driven idle_children
|
||||
body — no ``<system-reminder>`` envelope (the side-channel splice
|
||||
wraps it at the wire boundary).
|
||||
"""
|
||||
|
||||
def test_empty_list_returns_empty_string(self):
|
||||
# Caller short-circuits on `if not text: return` — so empty
|
||||
# input MUST produce empty output, not a header-only stub.
|
||||
assert format_idle_children_nudge([]) == ""
|
||||
|
||||
def test_single_child_renders(self):
|
||||
children = [{"ws_id": "ws-abc12345", "name": "research-task", "state": "running"}]
|
||||
text = format_idle_children_nudge(children)
|
||||
assert "ws-abc12" in text # short-id form (8 chars)
|
||||
assert "research-task" in text
|
||||
assert "running" in text
|
||||
assert "wait_for_workstream" in text
|
||||
assert "ws-abc12345" in text # full id appears in the suggestion's ws_ids list
|
||||
|
||||
def test_under_display_cap_no_overflow_line(self):
|
||||
children = [
|
||||
{"ws_id": f"ws-{i:08d}", "name": f"task-{i}", "state": "running"} for i in range(3)
|
||||
]
|
||||
text = format_idle_children_nudge(children)
|
||||
assert "...and" not in text
|
||||
for i in range(3):
|
||||
assert f"task-{i}" in text
|
||||
|
||||
def test_over_display_cap_renders_overflow_line(self):
|
||||
n = NUDGE_IDLE_CHILDREN_DISPLAY_CAP + 4
|
||||
children = [
|
||||
{"ws_id": f"ws-{i:08d}", "name": f"task-{i}", "state": "thinking"} for i in range(n)
|
||||
]
|
||||
text = format_idle_children_nudge(children)
|
||||
assert f"...and {n - NUDGE_IDLE_CHILDREN_DISPLAY_CAP} more" in text
|
||||
# First N children are inline; later ones are folded into "...and N more".
|
||||
for i in range(NUDGE_IDLE_CHILDREN_DISPLAY_CAP):
|
||||
assert f"task-{i}" in text
|
||||
for i in range(NUDGE_IDLE_CHILDREN_DISPLAY_CAP, n):
|
||||
# Names beyond the display cap aren't visible; only counted.
|
||||
assert f"task-{i}" not in text
|
||||
|
||||
def test_over_wait_cap_truncates_suggestion_ws_ids(self):
|
||||
n = NUDGE_IDLE_CHILDREN_WAIT_CAP + 5
|
||||
children = [
|
||||
{"ws_id": f"ws-{i:08d}", "name": f"task-{i}", "state": "running"} for i in range(n)
|
||||
]
|
||||
text = format_idle_children_nudge(children)
|
||||
# The first WAIT_CAP ids appear in the suggestion; later ones don't.
|
||||
first_in_suggestion = f"ws-{NUDGE_IDLE_CHILDREN_WAIT_CAP - 1:08d}"
|
||||
first_excluded = f"ws-{NUDGE_IDLE_CHILDREN_WAIT_CAP:08d}"
|
||||
assert first_in_suggestion in text
|
||||
assert first_excluded not in text
|
||||
|
||||
def test_unnamed_child_falls_back(self):
|
||||
children = [{"ws_id": "ws-deadbeef", "name": "", "state": "attention"}]
|
||||
text = format_idle_children_nudge(children)
|
||||
assert "(unnamed)" in text
|
||||
assert "attention" in text
|
||||
|
||||
def test_newline_in_name_does_not_forge_extra_bullet(self):
|
||||
"""A workstream name with embedded ``\\n`` / ``\\t`` / ``\\r`` MUST
|
||||
NOT break the bullet structure — :func:`sanitize_name`'s strict
|
||||
regex strips control chars (incl. TAB/LF/CR) so the name stays
|
||||
on a single line under its own bullet. Without this, a
|
||||
malicious child name like ``"foo\\n - ws-fake (running): bar"``
|
||||
would forge a fake sibling row in the rendered list.
|
||||
"""
|
||||
children = [
|
||||
{"ws_id": "ws-real0001", "name": "real", "state": "running"},
|
||||
{
|
||||
"ws_id": "ws-evil0002",
|
||||
"name": "evil\n - ws-fake (running): forged",
|
||||
"state": "thinking",
|
||||
},
|
||||
{"ws_id": "ws-real0003", "name": "tail", "state": "running"},
|
||||
]
|
||||
text = format_idle_children_nudge(children)
|
||||
bullet_rows = [ln for ln in text.splitlines() if ln.startswith(" - ")]
|
||||
assert len(bullet_rows) == 3, (
|
||||
f"expected 3 bullet rows; got {len(bullet_rows)}: {bullet_rows!r}"
|
||||
)
|
||||
evil_row = next(row for row in bullet_rows if "ws-evil" in row)
|
||||
assert "\n" not in evil_row
|
||||
assert "\t" not in evil_row
|
||||
assert "\r" not in evil_row
|
||||
assert "evil" in evil_row
|
||||
assert "ws-real" in bullet_rows[2]
|
||||
assert "tail" in bullet_rows[2]
|
||||
|
||||
def test_missing_state_renders_question_mark(self):
|
||||
children = [{"ws_id": "ws-12345678", "name": "x"}]
|
||||
text = format_idle_children_nudge(children)
|
||||
# Defensive default — exotic state keys / partial dicts shouldn't crash.
|
||||
assert "?" in text
|
||||
|
||||
def test_no_system_reminder_envelope(self):
|
||||
# The side-channel ``_apply_reminders_for_provider`` splice
|
||||
# adds ``<system-reminder>`` at the wire boundary; the formatter
|
||||
# MUST NOT wrap, or the model would see a doubled envelope.
|
||||
text = format_idle_children_nudge([{"ws_id": "ws-x", "name": "y", "state": "running"}])
|
||||
assert "<system-reminder>" not in text
|
||||
assert "</system-reminder>" not in text
|
||||
|
||||
def test_format_nudge_returns_empty_for_idle_children(self):
|
||||
# The static map's idle_children entry is the empty string by
|
||||
# design — format_idle_children_nudge produces the real body.
|
||||
assert format_nudge("idle_children") == ""
|
||||
|
||||
def test_should_nudge_recognises_idle_children_type(self, monkeypatch):
|
||||
# Type registration in ``_NUDGE_MAP`` makes ``should_nudge``
|
||||
# recognise it for cooldown gating; without the entry it would
|
||||
# silently return False on every call.
|
||||
state: dict[str, float] = {}
|
||||
# message_count > 1 to clear the first-message gate.
|
||||
assert should_nudge("idle_children", state, message_count=4, memory_count=0) is True
|
||||
# Cooldown set on success → second immediate call returns False.
|
||||
assert should_nudge("idle_children", state, message_count=5, memory_count=0) is False
|
||||
|
||||
|
||||
class TestSanitizeName:
|
||||
"""Strict sanitiser for single-line user-controlled name fields
|
||||
(used by :func:`format_idle_children_nudge` for the workstream
|
||||
``name``). Strips ASCII control chars **including** TAB/LF/CR
|
||||
plus Unicode steering vectors and angle-bracket tag breakers.
|
||||
"""
|
||||
|
||||
def test_empty_input_returns_empty(self):
|
||||
from turnstone.core.metacognition import sanitize_name
|
||||
|
||||
assert sanitize_name("") == ""
|
||||
|
||||
def test_strips_tab_lf_cr(self):
|
||||
"""Strict variant: TAB/LF/CR are stripped so a hostile name with
|
||||
an embedded newline can't break a bullet's one-line structure.
|
||||
"""
|
||||
from turnstone.core.metacognition import sanitize_name
|
||||
|
||||
# All three become spaces (then collapsed to one inline space
|
||||
# by the trailing ``strip()``-on-leading/trailing-only step
|
||||
# — interior runs stay as multiple spaces, that's fine for a
|
||||
# one-line name).
|
||||
assert sanitize_name("a\tb") == "a b"
|
||||
assert sanitize_name("a\nb") == "a b"
|
||||
assert sanitize_name("a\rb") == "a b"
|
||||
|
||||
def test_strips_other_ascii_control_chars(self):
|
||||
from turnstone.core.metacognition import sanitize_name
|
||||
|
||||
assert sanitize_name("a\x07b\x0bc\x0cd") == "a b c d"
|
||||
assert sanitize_name("a\x7fb") == "a b"
|
||||
|
||||
def test_strips_angle_bracket_tag_breakers(self):
|
||||
from turnstone.core.metacognition import sanitize_name
|
||||
|
||||
assert sanitize_name("a</thinking>b") == "a/thinkingb"
|
||||
|
||||
|
||||
class TestSanitizePayload:
|
||||
"""Permissive sanitiser used by the ``watch_triggered`` producer.
|
||||
Strips ASCII control chars (except TAB/LF/CR), Unicode steering
|
||||
vectors (bidi, zero-width, BOM, tag chars), and angle-bracket
|
||||
tag breakers — keeps everything else intact, so multi-line shell
|
||||
output retains its line structure.
|
||||
"""
|
||||
|
||||
def test_empty_input_returns_empty(self):
|
||||
from turnstone.core.metacognition import sanitize_payload
|
||||
|
||||
assert sanitize_payload("") == ""
|
||||
|
||||
def test_strips_ascii_control_chars(self):
|
||||
"""``\\x00``-``\\x1f`` minus TAB/LF/CR plus ``\\x7f`` (DEL) become spaces."""
|
||||
from turnstone.core.metacognition import sanitize_payload
|
||||
|
||||
# BEL (0x07), VT (0x0b), FF (0x0c) — all in strip set.
|
||||
assert sanitize_payload("a\x07b\x0bc\x0cd") == "a b c d"
|
||||
# DEL (0x7f).
|
||||
assert sanitize_payload("a\x7fb") == "a b"
|
||||
|
||||
def test_preserves_tab_lf_cr(self):
|
||||
"""TAB / LF / CR are intentionally preserved so multi-line shell
|
||||
output keeps its line structure when sanitised as a watch payload.
|
||||
"""
|
||||
from turnstone.core.metacognition import sanitize_payload
|
||||
|
||||
# Newlines kept; only the leading + trailing strip happens.
|
||||
out = sanitize_payload("line1\nline2\n\tindented\rline3")
|
||||
assert out == "line1\nline2\n\tindented\rline3"
|
||||
|
||||
def test_strips_bidi_and_zero_width(self):
|
||||
from turnstone.core.metacognition import sanitize_payload
|
||||
|
||||
# U+202E RIGHT-TO-LEFT OVERRIDE; U+200B ZERO WIDTH SPACE.
|
||||
assert sanitize_payload("abc") == "a b c"
|
||||
|
||||
def test_strips_angle_bracket_tag_breakers(self):
|
||||
from turnstone.core.metacognition import sanitize_payload
|
||||
|
||||
# "<" / ">" go away entirely (not replaced with space) so a name
|
||||
# like "</thinking>" doesn't leave a hole the model can read as
|
||||
# a structural marker.
|
||||
assert sanitize_payload("a</thinking>b") == "a/thinkingb"
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
"""Tests for alembic migration 049 (OAuth-MCP schema).
|
||||
|
||||
Drives ``command.upgrade`` from a programmatic Alembic config against
|
||||
an isolated SQLite database per test, then asserts:
|
||||
|
||||
* the two new tables (``mcp_user_tokens``, ``mcp_oauth_pending``) exist,
|
||||
* the eight new ``mcp_servers`` columns exist,
|
||||
* the post-upgrade ``UPDATE mcp_servers`` normalization rewrites rows
|
||||
with empty / missing headers to ``auth_type='none'`` while leaving
|
||||
rows with non-empty headers at ``auth_type='static'``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
_MIGRATIONS_DIR = str(
|
||||
Path(__file__).resolve().parent.parent / "turnstone" / "core" / "storage" / "migrations"
|
||||
)
|
||||
|
||||
|
||||
def _alembic_cfg(db_path: Path) -> Config:
|
||||
cfg = Config()
|
||||
cfg.set_main_option("script_location", _MIGRATIONS_DIR)
|
||||
cfg.set_main_option("sqlalchemy.url", f"sqlite:///{db_path}")
|
||||
return cfg
|
||||
|
||||
|
||||
class TestMigration049:
|
||||
def test_creates_new_tables_and_columns(self, tmp_path: Path) -> None:
|
||||
db_path = tmp_path / "049.db"
|
||||
cfg = _alembic_cfg(db_path)
|
||||
|
||||
# Walk forward through 048 first, then explicitly to 049 so we
|
||||
# exercise the *upgrade* function (not just the schema's `head`).
|
||||
command.upgrade(cfg, "048")
|
||||
command.upgrade(cfg, "049")
|
||||
|
||||
engine = sa.create_engine(f"sqlite:///{db_path}")
|
||||
try:
|
||||
inspector = sa.inspect(engine)
|
||||
tables = set(inspector.get_table_names())
|
||||
assert "mcp_user_tokens" in tables
|
||||
assert "mcp_oauth_pending" in tables
|
||||
|
||||
mcp_cols = {c["name"] for c in inspector.get_columns("mcp_servers")}
|
||||
new_cols = {
|
||||
"auth_type",
|
||||
"oauth_client_id",
|
||||
"oauth_client_secret_ct",
|
||||
"oauth_scopes",
|
||||
"oauth_audience",
|
||||
"oauth_registration_mode",
|
||||
"oauth_authorization_server_url",
|
||||
"oauth_as_issuer_cached",
|
||||
}
|
||||
assert new_cols.issubset(mcp_cols), new_cols - mcp_cols
|
||||
|
||||
# Index check on mcp_oauth_pending.
|
||||
indexes = {ix["name"] for ix in inspector.get_indexes("mcp_oauth_pending")}
|
||||
assert "idx_mcp_pending_created" in indexes
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
def test_normalizes_empty_headers_to_none(self, tmp_path: Path) -> None:
|
||||
"""Streamable-http rows with NULL / '' / '{}' headers become
|
||||
auth_type='none'; rows with non-empty headers stay 'static'.
|
||||
Stdio rows always stay 'static' regardless of headers — the
|
||||
column value is opaque when there is no HTTP transport."""
|
||||
db_path = tmp_path / "049-norm.db"
|
||||
cfg = _alembic_cfg(db_path)
|
||||
|
||||
# Apply everything up to 048, seed rows, then apply 049.
|
||||
command.upgrade(cfg, "048")
|
||||
|
||||
engine = sa.create_engine(f"sqlite:///{db_path}")
|
||||
try:
|
||||
with engine.begin() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO mcp_servers (
|
||||
server_id, name, transport, command, args, url,
|
||||
headers, env, auto_approve, enabled, created_by,
|
||||
registry_name, registry_version, registry_meta,
|
||||
created, updated
|
||||
) VALUES (
|
||||
:sid, :name, :transport, '', '[]',
|
||||
'https://x', :headers, '{}', 0, 1, '', NULL, '',
|
||||
'{}', '2026-05-04T11:00:00', '2026-05-04T11:00:00'
|
||||
)
|
||||
"""
|
||||
),
|
||||
[
|
||||
{
|
||||
"sid": "s-empty-str",
|
||||
"name": "empty-str",
|
||||
"transport": "streamable-http",
|
||||
"headers": "",
|
||||
},
|
||||
{
|
||||
"sid": "s-empty-obj",
|
||||
"name": "empty-obj",
|
||||
"transport": "streamable-http",
|
||||
"headers": "{}",
|
||||
},
|
||||
{
|
||||
"sid": "s-with-headers",
|
||||
"name": "with-headers",
|
||||
"transport": "streamable-http",
|
||||
"headers": '{"Authorization":"Bearer x"}',
|
||||
},
|
||||
# Stdio rows must keep the 'static' default, even
|
||||
# though their headers are empty — auth_type is
|
||||
# opaque for stdio.
|
||||
{
|
||||
"sid": "s-stdio-empty",
|
||||
"name": "stdio-empty",
|
||||
"transport": "stdio",
|
||||
"headers": "{}",
|
||||
},
|
||||
{
|
||||
"sid": "s-stdio-null",
|
||||
"name": "stdio-null",
|
||||
"transport": "stdio",
|
||||
"headers": "",
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
command.upgrade(cfg, "049")
|
||||
|
||||
with engine.connect() as conn:
|
||||
rows = dict(conn.execute(sa.text("SELECT name, auth_type FROM mcp_servers")).all())
|
||||
assert rows["empty-str"] == "none"
|
||||
assert rows["empty-obj"] == "none"
|
||||
assert rows["with-headers"] == "static"
|
||||
# Stdio rows must remain at the 'static' column default even
|
||||
# when headers are empty — the migration only touches HTTP
|
||||
# rows where auth_type is semantically meaningful.
|
||||
assert rows["stdio-empty"] == "static"
|
||||
assert rows["stdio-null"] == "static"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
def test_full_chain_to_head(self, tmp_path: Path) -> None:
|
||||
"""Sanity: running ``upgrade head`` on a fresh DB yields the
|
||||
same end-state column set as ``_schema.metadata``."""
|
||||
db_path = tmp_path / "049-head.db"
|
||||
cfg = _alembic_cfg(db_path)
|
||||
command.upgrade(cfg, "head")
|
||||
|
||||
engine = sa.create_engine(f"sqlite:///{db_path}")
|
||||
try:
|
||||
from turnstone.core.storage._schema import mcp_servers
|
||||
|
||||
inspector = sa.inspect(engine)
|
||||
actual = {c["name"] for c in inspector.get_columns("mcp_servers")}
|
||||
expected = {c.name for c in mcp_servers.columns}
|
||||
assert expected.issubset(actual), expected - actual
|
||||
finally:
|
||||
engine.dispose()
|
||||
@@ -212,3 +212,73 @@ class TestModelDefinitionStorage:
|
||||
m = db.get_model_definition(did)
|
||||
assert m is not None
|
||||
assert m["temperature"] is None
|
||||
|
||||
def test_reasoning_flags_default(self, db: SQLiteBackend) -> None:
|
||||
"""surface_persisted_reasoning defaults True; replay_reasoning_to_model defaults False."""
|
||||
did = _make_id()
|
||||
db.create_model_definition(definition_id=did, alias="reason-default", model="gpt-5")
|
||||
m = db.get_model_definition(did)
|
||||
assert m is not None
|
||||
assert m["surface_persisted_reasoning"] is True
|
||||
assert m["replay_reasoning_to_model"] is False
|
||||
|
||||
def test_create_with_explicit_reasoning_flags(self, db: SQLiteBackend) -> None:
|
||||
did = _make_id()
|
||||
db.create_model_definition(
|
||||
definition_id=did,
|
||||
alias="reason-explicit",
|
||||
model="claude-opus-4-7",
|
||||
surface_persisted_reasoning=False,
|
||||
replay_reasoning_to_model=True,
|
||||
)
|
||||
m = db.get_model_definition(did)
|
||||
assert m is not None
|
||||
assert m["surface_persisted_reasoning"] is False
|
||||
assert m["replay_reasoning_to_model"] is True
|
||||
# Same values must round-trip via the alias lookup too.
|
||||
m_alias = db.get_model_definition_by_alias("reason-explicit")
|
||||
assert m_alias is not None
|
||||
assert m_alias["surface_persisted_reasoning"] is False
|
||||
assert m_alias["replay_reasoning_to_model"] is True
|
||||
|
||||
def test_update_surface_persisted_reasoning(self, db: SQLiteBackend) -> None:
|
||||
did = _make_id()
|
||||
db.create_model_definition(definition_id=did, alias="upd-persist", model="gpt-5")
|
||||
ok = db.update_model_definition(did, surface_persisted_reasoning=False)
|
||||
assert ok is True
|
||||
m = db.get_model_definition(did)
|
||||
assert m is not None
|
||||
assert m["surface_persisted_reasoning"] is False
|
||||
assert m["replay_reasoning_to_model"] is False # untouched
|
||||
|
||||
def test_update_replay_reasoning_to_model(self, db: SQLiteBackend) -> None:
|
||||
did = _make_id()
|
||||
db.create_model_definition(definition_id=did, alias="upd-replay", model="gpt-5")
|
||||
ok = db.update_model_definition(did, replay_reasoning_to_model=True)
|
||||
assert ok is True
|
||||
m = db.get_model_definition(did)
|
||||
assert m is not None
|
||||
assert m["surface_persisted_reasoning"] is True # untouched
|
||||
assert m["replay_reasoning_to_model"] is True
|
||||
|
||||
def test_list_returns_reasoning_flags(self, db: SQLiteBackend) -> None:
|
||||
db.create_model_definition(
|
||||
definition_id=_make_id(),
|
||||
alias="list-a",
|
||||
model="gpt-5",
|
||||
surface_persisted_reasoning=True,
|
||||
replay_reasoning_to_model=False,
|
||||
)
|
||||
db.create_model_definition(
|
||||
definition_id=_make_id(),
|
||||
alias="list-b",
|
||||
model="claude-opus-4-7",
|
||||
surface_persisted_reasoning=False,
|
||||
replay_reasoning_to_model=True,
|
||||
)
|
||||
models = db.list_model_definitions()
|
||||
by_alias = {m["alias"]: m for m in models}
|
||||
assert by_alias["list-a"]["surface_persisted_reasoning"] is True
|
||||
assert by_alias["list-a"]["replay_reasoning_to_model"] is False
|
||||
assert by_alias["list-b"]["surface_persisted_reasoning"] is False
|
||||
assert by_alias["list-b"]["replay_reasoning_to_model"] is True
|
||||
|
||||
+277
-16
@@ -76,6 +76,23 @@ class TestModelConfig:
|
||||
assert cfg.temperature == 0.0
|
||||
assert cfg.temperature is not None
|
||||
|
||||
def test_reasoning_flags_default(self) -> None:
|
||||
cfg = ModelConfig(alias="x", base_url="x", api_key="x", model="x")
|
||||
assert cfg.surface_persisted_reasoning is True
|
||||
assert cfg.replay_reasoning_to_model is False
|
||||
|
||||
def test_reasoning_flags_set(self) -> None:
|
||||
cfg = ModelConfig(
|
||||
alias="x",
|
||||
base_url="x",
|
||||
api_key="x",
|
||||
model="x",
|
||||
surface_persisted_reasoning=False,
|
||||
replay_reasoning_to_model=True,
|
||||
)
|
||||
assert cfg.surface_persisted_reasoning is False
|
||||
assert cfg.replay_reasoning_to_model is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelRegistry
|
||||
@@ -314,7 +331,11 @@ class TestLoadModelRegistry:
|
||||
api_key="dummy",
|
||||
model="local-model",
|
||||
)
|
||||
assert reg.count == 2 # "openai" + "default"
|
||||
# The CLI ``"default"`` shim is suppressed once ``[models.*]``
|
||||
# populates configs — only the explicit alias survives.
|
||||
assert reg.count == 1
|
||||
assert reg.has_alias("openai")
|
||||
assert not reg.has_alias("default")
|
||||
assert reg.default == "openai"
|
||||
_, model, _ = reg.resolve()
|
||||
assert model == "gpt-4o"
|
||||
@@ -545,7 +566,12 @@ class TestLoadModelRegistryWithDB:
|
||||
assert cfg.source == "config"
|
||||
|
||||
def test_db_only_models_coexist(self) -> None:
|
||||
"""DB models coexist alongside config.toml models."""
|
||||
"""DB models coexist alongside config.toml models.
|
||||
|
||||
The CLI ``"default"`` shim is suppressed when DB / config models
|
||||
already populate the registry — see
|
||||
``test_cli_default_shim_skipped_when_db_models_present``.
|
||||
"""
|
||||
storage = _MockStorage(
|
||||
[
|
||||
{
|
||||
@@ -569,7 +595,7 @@ class TestLoadModelRegistryWithDB:
|
||||
reg = load_model_registry("http://x/v1", "x", "x", storage=storage)
|
||||
assert reg.has_alias("db-only")
|
||||
assert reg.has_alias("config-only")
|
||||
assert reg.has_alias("default")
|
||||
assert not reg.has_alias("default")
|
||||
assert reg.get_config("db-only").source == "db"
|
||||
assert reg.get_config("config-only").source == "config"
|
||||
|
||||
@@ -589,10 +615,12 @@ class TestLoadModelRegistryWithDB:
|
||||
}
|
||||
]
|
||||
)
|
||||
# The CLI default shim is suppressed when the DB row populates
|
||||
# configs, so only the DB-sourced alias exists here.
|
||||
with patch("turnstone.core.model_registry.load_config", return_value={}):
|
||||
reg = load_model_registry("http://x/v1", "x", "x", storage=storage)
|
||||
assert reg.get_config("from-db").source == "db"
|
||||
assert reg.get_config("default").source == ""
|
||||
assert not reg.has_alias("default")
|
||||
|
||||
def test_disabled_db_models_excluded(self) -> None:
|
||||
"""Disabled DB models are not loaded."""
|
||||
@@ -686,6 +714,53 @@ class TestLoadModelRegistryWithDB:
|
||||
assert cfg.max_tokens is None
|
||||
assert cfg.reasoning_effort is None
|
||||
|
||||
def test_db_reasoning_flags_loaded(self) -> None:
|
||||
"""Per-model reasoning flags from DB are carried in ModelConfig."""
|
||||
storage = _MockStorage(
|
||||
[
|
||||
{
|
||||
"alias": "anth-thinking",
|
||||
"model": "claude-opus-4-7",
|
||||
"provider": "anthropic",
|
||||
"base_url": "",
|
||||
"api_key": "sk-anth",
|
||||
"context_window": 200000,
|
||||
"capabilities": "{}",
|
||||
"enabled": True,
|
||||
"surface_persisted_reasoning": False,
|
||||
"replay_reasoning_to_model": True,
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch("turnstone.core.model_registry.load_config", return_value={}):
|
||||
reg = load_model_registry("http://x/v1", "x", "x", storage=storage)
|
||||
cfg = reg.get_config("anth-thinking")
|
||||
assert cfg.surface_persisted_reasoning is False
|
||||
assert cfg.replay_reasoning_to_model is True
|
||||
|
||||
def test_db_reasoning_flags_default_when_absent(self) -> None:
|
||||
"""Pre-052 rows without the columns degrade to dataclass defaults."""
|
||||
storage = _MockStorage(
|
||||
[
|
||||
{
|
||||
"alias": "legacy-row",
|
||||
"model": "gpt-5",
|
||||
"provider": "openai",
|
||||
"base_url": "",
|
||||
"api_key": "",
|
||||
"context_window": 32768,
|
||||
"capabilities": "{}",
|
||||
"enabled": True,
|
||||
# surface_persisted_reasoning + replay_reasoning_to_model intentionally absent
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch("turnstone.core.model_registry.load_config", return_value={}):
|
||||
reg = load_model_registry("http://x/v1", "x", "x", storage=storage)
|
||||
cfg = reg.get_config("legacy-row")
|
||||
assert cfg.surface_persisted_reasoning is True
|
||||
assert cfg.replay_reasoning_to_model is False
|
||||
|
||||
def test_db_default_alias_not_clobbered(self) -> None:
|
||||
"""DB model with alias='default' is not overwritten by CLI args."""
|
||||
storage = _MockStorage(
|
||||
@@ -770,16 +845,76 @@ class TestRegistryReload:
|
||||
assert reg.has_alias("b")
|
||||
assert reg.default == "b"
|
||||
|
||||
def test_reload_clears_clients(self) -> None:
|
||||
models = {"a": ModelConfig("a", "http://x/v1", "key", "m")}
|
||||
def test_reload_keeps_clients_when_connection_target_unchanged(self) -> None:
|
||||
"""Selective teardown: a model edit that leaves base_url / api_key /
|
||||
provider intact (e.g. admin tweaks the underlying ``model`` name or
|
||||
``temperature``) keeps the cached HTTP client warm — no need to
|
||||
re-establish TLS+pool when the endpoint is the same."""
|
||||
models = {"a": ModelConfig("a", "http://x/v1", "key", "m1", provider="openai")}
|
||||
reg = ModelRegistry(models=models, default="a")
|
||||
# Force client creation
|
||||
reg.get_client("a")
|
||||
assert "a" in reg._clients
|
||||
client_before = reg._clients["a"]
|
||||
provider_before = reg.get_provider("a")
|
||||
|
||||
# Same endpoint (base_url, api_key, provider), only ``model`` changed.
|
||||
new_models = {"a": ModelConfig("a", "http://x/v1", "key", "m2", provider="openai")}
|
||||
reg.reload(new_models, "a")
|
||||
|
||||
assert "a" in reg._clients
|
||||
assert reg._clients["a"] is client_before
|
||||
assert "a" in reg._providers
|
||||
assert reg._providers["a"] is provider_before
|
||||
|
||||
def test_reload_drops_client_when_base_url_changes(self) -> None:
|
||||
"""A ``base_url`` change drops the cached client (different
|
||||
endpoint = new connection) but keeps the cached provider —
|
||||
``LLMProvider`` is keyed only on the provider string, which
|
||||
didn't change."""
|
||||
models = {"a": ModelConfig("a", "http://x/v1", "key", "m", provider="openai")}
|
||||
reg = ModelRegistry(models=models, default="a")
|
||||
reg.get_client("a")
|
||||
provider_before = reg.get_provider("a")
|
||||
|
||||
new_models = {"a": ModelConfig("a", "http://y/v1", "key", "m", provider="openai")}
|
||||
reg.reload(new_models, "a")
|
||||
|
||||
# Reload with same models — clients should be cleared
|
||||
reg.reload(dict(models), "a")
|
||||
assert "a" not in reg._clients
|
||||
assert "a" in reg._providers
|
||||
assert reg._providers["a"] is provider_before
|
||||
|
||||
def test_reload_drops_provider_when_provider_string_changes(self) -> None:
|
||||
"""A provider-type swap (e.g. openai → anthropic) drops both the
|
||||
client AND the provider so the next resolve picks up the right
|
||||
``LLMProvider`` implementation against the new SDK."""
|
||||
models = {"a": ModelConfig("a", "http://x/v1", "key", "m", provider="openai")}
|
||||
reg = ModelRegistry(models=models, default="a")
|
||||
reg.get_client("a")
|
||||
reg.get_provider("a")
|
||||
|
||||
new_models = {"a": ModelConfig("a", "http://x/v1", "key", "m", provider="anthropic")}
|
||||
reg.reload(new_models, "a")
|
||||
|
||||
assert "a" not in reg._clients
|
||||
assert "a" not in reg._providers
|
||||
|
||||
def test_reload_drops_clients_for_removed_aliases(self) -> None:
|
||||
"""Aliases removed from the registry must release their cached
|
||||
clients — otherwise a deleted endpoint's connection pool would
|
||||
outlive the alias indefinitely."""
|
||||
models = {
|
||||
"a": ModelConfig("a", "http://x/v1", "key", "m"),
|
||||
"b": ModelConfig("b", "http://y/v1", "key", "m"),
|
||||
}
|
||||
reg = ModelRegistry(models=models, default="a")
|
||||
reg.get_client("a")
|
||||
reg.get_client("b")
|
||||
|
||||
# Drop "b" entirely.
|
||||
new_models = {"a": ModelConfig("a", "http://x/v1", "key", "m")}
|
||||
reg.reload(new_models, "a")
|
||||
|
||||
assert "a" in reg._clients # unchanged endpoint, kept warm
|
||||
assert "b" not in reg._clients
|
||||
|
||||
def test_reload_validates_default(self) -> None:
|
||||
models_a = {"a": ModelConfig("a", "x", "x", "m")}
|
||||
@@ -809,6 +944,8 @@ class _FakeUI:
|
||||
self.infos: list[str] = []
|
||||
self.errors: list[str] = []
|
||||
|
||||
def on_turn_start(self) -> None: ...
|
||||
def on_turn_committed(self) -> None: ...
|
||||
def on_thinking_start(self) -> None: ...
|
||||
def on_thinking_stop(self) -> None: ...
|
||||
def on_reasoning_token(self, text: str) -> None: ...
|
||||
@@ -1066,24 +1203,57 @@ class TestSessionAgentModel:
|
||||
def _captured_effort(captured: dict[str, Any]) -> str | None:
|
||||
"""Pull reasoning_effort out of provider-specific shapes.
|
||||
|
||||
openai-compatible servers receive it via extra_body.chat_template_kwargs;
|
||||
commercial providers receive it as a top-level kwarg.
|
||||
Chat Completions delivers it as a top-level ``reasoning_effort`` kwarg
|
||||
(when the model's caps permit it). Operators who route reasoning_effort
|
||||
through ``chat_template_kwargs`` (gpt-oss-style local templates) get
|
||||
it inside ``extra_body.chat_template_kwargs``.
|
||||
"""
|
||||
if "reasoning_effort" in captured:
|
||||
return captured["reasoning_effort"]
|
||||
eb = captured.get("extra_body") or {}
|
||||
ctk = eb.get("chat_template_kwargs") or {}
|
||||
return ctk.get("reasoning_effort") or captured.get("reasoning_effort")
|
||||
return ctk.get("reasoning_effort")
|
||||
|
||||
@staticmethod
|
||||
def _effort_caps() -> dict[str, Any]:
|
||||
"""Capabilities that allow Chat-Completions reasoning_effort to flow."""
|
||||
return {
|
||||
"reasoning_effort_values": [
|
||||
"minimal",
|
||||
"low",
|
||||
"medium",
|
||||
"high",
|
||||
"max",
|
||||
],
|
||||
}
|
||||
|
||||
def _three_model_registry(self, **kwargs: Any) -> ModelRegistry:
|
||||
caps = self._effort_caps()
|
||||
return ModelRegistry(
|
||||
models={
|
||||
"main": ModelConfig(
|
||||
"main", "http://m/v1", "k", "main-model", provider="openai-compatible"
|
||||
"main",
|
||||
"http://m/v1",
|
||||
"k",
|
||||
"main-model",
|
||||
provider="openai-compatible",
|
||||
capabilities=dict(caps),
|
||||
),
|
||||
"smart": ModelConfig(
|
||||
"smart", "http://s/v1", "k", "smart-model", provider="openai-compatible"
|
||||
"smart",
|
||||
"http://s/v1",
|
||||
"k",
|
||||
"smart-model",
|
||||
provider="openai-compatible",
|
||||
capabilities=dict(caps),
|
||||
),
|
||||
"fast": ModelConfig(
|
||||
"fast", "http://f/v1", "k", "fast-model", provider="openai-compatible"
|
||||
"fast",
|
||||
"http://f/v1",
|
||||
"k",
|
||||
"fast-model",
|
||||
provider="openai-compatible",
|
||||
capabilities=dict(caps),
|
||||
),
|
||||
},
|
||||
default="main",
|
||||
@@ -1189,6 +1359,45 @@ class TestSessionAgentModel:
|
||||
session._run_agent([{"role": "user", "content": "x"}], label="plan", agent_alias="fast")
|
||||
assert captured["model"] == "fast-model"
|
||||
|
||||
def test_session_fallback_inherits_primary_alias_for_caps(self) -> None:
|
||||
"""When _run_agent has no registry agent route, it must fall back to
|
||||
the session's primary alias for capability and server_compat lookup —
|
||||
otherwise per-model caps (reasoning_effort_values, server_compat) get
|
||||
silently dropped on the agent path."""
|
||||
reg = self._three_model_registry() # no agent_model / plan_model set
|
||||
session = _make_session(registry=reg, model_alias="main")
|
||||
# Probe what _run_agent passes to _provider_extra_params and
|
||||
# _resolve_capabilities by recording the model_alias on each call.
|
||||
captured_extra_alias: list[str | None] = []
|
||||
captured_resolve_alias: list[str | None] = []
|
||||
original_extra = session._provider_extra_params
|
||||
original_resolve = session._resolve_capabilities
|
||||
|
||||
def spy_extra(*args: Any, **kwargs: Any) -> Any:
|
||||
captured_extra_alias.append(kwargs.get("model_alias"))
|
||||
return original_extra(*args, **kwargs)
|
||||
|
||||
def spy_resolve(*args: Any, **kwargs: Any) -> Any:
|
||||
# _resolve_capabilities(provider, model, alias)
|
||||
alias = args[2] if len(args) >= 3 else kwargs.get("alias")
|
||||
captured_resolve_alias.append(alias)
|
||||
return original_resolve(*args, **kwargs)
|
||||
|
||||
session._provider_extra_params = spy_extra # type: ignore[method-assign]
|
||||
session._resolve_capabilities = spy_resolve # type: ignore[method-assign]
|
||||
|
||||
self._capture_on(session.client) # patch client.chat.completions.create
|
||||
session._run_agent([{"role": "user", "content": "x"}], label="plan")
|
||||
|
||||
assert captured_extra_alias and captured_extra_alias[-1] == "main", (
|
||||
f"agent fallback path did not inherit primary alias for extra_params: "
|
||||
f"{captured_extra_alias!r}"
|
||||
)
|
||||
assert captured_resolve_alias and captured_resolve_alias[-1] == "main", (
|
||||
f"agent fallback path did not inherit primary alias for caps: "
|
||||
f"{captured_resolve_alias!r}"
|
||||
)
|
||||
|
||||
def test_invalid_alias_raises_in_run_agent(self) -> None:
|
||||
"""Defence-in-depth: _prepare_* validates first, but _run_agent
|
||||
rejects unknown aliases too rather than silently falling back."""
|
||||
@@ -1477,6 +1686,58 @@ class TestLoadModelRegistryDBOnly:
|
||||
reg = load_model_registry(model="", storage=storage)
|
||||
assert not reg.has_alias("default")
|
||||
|
||||
def test_cli_default_shim_skipped_when_db_models_present(self) -> None:
|
||||
"""An auto-detected ``--model`` does NOT synthesise a ``default``
|
||||
alias when the DB already contributes models.
|
||||
|
||||
Regression for the silent bypass of ``model.task_alias`` /
|
||||
``model.plan_alias``: a synthesised ``default`` aliased to whatever
|
||||
``--base-url`` was at boot leaks into the LLM-visible alias list,
|
||||
and the LLM picks it for ``task_agent(model="default")`` — which
|
||||
then routes around the operator-configured per-role default.
|
||||
"""
|
||||
storage = _MockStorage(
|
||||
[
|
||||
{
|
||||
"alias": "gh200",
|
||||
"model": "deepseek-ai/DeepSeek-V4-Flash",
|
||||
"provider": "openai",
|
||||
"base_url": "http://gh200:8000/v1",
|
||||
"api_key": "sk-gh200",
|
||||
"context_window": 1048576,
|
||||
"capabilities": "{}",
|
||||
"enabled": True,
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch("turnstone.core.model_registry.load_config", return_value={}):
|
||||
reg = load_model_registry(
|
||||
base_url="http://flatspark:8000/v1",
|
||||
api_key="sk-flatspark",
|
||||
model="qwen3.6-35B-A3B", # populated by ``detect_model``
|
||||
storage=storage,
|
||||
)
|
||||
assert reg.has_alias("gh200")
|
||||
assert not reg.has_alias("default")
|
||||
|
||||
def test_cli_default_shim_skipped_when_config_models_present(self) -> None:
|
||||
"""Same shim suppression when only ``[models.*]`` populates configs."""
|
||||
fake_cfg: dict[str, Any] = {
|
||||
"models": {"local": {"model": "qwen3-32b"}},
|
||||
}
|
||||
with patch("turnstone.core.model_registry.load_config", return_value=fake_cfg):
|
||||
reg = load_model_registry("http://x/v1", "x", "fallback-model")
|
||||
assert reg.has_alias("local")
|
||||
assert not reg.has_alias("default")
|
||||
|
||||
def test_cli_default_shim_still_fires_when_registry_empty(self) -> None:
|
||||
"""Single-model CLI mode (no DB, no config.toml [models.*]) keeps
|
||||
the back-compat ``default`` alias."""
|
||||
with patch("turnstone.core.model_registry.load_config", return_value={}):
|
||||
reg = load_model_registry("http://x/v1", "x", "lone-model")
|
||||
assert reg.has_alias("default")
|
||||
assert reg.get_config("default").model == "lone-model"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# server._effective_routing / _apply_routing_overrides
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
"""``models_changed`` SSE fanout coverage.
|
||||
|
||||
The console pushes a ``models_changed`` cluster event whenever a model
|
||||
definition is created / updated / deleted / reloaded, or whenever a
|
||||
setting in :data:`turnstone.console.server._MODEL_AFFECTING_SETTING_KEYS`
|
||||
is updated or reset. Connected browsers refetch ``/v1/api/models`` on
|
||||
receipt so the home composer dropdown + admin Models → Roles sub-tab
|
||||
reflect alias edits without a manual reload.
|
||||
|
||||
These tests pin two contracts:
|
||||
|
||||
- every model-definition CRUD path emits exactly one ``models_changed``
|
||||
fanout (so the browser stays in sync with the DB);
|
||||
- settings PUT / DELETE only emit the fanout when the key is
|
||||
model-affecting — unrelated keys (e.g. ``session.retention_days``)
|
||||
must not trigger spurious dropdown re-renders across the cluster.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.routing import Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from tests._coord_test_helpers import _AuthMiddleware
|
||||
from turnstone.console.server import (
|
||||
_MODEL_AFFECTING_SETTING_KEYS,
|
||||
admin_create_model_definition,
|
||||
admin_delete_model_definition,
|
||||
admin_delete_setting,
|
||||
admin_model_reload,
|
||||
admin_update_model_definition,
|
||||
admin_update_setting,
|
||||
)
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(tmp_path: Any) -> SQLiteBackend:
|
||||
return SQLiteBackend(str(tmp_path / "models_changed.db"))
|
||||
|
||||
|
||||
def _seed(storage: SQLiteBackend, *, definition_id: str, alias: str) -> None:
|
||||
storage.create_model_definition(
|
||||
definition_id=definition_id,
|
||||
alias=alias,
|
||||
model="model-x",
|
||||
provider="openai-compatible",
|
||||
base_url="http://localhost:8000/v1",
|
||||
api_key="sk-test",
|
||||
context_window=8192,
|
||||
capabilities="{}",
|
||||
enabled=True,
|
||||
created_by="admin",
|
||||
)
|
||||
|
||||
|
||||
def _make_client(storage: SQLiteBackend) -> tuple[TestClient, MagicMock]:
|
||||
"""Build a TestClient + return the stub collector for assertion.
|
||||
|
||||
Wires the four model-definition CRUD/reload routes plus the two
|
||||
settings mutation routes. Collector is a MagicMock so each
|
||||
``emit_models_changed`` call lands as a recorded call without
|
||||
spinning up the full SSE listener queue.
|
||||
"""
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route(
|
||||
"/v1/api/admin/model-definitions",
|
||||
admin_create_model_definition,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/v1/api/admin/model-definitions/reload",
|
||||
admin_model_reload,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/v1/api/admin/model-definitions/{definition_id}",
|
||||
admin_update_model_definition,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route(
|
||||
"/v1/api/admin/model-definitions/{definition_id}",
|
||||
admin_delete_model_definition,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
Route(
|
||||
"/v1/api/admin/settings/{key:path}",
|
||||
admin_update_setting,
|
||||
methods=["PUT"],
|
||||
),
|
||||
Route(
|
||||
"/v1/api/admin/settings/{key:path}",
|
||||
admin_delete_setting,
|
||||
methods=["DELETE"],
|
||||
),
|
||||
],
|
||||
middleware=[Middleware(_AuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
app.state.coord_registry = None # CRUD endpoints handle this gracefully
|
||||
collector = MagicMock()
|
||||
collector.get_all_nodes.return_value = []
|
||||
app.state.collector = collector
|
||||
app.state.proxy_client = MagicMock()
|
||||
app.state.config_store = MagicMock()
|
||||
client = TestClient(app)
|
||||
client.headers.update(
|
||||
{
|
||||
"X-Test-User": "admin",
|
||||
"X-Test-Perms": "admin.models,admin.settings",
|
||||
}
|
||||
)
|
||||
return client, collector
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Model-definition CRUD endpoints fan out ``models_changed``
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_emits_models_changed(storage: SQLiteBackend) -> None:
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.post(
|
||||
"/v1/api/admin/model-definitions",
|
||||
json={
|
||||
"alias": "fast",
|
||||
"model": "fast-model",
|
||||
"provider": "openai-compatible",
|
||||
"base_url": "http://localhost:9000/v1",
|
||||
"api_key": "sk-x",
|
||||
"context_window": 4096,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
def test_update_emits_models_changed(storage: SQLiteBackend) -> None:
|
||||
_seed(storage, definition_id="m1", alias="local")
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.put(
|
||||
"/v1/api/admin/model-definitions/m1",
|
||||
json={"model": "swapped-model"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
def test_update_with_empty_body_does_not_emit(storage: SQLiteBackend) -> None:
|
||||
"""Empty-body PUT writes no rows + skips the registry refresh — no
|
||||
SSE fanout either, since nothing actually changed."""
|
||||
_seed(storage, definition_id="m1", alias="local")
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.put("/v1/api/admin/model-definitions/m1", json={})
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 0
|
||||
|
||||
|
||||
def test_delete_emits_models_changed(storage: SQLiteBackend) -> None:
|
||||
_seed(storage, definition_id="m1", alias="local")
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.delete("/v1/api/admin/model-definitions/m1")
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
def test_reload_emits_models_changed(storage: SQLiteBackend) -> None:
|
||||
_seed(storage, definition_id="m1", alias="local")
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.post("/v1/api/admin/model-definitions/reload")
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Settings PUT / DELETE only emit for model-affecting keys
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# Pinned snapshot of the role-related keys we expect the allowlist to
|
||||
# cover today. The frozenset itself is asserted further down so a
|
||||
# stray addition doesn't silently bypass coverage.
|
||||
_EXPECTED_AFFECTING_KEYS = frozenset(
|
||||
{
|
||||
"model.default_alias",
|
||||
"model.plan_alias",
|
||||
"model.plan_effort",
|
||||
"model.task_alias",
|
||||
"model.task_effort",
|
||||
"coordinator.model_alias",
|
||||
"coordinator.reasoning_effort",
|
||||
"judge.model",
|
||||
"channels.default_model_alias",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _value_for_key(key: str) -> str:
|
||||
"""Return a registry-valid value for ``key``.
|
||||
|
||||
``reasoning_effort`` keys have a fixed choice list; alias-shaped
|
||||
keys accept arbitrary strings. Avoids per-key custom payloads.
|
||||
"""
|
||||
if (
|
||||
key.endswith("reasoning_effort")
|
||||
or key.endswith("plan_effort")
|
||||
or key.endswith("task_effort")
|
||||
):
|
||||
return "low"
|
||||
return "anything"
|
||||
|
||||
|
||||
def test_affecting_keys_set_matches_expected() -> None:
|
||||
"""Lock in the allowlist so an unintentional removal is caught."""
|
||||
assert _MODEL_AFFECTING_SETTING_KEYS == _EXPECTED_AFFECTING_KEYS
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key", sorted(_EXPECTED_AFFECTING_KEYS))
|
||||
def test_settings_put_emits_for_model_affecting_key(storage: SQLiteBackend, key: str) -> None:
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.put(
|
||||
f"/v1/api/admin/settings/{key}",
|
||||
json={"value": _value_for_key(key)},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key", sorted(_EXPECTED_AFFECTING_KEYS))
|
||||
def test_settings_delete_emits_for_model_affecting_key(storage: SQLiteBackend, key: str) -> None:
|
||||
client, collector = _make_client(storage)
|
||||
# Seed a row so DELETE has something to remove (otherwise 404).
|
||||
client.put(
|
||||
f"/v1/api/admin/settings/{key}",
|
||||
json={"value": _value_for_key(key)},
|
||||
)
|
||||
collector.emit_models_changed.reset_mock()
|
||||
resp = client.delete(f"/v1/api/admin/settings/{key}")
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 1
|
||||
|
||||
|
||||
def test_settings_put_does_not_emit_for_unrelated_key(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Updating a non-model setting (here: a session retention knob)
|
||||
must not trigger a cluster-wide dropdown refresh."""
|
||||
client, collector = _make_client(storage)
|
||||
resp = client.put(
|
||||
"/v1/api/admin/settings/session.retention_days",
|
||||
json={"value": 30},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 0
|
||||
|
||||
|
||||
def test_settings_delete_does_not_emit_for_unrelated_key(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
client, collector = _make_client(storage)
|
||||
client.put(
|
||||
"/v1/api/admin/settings/session.retention_days",
|
||||
json={"value": 30},
|
||||
)
|
||||
collector.emit_models_changed.reset_mock()
|
||||
resp = client.delete("/v1/api/admin/settings/session.retention_days")
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert collector.emit_models_changed.call_count == 0
|
||||
@@ -0,0 +1,400 @@
|
||||
"""Tests for the console-side ``NotifyDispatcher``.
|
||||
|
||||
Exercises the dispatcher against the SQLite synthetic-sweep path so the
|
||||
suite runs without a Postgres dependency. The PG path is shaped the
|
||||
same way (same handler invocation semantics) — the only difference is
|
||||
the underlying stream's wake-up source, which is covered separately in
|
||||
``test_storage_notify.py::TestPostgresNotify``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def dispatcher_factory(storage):
|
||||
"""Yield a factory that constructs + tracks dispatchers for teardown."""
|
||||
from turnstone.console.notify_dispatcher import NotifyDispatcher
|
||||
|
||||
created: list[NotifyDispatcher] = []
|
||||
|
||||
def _make(*, channels: list[str]) -> NotifyDispatcher:
|
||||
d = NotifyDispatcher(storage, channels=channels)
|
||||
created.append(d)
|
||||
return d
|
||||
|
||||
yield _make
|
||||
|
||||
for d in created:
|
||||
d.stop(timeout=2.0)
|
||||
|
||||
|
||||
def _wait_for(predicate, deadline_sec: float = 3.0) -> bool:
|
||||
"""Poll ``predicate`` until True or timeout. Returns bool."""
|
||||
deadline = time.monotonic() + deadline_sec
|
||||
while time.monotonic() < deadline:
|
||||
if predicate():
|
||||
return True
|
||||
time.sleep(0.02)
|
||||
return False
|
||||
|
||||
|
||||
def _start_ready(d, *, timeout: float = 5.0) -> None:
|
||||
"""``d.start()`` + assert the listener is actually listening.
|
||||
|
||||
Closes the start-vs-notify race for backends where ``storage.listen``
|
||||
blocks on the network (Postgres ``LISTEN`` over a fresh psycopg
|
||||
connection): without the sync, a same-thread ``storage.notify`` can
|
||||
fire before the LISTEN registers and the notification is lost.
|
||||
"""
|
||||
d.start()
|
||||
if not d.wait_until_ready(timeout=timeout):
|
||||
msg = f"dispatcher listener did not open within {timeout}s"
|
||||
raise AssertionError(msg)
|
||||
|
||||
|
||||
class TestSubscribe:
|
||||
def test_subscribe_registers_handler(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
seen: list = []
|
||||
d.subscribe("alpha", lambda n: seen.append(n))
|
||||
_start_ready(d)
|
||||
# Fire a notify via the storage layer — dispatcher delivers to handler.
|
||||
storage.notify("alpha", "hello")
|
||||
assert _wait_for(lambda: any(n.payload == "hello" for n in seen))
|
||||
|
||||
def test_subscribe_undeclared_channel_raises(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
with pytest.raises(ValueError, match="not declared"):
|
||||
d.subscribe("beta", lambda n: None)
|
||||
|
||||
def test_subscribe_returns_unsubscribe_callable(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
seen: list = []
|
||||
unsub = d.subscribe("alpha", lambda n: seen.append(n))
|
||||
_start_ready(d)
|
||||
storage.notify("alpha", "first")
|
||||
assert _wait_for(lambda: any(n.payload == "first" for n in seen))
|
||||
unsub()
|
||||
# After unsubscribe, the handler no longer fires. Drain old hits
|
||||
# so the next notify-vs-handler-count check is unambiguous.
|
||||
seen.clear()
|
||||
storage.notify("alpha", "second")
|
||||
# Give the dispatcher a beat to deliver if it were going to.
|
||||
time.sleep(0.2)
|
||||
assert not any(n.payload == "second" for n in seen)
|
||||
|
||||
def test_construction_requires_at_least_one_channel(self, storage):
|
||||
from turnstone.console.notify_dispatcher import NotifyDispatcher
|
||||
|
||||
with pytest.raises(ValueError, match="at least one"):
|
||||
NotifyDispatcher(storage, channels=[])
|
||||
|
||||
def test_duplicate_channels_deduplicated(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha", "alpha", "beta"])
|
||||
assert d.channels == ["alpha", "beta"]
|
||||
|
||||
|
||||
class TestDispatch:
|
||||
def test_multiple_handlers_each_invoked(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
seen_a: list = []
|
||||
seen_b: list = []
|
||||
d.subscribe("alpha", lambda n: seen_a.append(n))
|
||||
d.subscribe("alpha", lambda n: seen_b.append(n))
|
||||
_start_ready(d)
|
||||
storage.notify("alpha", "shared")
|
||||
assert _wait_for(lambda: seen_a and seen_b)
|
||||
assert seen_a[0].payload == "shared"
|
||||
assert seen_b[0].payload == "shared"
|
||||
|
||||
def test_handler_exception_does_not_break_dispatch(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
survived: list = []
|
||||
|
||||
def _broken(_n):
|
||||
msg = "boom"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
d.subscribe("alpha", _broken)
|
||||
d.subscribe("alpha", lambda n: survived.append(n))
|
||||
_start_ready(d)
|
||||
storage.notify("alpha", "after_broken")
|
||||
# The second handler runs even though the first raised.
|
||||
assert _wait_for(lambda: any(n.payload == "after_broken" for n in survived))
|
||||
|
||||
def test_dispatch_filters_by_channel(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha", "beta"])
|
||||
seen_a: list = []
|
||||
seen_b: list = []
|
||||
d.subscribe("alpha", lambda n: seen_a.append(n))
|
||||
d.subscribe("beta", lambda n: seen_b.append(n))
|
||||
_start_ready(d)
|
||||
storage.notify("alpha", "for_a")
|
||||
storage.notify("beta", "for_b")
|
||||
assert _wait_for(lambda: seen_a and seen_b)
|
||||
assert all(n.payload == "for_a" for n in seen_a)
|
||||
assert all(n.payload == "for_b" for n in seen_b)
|
||||
|
||||
|
||||
class TestReconnect:
|
||||
"""Reconnect + synthetic ``reconcile`` notify on stream-open success.
|
||||
|
||||
Uses a stub storage that owns its own listen stream so the test can
|
||||
drive a controlled stream-error sequence — the SQLite path can't
|
||||
raise :class:`NotifyConnectionError`, and the PG path requires a
|
||||
real database outage to exercise this code, neither of which fits a
|
||||
unit test. The dispatcher's threading and reconcile-pending logic
|
||||
are storage-agnostic — the dispatcher sees the same
|
||||
:class:`NotifyStream` Protocol regardless of backend.
|
||||
"""
|
||||
|
||||
def test_reconcile_fires_after_reopen_not_before(self):
|
||||
from turnstone.console.notify_dispatcher import NotifyDispatcher
|
||||
from turnstone.core.storage._notify import Notify, NotifyConnectionError
|
||||
|
||||
# State machine: open -> first poll raises NotifyConnectionError
|
||||
# -> dispatcher waits backoff then reopens -> second open's first
|
||||
# poll blocks forever (test stops the dispatcher before then).
|
||||
# The fix: synthetic reconcile fires AFTER the second open
|
||||
# succeeds, not after the first open fails.
|
||||
sequence: list[str] = []
|
||||
reopen_event = threading.Event()
|
||||
|
||||
class _StubStream:
|
||||
def __init__(self, fail_first_poll: bool):
|
||||
self._fail = fail_first_poll
|
||||
self._closed = False
|
||||
|
||||
def poll(self, _timeout):
|
||||
if self._closed:
|
||||
return []
|
||||
if self._fail:
|
||||
self._fail = False
|
||||
sequence.append("poll_raises")
|
||||
msg = "fake-disconnect"
|
||||
raise NotifyConnectionError(msg)
|
||||
sequence.append("poll_returns")
|
||||
# Block until close to simulate a quiet steady-state.
|
||||
time.sleep(0.5)
|
||||
return []
|
||||
|
||||
def close(self):
|
||||
self._closed = True
|
||||
|
||||
class _StubStorage:
|
||||
def __init__(self):
|
||||
self._open_count = 0
|
||||
|
||||
def listen(self, _channels):
|
||||
import contextlib as _contextlib
|
||||
|
||||
@_contextlib.contextmanager
|
||||
def _cm():
|
||||
self._open_count += 1
|
||||
sequence.append(f"open_{self._open_count}")
|
||||
if self._open_count == 2:
|
||||
reopen_event.set()
|
||||
stream = _StubStream(fail_first_poll=(self._open_count == 1))
|
||||
try:
|
||||
yield stream
|
||||
finally:
|
||||
stream.close()
|
||||
|
||||
return _cm()
|
||||
|
||||
# Speed up backoff so the reopen happens promptly in the test.
|
||||
import turnstone.console.notify_dispatcher as nd_mod
|
||||
|
||||
original_backoff = nd_mod._RECONNECT_BACKOFF_INITIAL
|
||||
nd_mod._RECONNECT_BACKOFF_INITIAL = 0.05
|
||||
try:
|
||||
d = NotifyDispatcher(_StubStorage(), channels=["alpha"])
|
||||
got: list[Notify] = []
|
||||
d.subscribe("alpha", lambda n: got.append(n))
|
||||
d.start()
|
||||
try:
|
||||
# Wait for the second open (post-reconnect).
|
||||
assert reopen_event.wait(3.0), "dispatcher did not reopen after disconnect"
|
||||
# Reconcile should be delivered shortly after the reopen.
|
||||
deadline = time.monotonic() + 2.0
|
||||
while time.monotonic() < deadline:
|
||||
if any(n.payload == "reconcile" for n in got):
|
||||
break
|
||||
time.sleep(0.02)
|
||||
assert any(n.payload == "reconcile" for n in got), (
|
||||
f"no reconcile delivered; sequence={sequence}, got={got}"
|
||||
)
|
||||
# The reconcile must NOT fire before the second open —
|
||||
# if it did, the index of 'open_2' in sequence would
|
||||
# come after any reconcile-emitting work. Check ordering:
|
||||
# 'open_1' < 'poll_raises' < 'open_2' (synthesize happens
|
||||
# inside the with-block of the SECOND open).
|
||||
ix_open_1 = sequence.index("open_1")
|
||||
ix_raises = sequence.index("poll_raises")
|
||||
ix_open_2 = sequence.index("open_2")
|
||||
assert ix_open_1 < ix_raises < ix_open_2
|
||||
finally:
|
||||
d.stop(timeout=2.0)
|
||||
finally:
|
||||
nd_mod._RECONNECT_BACKOFF_INITIAL = original_backoff
|
||||
|
||||
def test_generic_exception_path_also_synthesizes_reconcile(self):
|
||||
"""Exceptions thrown during ``listen()`` (not via stream.poll) still trigger reconcile.
|
||||
|
||||
Models the ``psycopg.connect()`` / initial ``LISTEN`` failure
|
||||
shape, which doesn't go through the stream's exception
|
||||
translator and would hit the generic ``except Exception``
|
||||
branch. Pre-fix, that branch emitted no reconcile.
|
||||
"""
|
||||
from turnstone.console.notify_dispatcher import NotifyDispatcher
|
||||
|
||||
reopen_event = threading.Event()
|
||||
|
||||
class _StubStream:
|
||||
def __init__(self):
|
||||
self._closed = False
|
||||
|
||||
def poll(self, _timeout):
|
||||
if self._closed:
|
||||
return []
|
||||
time.sleep(0.5)
|
||||
return []
|
||||
|
||||
def close(self):
|
||||
self._closed = True
|
||||
|
||||
class _StubStorage:
|
||||
def __init__(self):
|
||||
self._open_count = 0
|
||||
|
||||
def listen(self, _channels):
|
||||
import contextlib as _contextlib
|
||||
|
||||
self._open_count += 1
|
||||
if self._open_count == 1:
|
||||
# First open raises a generic exception (e.g.
|
||||
# ``psycopg.OperationalError`` from a failed connect)
|
||||
# — landing in the dispatcher's generic except branch.
|
||||
msg = "fake-connect-failure"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
@_contextlib.contextmanager
|
||||
def _cm():
|
||||
reopen_event.set()
|
||||
stream = _StubStream()
|
||||
try:
|
||||
yield stream
|
||||
finally:
|
||||
stream.close()
|
||||
|
||||
return _cm()
|
||||
|
||||
import turnstone.console.notify_dispatcher as nd_mod
|
||||
|
||||
original_backoff = nd_mod._RECONNECT_BACKOFF_INITIAL
|
||||
nd_mod._RECONNECT_BACKOFF_INITIAL = 0.05
|
||||
try:
|
||||
d = NotifyDispatcher(_StubStorage(), channels=["alpha"])
|
||||
got: list = []
|
||||
d.subscribe("alpha", lambda n: got.append(n))
|
||||
d.start()
|
||||
try:
|
||||
assert reopen_event.wait(3.0), "dispatcher did not reopen after generic exception"
|
||||
deadline = time.monotonic() + 2.0
|
||||
while time.monotonic() < deadline:
|
||||
if any(n.payload == "reconcile" for n in got):
|
||||
break
|
||||
time.sleep(0.02)
|
||||
assert any(n.payload == "reconcile" for n in got), (
|
||||
"no reconcile delivered after generic-exception recovery"
|
||||
)
|
||||
finally:
|
||||
d.stop(timeout=2.0)
|
||||
finally:
|
||||
nd_mod._RECONNECT_BACKOFF_INITIAL = original_backoff
|
||||
|
||||
|
||||
class TestCoalescing:
|
||||
"""Same-channel burst collapses to one handler invocation per batch."""
|
||||
|
||||
def test_burst_coalesces_to_one_handler_call_per_channel(self, dispatcher_factory, storage):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
invocations: list = []
|
||||
# Slow handler to ensure all bursts queue up before the first
|
||||
# call returns — gives the dispatch loop time to coalesce.
|
||||
coalesce_gate = threading.Event()
|
||||
|
||||
def _slow_handler(n):
|
||||
invocations.append(n)
|
||||
coalesce_gate.wait(0.05)
|
||||
|
||||
d.subscribe("alpha", _slow_handler)
|
||||
_start_ready(d)
|
||||
# Burst of 10 notifies on the same channel — should coalesce
|
||||
# down to many fewer handler invocations.
|
||||
for i in range(10):
|
||||
storage.notify("alpha", str(i))
|
||||
# Wait until the dispatch settles (handler is called at least once
|
||||
# and the queue empties).
|
||||
deadline = time.monotonic() + 2.0
|
||||
while time.monotonic() < deadline:
|
||||
if invocations and d._dispatch_queue.empty():
|
||||
time.sleep(0.1) # allow any final coalesced call to land
|
||||
break
|
||||
time.sleep(0.02)
|
||||
coalesce_gate.set()
|
||||
# At least one handler call; well fewer than 10 (coalescing
|
||||
# collapsed the burst). Exact count depends on timing — typical
|
||||
# is 1-2 invocations per burst on a fast machine.
|
||||
assert invocations, "handler never fired"
|
||||
assert len(invocations) < 10, (
|
||||
f"expected coalescing to collapse burst of 10; got {len(invocations)} invocations"
|
||||
)
|
||||
|
||||
|
||||
class TestLifecycle:
|
||||
def test_start_is_idempotent(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
d.start()
|
||||
d.start() # No-op, no thread doubling
|
||||
# Single listener + single dispatch thread are spawned regardless.
|
||||
# Inspect by name so we don't depend on the exact thread count of
|
||||
# the test runner.
|
||||
listener_threads = [
|
||||
t for t in threading.enumerate() if t.name == "notify-dispatcher-listener"
|
||||
]
|
||||
dispatch_threads = [
|
||||
t for t in threading.enumerate() if t.name == "notify-dispatcher-dispatch"
|
||||
]
|
||||
assert len(listener_threads) == 1
|
||||
assert len(dispatch_threads) == 1
|
||||
|
||||
def test_stop_is_idempotent(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
d.start()
|
||||
d.stop(timeout=2.0)
|
||||
d.stop(timeout=2.0) # No-op, no error
|
||||
|
||||
def test_stop_without_start_is_noop(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
d.stop(timeout=1.0) # No-op, no thread to join
|
||||
|
||||
def test_stop_joins_threads(self, dispatcher_factory):
|
||||
d = dispatcher_factory(channels=["alpha"])
|
||||
d.start()
|
||||
# Capture thread references then stop and assert they exited.
|
||||
threads_before = [
|
||||
t
|
||||
for t in threading.enumerate()
|
||||
if t.name in {"notify-dispatcher-listener", "notify-dispatcher-dispatch"}
|
||||
]
|
||||
assert threads_before
|
||||
d.stop(timeout=3.0)
|
||||
time.sleep(0.05)
|
||||
for t in threads_before:
|
||||
assert not t.is_alive(), f"{t.name} still alive after stop"
|
||||
@@ -0,0 +1,562 @@
|
||||
"""Unit tests for :class:`NudgeQueue`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.nudge_queue import TOOL_DRAIN, USER_DRAIN, NudgeQueue
|
||||
|
||||
|
||||
class TestEnqueueDrain:
|
||||
def test_enqueue_drain_fifo_order(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
q.enqueue("c", "3", "any")
|
||||
# Drain everything regardless of channel — preserves insertion order.
|
||||
out = q.drain({"user", "tool", "any"})
|
||||
# Drain returns ``(nudge_type, text, metadata)``; producers
|
||||
# without ``metadata`` see ``None`` in the third slot.
|
||||
assert out == [("a", "1", None), ("b", "2", None), ("c", "3", None)]
|
||||
assert len(q) == 0
|
||||
|
||||
def test_drain_filter_keeps_non_matching(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "x", "user")
|
||||
q.enqueue("b", "y", "tool")
|
||||
# Drain only user → tool entry stays.
|
||||
out = q.drain(USER_DRAIN)
|
||||
assert out == [("a", "x", None)]
|
||||
assert len(q) == 1
|
||||
# Now drain tool — gets the remaining entry.
|
||||
out = q.drain(TOOL_DRAIN)
|
||||
assert out == [("b", "y", None)]
|
||||
assert len(q) == 0
|
||||
|
||||
def test_any_channel_drains_on_either_seam(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("c", "z", "any")
|
||||
# User-seam drain pulls "any".
|
||||
assert q.drain(USER_DRAIN) == [("c", "z", None)]
|
||||
assert len(q) == 0
|
||||
# Re-enqueue and prove tool-seam also drains "any".
|
||||
q.enqueue("d", "w", "any")
|
||||
assert q.drain(TOOL_DRAIN) == [("d", "w", None)]
|
||||
assert len(q) == 0
|
||||
|
||||
def test_drain_empty_filter_no_op(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
# Empty filter drains nothing.
|
||||
assert q.drain(set()) == []
|
||||
assert len(q) == 1
|
||||
|
||||
def test_drain_empty_queue_returns_empty_list(self):
|
||||
q = NudgeQueue()
|
||||
# Fast-path: no items → no kept-deque allocation, just `[]`.
|
||||
assert q.drain(USER_DRAIN) == []
|
||||
assert q.drain({"user", "tool", "any"}) == []
|
||||
assert len(q) == 0
|
||||
|
||||
def test_drain_preserves_order_across_partial_drain(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
q.enqueue("c", "3", "user")
|
||||
q.enqueue("d", "4", "tool")
|
||||
# Drain user — should get "a" then "c" in order; "b","d" stay.
|
||||
assert q.drain({"user"}) == [("a", "1", None), ("c", "3", None)]
|
||||
# Tool drain follows insertion order on remaining.
|
||||
assert q.drain({"tool"}) == [("b", "2", None), ("d", "4", None)]
|
||||
|
||||
|
||||
class TestLenAndClear:
|
||||
def test_len_does_not_mutate(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
assert len(q) == 1
|
||||
assert len(q) == 1 # second call still 1; not consumed
|
||||
assert q.pending() == [("a", "1")]
|
||||
|
||||
def test_len_empty_is_zero(self):
|
||||
q = NudgeQueue()
|
||||
assert len(q) == 0
|
||||
|
||||
def test_clear_returns_count(self):
|
||||
q = NudgeQueue()
|
||||
assert q.clear() == 0
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
q.enqueue("c", "3", "any")
|
||||
assert q.clear() == 3
|
||||
assert len(q) == 0
|
||||
|
||||
def test_clear_empty_returns_zero(self):
|
||||
q = NudgeQueue()
|
||||
assert q.clear() == 0
|
||||
|
||||
|
||||
class TestDropOldestByType:
|
||||
def test_drop_oldest_by_type_removes_earliest_match(self):
|
||||
"""Drop the FIRST entry of the matching type; later matches stay."""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("other", "first", "any")
|
||||
q.enqueue("target", "older", "any")
|
||||
q.enqueue("target", "newer", "any")
|
||||
# "older" is the earliest target — drop it.
|
||||
assert q.drop_oldest_by_type("target") is True
|
||||
assert q.pending() == [("other", "first"), ("target", "newer")]
|
||||
|
||||
def test_drop_oldest_by_type_no_match_returns_false(self):
|
||||
"""Empty queue and unmatched-type cases both return False."""
|
||||
q = NudgeQueue()
|
||||
# Empty.
|
||||
assert q.drop_oldest_by_type("target") is False
|
||||
# Non-matching items only.
|
||||
q.enqueue("other", "1", "any")
|
||||
q.enqueue("other", "2", "tool")
|
||||
assert q.drop_oldest_by_type("target") is False
|
||||
# Queue is unaffected.
|
||||
assert q.pending() == [("other", "1"), ("other", "2")]
|
||||
|
||||
def test_drop_oldest_by_type_only_drops_one(self):
|
||||
"""Multiple matching entries → only the first is removed."""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("target", "1", "any")
|
||||
q.enqueue("target", "2", "any")
|
||||
q.enqueue("target", "3", "any")
|
||||
assert q.drop_oldest_by_type("target") is True
|
||||
assert q.pending() == [("target", "2"), ("target", "3")]
|
||||
|
||||
def test_drop_oldest_by_type_channel_filter(self):
|
||||
"""With ``channel`` set, drop walks only that channel. Pairs with
|
||||
:meth:`count_by_type(..., channel=...)` so producer-side soft caps
|
||||
operate on a consistent entry set.
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("target", "user-1", "user")
|
||||
q.enqueue("target", "any-1", "any")
|
||||
q.enqueue("target", "any-2", "any")
|
||||
# Drop the oldest "any"-channel target — leaves the user one
|
||||
# untouched even though it's earlier in insertion order.
|
||||
assert q.drop_oldest_by_type("target", channel="any") is True
|
||||
assert q.pending() == [
|
||||
("target", "user-1"),
|
||||
("target", "any-2"),
|
||||
]
|
||||
# And a channel with no matches returns False without touching
|
||||
# the queue.
|
||||
assert q.drop_oldest_by_type("target", channel="tool") is False
|
||||
assert q.pending() == [
|
||||
("target", "user-1"),
|
||||
("target", "any-2"),
|
||||
]
|
||||
|
||||
|
||||
class TestCapAtOrDropOldest:
|
||||
def test_below_cap_no_drop(self):
|
||||
q = NudgeQueue()
|
||||
for i in range(3):
|
||||
q.enqueue("target", f"t-{i}", "any")
|
||||
# 3 entries, cap=5 → no drop.
|
||||
assert q.cap_at_or_drop_oldest("target", 5, channel="any") is False
|
||||
assert q.count_by_type("target") == 3
|
||||
|
||||
def test_at_cap_drops_oldest(self):
|
||||
q = NudgeQueue()
|
||||
for i in range(5):
|
||||
q.enqueue("target", f"t-{i}", "any")
|
||||
# 5 entries, cap=5 → drop the oldest ("t-0"), leaving 4.
|
||||
assert q.cap_at_or_drop_oldest("target", 5, channel="any") is True
|
||||
remaining = q.pending()
|
||||
assert ("target", "t-0") not in remaining
|
||||
assert len(remaining) == 4
|
||||
assert remaining[0] == ("target", "t-1") # FIFO drop-oldest preserved
|
||||
|
||||
def test_above_cap_drops_only_one(self):
|
||||
q = NudgeQueue()
|
||||
for i in range(7):
|
||||
q.enqueue("target", f"t-{i}", "any")
|
||||
# 7 entries, cap=5 → drop only ONE per call (soft-cap regulates over time).
|
||||
assert q.cap_at_or_drop_oldest("target", 5, channel="any") is True
|
||||
assert q.count_by_type("target") == 6
|
||||
|
||||
def test_channel_filter_respected(self):
|
||||
q = NudgeQueue()
|
||||
for i in range(3):
|
||||
q.enqueue("target", f"any-{i}", "any")
|
||||
for i in range(3):
|
||||
q.enqueue("target", f"user-{i}", "user")
|
||||
# 3 "any"-channel entries; cap=3 on channel="any" → drop oldest "any" only.
|
||||
assert q.cap_at_or_drop_oldest("target", 3, channel="any") is True
|
||||
# User-channel entries untouched.
|
||||
assert q.count_by_type("target", channel="user") == 3
|
||||
assert q.count_by_type("target", channel="any") == 2
|
||||
|
||||
def test_other_types_ignored(self):
|
||||
q = NudgeQueue()
|
||||
for i in range(5):
|
||||
q.enqueue("other", f"o-{i}", "any")
|
||||
q.enqueue("target", "t-0", "any")
|
||||
# Only one "target" entry; cap=1 on "target" → drop it. "other"
|
||||
# entries are untouched even though queue holds 6 total.
|
||||
assert q.cap_at_or_drop_oldest("target", 1, channel="any") is True
|
||||
assert q.count_by_type("target") == 0
|
||||
assert q.count_by_type("other") == 5
|
||||
|
||||
def test_zero_or_negative_cap_no_op(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("target", "t-0", "any")
|
||||
assert q.cap_at_or_drop_oldest("target", 0, channel="any") is False
|
||||
assert q.cap_at_or_drop_oldest("target", -1, channel="any") is False
|
||||
assert q.count_by_type("target") == 1
|
||||
|
||||
def test_no_match_returns_false(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("other", "o-0", "any")
|
||||
assert q.cap_at_or_drop_oldest("target", 1, channel="any") is False
|
||||
assert q.count_by_type("other") == 1
|
||||
|
||||
|
||||
class TestCountByType:
|
||||
def test_count_by_type_no_channel(self):
|
||||
"""Count across all channels with ``channel=None``."""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("target", "1", "user")
|
||||
q.enqueue("other", "x", "any")
|
||||
q.enqueue("target", "2", "any")
|
||||
q.enqueue("target", "3", "tool")
|
||||
assert q.count_by_type("target") == 3
|
||||
assert q.count_by_type("other") == 1
|
||||
assert q.count_by_type("missing") == 0
|
||||
|
||||
def test_count_by_type_with_channel_filter(self):
|
||||
"""Filter narrows the count to one channel — used by producer-side
|
||||
soft caps that pair with ``drop_oldest_by_type(..., channel=...)``.
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("target", "u-1", "user")
|
||||
q.enqueue("target", "a-1", "any")
|
||||
q.enqueue("target", "a-2", "any")
|
||||
q.enqueue("target", "t-1", "tool")
|
||||
assert q.count_by_type("target", channel="any") == 2
|
||||
assert q.count_by_type("target", channel="user") == 1
|
||||
assert q.count_by_type("target", channel="tool") == 1
|
||||
|
||||
def test_count_by_type_empty_queue(self):
|
||||
q = NudgeQueue()
|
||||
assert q.count_by_type("anything") == 0
|
||||
assert q.count_by_type("anything", channel="any") == 0
|
||||
|
||||
|
||||
class TestPending:
|
||||
def test_pending_no_filter_returns_all_in_order(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
q.enqueue("c", "3", "any")
|
||||
# All three, in insertion order, as (nudge_type, text) tuples.
|
||||
assert q.pending() == [("a", "1"), ("b", "2"), ("c", "3")]
|
||||
|
||||
def test_pending_channel_filter(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
q.enqueue("c", "3", "user")
|
||||
q.enqueue("d", "4", "any")
|
||||
assert q.pending("user") == [("a", "1"), ("c", "3")]
|
||||
assert q.pending("tool") == [("b", "2")]
|
||||
assert q.pending("any") == [("d", "4")]
|
||||
|
||||
def test_pending_does_not_mutate(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
# Two pending calls return same content; nothing consumed.
|
||||
first = q.pending()
|
||||
second = q.pending()
|
||||
assert first == second
|
||||
assert len(q) == 2
|
||||
|
||||
|
||||
class TestMetadata:
|
||||
"""Producer-supplied ``metadata`` rides alongside ``(type, text)`` on
|
||||
drain. Today only ``watch_triggered`` populates it; the wire shape
|
||||
accommodates future producers (e.g. structured tool_error context)
|
||||
without another schema bump.
|
||||
"""
|
||||
|
||||
def test_drain_returns_metadata_when_set(self):
|
||||
q = NudgeQueue()
|
||||
meta = {"watch_name": "w1", "command": "ls", "poll_count": 2}
|
||||
q.enqueue("watch_triggered", "$ ls\nfile.txt", "any", metadata=meta)
|
||||
out = q.drain({"any"})
|
||||
assert out == [("watch_triggered", "$ ls\nfile.txt", meta)]
|
||||
|
||||
def test_drain_returns_none_when_metadata_unset(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("idle_children", "kids", "any") # no metadata kwarg
|
||||
out = q.drain({"any"})
|
||||
assert out == [("idle_children", "kids", None)]
|
||||
|
||||
def test_pending_with_metadata_projects_third_field(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("watch_triggered", "out", "any", metadata={"watch_name": "w"})
|
||||
snapshot = q.pending_with_metadata()
|
||||
assert snapshot == [
|
||||
("a", "1", None),
|
||||
("watch_triggered", "out", {"watch_name": "w"}),
|
||||
]
|
||||
# ``pending`` (without metadata) keeps the legacy 2-tuple shape.
|
||||
assert q.pending() == [("a", "1"), ("watch_triggered", "out")]
|
||||
|
||||
def test_metadata_survives_partial_drain(self):
|
||||
"""A ``user``-channel drain leaves an unaffected ``tool``-channel
|
||||
entry — its metadata must still be present on the next drain."""
|
||||
q = NudgeQueue()
|
||||
q.enqueue("user_thing", "u", "user")
|
||||
q.enqueue("watch_triggered", "w-out", "tool", metadata={"watch_name": "w1"})
|
||||
# User drain doesn't touch the tool entry.
|
||||
assert q.drain({"user"}) == [("user_thing", "u", None)]
|
||||
# Tool drain still has the metadata.
|
||||
assert q.drain({"tool"}) == [
|
||||
("watch_triggered", "w-out", {"watch_name": "w1"}),
|
||||
]
|
||||
|
||||
def test_metadata_with_valid_until_predicate(self):
|
||||
"""Metadata + ``valid_until`` co-exist on the same entry; the
|
||||
predicate gate runs as before, and on a True result the metadata
|
||||
rides the drained tuple.
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
q.enqueue(
|
||||
"watch_triggered",
|
||||
"w-out",
|
||||
"any",
|
||||
valid_until=lambda: True,
|
||||
metadata={"watch_name": "w1", "is_final": True},
|
||||
)
|
||||
out = q.drain({"any"})
|
||||
assert out == [
|
||||
("watch_triggered", "w-out", {"watch_name": "w1", "is_final": True}),
|
||||
]
|
||||
|
||||
|
||||
class TestHasPending:
|
||||
def test_has_pending_returns_false_on_empty_queue(self):
|
||||
q = NudgeQueue()
|
||||
assert q.has_pending({"user", "any"}) is False
|
||||
assert q.has_pending({"tool"}) is False
|
||||
|
||||
def test_has_pending_short_circuits_on_first_match(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "tool")
|
||||
q.enqueue("b", "2", "user")
|
||||
# First entry doesn't match, second does — true after walking 2.
|
||||
assert q.has_pending({"user"}) is True
|
||||
|
||||
def test_has_pending_returns_false_when_no_match(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "tool")
|
||||
q.enqueue("b", "2", "tool")
|
||||
assert q.has_pending({"user", "any"}) is False
|
||||
|
||||
def test_has_pending_matches_any_channel(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "any")
|
||||
# USER_DRAIN-shaped filter pulls "any" entries.
|
||||
assert q.has_pending(USER_DRAIN) is True
|
||||
# TOOL_DRAIN-shaped filter also pulls "any" entries.
|
||||
assert q.has_pending(TOOL_DRAIN) is True
|
||||
|
||||
def test_has_pending_does_not_mutate(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "user")
|
||||
q.enqueue("b", "2", "tool")
|
||||
before = q.pending()
|
||||
q.has_pending({"user"})
|
||||
q.has_pending({"tool"})
|
||||
q.has_pending(set())
|
||||
assert q.pending() == before
|
||||
|
||||
|
||||
class TestValidation:
|
||||
def test_invalid_channel_raises(self):
|
||||
q = NudgeQueue()
|
||||
with pytest.raises(ValueError, match="channel"):
|
||||
q.enqueue("a", "1", "wake") # type: ignore[arg-type]
|
||||
with pytest.raises(ValueError):
|
||||
q.enqueue("b", "2", "") # type: ignore[arg-type]
|
||||
# Queue is unaffected by the failed enqueues.
|
||||
assert len(q) == 0
|
||||
|
||||
def test_channel_is_required(self):
|
||||
q = NudgeQueue()
|
||||
# No default — caller MUST pick a seam consciously.
|
||||
with pytest.raises(TypeError):
|
||||
q.enqueue("a", "1") # type: ignore[call-arg]
|
||||
|
||||
|
||||
class TestValidUntil:
|
||||
"""``valid_until`` predicate: drain re-checks freshness. Falsy
|
||||
predicates drop the entry without delivery and log at ``info``
|
||||
(normal lifecycle outcome); raising predicates drop the entry and
|
||||
log at ``warning`` with ``exc_info`` (misbehaving predicate).
|
||||
"""
|
||||
|
||||
def test_valid_until_true_delivers(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "any", valid_until=lambda: True)
|
||||
out = q.drain({"any"})
|
||||
assert out == [("a", "1", None)]
|
||||
|
||||
def test_valid_until_false_drops_with_info_log(self, caplog: pytest.LogCaptureFixture):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "any", valid_until=lambda: False)
|
||||
with caplog.at_level(logging.INFO, logger="turnstone.core.nudge_queue"):
|
||||
out = q.drain({"any"})
|
||||
assert out == []
|
||||
# Already removed from queue (drain partition removes BEFORE
|
||||
# predicate check — falsy doesn't return to queue).
|
||||
assert len(q) == 0
|
||||
# The drop emits a structured info record so a wiring
|
||||
# regression (a predicate that always returns False) is still
|
||||
# observable, without spamming ``warning`` for the routine
|
||||
# lifecycle case where ``valid_until`` is doing its job.
|
||||
# structlog renders the event name + extras into ``msg`` as a
|
||||
# single rendered string, so substring-match like the
|
||||
# ``watch_dispatch.queue_full`` assertion in
|
||||
# tests/test_watch_dispatch.py.
|
||||
drops = [r for r in caplog.records if "nudge_queue.predicate_dropped" in r.getMessage()]
|
||||
assert len(drops) == 1
|
||||
assert drops[0].levelno == logging.INFO
|
||||
assert "predicate_false" in drops[0].getMessage()
|
||||
assert "'nudge_type': 'a'" in drops[0].getMessage()
|
||||
assert "'channel': 'any'" in drops[0].getMessage()
|
||||
assert "'text_len': 1" in drops[0].getMessage()
|
||||
|
||||
def test_valid_until_exception_drops_with_warning(self, caplog: pytest.LogCaptureFixture):
|
||||
q = NudgeQueue()
|
||||
|
||||
def boom() -> bool:
|
||||
raise RuntimeError("predicate crash")
|
||||
|
||||
q.enqueue("a", "1", "any", valid_until=boom)
|
||||
with caplog.at_level(logging.WARNING, logger="turnstone.core.nudge_queue"):
|
||||
out = q.drain({"any"})
|
||||
assert out == []
|
||||
# Crash-on-predicate is treated as "no longer valid" — drop, not propagate.
|
||||
assert len(q) == 0
|
||||
# Stays at ``warning`` (with ``exc_info``) because a raising
|
||||
# predicate is a bug, not a normal lifecycle outcome.
|
||||
drops = [r for r in caplog.records if "nudge_queue.predicate_dropped" in r.getMessage()]
|
||||
assert len(drops) == 1
|
||||
assert drops[0].levelno == logging.WARNING
|
||||
rendered = drops[0].getMessage()
|
||||
assert "predicate_raised" in rendered
|
||||
assert "RuntimeError" in rendered
|
||||
assert "predicate crash" in rendered
|
||||
|
||||
def test_valid_until_evaluated_outside_lock(self):
|
||||
"""The predicate may do non-trivial work (e.g. storage I/O)
|
||||
without blocking other producers. Verify the predicate runs
|
||||
outside the queue's internal lock by enqueueing from inside
|
||||
the predicate — would deadlock if the lock was still held.
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
|
||||
def reentrant() -> bool:
|
||||
# If the lock is held during predicate eval, this enqueue
|
||||
# would block forever (RLock would let it through, but the
|
||||
# queue uses a plain Lock).
|
||||
q.enqueue("inner", "from-predicate", "any", valid_until=lambda: True)
|
||||
return True
|
||||
|
||||
q.enqueue("outer", "1", "any", valid_until=reentrant)
|
||||
out = q.drain({"any"})
|
||||
# Outer's predicate ran outside the lock, enqueued "inner";
|
||||
# outer's True return delivered "outer". "inner" was enqueued
|
||||
# AFTER the partition snapshot, so it stays in the queue.
|
||||
assert out == [("outer", "1", None)]
|
||||
assert q.pending() == [("inner", "from-predicate")]
|
||||
|
||||
def test_valid_until_only_evaluated_for_matching_channel(self):
|
||||
"""A non-matching entry's predicate must NOT fire — that would
|
||||
be wasted work (or worse, a side-effecting predicate would run
|
||||
when the entry is supposed to stay queued).
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
calls = []
|
||||
|
||||
def track() -> bool:
|
||||
calls.append(1)
|
||||
return True
|
||||
|
||||
# Tool-channel entry; we drain user-channel. Predicate must not run.
|
||||
q.enqueue("a", "1", "tool", valid_until=track)
|
||||
q.drain({"user", "any"})
|
||||
assert calls == []
|
||||
# Entry stays queued.
|
||||
assert q.pending("tool") == [("a", "1")]
|
||||
|
||||
def test_valid_until_default_none_always_delivers(self):
|
||||
# No predicate → entry behaves identically to pre-PR-3 entries.
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "any") # no valid_until kwarg
|
||||
assert q.drain({"any"}) == [("a", "1", None)]
|
||||
|
||||
|
||||
class TestConcurrency:
|
||||
def test_concurrent_enqueue_drain_no_loss(self):
|
||||
"""16 producer threads × 64 nudges = 1024 total; one consumer
|
||||
drains in a loop until producers finish + queue empty. Verify
|
||||
every produced item is observed exactly once.
|
||||
"""
|
||||
q = NudgeQueue()
|
||||
producers = 16
|
||||
per_producer = 64
|
||||
total = producers * per_producer
|
||||
|
||||
produced: set[tuple[str, str]] = set()
|
||||
produced_lock = threading.Lock()
|
||||
observed: list[tuple[str, str]] = []
|
||||
observed_lock = threading.Lock()
|
||||
done_event = threading.Event()
|
||||
|
||||
def produce(pid: int) -> None:
|
||||
for i in range(per_producer):
|
||||
key = (f"p{pid}", f"i{i}")
|
||||
with produced_lock:
|
||||
produced.add(key)
|
||||
q.enqueue(key[0], key[1], "user")
|
||||
|
||||
def consume() -> None:
|
||||
while not done_event.is_set() or len(q) > 0:
|
||||
drained = q.drain({"user"})
|
||||
if drained:
|
||||
with observed_lock:
|
||||
# Drop the trailing ``metadata`` slot — every
|
||||
# entry here was enqueued without metadata, so
|
||||
# the comparison set / count matches the produced
|
||||
# ``(type, text)`` shape.
|
||||
observed.extend((nt, txt) for nt, txt, _meta in drained)
|
||||
|
||||
consumer = threading.Thread(target=consume, daemon=True)
|
||||
consumer.start()
|
||||
threads = [threading.Thread(target=produce, args=(i,)) for i in range(producers)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
done_event.set()
|
||||
consumer.join(timeout=5.0)
|
||||
assert not consumer.is_alive(), "consumer didn't finish in time"
|
||||
|
||||
# Every produced key observed; no duplicates.
|
||||
assert set(observed) == produced
|
||||
assert len(observed) == total
|
||||
assert len(q) == 0
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user