mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-13 23:42:25 -06:00
Compare commits
80 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 |
@@ -546,11 +546,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 +570,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
|
||||
|
||||
@@ -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>>
|
||||
}
|
||||
|
||||
@@ -253,7 +253,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.
|
||||
--
|
||||
|
||||
@@ -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:25b5448bbb7da8ddafe4f65c6c5e6cbcaa9cb9f31746ca46d3a2241bc47b1956
|
||||
size 259687
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7623df33be9baf7647ca1c2450640df57e1cd73e8be1f8168aae16e546ad683c
|
||||
size 459941
|
||||
oid sha256:d6aff446a062aa08f316985d00c2183148694f786d7f22172bc50b30046c728b
|
||||
size 379259
|
||||
|
||||
+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"
|
||||
|
||||
|
||||
+1
-1
@@ -100,7 +100,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 |
|
||||
|
||||
+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:
|
||||
|
||||
+3
-2
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "turnstone"
|
||||
version = "1.5.7"
|
||||
version = "1.5.9"
|
||||
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",
|
||||
]
|
||||
|
||||
|
||||
@@ -26,3 +26,28 @@ def make_chat_session(**overrides: Any) -> Any:
|
||||
}
|
||||
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
|
||||
|
||||
@@ -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())
|
||||
@@ -22,6 +22,8 @@ to lock in:
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -33,6 +35,7 @@ from starlette.testclient import TestClient
|
||||
|
||||
from tests._coord_test_helpers import _AuthMiddleware
|
||||
from turnstone.console.server import (
|
||||
_maybe_bootstrap_coord_subsystem,
|
||||
_refresh_coord_registry,
|
||||
admin_create_model_definition,
|
||||
admin_delete_model_definition,
|
||||
@@ -42,6 +45,29 @@ from turnstone.console.server import (
|
||||
from turnstone.core.model_registry import ModelConfig, ModelRegistry
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
|
||||
def _bootstrap_app(**overrides: Any) -> Any:
|
||||
"""Build a fake ``app`` with the ``state`` attrs the bootstrap helper
|
||||
inspects. Defaults match a freshly-installed console (no coord
|
||||
subsystem yet) with all required prereqs (collector, console_metrics,
|
||||
config_store) populated as MagicMocks. Tests pass overrides to
|
||||
suppress individual prereqs or pre-set ``coord_mgr`` etc.
|
||||
"""
|
||||
state_kwargs: dict[str, Any] = {
|
||||
"coord_mgr": None,
|
||||
"coord_adapter": None,
|
||||
"coord_registry": None,
|
||||
"coord_registry_error": "",
|
||||
"coord_state_writer": None,
|
||||
"coord_idle_observer": None,
|
||||
"config_store": MagicMock(),
|
||||
"collector": MagicMock(),
|
||||
"console_metrics": MagicMock(),
|
||||
}
|
||||
state_kwargs.update(overrides)
|
||||
return SimpleNamespace(state=SimpleNamespace(**state_kwargs))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -213,6 +239,538 @@ def test_helper_preserves_registry_when_no_enabled_rows(storage: SQLiteBackend)
|
||||
assert state.coord_registry.get_config("local").model == "cached-model"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# First-row bootstrap tests — ``_maybe_bootstrap_coord_subsystem`` semantics.
|
||||
# A console booted with no model rows leaves coord_mgr = None; the operator
|
||||
# adding the first row at runtime must promote the subsystem to ready
|
||||
# without a console restart.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_bootstrap_noop_when_coord_mgr_already_built(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Idempotent fast-path — already-bootstrapped subsystem must not
|
||||
re-stand-up a second SessionManager / StateWriter pair."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app(coord_mgr=MagicMock()) # subsystem already built
|
||||
|
||||
calls: list[Any] = []
|
||||
monkeypatch.setattr(
|
||||
server_module,
|
||||
"_bootstrap_coord_subsystem",
|
||||
lambda *a, **kw: calls.append(a),
|
||||
)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
assert calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("missing_attr", ["config_store", "collector", "console_metrics"])
|
||||
def test_bootstrap_noop_when_prerequisites_missing(
|
||||
storage: SQLiteBackend,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
missing_attr: str,
|
||||
) -> None:
|
||||
"""Each strictly-required ``app.state`` attr (config_store, collector,
|
||||
console_metrics) must individually short-circuit the bootstrap to a
|
||||
no-op — partial init / test harnesses don't have the full set, and a
|
||||
CRUD write that already landed mustn't 500 on a missing prereq."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app(**{missing_attr: None})
|
||||
calls: list[Any] = []
|
||||
monkeypatch.setattr(
|
||||
server_module,
|
||||
"_bootstrap_coord_subsystem",
|
||||
lambda *a, **kw: calls.append(a),
|
||||
)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
assert calls == []
|
||||
assert app.state.coord_mgr is None
|
||||
|
||||
|
||||
def test_bootstrap_records_error_when_no_rows(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""All rows disabled (or none seeded) — load_model_registry raises
|
||||
ValueError. Helper records the message on app.state so the
|
||||
coord-endpoint 503 surfaces a current diagnosis instead of a stale
|
||||
one from boot."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
app = _bootstrap_app()
|
||||
calls: list[Any] = []
|
||||
monkeypatch.setattr(
|
||||
server_module,
|
||||
"_bootstrap_coord_subsystem",
|
||||
lambda *a, **kw: calls.append(a),
|
||||
)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
assert calls == []
|
||||
assert "No model definitions found" in app.state.coord_registry_error
|
||||
|
||||
|
||||
def test_bootstrap_calls_subsystem_builder_on_first_row(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""A row exists ⇒ helper loads the registry, hands it to the
|
||||
subsystem builder, and the builder stamps it on app.state. Mirrors
|
||||
the post-build invariant the real ``_bootstrap_coord_subsystem``
|
||||
establishes (coord_registry set iff coord_mgr set) so the stale
|
||||
boot-time error string clears as part of the same commit step."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app(coord_registry_error="stale boot-time message")
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _fake_build(app_arg: Any, _storage: Any, _cfg: Any, registry_arg: Any) -> None:
|
||||
captured["app"] = app_arg
|
||||
captured["registry"] = registry_arg
|
||||
# Simulate the real builder's final commit step: stamp registry
|
||||
# + clear stale error + set coord_mgr atomically.
|
||||
app_arg.state.coord_registry = registry_arg
|
||||
app_arg.state.coord_registry_error = ""
|
||||
app_arg.state.coord_mgr = MagicMock()
|
||||
|
||||
monkeypatch.setattr(server_module, "_bootstrap_coord_subsystem", _fake_build)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
assert captured["app"] is app
|
||||
assert captured["registry"].has_alias("local")
|
||||
assert app.state.coord_registry is captured["registry"]
|
||||
assert app.state.coord_registry_error == ""
|
||||
|
||||
|
||||
def test_bootstrap_replaces_stale_error_on_builder_failure(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""A builder failure after a successful registry load must not leave
|
||||
the stale "no model definitions" message on app.state — that
|
||||
diagnosis is demonstrably wrong (rows ARE present, the build failed
|
||||
for a different reason). Replacement message must surface the
|
||||
actual exception type so operators can correlate with logs."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app(
|
||||
coord_registry_error=(
|
||||
"No model definitions found. Provide --model, configure [models.*] "
|
||||
"in config.toml, or add model definitions in the admin panel."
|
||||
)
|
||||
)
|
||||
|
||||
def _boom(*_a: Any, **_kw: Any) -> None:
|
||||
raise RuntimeError("simulated builder failure")
|
||||
|
||||
monkeypatch.setattr(server_module, "_bootstrap_coord_subsystem", _boom)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage) # must not raise
|
||||
assert app.state.coord_mgr is None
|
||||
# Stale "no models" message replaced.
|
||||
assert "No model definitions found" not in app.state.coord_registry_error
|
||||
# New message mentions the actual failure class so the 503 banner
|
||||
# gives operators something actionable beyond "look at logs".
|
||||
assert "RuntimeError" in app.state.coord_registry_error
|
||||
assert "failed to initialise" in app.state.coord_registry_error
|
||||
|
||||
|
||||
def test_bootstrap_tears_down_partial_state_on_builder_failure(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""If the builder partially stamps handles on app.state and then
|
||||
raises, the helper must call the teardown path so a subsequent
|
||||
retry doesn't leak a StateWriter daemon / observer subscription."""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app()
|
||||
|
||||
state_writer = MagicMock()
|
||||
idle_observer = MagicMock()
|
||||
coord_adapter = MagicMock()
|
||||
|
||||
def _partial_then_boom(app_arg: Any, *_a: Any, **_kw: Any) -> None:
|
||||
# Mirror the real builder's stamp-immediately-after-start order:
|
||||
# StateWriter spawned + stamped before SessionManager validates.
|
||||
app_arg.state.coord_state_writer = state_writer
|
||||
app_arg.state.coord_idle_observer = idle_observer
|
||||
app_arg.state.coord_adapter = coord_adapter
|
||||
raise RuntimeError("simulated mid-build failure")
|
||||
|
||||
monkeypatch.setattr(server_module, "_bootstrap_coord_subsystem", _partial_then_boom)
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
# Teardown ran for each partially-stamped handle.
|
||||
state_writer.shutdown.assert_called_once()
|
||||
idle_observer.shutdown.assert_called_once()
|
||||
coord_adapter.shutdown.assert_called_once()
|
||||
# And the app.state slots are reset so a retry sees a clean field.
|
||||
assert app.state.coord_state_writer is None
|
||||
assert app.state.coord_idle_observer is None
|
||||
assert app.state.coord_adapter is None
|
||||
assert app.state.coord_mgr is None
|
||||
assert app.state.coord_registry is None
|
||||
|
||||
|
||||
def test_real_bootstrap_stands_up_subsystem_end_to_end(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""End-to-end: the real ``_bootstrap_coord_subsystem`` constructs a
|
||||
working ``SessionManager`` against a real ``ConfigStore`` + real
|
||||
``ClusterCollector`` when an operator adds the first model row to
|
||||
a freshly-installed console.
|
||||
|
||||
This is the test that reproduces the user-reported bug — without it,
|
||||
all the bootstrap helper-level tests can pass even if the real
|
||||
builder never actually completes (the helper-level tests
|
||||
monkeypatch the builder out). Asserts the post-bootstrap invariant
|
||||
that ``_require_coord_mgr`` relies on: ``coord_mgr`` is a real
|
||||
SessionManager and ``coord_registry_error`` has been cleared.
|
||||
"""
|
||||
from turnstone.console import server as server_module
|
||||
from turnstone.console.collector import ClusterCollector
|
||||
from turnstone.console.coordinator_ui import ConsoleCoordinatorUI
|
||||
from turnstone.console.metrics import ConsoleMetrics
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
from turnstone.core.session_manager import SessionManager
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
config_store = ConfigStore(storage)
|
||||
# Disable the idle-cleanup daemon for this test — it has no
|
||||
# stop_event hook in the bootstrap (the loop runs until process
|
||||
# termination) so leaving the default 120-minute timeout would
|
||||
# leak a daemon thread across every test run.
|
||||
config_store.set("server.workstream_idle_timeout", 0)
|
||||
# ClusterCollector is constructed but NOT started — start() spawns
|
||||
# network discovery + SSE manager threads we don't need for this
|
||||
# test. ensure_console_pseudo_node() (called by the bootstrap via
|
||||
# start_child_event_fanout) operates on the in-memory snapshot map
|
||||
# without requiring the discovery loop to be live.
|
||||
collector = ClusterCollector(storage=storage)
|
||||
# Snapshot ConsoleCoordinatorUI's class attrs so the test can
|
||||
# restore them on teardown — the bootstrap mutates them and they
|
||||
# persist across tests at process scope.
|
||||
saved_coord_mgr = ConsoleCoordinatorUI._coord_mgr
|
||||
saved_collector = ConsoleCoordinatorUI._collector
|
||||
saved_metrics = ConsoleCoordinatorUI._console_metrics
|
||||
|
||||
app = SimpleNamespace(
|
||||
state=SimpleNamespace(
|
||||
coord_mgr=None,
|
||||
coord_adapter=None,
|
||||
coord_registry=None,
|
||||
coord_registry_error=(
|
||||
"No model definitions found. Provide --model, configure [models.*] "
|
||||
"in config.toml, or add model definitions in the admin panel."
|
||||
),
|
||||
coord_state_writer=None,
|
||||
coord_idle_observer=None,
|
||||
config_store=config_store,
|
||||
collector=collector,
|
||||
console_metrics=ConsoleMetrics(),
|
||||
jwt_secret="x" * 32,
|
||||
console_url="http://127.0.0.1:8001",
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
# The real builder ran and produced a working SessionManager.
|
||||
assert isinstance(app.state.coord_mgr, SessionManager)
|
||||
assert app.state.coord_adapter is not None
|
||||
# Registry stamped with the seeded alias.
|
||||
assert app.state.coord_registry is not None
|
||||
assert app.state.coord_registry.has_alias("local")
|
||||
# Stale boot-time error string cleared as part of the commit.
|
||||
assert app.state.coord_registry_error == ""
|
||||
# StateWriter daemon is alive — it's the load-bearing async
|
||||
# persistence layer for SessionManager state transitions.
|
||||
assert app.state.coord_state_writer is not None
|
||||
# Class-level wiring on ConsoleCoordinatorUI is the path
|
||||
# on_state_change / on_rename use to fan out to the dashboard.
|
||||
assert ConsoleCoordinatorUI._coord_mgr is app.state.coord_mgr
|
||||
assert ConsoleCoordinatorUI._collector is collector
|
||||
finally:
|
||||
# Tear down threads + subscriptions spawned by the bootstrap.
|
||||
# ``_teardown_partial_coord_subsystem`` does the same work the
|
||||
# runtime-bootstrap failure path does, so reusing it here also
|
||||
# exercises that helper end-to-end.
|
||||
server_module._teardown_partial_coord_subsystem(app)
|
||||
# Restore ConsoleCoordinatorUI class attrs so other tests in
|
||||
# the suite see them as they were before this test ran.
|
||||
ConsoleCoordinatorUI._coord_mgr = saved_coord_mgr
|
||||
ConsoleCoordinatorUI._collector = saved_collector
|
||||
ConsoleCoordinatorUI._console_metrics = saved_metrics
|
||||
|
||||
|
||||
def test_real_bootstrap_rolls_back_partial_state_on_side_effect_failure(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""The real ``_bootstrap_coord_subsystem`` must roll back from
|
||||
locally-held handles when a side-effect step fails mid-build, so
|
||||
``app.state`` is never stamped (no half-built subsystem visible)
|
||||
and the started ``StateWriter`` daemon is shut down (no leaked
|
||||
thread across retries).
|
||||
|
||||
Exercises the bug-2 fix end-to-end: monkeypatches
|
||||
``install_idle_nudge_watcher`` to raise, drives the real builder,
|
||||
and asserts (a) the exception propagates, (b) ``app.state`` shows
|
||||
a clean fresh-install state, (c) the started ``StateWriter`` is
|
||||
no longer alive.
|
||||
"""
|
||||
from turnstone.console import server as server_module
|
||||
from turnstone.console.collector import ClusterCollector
|
||||
from turnstone.console.coordinator_ui import ConsoleCoordinatorUI
|
||||
from turnstone.console.metrics import ConsoleMetrics
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
config_store = ConfigStore(storage)
|
||||
config_store.set("server.workstream_idle_timeout", 0)
|
||||
collector = ClusterCollector(storage=storage)
|
||||
saved_coord_mgr = ConsoleCoordinatorUI._coord_mgr
|
||||
saved_collector = ConsoleCoordinatorUI._collector
|
||||
saved_metrics = ConsoleCoordinatorUI._console_metrics
|
||||
|
||||
app = SimpleNamespace(
|
||||
state=SimpleNamespace(
|
||||
coord_mgr=None,
|
||||
coord_adapter=None,
|
||||
coord_registry=None,
|
||||
coord_registry_error="boot-time stale message",
|
||||
coord_state_writer=None,
|
||||
coord_idle_observer=None,
|
||||
config_store=config_store,
|
||||
collector=collector,
|
||||
console_metrics=ConsoleMetrics(),
|
||||
jwt_secret="x" * 32,
|
||||
console_url="http://127.0.0.1:8001",
|
||||
)
|
||||
)
|
||||
|
||||
# Monkeypatch a mid-build side-effect to fail AFTER StateWriter +
|
||||
# observer have started but BEFORE the atomic commit. This is the
|
||||
# exact failure shape the new local-rollback path is designed to
|
||||
# handle cleanly.
|
||||
def _boom(*_a: Any, **_kw: Any) -> Any:
|
||||
raise RuntimeError("simulated mid-build subscription failure")
|
||||
|
||||
monkeypatch.setattr("turnstone.console.server.install_idle_nudge_watcher", _boom, raising=False)
|
||||
# The bootstrap helper imports install_idle_nudge_watcher locally
|
||||
# at call time (inside the function), so we need to patch the
|
||||
# source module too — server.py's import is a name lookup against
|
||||
# the module each call.
|
||||
monkeypatch.setattr(
|
||||
"turnstone.core.idle_nudge_watcher.install_idle_nudge_watcher",
|
||||
_boom,
|
||||
)
|
||||
|
||||
try:
|
||||
# ``_maybe_bootstrap_coord_subsystem`` swallows the exception,
|
||||
# logs it, and replaces the stale boot-time error string with
|
||||
# a builder-failure-specific one — but the underlying invariant
|
||||
# we're testing here is that the real builder cleaned up its
|
||||
# own partial side-effects so ``app.state`` is left clean.
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
# No state stamped — atomic commit never reached.
|
||||
assert app.state.coord_mgr is None
|
||||
assert app.state.coord_registry is None
|
||||
assert app.state.coord_state_writer is None
|
||||
assert app.state.coord_idle_observer is None
|
||||
assert app.state.coord_adapter is None
|
||||
# ConsoleCoordinatorUI class attrs were never stamped because
|
||||
# they sit AFTER the side-effect phase — local-rollback never
|
||||
# had to touch them, but the post-failure state still matches
|
||||
# the lifespan's clean state.
|
||||
assert ConsoleCoordinatorUI._coord_mgr is None
|
||||
# The error string surfaces the actual failure cause, not the
|
||||
# stale boot-time "no models" message.
|
||||
assert "RuntimeError" in app.state.coord_registry_error
|
||||
assert "failed to initialise" in app.state.coord_registry_error
|
||||
finally:
|
||||
# Defensive — _maybe_bootstrap should already have torn down,
|
||||
# but call once more in case future drift introduces a leak.
|
||||
server_module._teardown_partial_coord_subsystem(app)
|
||||
ConsoleCoordinatorUI._coord_mgr = saved_coord_mgr
|
||||
ConsoleCoordinatorUI._collector = saved_collector
|
||||
ConsoleCoordinatorUI._console_metrics = saved_metrics
|
||||
|
||||
|
||||
def test_bootstrap_atomic_commit_no_partial_visibility(
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""A concurrent reader scanning ``app.state`` while the bootstrap
|
||||
runs must never observe ``coord_mgr`` set with ``coord_registry``
|
||||
still ``None`` — that combination would surface a misleading
|
||||
"Restart the console after adding a model definition" 503 from
|
||||
:func:`_require_coord_mgr` even though the operator just
|
||||
successfully added a model.
|
||||
|
||||
Drives the real builder while a separate thread polls
|
||||
``coord_mgr`` / ``coord_registry`` in tight loops; if the bootstrap
|
||||
ever stamps ``coord_mgr`` before ``coord_registry``, the polling
|
||||
thread will catch it.
|
||||
"""
|
||||
from turnstone.console import server as server_module
|
||||
from turnstone.console.collector import ClusterCollector
|
||||
from turnstone.console.coordinator_ui import ConsoleCoordinatorUI
|
||||
from turnstone.console.metrics import ConsoleMetrics
|
||||
from turnstone.core.config_store import ConfigStore
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
config_store = ConfigStore(storage)
|
||||
config_store.set("server.workstream_idle_timeout", 0)
|
||||
collector = ClusterCollector(storage=storage)
|
||||
saved_coord_mgr = ConsoleCoordinatorUI._coord_mgr
|
||||
saved_collector = ConsoleCoordinatorUI._collector
|
||||
saved_metrics = ConsoleCoordinatorUI._console_metrics
|
||||
|
||||
app = SimpleNamespace(
|
||||
state=SimpleNamespace(
|
||||
coord_mgr=None,
|
||||
coord_adapter=None,
|
||||
coord_registry=None,
|
||||
coord_registry_error="",
|
||||
coord_state_writer=None,
|
||||
coord_idle_observer=None,
|
||||
config_store=config_store,
|
||||
collector=collector,
|
||||
console_metrics=ConsoleMetrics(),
|
||||
jwt_secret="x" * 32,
|
||||
console_url="http://127.0.0.1:8001",
|
||||
)
|
||||
)
|
||||
|
||||
stop_polling = threading.Event()
|
||||
violations: list[str] = []
|
||||
|
||||
def _poll_for_partial_state() -> None:
|
||||
# Tight loop emulating ``_require_coord_mgr``'s read pattern
|
||||
# (coord_mgr first, then coord_registry). Any iteration that
|
||||
# observes coord_mgr set with coord_registry still None is the
|
||||
# exact bug Copilot's first finding pointed at.
|
||||
while not stop_polling.is_set():
|
||||
mgr = app.state.coord_mgr
|
||||
reg = app.state.coord_registry
|
||||
if mgr is not None and reg is None:
|
||||
violations.append(f"mgr={mgr!r} reg={reg!r}")
|
||||
return
|
||||
|
||||
poller = threading.Thread(target=_poll_for_partial_state, name="partial-state-poller")
|
||||
poller.start()
|
||||
try:
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
finally:
|
||||
stop_polling.set()
|
||||
poller.join(timeout=2.0)
|
||||
server_module._teardown_partial_coord_subsystem(app)
|
||||
ConsoleCoordinatorUI._coord_mgr = saved_coord_mgr
|
||||
ConsoleCoordinatorUI._collector = saved_collector
|
||||
ConsoleCoordinatorUI._console_metrics = saved_metrics
|
||||
|
||||
assert violations == [], (
|
||||
"concurrent reader observed coord_mgr set with coord_registry still None — "
|
||||
f"atomic commit invariant violated: {violations}"
|
||||
)
|
||||
|
||||
|
||||
def test_bootstrap_lock_serialises_concurrent_calls(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Two simultaneous CRUD writes both seeing ``coord_mgr is None``
|
||||
must serialise via ``_COORD_BOOTSTRAP_LOCK`` and the second caller
|
||||
must observe the post-build state on its inside-the-lock re-check —
|
||||
so the builder runs exactly once. Without the lock + double-check,
|
||||
both threads enter the build and stamp duplicate SessionManager /
|
||||
StateWriter / observer triples on app.state.
|
||||
|
||||
The synchronisation is deterministic, not wall-clock-based: an
|
||||
instrumented lock wrapper signals when a second acquirer arrives,
|
||||
so the test fails fast and reproducibly on slow CI rather than
|
||||
relying on a sleep long enough to "probably" let thread 2 reach
|
||||
the lock — a dependence the previous version was rightly criticised
|
||||
for.
|
||||
"""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
_seed_model_def(storage, definition_id="m1", alias="local", model="m")
|
||||
app = _bootstrap_app()
|
||||
build_count = 0
|
||||
count_lock = threading.Lock()
|
||||
in_build = threading.Event()
|
||||
release_build = threading.Event()
|
||||
|
||||
def _slow_build(app_arg: Any, *_a: Any, **_kw: Any) -> None:
|
||||
nonlocal build_count
|
||||
with count_lock:
|
||||
build_count += 1
|
||||
is_first = build_count == 1
|
||||
if is_first:
|
||||
# Hold inside the build so the second thread is forced to
|
||||
# queue at the lock — without the lock it would race ahead
|
||||
# and increment build_count to 2.
|
||||
in_build.set()
|
||||
release_build.wait(timeout=2.0)
|
||||
# Mirror the real builder's commit step.
|
||||
app_arg.state.coord_mgr = MagicMock()
|
||||
app_arg.state.coord_registry = MagicMock()
|
||||
|
||||
monkeypatch.setattr(server_module, "_bootstrap_coord_subsystem", _slow_build)
|
||||
|
||||
# Instrumented wrapper: delegates to a real ``threading.Lock`` so
|
||||
# the production ``with _COORD_BOOTSTRAP_LOCK:`` block keeps doing
|
||||
# genuine serialisation work, but counts arrivals so the main
|
||||
# thread can wait deterministically until thread 2 is at the lock
|
||||
# before releasing thread 1. If the production code drops the
|
||||
# ``with`` block entirely, the wrapper is never entered, the
|
||||
# arrival event never fires, and the assertion below times out
|
||||
# with a clear error rather than the subtler false-pass a sleep
|
||||
# would allow.
|
||||
real_lock = threading.Lock()
|
||||
arrivals_lock = threading.Lock()
|
||||
arrivals = 0
|
||||
second_waiter_arrived = threading.Event()
|
||||
|
||||
class _InstrumentedLock:
|
||||
def __enter__(self) -> Any:
|
||||
nonlocal arrivals
|
||||
with arrivals_lock:
|
||||
arrivals += 1
|
||||
arrival_index = arrivals
|
||||
if arrival_index >= 2:
|
||||
second_waiter_arrived.set()
|
||||
real_lock.acquire()
|
||||
return self
|
||||
|
||||
def __exit__(self, *_exc: Any) -> None:
|
||||
real_lock.release()
|
||||
|
||||
monkeypatch.setattr(server_module, "_COORD_BOOTSTRAP_LOCK", _InstrumentedLock())
|
||||
|
||||
def _run() -> None:
|
||||
_maybe_bootstrap_coord_subsystem(app, storage)
|
||||
|
||||
t1 = threading.Thread(target=_run, name="bootstrap-thread-1")
|
||||
t2 = threading.Thread(target=_run, name="bootstrap-thread-2")
|
||||
t1.start()
|
||||
assert in_build.wait(timeout=2.0), "thread 1 never entered the builder"
|
||||
t2.start()
|
||||
# Deterministic: block here until thread 2 has reached the lock
|
||||
# (or the wait times out, signalling the lock was bypassed entirely).
|
||||
assert second_waiter_arrived.wait(timeout=2.0), (
|
||||
"thread 2 never reached the lock — concurrency was not exercised, "
|
||||
"production code may be skipping the lock"
|
||||
)
|
||||
release_build.set()
|
||||
t1.join(timeout=5.0)
|
||||
t2.join(timeout=5.0)
|
||||
assert not t1.is_alive() and not t2.is_alive()
|
||||
assert build_count == 1, (
|
||||
f"builder ran {build_count} times — lock failed to serialise concurrent calls"
|
||||
)
|
||||
|
||||
|
||||
def test_helper_preserves_registry_on_reload_validation_error(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
@@ -309,6 +867,79 @@ def test_create_endpoint_refreshes_registry(storage: SQLiteBackend) -> None:
|
||||
assert registry.get_config("fast").model == "fast-model"
|
||||
|
||||
|
||||
def test_create_endpoint_bootstraps_subsystem_on_fresh_install(
|
||||
storage: SQLiteBackend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""User-visible regression: a console booted with no model rows leaves
|
||||
coord_mgr unbuilt; the operator adding their first model via the
|
||||
admin panel must promote the subsystem to ready (no console restart).
|
||||
Before the fix, ``_refresh_coord_registry`` short-circuited on
|
||||
``coord_registry is None`` and the dashboard's 503 banner persisted
|
||||
until the user restarted.
|
||||
"""
|
||||
from turnstone.console import server as server_module
|
||||
|
||||
# Fresh-install state: registry=None, coord_mgr=None, boot-time
|
||||
# error string set by the lifespan's ValueError catch. Build the
|
||||
# app explicitly so the test can inspect ``app.state`` after the
|
||||
# request completes (TestClient's ``.app`` attribute is typed as
|
||||
# ASGIApp, which loses the ``.state`` accessor).
|
||||
app = Starlette(
|
||||
routes=[
|
||||
Route(
|
||||
"/v1/api/admin/model-definitions",
|
||||
admin_create_model_definition,
|
||||
methods=["POST"],
|
||||
),
|
||||
],
|
||||
middleware=[Middleware(_AuthMiddleware)],
|
||||
)
|
||||
app.state.auth_storage = storage
|
||||
app.state.coord_registry = None
|
||||
app.state.coord_mgr = None
|
||||
app.state.coord_registry_error = (
|
||||
"No model definitions found. Provide --model, configure [models.*] "
|
||||
"in config.toml, or add model definitions in the admin panel."
|
||||
)
|
||||
app.state.collector = MagicMock()
|
||||
app.state.collector.get_all_nodes.return_value = []
|
||||
app.state.config_store = MagicMock()
|
||||
app.state.console_metrics = MagicMock()
|
||||
client = TestClient(app)
|
||||
client.headers.update({"X-Test-User": "admin", "X-Test-Perms": "admin.models"})
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _fake_build(app_arg: Any, _storage: Any, _cfg: Any, registry_arg: Any) -> None:
|
||||
captured["registry"] = registry_arg
|
||||
# Mirror the real builder's commit step so the post-call asserts
|
||||
# see the same invariant a successful real bootstrap establishes.
|
||||
app_arg.state.coord_registry = registry_arg
|
||||
app_arg.state.coord_registry_error = ""
|
||||
app_arg.state.coord_mgr = MagicMock()
|
||||
|
||||
monkeypatch.setattr(server_module, "_bootstrap_coord_subsystem", _fake_build)
|
||||
|
||||
resp = client.post(
|
||||
"/v1/api/admin/model-definitions",
|
||||
json={
|
||||
"alias": "first",
|
||||
"model": "first-model",
|
||||
"provider": "openai-compatible",
|
||||
"base_url": "http://localhost:9000/v1",
|
||||
"api_key": "sk-x",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
# Bootstrap fired with a registry holding the just-added alias.
|
||||
assert "registry" in captured and captured["registry"].has_alias("first")
|
||||
# coord_mgr is now non-None (bootstrap completed) and the stale
|
||||
# boot-time error message has been cleared so subsequent 503s
|
||||
# don't lie about current state.
|
||||
assert app.state.coord_mgr is not None
|
||||
assert app.state.coord_registry_error == ""
|
||||
|
||||
|
||||
def test_update_endpoint_refreshes_registry(storage: SQLiteBackend) -> None:
|
||||
"""PUT swaps the underlying model name behind a stable alias — the
|
||||
user's reported regression."""
|
||||
|
||||
+296
-1
@@ -124,7 +124,14 @@ def test_replay_history_renders_content_before_tool_block() -> None:
|
||||
asst_start = fn.index('msg.role === "assistant"')
|
||||
asst_end = fn.index('msg.role === "tool"', asst_start)
|
||||
asst = fn[asst_start:asst_end]
|
||||
content_idx = asst.index("if (msg.content)")
|
||||
# ``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 "
|
||||
@@ -160,3 +167,291 @@ def test_replay_history_renders_persisted_verdict_badge() -> None:
|
||||
"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 panel 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="openSettingsPanel()"' in body, (
|
||||
"settings-btn must wire onclick=openSettingsPanel() — losing "
|
||||
"the binding leaves the panel 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_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."
|
||||
)
|
||||
|
||||
@@ -207,6 +207,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,143 @@
|
||||
"""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"}]
|
||||
@@ -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
|
||||
@@ -167,6 +167,28 @@ def test_coordinator_js_exposes_inline_approval_helpers():
|
||||
# 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():
|
||||
@@ -280,3 +302,43 @@ def test_coordinator_js_handle_child_state_no_longer_reads_sse_pending_approval_
|
||||
"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."
|
||||
)
|
||||
|
||||
@@ -167,7 +167,7 @@ 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_and_marks_truncated(self) -> None:
|
||||
def test_decorates_tool_calls_with_verdict_and_assessment(self) -> None:
|
||||
verdicts = {
|
||||
"call_a": {
|
||||
"risk_level": "high",
|
||||
@@ -181,13 +181,6 @@ class TestDecorateHistoryMessages:
|
||||
assessments = {
|
||||
"call_a": {"risk_level": "high", "flags": '["secret"]', "redacted": 1},
|
||||
}
|
||||
# Tool result content of exactly TOOL_RESULT_STORAGE_CAP chars
|
||||
# hits the storage cap (longer is impossible — storage clamps
|
||||
# at the cap). Reference the constant rather than a literal so
|
||||
# this test stays correct if the cap moves again.
|
||||
from turnstone.core.history_decoration import TOOL_RESULT_STORAGE_CAP
|
||||
|
||||
truncated_content = "x" * TOOL_RESULT_STORAGE_CAP
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
@@ -200,7 +193,7 @@ class TestDecorateHistoryMessages:
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_a", "content": truncated_content},
|
||||
{"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)
|
||||
@@ -211,9 +204,12 @@ class TestDecorateHistoryMessages:
|
||||
assert "reasoning" in tc["verdict"]
|
||||
assert tc["output_assessment"]["flags"] == ["secret"]
|
||||
assert tc["output_assessment"]["redacted"] is True
|
||||
# Truncated tool message got the flag; the short one did not.
|
||||
assert messages[2].get("truncated") is True
|
||||
assert "truncated" not in messages[3]
|
||||
# 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
|
||||
@@ -230,3 +226,230 @@ class TestDecorateHistoryMessages:
|
||||
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"}
|
||||
]
|
||||
|
||||
@@ -0,0 +1,385 @@
|
||||
"""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_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
|
||||
@@ -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
+786
-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,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 == {}
|
||||
@@ -4,6 +4,8 @@ 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,
|
||||
@@ -11,6 +13,7 @@ from turnstone.core.metacognition import (
|
||||
RepeatDetector,
|
||||
detect_completion,
|
||||
detect_correction,
|
||||
format_idle_children_nudge,
|
||||
format_nudge,
|
||||
should_nudge,
|
||||
)
|
||||
@@ -376,3 +379,209 @@ class TestRepeatDetector:
|
||||
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()
|
||||
@@ -0,0 +1,533 @@
|
||||
"""Unit tests for :class:`NudgeQueue`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
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 /
|
||||
raising predicates drop the entry without delivery.
|
||||
"""
|
||||
|
||||
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_silently(self):
|
||||
q = NudgeQueue()
|
||||
q.enqueue("a", "1", "any", valid_until=lambda: False)
|
||||
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
|
||||
|
||||
def test_valid_until_exception_drops_silently(self):
|
||||
q = NudgeQueue()
|
||||
|
||||
def boom() -> bool:
|
||||
raise RuntimeError("predicate crash")
|
||||
|
||||
q.enqueue("a", "1", "any", valid_until=boom)
|
||||
out = q.drain({"any"})
|
||||
assert out == []
|
||||
# Crash-on-predicate is treated as "no longer valid" — drop, not propagate.
|
||||
assert len(q) == 0
|
||||
|
||||
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
|
||||
@@ -0,0 +1,172 @@
|
||||
"""Direct tests for the shared SSRF helpers in :mod:`turnstone.core.oauth_ssrf`.
|
||||
|
||||
The OIDC test suite already exercises these via the OIDC adapter
|
||||
(``OIDCError`` re-raises). This file pins the canonical
|
||||
:class:`OAuthSSRFError` exception so callers that don't go through OIDC
|
||||
(notably ``mcp_oauth``) can rely on a stable contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import urllib.parse
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.oauth_ssrf import (
|
||||
OAuthSSRFError,
|
||||
effective_port,
|
||||
is_localhost,
|
||||
validate_discovered_endpoint,
|
||||
validate_url_no_ssrf,
|
||||
)
|
||||
|
||||
|
||||
class TestIsLocalhost:
|
||||
def test_loopback_names(self) -> None:
|
||||
assert is_localhost("localhost")
|
||||
assert is_localhost("127.0.0.1")
|
||||
assert is_localhost("::1")
|
||||
assert is_localhost("foo.localhost")
|
||||
|
||||
def test_non_loopback(self) -> None:
|
||||
assert not is_localhost("example.com")
|
||||
assert not is_localhost("internal.corp")
|
||||
|
||||
|
||||
class TestEffectivePort:
|
||||
def test_explicit_port(self) -> None:
|
||||
p = urllib.parse.urlparse("https://idp.example.com:9443/foo")
|
||||
assert effective_port(p) == 9443
|
||||
|
||||
def test_default_https(self) -> None:
|
||||
p = urllib.parse.urlparse("https://idp.example.com/foo")
|
||||
assert effective_port(p) == 443
|
||||
|
||||
def test_default_http(self) -> None:
|
||||
p = urllib.parse.urlparse("http://idp.example.com/foo")
|
||||
assert effective_port(p) == 80
|
||||
|
||||
def test_unknown_scheme(self) -> None:
|
||||
p = urllib.parse.urlparse("ftp://idp.example.com/foo")
|
||||
assert effective_port(p) is None
|
||||
|
||||
|
||||
class TestValidateUrlNoSSRF:
|
||||
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
||||
_PRIVATE_ADDR = [(2, 1, 6, "", ("10.0.0.1", 0))]
|
||||
_LOOPBACK_ADDR = [(2, 1, 6, "", ("127.0.0.1", 0))]
|
||||
|
||||
def test_valid_https(self) -> None:
|
||||
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
||||
parsed = validate_url_no_ssrf("https://idp.example.com/foo", allow_http=False)
|
||||
assert parsed.scheme == "https"
|
||||
assert parsed.hostname == "idp.example.com"
|
||||
|
||||
def test_rejects_http_when_not_allowed(self) -> None:
|
||||
with pytest.raises(OAuthSSRFError, match="must use HTTPS"):
|
||||
validate_url_no_ssrf("http://idp.example.com", allow_http=False)
|
||||
|
||||
def test_allows_http_localhost_with_flag(self) -> None:
|
||||
with patch("socket.getaddrinfo", return_value=self._LOOPBACK_ADDR):
|
||||
validate_url_no_ssrf("http://localhost:8080", allow_http=True)
|
||||
|
||||
def test_rejects_http_non_localhost_even_with_flag(self) -> None:
|
||||
with pytest.raises(OAuthSSRFError, match="must use HTTPS"):
|
||||
validate_url_no_ssrf("http://idp.example.com", allow_http=True)
|
||||
|
||||
def test_rejects_userinfo(self) -> None:
|
||||
with pytest.raises(OAuthSSRFError, match="embedded credentials"):
|
||||
validate_url_no_ssrf("https://user:pass@idp.example.com", allow_http=False)
|
||||
|
||||
def test_rejects_private_address(self) -> None:
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PRIVATE_ADDR),
|
||||
pytest.raises(OAuthSSRFError, match="non-public address"),
|
||||
):
|
||||
validate_url_no_ssrf("https://corp.example.com", allow_http=False)
|
||||
|
||||
def test_rejects_unresolvable(self) -> None:
|
||||
import socket
|
||||
|
||||
with (
|
||||
patch("socket.getaddrinfo", side_effect=socket.gaierror("fail")),
|
||||
pytest.raises(OAuthSSRFError, match="cannot be resolved"),
|
||||
):
|
||||
validate_url_no_ssrf("https://no.such.host.invalid", allow_http=False)
|
||||
|
||||
|
||||
class TestValidateDiscoveredEndpoint:
|
||||
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
||||
|
||||
def test_same_origin_passes(self) -> None:
|
||||
issuer = urllib.parse.urlparse("https://idp.example.com")
|
||||
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
||||
validate_discovered_endpoint(
|
||||
"https://idp.example.com/token",
|
||||
issuer,
|
||||
allow_http=False,
|
||||
trusted_endpoint_hosts=frozenset(),
|
||||
)
|
||||
|
||||
def test_third_party_host_rejected(self) -> None:
|
||||
issuer = urllib.parse.urlparse("https://idp.example.com")
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
||||
pytest.raises(OAuthSSRFError, match="not trusted"),
|
||||
):
|
||||
validate_discovered_endpoint(
|
||||
"https://attacker.example.com/token",
|
||||
issuer,
|
||||
allow_http=False,
|
||||
trusted_endpoint_hosts=frozenset(),
|
||||
)
|
||||
|
||||
def test_trusted_endpoint_host_passes(self) -> None:
|
||||
issuer = urllib.parse.urlparse("https://idp.example.com")
|
||||
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
||||
validate_discovered_endpoint(
|
||||
"https://shard.example.com/token",
|
||||
issuer,
|
||||
allow_http=False,
|
||||
trusted_endpoint_hosts=frozenset({"shard.example.com"}),
|
||||
)
|
||||
|
||||
def test_known_google_alias_passes(self) -> None:
|
||||
"""The hard-coded Google alias map covers oauth2.googleapis.com."""
|
||||
issuer = urllib.parse.urlparse("https://accounts.google.com")
|
||||
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
||||
validate_discovered_endpoint(
|
||||
"https://oauth2.googleapis.com/token",
|
||||
issuer,
|
||||
allow_http=False,
|
||||
trusted_endpoint_hosts=frozenset(),
|
||||
)
|
||||
|
||||
def test_scheme_mismatch_rejected(self) -> None:
|
||||
# When the issuer is http://localhost (allow_http=True), an
|
||||
# https:// endpoint must still be rejected as a scheme mismatch.
|
||||
issuer = urllib.parse.urlparse("http://localhost:8080")
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("127.0.0.1", 0))]),
|
||||
pytest.raises(OAuthSSRFError, match="scheme"),
|
||||
):
|
||||
validate_discovered_endpoint(
|
||||
"https://localhost:8080/token",
|
||||
issuer,
|
||||
allow_http=True,
|
||||
trusted_endpoint_hosts=frozenset(),
|
||||
)
|
||||
|
||||
def test_port_mismatch_rejected(self) -> None:
|
||||
issuer = urllib.parse.urlparse("https://idp.example.com")
|
||||
with (
|
||||
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
||||
pytest.raises(OAuthSSRFError, match="port"),
|
||||
):
|
||||
validate_discovered_endpoint(
|
||||
"https://idp.example.com:9443/token",
|
||||
issuer,
|
||||
allow_http=False,
|
||||
trusted_endpoint_hosts=frozenset(),
|
||||
)
|
||||
+1432
-149
File diff suppressed because it is too large
Load Diff
+259
-45
@@ -18,10 +18,7 @@ from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.routing import Mount, Route
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from tests.conftest import make_oidc_test_config as _make_oidc_config
|
||||
from turnstone.console.server import (
|
||||
admin_delete_oidc_identity,
|
||||
admin_list_oidc_identities,
|
||||
@@ -32,34 +29,12 @@ from turnstone.core.auth import (
|
||||
handle_oidc_authorize,
|
||||
handle_oidc_callback,
|
||||
)
|
||||
from turnstone.core.oidc import OIDCConfig, OIDCError
|
||||
from turnstone.core.oidc import OIDCConfig, OIDCError, OIDCKeyNotFoundError
|
||||
from turnstone.core.storage._sqlite import SQLiteBackend
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_oidc_config(**overrides: Any) -> OIDCConfig:
|
||||
"""Build a test OIDCConfig with sensible defaults."""
|
||||
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,
|
||||
"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)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thin handler wrappers — match the pattern used in server.py / console
|
||||
@@ -290,9 +265,9 @@ class TestOIDCCallback:
|
||||
) -> None:
|
||||
storage.create_oidc_pending_state(state, nonce, code_verifier, audience)
|
||||
|
||||
@patch("turnstone.core.oidc.provision_oidc_user")
|
||||
@patch("turnstone.core.oidc.validate_id_token")
|
||||
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_happy_path(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
@@ -379,7 +354,7 @@ class TestOIDCCallback:
|
||||
assert resp.status_code == 302
|
||||
assert "Login+session+expired" in resp.headers["location"]
|
||||
|
||||
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_code_exchange_failure(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
@@ -396,8 +371,8 @@ class TestOIDCCallback:
|
||||
assert resp.status_code == 302
|
||||
assert "Authentication+failed" in resp.headers["location"]
|
||||
|
||||
@patch("turnstone.core.oidc.validate_id_token")
|
||||
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_token_validation_failure(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
@@ -416,10 +391,10 @@ class TestOIDCCallback:
|
||||
assert resp.status_code == 302
|
||||
assert "Authentication+failed" in resp.headers["location"]
|
||||
|
||||
@patch("turnstone.core.oidc.provision_oidc_user")
|
||||
@patch("turnstone.core.oidc.validate_id_token")
|
||||
@patch("turnstone.core.oidc.fetch_jwks", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.fetch_jwks", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_jwks_key_rotation_retry(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
@@ -429,13 +404,13 @@ class TestOIDCCallback:
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""First validate raises 'kid not found in JWKS', fetch_jwks retried, second validate succeeds."""
|
||||
"""First validate raises kid-not-found, fetch_jwks retried, second validate succeeds."""
|
||||
self._seed_pending_state(storage)
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
|
||||
# First call raises kid-not-found; second call (after JWKS refresh) succeeds
|
||||
mock_validate.side_effect = [
|
||||
OIDCError("Signing key 'new-kid' not found in JWKS"),
|
||||
OIDCKeyNotFoundError("Signing key 'new-kid' not found in JWKS"),
|
||||
{"sub": "user123", "email": "u@example.com", "nonce": "test-nonce"},
|
||||
]
|
||||
mock_fetch_jwks.return_value = {"keys": [{"kid": "new-kid", "kty": "RSA"}]}
|
||||
@@ -450,9 +425,61 @@ class TestOIDCCallback:
|
||||
mock_fetch_jwks.assert_called_once()
|
||||
assert mock_validate.call_count == 2
|
||||
|
||||
@patch("turnstone.core.oidc.provision_oidc_user")
|
||||
@patch("turnstone.core.oidc.validate_id_token")
|
||||
@patch("turnstone.core.oidc.exchange_code", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.fetch_jwks", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_callback_uses_keynotfound_for_jwks_retry(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
mock_fetch_jwks: AsyncMock,
|
||||
mock_validate: Any,
|
||||
mock_provision: Any,
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Retry path keys off the OIDCKeyNotFoundError type, not message substring."""
|
||||
self._seed_pending_state(storage)
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
|
||||
# First raises subclass; rephrased message must not affect retry behaviour.
|
||||
mock_validate.side_effect = [
|
||||
OIDCKeyNotFoundError("rotated key absent from cached set"),
|
||||
{"sub": "user123", "email": "u@example.com", "nonce": "test-nonce"},
|
||||
]
|
||||
mock_fetch_jwks.return_value = {"keys": [{"kid": "new-kid", "kty": "RSA"}]}
|
||||
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
|
||||
|
||||
resp = authorize_client.get(
|
||||
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
|
||||
follow_redirects=False,
|
||||
)
|
||||
assert resp.status_code == 302
|
||||
assert "oidc_success=1" in resp.headers["location"]
|
||||
mock_fetch_jwks.assert_called_once()
|
||||
assert mock_validate.call_count == 2
|
||||
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_callback_returns_authentication_failed_on_missing_id_token(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""A token endpoint response without id_token must redirect with auth-failed."""
|
||||
self._seed_pending_state(storage)
|
||||
mock_exchange.return_value = {"access_token": "x"}
|
||||
|
||||
resp = authorize_client.get(
|
||||
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
|
||||
follow_redirects=False,
|
||||
)
|
||||
assert resp.status_code == 302
|
||||
assert "oidc_error=Authentication+failed" in resp.headers["location"]
|
||||
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_no_users_after_oidc_success_redirects_setup(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
@@ -507,6 +534,193 @@ class TestOIDCCallback:
|
||||
assert "oidc_error" in resp.headers["location"]
|
||||
assert "Too+many" in resp.headers["location"]
|
||||
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_setup_gate_uses_count_users_not_full_scan(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
mock_validate: Any,
|
||||
mock_provision: Any,
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Callback's setup-complete gate must call count_users, not list_users."""
|
||||
from unittest.mock import patch as obj_patch
|
||||
|
||||
self._seed_pending_state(storage)
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
mock_validate.return_value = {
|
||||
"sub": "u1",
|
||||
"email": "u@example.com",
|
||||
"nonce": "test-nonce",
|
||||
}
|
||||
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
|
||||
|
||||
with (
|
||||
obj_patch.object(storage, "count_users", wraps=storage.count_users) as count_spy,
|
||||
obj_patch.object(storage, "list_users", wraps=storage.list_users) as list_spy,
|
||||
):
|
||||
resp = authorize_client.get(
|
||||
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
|
||||
follow_redirects=False,
|
||||
)
|
||||
|
||||
assert resp.status_code == 302
|
||||
assert "oidc_success=1" in resp.headers["location"]
|
||||
count_spy.assert_called_once_with()
|
||||
list_spy.assert_not_called()
|
||||
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_state_cleanup_is_gated(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
mock_validate: Any,
|
||||
mock_provision: Any,
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""Cleanup runs once per cleanup-interval window, not every callback."""
|
||||
from unittest.mock import patch as obj_patch
|
||||
|
||||
# First call seeds the cleanup timestamp; subsequent calls within
|
||||
# _OIDC_STATE_CLEANUP_INTERVAL_S must NOT trigger cleanup again.
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
mock_validate.return_value = {
|
||||
"sub": "u1",
|
||||
"email": "u@example.com",
|
||||
"nonce": "test-nonce",
|
||||
}
|
||||
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
|
||||
|
||||
with obj_patch.object(
|
||||
storage, "cleanup_expired_oidc_states", wraps=storage.cleanup_expired_oidc_states
|
||||
) as cleanup_spy:
|
||||
for state in ("s1", "s2", "s3"):
|
||||
self._seed_pending_state(storage, state=state, nonce="test-nonce")
|
||||
authorize_client.get(
|
||||
f"/v1/api/auth/oidc/callback?code=c&state={state}",
|
||||
follow_redirects=False,
|
||||
)
|
||||
|
||||
assert cleanup_spy.call_count == 1
|
||||
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_callback_uses_pending_audience_not_handler_audience(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
mock_validate: Any,
|
||||
mock_provision: Any,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""JWT ``aud`` claim must come from the audience stored at /authorize,
|
||||
not the audience the callback handler was invoked with.
|
||||
|
||||
Regression for the cross-service audience-confusion concern: a
|
||||
login flow opened against the server (audience ``"turnstone-server"``)
|
||||
must not be silently re-targeted to ``"turnstone-console"`` when
|
||||
the callback runs through the console's handler wrapper.
|
||||
"""
|
||||
import jwt as pyjwt
|
||||
|
||||
# Seed pending state with the SERVER audience.
|
||||
storage.create_oidc_pending_state(
|
||||
"audience-state",
|
||||
"audience-nonce",
|
||||
"audience-verifier",
|
||||
"turnstone-server",
|
||||
)
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
mock_validate.return_value = {
|
||||
"sub": "user-aud",
|
||||
"email": "u@example.com",
|
||||
"nonce": "audience-nonce",
|
||||
}
|
||||
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
|
||||
|
||||
# Wire a callback bound to the CONSOLE audience. After bug-3 the
|
||||
# stored audience must take precedence.
|
||||
async def _console_callback(request: Request) -> Response:
|
||||
return await handle_oidc_callback(request, "turnstone-console")
|
||||
|
||||
jwt_secret = "test-jwt-secret-key-padded-32b!!"
|
||||
app = Starlette(
|
||||
routes=[Mount("/v1", routes=[Route("/api/auth/oidc/callback", _console_callback)])]
|
||||
)
|
||||
app.state.oidc_config = _make_oidc_config()
|
||||
app.state.auth_storage = storage
|
||||
app.state.jwt_secret = jwt_secret
|
||||
app.state.jwks_data = {"keys": []}
|
||||
app.state.login_limiter = None
|
||||
client = TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
resp = client.get(
|
||||
"/v1/api/auth/oidc/callback?code=authcode&state=audience-state",
|
||||
follow_redirects=False,
|
||||
)
|
||||
|
||||
assert resp.status_code == 302
|
||||
assert "oidc_success=1" in resp.headers["location"]
|
||||
|
||||
# Extract the JWT from the Set-Cookie header and decode it.
|
||||
set_cookie = resp.headers["set-cookie"]
|
||||
cookie_kv = set_cookie.split(";", 1)[0]
|
||||
name, _, token = cookie_kv.partition("=")
|
||||
assert name == "turnstone_auth"
|
||||
assert token
|
||||
|
||||
# Decoding without audience verification first to inspect the claim.
|
||||
claims = pyjwt.decode(
|
||||
token, jwt_secret, algorithms=["HS256"], options={"verify_aud": False}
|
||||
)
|
||||
assert claims["aud"] == "turnstone-server"
|
||||
assert claims["aud"] != "turnstone-console"
|
||||
|
||||
@patch("turnstone.core.auth.provision_oidc_user")
|
||||
@patch("turnstone.core.auth.validate_id_token")
|
||||
@patch("turnstone.core.auth.fetch_jwks", new_callable=AsyncMock)
|
||||
@patch("turnstone.core.auth.exchange_code", new_callable=AsyncMock)
|
||||
def test_jwks_refetch_dedup_when_kid_appears(
|
||||
self,
|
||||
mock_exchange: AsyncMock,
|
||||
mock_fetch_jwks: AsyncMock,
|
||||
mock_validate: Any,
|
||||
mock_provision: Any,
|
||||
authorize_client: TestClient,
|
||||
storage: SQLiteBackend,
|
||||
) -> None:
|
||||
"""If a concurrent caller already refreshed JWKS, second caller skips fetch."""
|
||||
from unittest.mock import patch as obj_patch
|
||||
|
||||
self._seed_pending_state(storage)
|
||||
mock_exchange.return_value = {"id_token": "fake.jwt.token"}
|
||||
mock_provision.return_value = {"user_id": "test-admin", "username": "testadmin"}
|
||||
|
||||
# First validate raises kid-not-found; second succeeds.
|
||||
mock_validate.side_effect = [
|
||||
OIDCKeyNotFoundError("Signing key 'k-rotated' not found"),
|
||||
{"sub": "u1", "email": "u@example.com", "nonce": "test-nonce"},
|
||||
]
|
||||
|
||||
# Pre-populate the JWKS cache so the rotated kid is already
|
||||
# present — analog of a concurrent caller having won the lock.
|
||||
# The retry path must short-circuit and skip the network fetch.
|
||||
authorize_client.app.state.jwks_data = {"keys": [{"kid": "k-rotated", "kty": "RSA"}]}
|
||||
|
||||
with obj_patch("jwt.get_unverified_header", return_value={"kid": "k-rotated"}):
|
||||
resp = authorize_client.get(
|
||||
"/v1/api/auth/oidc/callback?code=authcode&state=valid-state",
|
||||
follow_redirects=False,
|
||||
)
|
||||
|
||||
assert resp.status_code == 302
|
||||
assert "oidc_success=1" in resp.headers["location"]
|
||||
mock_fetch_jwks.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Admin OIDC identity endpoint tests
|
||||
|
||||
@@ -6,6 +6,85 @@ import time
|
||||
|
||||
import pytest
|
||||
|
||||
from turnstone.core.storage import StorageConflictError
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Atomic OIDC user provisioning
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateOIDCUser:
|
||||
def test_create_oidc_user_success(self, db):
|
||||
"""Both rows present after one atomic call."""
|
||||
db.create_oidc_user(
|
||||
user_id="u-new",
|
||||
username="alice",
|
||||
display_name="Alice",
|
||||
password_hash="!oidc",
|
||||
issuer="https://idp.example.com",
|
||||
subject="sub-1",
|
||||
email="alice@example.com",
|
||||
)
|
||||
|
||||
user = db.get_user("u-new")
|
||||
assert user is not None
|
||||
assert user["username"] == "alice"
|
||||
assert user["password_hash"] == "!oidc"
|
||||
|
||||
identity = db.get_oidc_identity("https://idp.example.com", "sub-1")
|
||||
assert identity is not None
|
||||
assert identity["user_id"] == "u-new"
|
||||
assert identity["email"] == "alice@example.com"
|
||||
|
||||
def test_create_oidc_user_username_conflict_rolls_back(self, db):
|
||||
"""Pre-existing username -> StorageConflictError; identity NOT inserted."""
|
||||
db.create_user("u-existing", "alice", "Alice", "$2b$12$hash")
|
||||
|
||||
with pytest.raises(StorageConflictError, match="username"):
|
||||
db.create_oidc_user(
|
||||
user_id="u-new",
|
||||
username="alice",
|
||||
display_name="Alice2",
|
||||
password_hash="!oidc",
|
||||
issuer="https://idp.example.com",
|
||||
subject="sub-1",
|
||||
email="alice2@example.com",
|
||||
)
|
||||
|
||||
# The new user_id row must not exist.
|
||||
assert db.get_user("u-new") is None
|
||||
# The identity row must not exist.
|
||||
assert db.get_oidc_identity("https://idp.example.com", "sub-1") is None
|
||||
# The pre-existing user is untouched.
|
||||
existing = db.get_user("u-existing")
|
||||
assert existing is not None
|
||||
assert existing["password_hash"] == "$2b$12$hash"
|
||||
|
||||
def test_create_oidc_user_identity_conflict_rolls_back(self, db):
|
||||
"""Pre-existing (issuer, subject) -> StorageConflictError; user row rolled back."""
|
||||
db.create_user("u-other", "other", "Other", "!oidc")
|
||||
db.create_oidc_identity("https://idp.example.com", "sub-1", "u-other", "other@example.com")
|
||||
|
||||
with pytest.raises(StorageConflictError, match="OIDC identity"):
|
||||
db.create_oidc_user(
|
||||
user_id="u-new",
|
||||
username="bob",
|
||||
display_name="Bob",
|
||||
password_hash="!oidc",
|
||||
issuer="https://idp.example.com",
|
||||
subject="sub-1",
|
||||
email="bob@example.com",
|
||||
)
|
||||
|
||||
# The candidate user row was rolled back.
|
||||
assert db.get_user("u-new") is None
|
||||
assert db.get_user_by_username("bob") is None
|
||||
# The pre-existing identity still points at the original user.
|
||||
identity = db.get_oidc_identity("https://idp.example.com", "sub-1")
|
||||
assert identity is not None
|
||||
assert identity["user_id"] == "u-other"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OIDC Identity CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -304,3 +383,217 @@ class TestOIDCPendingState:
|
||||
.where(oidc_pending_states.c.state == "state-cleanup")
|
||||
).scalar()
|
||||
assert count == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# count_users / find_existing_usernames
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCountUsers:
|
||||
def test_count_users_empty(self, db):
|
||||
assert db.count_users() == 0
|
||||
|
||||
def test_count_users_after_inserts(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h1")
|
||||
db.create_user("u2", "bob", "Bob", "h2")
|
||||
db.create_user("u3", "carol", "Carol", "h3")
|
||||
assert db.count_users() == 3
|
||||
|
||||
|
||||
class TestFindExistingUsernames:
|
||||
def test_empty_input_returns_empty_set(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h1")
|
||||
assert db.find_existing_usernames([]) == set()
|
||||
|
||||
def test_returns_subset_present_in_db(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h1")
|
||||
db.create_user("u2", "bob", "Bob", "h2")
|
||||
|
||||
existing = db.find_existing_usernames(["alice", "bob", "carol", "dave"])
|
||||
assert existing == {"alice", "bob"}
|
||||
|
||||
def test_no_matches_returns_empty_set(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h1")
|
||||
assert db.find_existing_usernames(["bob", "carol"]) == set()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# replace_oidc_roles
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestReplaceOIDCRoles:
|
||||
def _seed_role(self, db, role_id):
|
||||
db.create_role(role_id, role_id, role_id, "perm.read", False, "")
|
||||
|
||||
def test_inserts_added_roles(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
self._seed_role(db, "role-b")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a", "role-b"})
|
||||
|
||||
assert added == {"role-a", "role-b"}
|
||||
assert removed == set()
|
||||
roles = {r["role_id"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a", "role-b"}
|
||||
|
||||
def test_removes_stale_oidc_roles(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
self._seed_role(db, "role-b")
|
||||
db.assign_role("u1", "role-a", "oidc")
|
||||
db.assign_role("u1", "role-b", "oidc")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == {"role-b"}
|
||||
roles = {r["role_id"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a"}
|
||||
|
||||
def test_preserves_non_oidc_roles(self, db):
|
||||
"""Manually-assigned and oidc-default rows are NOT touched."""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-manual")
|
||||
self._seed_role(db, "role-default")
|
||||
self._seed_role(db, "role-oidc-old")
|
||||
db.assign_role("u1", "role-manual", "admin-ui")
|
||||
db.assign_role("u1", "role-default", "oidc-default")
|
||||
db.assign_role("u1", "role-oidc-old", "oidc")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", set())
|
||||
|
||||
# Only the oidc-assigned row was diffed
|
||||
assert added == set()
|
||||
assert removed == {"role-oidc-old"}
|
||||
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {
|
||||
"role-manual": "admin-ui",
|
||||
"role-default": "oidc-default",
|
||||
}
|
||||
|
||||
def test_no_op_when_desired_matches_current(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
db.assign_role("u1", "role-a", "oidc")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
assert {r["role_id"] for r in db.list_user_roles("u1")} == {"role-a"}
|
||||
|
||||
def test_empty_user_no_oidc_history(self, db):
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", set())
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
|
||||
def test_desired_role_blocked_by_admin_ui_assignment(self, db):
|
||||
"""Desired role already held via admin-ui: untouched, no PK conflict."""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
db.assign_role("u1", "role-a", "admin-ui")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a": "admin-ui"}
|
||||
|
||||
def test_desired_role_blocked_by_oidc_default_assignment(self, db):
|
||||
"""Desired role already held via oidc-default fallback: untouched."""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
db.assign_role("u1", "role-a", "oidc-default")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a": "oidc-default"}
|
||||
|
||||
def test_desired_role_added_alongside_blocked_role(self, db):
|
||||
"""Mixed case: one desired role is blocked (admin-ui), the other inserts cleanly."""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
self._seed_role(db, "role-b")
|
||||
db.assign_role("u1", "role-a", "admin-ui")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a", "role-b"})
|
||||
|
||||
assert added == {"role-b"}
|
||||
assert removed == set()
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a": "admin-ui", "role-b": "oidc"}
|
||||
|
||||
def test_revoke_only_oidc_assigned_roles(self, db):
|
||||
"""OIDC-assigned roles get revoked when not in desired; admin-ui rows survive."""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-manual")
|
||||
self._seed_role(db, "role-oidc-old")
|
||||
self._seed_role(db, "role-default")
|
||||
db.assign_role("u1", "role-manual", "admin-ui")
|
||||
db.assign_role("u1", "role-oidc-old", "oidc")
|
||||
db.assign_role("u1", "role-default", "oidc-default")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", set())
|
||||
|
||||
assert added == set()
|
||||
assert removed == {"role-oidc-old"}
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-manual": "admin-ui", "role-default": "oidc-default"}
|
||||
|
||||
def test_replace_oidc_roles_no_op_steady_state(self, db):
|
||||
"""Steady-state re-login: claims unchanged, function must short-circuit.
|
||||
|
||||
This pins the contract that drives the SQLite optimistic-read fast
|
||||
path — the common case (token refresh with identical role claims)
|
||||
must not acquire a write lock.
|
||||
"""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
self._seed_role(db, "role-b")
|
||||
db.assign_role("u1", "role-a", "oidc")
|
||||
db.assign_role("u1", "role-b", "oidc")
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a", "role-b"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
# All rows still oidc-assigned with identical membership.
|
||||
roles = {r["role_id"]: r["assigned_by"] for r in db.list_user_roles("u1")}
|
||||
assert roles == {"role-a": "oidc", "role-b": "oidc"}
|
||||
|
||||
def test_replace_oidc_roles_returns_post_lock_diff(self, db):
|
||||
"""Returned (added, removed) reflects the post-lock state, not the optimistic read.
|
||||
|
||||
The SQLite implementation re-reads under the write lock to defend
|
||||
against races; the values returned must come from that re-read so
|
||||
callers (apply_role_mapping audit logs) see the actual transition
|
||||
that hit the table. Steady-state input must collapse to empty
|
||||
sets and leave row timestamps unchanged.
|
||||
"""
|
||||
db.create_user("u1", "alice", "Alice", "h")
|
||||
self._seed_role(db, "role-a")
|
||||
db.assign_role("u1", "role-a", "oidc")
|
||||
|
||||
before = db.list_user_roles("u1")
|
||||
assert len(before) == 1
|
||||
original_created = before[0]["assignment_created"]
|
||||
|
||||
added, removed = db.replace_oidc_roles("u1", {"role-a"})
|
||||
|
||||
assert added == set()
|
||||
assert removed == set()
|
||||
# No write occurred — the assignment row's timestamp is untouched.
|
||||
after = db.list_user_roles("u1")
|
||||
assert len(after) == 1
|
||||
assert after[0]["assignment_created"] == original_created
|
||||
|
||||
@@ -15,9 +15,26 @@ def _row(
|
||||
tc_id=None,
|
||||
pdata=None,
|
||||
tool_calls=None,
|
||||
source=None,
|
||||
reminders=None,
|
||||
):
|
||||
"""Build a 7-element conversation row tuple (id, role, ...)."""
|
||||
return (next(_row_ids), role, content, tool_name, tc_id, pdata, tool_calls)
|
||||
"""Build a 9-element conversation row tuple (id, role, ...).
|
||||
|
||||
Trailing ``source`` / ``reminders`` mirror the persisted twins of
|
||||
the in-memory ``_source`` / ``_reminders`` side-channels added in
|
||||
migration 050.
|
||||
"""
|
||||
return (
|
||||
next(_row_ids),
|
||||
role,
|
||||
content,
|
||||
tool_name,
|
||||
tc_id,
|
||||
pdata,
|
||||
tool_calls,
|
||||
source,
|
||||
reminders,
|
||||
)
|
||||
|
||||
|
||||
class TestAssistantWithToolCalls:
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Tests for ``initialize_mcp_crypto_state`` startup gate.
|
||||
|
||||
Phase 3 of the OAuth-MCP RFC: validates fail-loud behavior when an
|
||||
operator forgets the encryption key on a node that hosts OAuth-protected
|
||||
MCP server rows.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import types
|
||||
|
||||
import pytest
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
import turnstone.core.config as cfg_mod
|
||||
from turnstone.core.mcp_crypto import (
|
||||
MCPTokenCipher,
|
||||
MCPTokenStore,
|
||||
initialize_mcp_crypto_state,
|
||||
)
|
||||
|
||||
|
||||
def _patch_security(monkeypatch: pytest.MonkeyPatch, payload: dict) -> None:
|
||||
"""Override ``load_config('security')`` to return ``payload``."""
|
||||
|
||||
def fake(section: str | None = None) -> dict:
|
||||
if section == "security":
|
||||
return payload
|
||||
return {}
|
||||
|
||||
monkeypatch.setattr(cfg_mod, "load_config", fake)
|
||||
|
||||
|
||||
class TestInitializeMcpCryptoState:
|
||||
def test_startup_succeeds_with_no_oauth_user_rows_and_no_key(
|
||||
self, backend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Common case: no key, no oauth_user rows -> sentinels installed."""
|
||||
_patch_security(monkeypatch, {})
|
||||
|
||||
state = types.SimpleNamespace()
|
||||
initialize_mcp_crypto_state(state, node_id="n1")
|
||||
|
||||
assert state.mcp_token_cipher is None
|
||||
assert state.mcp_token_store is None
|
||||
|
||||
def test_startup_succeeds_with_key_and_oauth_user_row(
|
||||
self, backend, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Operator has wired a key and at least one oauth_user row.
|
||||
|
||||
Cipher + store should land on app_state.
|
||||
"""
|
||||
# Plant an oauth_user row.
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-1",
|
||||
name="oauth-srv",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
_patch_security(
|
||||
monkeypatch,
|
||||
{"mcp_token_encryption_key": Fernet.generate_key().decode()},
|
||||
)
|
||||
|
||||
state = types.SimpleNamespace()
|
||||
initialize_mcp_crypto_state(state, node_id="n1")
|
||||
|
||||
assert isinstance(state.mcp_token_cipher, MCPTokenCipher)
|
||||
assert isinstance(state.mcp_token_store, MCPTokenStore)
|
||||
|
||||
def test_startup_aborts_with_oauth_user_row_and_no_key(
|
||||
self, backend, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
"""Misconfiguration: oauth_user row exists, no key -> SystemExit(1)."""
|
||||
backend.create_mcp_server(
|
||||
server_id="srv-1",
|
||||
name="oauth-srv",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
)
|
||||
_patch_security(monkeypatch, {})
|
||||
|
||||
state = types.SimpleNamespace()
|
||||
with (
|
||||
caplog.at_level("ERROR", logger="turnstone.core.mcp_crypto"),
|
||||
pytest.raises(SystemExit) as exc_info,
|
||||
):
|
||||
initialize_mcp_crypto_state(state, node_id="n1")
|
||||
assert exc_info.value.code == 1
|
||||
# Operator-actionable error message names BOTH supported config-key
|
||||
# forms so an operator using the rotation list (plural) is not
|
||||
# misled into thinking only the singular form is valid.
|
||||
messages = " ".join(record.message for record in caplog.records)
|
||||
assert "mcp_token_encryption_keys" in messages
|
||||
assert re.search(r"mcp_token_encryption_key(?!s)", messages) is not None
|
||||
|
||||
def test_startup_aborts_with_invalid_key(
|
||||
self, backend, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
"""Malformed key material should fail loud at startup, not at first use."""
|
||||
_patch_security(monkeypatch, {"mcp_token_encryption_key": "###not-base64###"})
|
||||
|
||||
state = types.SimpleNamespace()
|
||||
with (
|
||||
caplog.at_level("ERROR", logger="turnstone.core.mcp_crypto"),
|
||||
pytest.raises(SystemExit) as exc_info,
|
||||
):
|
||||
initialize_mcp_crypto_state(state, node_id="n1")
|
||||
assert exc_info.value.code == 1
|
||||
+1221
-48
File diff suppressed because it is too large
Load Diff
@@ -163,6 +163,28 @@ class FakeAdapter:
|
||||
return [e for e in self.events if e.kind == kind]
|
||||
|
||||
|
||||
class _FakeRowMapping:
|
||||
"""SQLAlchemy-Row-like wrapper exposing ``_mapping`` over a ``_Row``.
|
||||
|
||||
The real backends return ``Row`` objects with a ``_mapping`` attribute;
|
||||
consumers (e.g. ``CoordinatorIdleObserver._active_children``) prefer
|
||||
``row._mapping[<col>]`` access. This shim mirrors that contract so
|
||||
fakes are interchangeable with real Rows in tests.
|
||||
"""
|
||||
|
||||
def __init__(self, row: _Row) -> None:
|
||||
self._mapping = {
|
||||
"ws_id": row.ws_id,
|
||||
"user_id": row.user_id,
|
||||
"name": row.name,
|
||||
"kind": row.kind,
|
||||
"state": row.state,
|
||||
"parent_ws_id": row.parent_ws_id,
|
||||
"updated": row.updated,
|
||||
"node_id": row.node_id,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Row:
|
||||
ws_id: str
|
||||
@@ -300,6 +322,47 @@ class FakeStorage:
|
||||
"parent_ws_id": row.parent_ws_id,
|
||||
}
|
||||
|
||||
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]:
|
||||
kind_str = kind.value if isinstance(kind, WorkstreamKind) else kind
|
||||
with self.lock:
|
||||
matched: list[_FakeRowMapping] = []
|
||||
for row in self.rows.values():
|
||||
if parent_ws_id is not None and row.parent_ws_id != parent_ws_id:
|
||||
continue
|
||||
if kind_str is not None and row.kind != kind_str:
|
||||
continue
|
||||
if user_id is not None and row.user_id != user_id:
|
||||
continue
|
||||
matched.append(_FakeRowMapping(row))
|
||||
# Order by updated DESC so the consumer's LIMIT semantics match
|
||||
# production (storage backends order this way).
|
||||
matched.sort(key=lambda r: r._mapping["updated"], reverse=True)
|
||||
return matched[:limit]
|
||||
|
||||
def count_workstreams_by_state(
|
||||
self,
|
||||
*,
|
||||
parent_ws_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, int]:
|
||||
counts: dict[str, int] = {}
|
||||
with self.lock:
|
||||
for row in self.rows.values():
|
||||
if parent_ws_id is not None and row.parent_ws_id != parent_ws_id:
|
||||
continue
|
||||
if user_id is not None and row.user_id != user_id:
|
||||
continue
|
||||
counts[row.state] = counts.get(row.state, 0) + 1
|
||||
return counts
|
||||
|
||||
def delete_workstream(self, ws_id: str) -> None:
|
||||
with self.lock:
|
||||
self.rows.pop(ws_id, None)
|
||||
|
||||
@@ -0,0 +1,236 @@
|
||||
"""Tests for ``_format_mcp_dispatch_error`` and the three MCP exec sites.
|
||||
|
||||
The Phase 7b pool dispatcher signals user-actionable failures (consent
|
||||
required, insufficient scope) via ``RuntimeError(json_str)`` where
|
||||
``json_str`` is the structured-error payload built by
|
||||
:func:`turnstone.core.mcp_client._structured_error`. The exec sites in
|
||||
:mod:`turnstone.core.session` previously wrapped that JSON in
|
||||
``f"MCP X error: {e}"``, destroying the structured shape the dashboard
|
||||
renderer keys on. The helper preserves the JSON when the exception
|
||||
text decodes to a structured-error envelope and prefixes otherwise.
|
||||
|
||||
Sibling-bug coverage: every exec site (tool / read_resource /
|
||||
use_prompt) gets two assertions — JSON preserved on a consent-required
|
||||
exception, JSON-prefixed on a generic transport failure.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from tests.test_session import _make_session
|
||||
from turnstone.core.session import _format_mcp_dispatch_error
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit tests for the helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatMcpDispatchError:
|
||||
def test_preserves_consent_required_payload(self) -> None:
|
||||
payload = json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"code": "mcp_consent_required",
|
||||
"server": "srv-x",
|
||||
"detail": "No token for user. Consent flow required.",
|
||||
"consent_url": "/v1/api/mcp/oauth/start?server=srv-x",
|
||||
}
|
||||
}
|
||||
)
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError(payload))
|
||||
assert out == payload
|
||||
|
||||
def test_preserves_insufficient_scope_payload(self) -> None:
|
||||
payload = json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"code": "mcp_insufficient_scope",
|
||||
"server": "srv-x",
|
||||
"detail": "Tool requires elevated scopes.",
|
||||
"scopes_required": ["read", "write"],
|
||||
"consent_url": "/v1/api/mcp/oauth/start?server=srv-x&scopes=read+write",
|
||||
}
|
||||
}
|
||||
)
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError(payload))
|
||||
assert out == payload
|
||||
|
||||
def test_prefixes_generic_runtime_error(self) -> None:
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError("connection lost"))
|
||||
assert out == "MCP tool error: connection lost"
|
||||
|
||||
def test_prefixes_value_error(self) -> None:
|
||||
out = _format_mcp_dispatch_error("MCP tool error", ValueError("bad input"))
|
||||
assert out == "MCP tool error: bad input"
|
||||
|
||||
def test_prefixes_random_json_without_mcp_code(self) -> None:
|
||||
# JSON that isn't a structured-error envelope must NOT be passed
|
||||
# through verbatim — the helper only opens the gate for codes
|
||||
# prefixed ``mcp_``.
|
||||
payload = json.dumps({"foo": "bar"})
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError(payload))
|
||||
assert out == f"MCP tool error: {payload}"
|
||||
|
||||
def test_prefixes_envelope_with_non_mcp_code(self) -> None:
|
||||
payload = json.dumps({"error": {"code": "other_error", "server": "x", "detail": "y"}})
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError(payload))
|
||||
assert out == f"MCP tool error: {payload}"
|
||||
|
||||
def test_prefixes_envelope_without_dict_error(self) -> None:
|
||||
payload = json.dumps({"error": "plain string"})
|
||||
out = _format_mcp_dispatch_error("MCP tool error", RuntimeError(payload))
|
||||
assert out == f"MCP tool error: {payload}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration tests against the three MCP exec sites
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_CONSENT_REQUIRED_JSON = json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"code": "mcp_consent_required",
|
||||
"server": "srv-oauth",
|
||||
"detail": "No token for user. Consent flow required.",
|
||||
"consent_url": "/v1/api/mcp/oauth/start?server=srv-oauth",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _record_outputs(session) -> list[tuple[str, str, str, bool]]:
|
||||
"""Patch ``_report_tool_result`` to capture (call_id, name, output, is_error)."""
|
||||
captures: list[tuple[str, str, str, bool]] = []
|
||||
|
||||
def _capture(call_id: str, name: str, output: str, *, is_error: bool = False) -> None:
|
||||
captures.append((call_id, name, output, is_error))
|
||||
|
||||
session._report_tool_result = _capture # type: ignore[method-assign]
|
||||
return captures
|
||||
|
||||
|
||||
class TestExecMcpToolDispatchError:
|
||||
def test_exec_mcp_tool_preserves_structured_error_json(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool_sync.side_effect = RuntimeError(_CONSENT_REQUIRED_JSON)
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "tc_1",
|
||||
"mcp_func_name": "mcp__srv-oauth__do",
|
||||
"mcp_args": {},
|
||||
}
|
||||
session._exec_mcp_tool(item)
|
||||
|
||||
assert len(captures) == 1
|
||||
_, _, output, is_error = captures[0]
|
||||
assert output == _CONSENT_REQUIRED_JSON
|
||||
assert is_error is True
|
||||
|
||||
def test_exec_mcp_tool_prefixes_non_structured_error(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool_sync.side_effect = RuntimeError("connection lost")
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "tc_2",
|
||||
"mcp_func_name": "mcp__srv-oauth__do",
|
||||
"mcp_args": {},
|
||||
}
|
||||
session._exec_mcp_tool(item)
|
||||
|
||||
assert captures[0][2] == "MCP tool error: connection lost"
|
||||
assert captures[0][3] is True
|
||||
|
||||
|
||||
class TestExecReadResourceDispatchError:
|
||||
def test_exec_read_resource_preserves_structured_error_json(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.read_resource_sync.side_effect = RuntimeError(_CONSENT_REQUIRED_JSON)
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "rc_1",
|
||||
"resource_uri": "https://example.com/r",
|
||||
}
|
||||
# The exec site emits a ``log.warning`` (no ``exc_info`` — bearer-leak
|
||||
# invariant) on failure. Patch the logger so the test doesn't emit
|
||||
# noise to the captured stderr — assertions don't depend on log
|
||||
# output.
|
||||
with patch("turnstone.core.session.log"):
|
||||
session._exec_read_resource(item)
|
||||
|
||||
assert len(captures) == 1
|
||||
assert captures[0][2] == _CONSENT_REQUIRED_JSON
|
||||
assert captures[0][3] is True
|
||||
|
||||
def test_exec_read_resource_prefixes_non_structured_error(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.read_resource_sync.side_effect = RuntimeError("connection lost")
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "rc_2",
|
||||
"resource_uri": "https://example.com/r",
|
||||
}
|
||||
with patch("turnstone.core.session.log"):
|
||||
session._exec_read_resource(item)
|
||||
|
||||
assert captures[0][2] == "MCP resource error: connection lost"
|
||||
assert captures[0][3] is True
|
||||
|
||||
|
||||
class TestExecUsePromptDispatchError:
|
||||
def test_exec_use_prompt_preserves_structured_error_json(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get_prompt_sync.side_effect = RuntimeError(_CONSENT_REQUIRED_JSON)
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "pc_1",
|
||||
"prompt_name": "mcp__srv-oauth__greet",
|
||||
"prompt_arguments": {},
|
||||
}
|
||||
with patch("turnstone.core.session.log"):
|
||||
session._exec_use_prompt(item)
|
||||
|
||||
assert len(captures) == 1
|
||||
assert captures[0][2] == _CONSENT_REQUIRED_JSON
|
||||
assert captures[0][3] is True
|
||||
|
||||
def test_exec_use_prompt_prefixes_non_structured_error(self, tmp_db) -> None:
|
||||
session = _make_session()
|
||||
captures = _record_outputs(session)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get_prompt_sync.side_effect = RuntimeError("connection lost")
|
||||
session._mcp_client = mock_client
|
||||
|
||||
item = {
|
||||
"call_id": "pc_2",
|
||||
"prompt_name": "mcp__srv-oauth__greet",
|
||||
"prompt_arguments": {},
|
||||
}
|
||||
with patch("turnstone.core.session.log"):
|
||||
session._exec_use_prompt(item)
|
||||
|
||||
assert captures[0][2] == "MCP prompt error: connection lost"
|
||||
assert captures[0][3] is True
|
||||
+117
-10
@@ -451,6 +451,71 @@ class TestInterruptedWorkstreamRepair:
|
||||
assert msgs[1]["role"] == "assistant"
|
||||
assert msgs[2]["role"] == "user"
|
||||
|
||||
def test_repair_false_preserves_partial_trailing_turn(self, tmp_db):
|
||||
"""``repair=False`` is the display-read contract for ``/history``.
|
||||
|
||||
The default repair pass strips the trailing
|
||||
``assistant(tool_calls)`` when not all tool results are persisted
|
||||
— correct for ``session.resume`` (LLM context), wrong for the
|
||||
REST display read. A user refreshing the coordinator page mid-
|
||||
tool-execution would otherwise lose the entire trailing turn
|
||||
from the UI. ``repair=False`` returns the raw persisted state.
|
||||
"""
|
||||
import json
|
||||
|
||||
tc_json = json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
||||
},
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
||||
},
|
||||
]
|
||||
)
|
||||
save_message("s1", "user", "hello")
|
||||
save_message("s1", "assistant", "Checking", tool_calls=tc_json)
|
||||
save_message("s1", "tool", "file.txt", tool_call_id="call_1")
|
||||
# No call_2 result persisted — mid-execution refresh.
|
||||
msgs = get_storage().load_messages("s1", repair=False)
|
||||
# All three rows survive — the trailing partial turn is what the
|
||||
# operator was actually watching live.
|
||||
assert [m["role"] for m in msgs] == ["user", "assistant", "tool"]
|
||||
assert msgs[1].get("tool_calls") and len(msgs[1]["tool_calls"]) == 2
|
||||
assert msgs[2]["tool_call_id"] == "call_1"
|
||||
|
||||
def test_repair_false_does_not_synthesize_orphan_results(self, tmp_db):
|
||||
"""``repair=False`` must NOT splice synthetic ``"Tool execution
|
||||
was cancelled."`` rows for mid-conversation orphans either —
|
||||
the operator never saw those rows, and showing them would
|
||||
invent UI content that doesn't reflect persisted state.
|
||||
"""
|
||||
import json
|
||||
|
||||
tc_json = json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
||||
},
|
||||
]
|
||||
)
|
||||
save_message("s1", "user", "first")
|
||||
save_message("s1", "assistant", "Working", tool_calls=tc_json)
|
||||
# Cancel landed before any tool result — next turn happens.
|
||||
save_message("s1", "user", "second")
|
||||
save_message("s1", "assistant", "ok")
|
||||
msgs = get_storage().load_messages("s1", repair=False)
|
||||
roles = [m["role"] for m in msgs]
|
||||
# No synthetic tool row spliced after the orphaned tool_calls.
|
||||
assert roles == ["user", "assistant", "user", "assistant"]
|
||||
assert all(m["role"] != "tool" for m in msgs)
|
||||
|
||||
|
||||
# ── Workstream config persistence ─────────────────────────────────────
|
||||
|
||||
@@ -914,8 +979,11 @@ class TestMCPToolGating:
|
||||
"""read_resource excluded when MCP client has no resources."""
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
mcp_client.resource_count = 0
|
||||
mcp_client.prompt_count = 2
|
||||
# Phase 7b: gating uses ``*_count_for_user`` so the test mocks
|
||||
# the per-user variant (the property remains for static-only
|
||||
# admin paths). Returning 0 / 2 mirrors the prior contract.
|
||||
mcp_client.resource_count_for_user.return_value = 0
|
||||
mcp_client.prompt_count_for_user.return_value = 2
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
@@ -937,8 +1005,8 @@ class TestMCPToolGating:
|
||||
"""use_prompt excluded when MCP client has no prompts."""
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
mcp_client.resource_count = 3
|
||||
mcp_client.prompt_count = 0
|
||||
mcp_client.resource_count_for_user.return_value = 3
|
||||
mcp_client.prompt_count_for_user.return_value = 0
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
@@ -960,8 +1028,8 @@ class TestMCPToolGating:
|
||||
"""Both tools present when MCP client has resources and prompts."""
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
mcp_client.resource_count = 1
|
||||
mcp_client.prompt_count = 1
|
||||
mcp_client.resource_count_for_user.return_value = 1
|
||||
mcp_client.prompt_count_for_user.return_value = 1
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
@@ -983,8 +1051,8 @@ class TestMCPToolGating:
|
||||
"""Gating applies even when tool_search is active (client-side path)."""
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
mcp_client.resource_count = 0
|
||||
mcp_client.prompt_count = 0
|
||||
mcp_client.resource_count_for_user.return_value = 0
|
||||
mcp_client.prompt_count_for_user.return_value = 0
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
@@ -1012,8 +1080,8 @@ class TestMCPToolGating:
|
||||
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
mcp_client.resource_count = 0
|
||||
mcp_client.prompt_count = 0
|
||||
mcp_client.resource_count_for_user.return_value = 0
|
||||
mcp_client.prompt_count_for_user.return_value = 0
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
@@ -1034,3 +1102,42 @@ class TestMCPToolGating:
|
||||
names = [t.get("function", {}).get("name") for t in tools]
|
||||
assert "read_resource" not in names
|
||||
assert "use_prompt" not in names
|
||||
|
||||
def test_pool_only_user_keeps_read_resource_and_use_prompt(self, tmp_db, mock_openai_client):
|
||||
"""Phase 7b canary: a pool-only user (static catalog empty) still
|
||||
sees ``read_resource`` and ``use_prompt`` because the gating
|
||||
consults ``*_count_for_user`` (scope decision 0.2).
|
||||
|
||||
Drives ``resource_count = prompt_count = 0`` (the static-only
|
||||
properties are zero) but ``*_count_for_user(uid) > 0`` because
|
||||
the user has pool entries; the tools must remain visible.
|
||||
"""
|
||||
mcp_client = MagicMock()
|
||||
mcp_client.get_tools.return_value = []
|
||||
# Static catalog is empty; admin-style legacy properties say 0.
|
||||
mcp_client.resource_count = 0
|
||||
mcp_client.prompt_count = 0
|
||||
# Per-user variant reports the user's pool entries.
|
||||
mcp_client.resource_count_for_user.return_value = 2
|
||||
mcp_client.prompt_count_for_user.return_value = 1
|
||||
|
||||
session = ChatSession(
|
||||
client=mock_openai_client,
|
||||
model="local-model",
|
||||
ui=MagicMock(),
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=1000,
|
||||
tool_timeout=10,
|
||||
mcp_client=mcp_client,
|
||||
user_id="pool-only-user",
|
||||
)
|
||||
|
||||
tools = session._get_active_tools()
|
||||
names = [t.get("function", {}).get("name") for t in tools]
|
||||
assert "read_resource" in names
|
||||
assert "use_prompt" in names
|
||||
# Verify the per-user gate was actually consulted with the
|
||||
# session's ``user_id`` (sanity-check on the wiring).
|
||||
mcp_client.resource_count_for_user.assert_any_call("pool-only-user")
|
||||
mcp_client.prompt_count_for_user.assert_any_call("pool-only-user")
|
||||
|
||||
@@ -0,0 +1,232 @@
|
||||
"""Tests for the SKILL.md parse admin API endpoint.
|
||||
|
||||
The endpoint is a thin permission-checked wrapper around
|
||||
``turnstone.core.skill_parser.parse_skill_md``. These tests cover the
|
||||
routing, auth, and error-handling layers — parser semantics live in
|
||||
``test_skill_parser.py``.
|
||||
"""
|
||||
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from turnstone.console.server import admin_parse_skill
|
||||
from turnstone.core.auth import AuthResult
|
||||
|
||||
|
||||
class _InjectAuthMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id="test-user",
|
||||
scopes=frozenset({"approve"}),
|
||||
token_source="config",
|
||||
permissions=frozenset({"read", "write", "approve", "admin.skills"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
class _InjectAuthNoSkillsMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: Any) -> Response:
|
||||
request.state.auth_result = AuthResult(
|
||||
user_id="test-user",
|
||||
scopes=frozenset({"approve"}),
|
||||
token_source="jwt",
|
||||
permissions=frozenset({"read", "write", "approve"}),
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
_ROUTES = [
|
||||
Mount(
|
||||
"/v1",
|
||||
routes=[
|
||||
Route("/api/admin/skills/parse", admin_parse_skill, methods=["POST"]),
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client() -> TestClient:
|
||||
app = Starlette(
|
||||
routes=_ROUTES,
|
||||
middleware=[Middleware(_InjectAuthMiddleware)],
|
||||
)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client_no_perm() -> TestClient:
|
||||
app = Starlette(
|
||||
routes=_ROUTES,
|
||||
middleware=[Middleware(_InjectAuthNoSkillsMiddleware)],
|
||||
)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
_FULL_SKILL = """\
|
||||
---
|
||||
name: code-review
|
||||
description: Automated code review skill
|
||||
author: Test Author
|
||||
version: 2.0.0
|
||||
tags: [python, review, quality]
|
||||
allowed-tools: [read_file, list_directory]
|
||||
license: MIT
|
||||
compatibility: ">=0.7"
|
||||
---
|
||||
|
||||
# Code Review
|
||||
|
||||
Review code for best practices.
|
||||
"""
|
||||
|
||||
_MINIMAL_SKILL = """\
|
||||
---
|
||||
name: minimal
|
||||
---
|
||||
|
||||
Just some content.
|
||||
"""
|
||||
|
||||
|
||||
class TestParseSkill:
|
||||
def test_parses_full_frontmatter(self, client: TestClient) -> None:
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": _FULL_SKILL})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "code-review"
|
||||
assert data["description"] == "Automated code review skill"
|
||||
assert data["author"] == "Test Author"
|
||||
assert data["version"] == "2.0.0"
|
||||
assert data["tags"] == ["python", "review", "quality"]
|
||||
assert data["allowed_tools"] == ["read_file", "list_directory"]
|
||||
assert data["license"] == "MIT"
|
||||
assert data["compatibility"] == ">=0.7"
|
||||
assert "# Code Review" in data["content"]
|
||||
# Frontmatter should not leak into the body.
|
||||
assert "name: code-review" not in data["content"]
|
||||
# ParsedSkill carries raw_frontmatter (the full YAML dict) but the
|
||||
# handler whitelists fields by hand to avoid leaking arbitrary keys.
|
||||
# Pin that contract — a future refactor to dataclasses.asdict would
|
||||
# silently break it without this assertion.
|
||||
assert "raw_frontmatter" not in data
|
||||
|
||||
def test_parses_minimal_frontmatter(self, client: TestClient) -> None:
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": _MINIMAL_SKILL})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "minimal"
|
||||
assert data["description"] == "Just some content."
|
||||
assert data["version"] == "1.0.0"
|
||||
assert data["tags"] == []
|
||||
assert data["allowed_tools"] == []
|
||||
assert data["license"] == ""
|
||||
|
||||
def test_anthropic_nested_metadata_tags(self, client: TestClient) -> None:
|
||||
# Anthropic-style skill puts tags under metadata.tags rather than
|
||||
# at the top level — the parser must handle both layouts.
|
||||
raw = """\
|
||||
---
|
||||
name: nested-meta
|
||||
description: A skill using nested metadata
|
||||
metadata:
|
||||
tags: [alpha, beta]
|
||||
author: Anthropic
|
||||
version: 3.1.4
|
||||
---
|
||||
|
||||
Body.
|
||||
"""
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": raw})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["tags"] == ["alpha", "beta"]
|
||||
assert data["author"] == "Anthropic"
|
||||
assert data["version"] == "3.1.4"
|
||||
|
||||
def test_unquoted_colon_in_description(self, client: TestClient) -> None:
|
||||
# Common cross-client mistake: ``description: Use when: the user...``
|
||||
# The parser retries with the description value quoted.
|
||||
raw = """\
|
||||
---
|
||||
name: colon-desc
|
||||
description: Use when: the user asks for a review
|
||||
---
|
||||
|
||||
Body.
|
||||
"""
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": raw})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "colon-desc"
|
||||
assert "Use when" in data["description"]
|
||||
|
||||
def test_missing_name_returns_400(self, client: TestClient) -> None:
|
||||
raw = """\
|
||||
---
|
||||
description: No name field
|
||||
---
|
||||
|
||||
Body.
|
||||
"""
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": raw})
|
||||
assert resp.status_code == 400
|
||||
assert "name" in resp.json()["error"].lower()
|
||||
|
||||
def test_missing_raw_returns_400(self, client: TestClient) -> None:
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={})
|
||||
assert resp.status_code == 400
|
||||
assert "raw" in resp.json()["error"].lower()
|
||||
|
||||
def test_blank_raw_returns_400(self, client: TestClient) -> None:
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": " \n"})
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_invalid_yaml_returns_400(self, client: TestClient) -> None:
|
||||
# YAML that the malformed-description retry can't fix.
|
||||
raw = "---\nname: [not, valid, here\n---\nBody.\n"
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": raw})
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_requires_admin_skills_permission(self, client_no_perm: TestClient) -> None:
|
||||
resp = client_no_perm.post("/v1/api/admin/skills/parse", json={"raw": _MINIMAL_SKILL})
|
||||
assert resp.status_code == 403
|
||||
|
||||
def test_oversized_content_length_returns_413(self, client: TestClient) -> None:
|
||||
# Content-Length pre-check rejects oversized bodies before they're
|
||||
# buffered into memory. Caps worker memory against an admin-token
|
||||
# holder spraying multi-GB JSON. The threshold is generous (~4×
|
||||
# the per-string cap) so payload here must clearly exceed it.
|
||||
oversized = "a" * 200_000
|
||||
resp = client.post("/v1/api/admin/skills/parse", json={"raw": oversized})
|
||||
assert resp.status_code == 413
|
||||
|
||||
def test_oversized_raw_chunked_returns_413(self, client: TestClient) -> None:
|
||||
# When the client sends Transfer-Encoding: chunked there is no
|
||||
# Content-Length header, so the pre-check is skipped and the
|
||||
# application-layer cap is the only line of defence. httpx switches
|
||||
# to chunked when the body is a generator.
|
||||
def _gen() -> Iterator[bytes]:
|
||||
yield b'{"raw":"' + b"a" * 33_000 + b'"}'
|
||||
|
||||
resp = client.post(
|
||||
"/v1/api/admin/skills/parse",
|
||||
content=_gen(),
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
assert resp.status_code == 413
|
||||
assert "raw" in resp.json()["error"].lower()
|
||||
@@ -0,0 +1,158 @@
|
||||
"""Tests for ``_source`` / ``_reminders`` round-tripping through both
|
||||
storage backends.
|
||||
|
||||
Persisting the in-memory side-channels lets multi-tab / multi-device
|
||||
replay show the same metacognitive bubble shape the originating tab
|
||||
saw live — see ``docs/design/watch-card-ux.md`` §1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from turnstone.core.storage._schema import conversations
|
||||
|
||||
|
||||
class TestSourceRoundtrip:
|
||||
def test_source_roundtrip(self, backend):
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "", source="system_nudge")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert len(msgs) == 1
|
||||
assert msgs[0]["role"] == "user"
|
||||
assert msgs[0]["content"] == ""
|
||||
assert msgs[0].get("_source") == "system_nudge"
|
||||
|
||||
def test_source_absent_when_not_set(self, backend):
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hello")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert "_source" not in msgs[0]
|
||||
|
||||
|
||||
class TestRemindersRoundtrip:
|
||||
def test_reminders_roundtrip(self, backend):
|
||||
backend.register_workstream("s1")
|
||||
payload = [
|
||||
{
|
||||
"type": "watch_triggered",
|
||||
"text": "$ ls\nfile.txt\n",
|
||||
"watch_name": "w1",
|
||||
"command": "ls",
|
||||
"poll_count": 2,
|
||||
"max_polls": 100,
|
||||
"is_final": False,
|
||||
}
|
||||
]
|
||||
backend.save_message(
|
||||
"s1",
|
||||
"user",
|
||||
"",
|
||||
source="system_nudge",
|
||||
reminders=json.dumps(payload, separators=(",", ":")),
|
||||
)
|
||||
msgs = backend.load_messages("s1")
|
||||
assert msgs[0].get("_reminders") == payload
|
||||
# Optional fields preserved verbatim.
|
||||
rem = msgs[0]["_reminders"][0]
|
||||
assert rem["watch_name"] == "w1"
|
||||
assert rem["command"] == "ls"
|
||||
assert rem["poll_count"] == 2
|
||||
assert rem["max_polls"] == 100
|
||||
assert rem["is_final"] is False
|
||||
|
||||
def test_reminders_null_renders_as_no_key(self, backend):
|
||||
"""Absent vs. empty-list should map to the same shape on the
|
||||
load side: ``_reminders`` simply not present in the dict.
|
||||
Mirrors the ``_attachments_meta`` precedent in
|
||||
``reconstruct_messages``.
|
||||
"""
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message("s1", "user", "hello")
|
||||
msgs = backend.load_messages("s1")
|
||||
assert "_reminders" not in msgs[0]
|
||||
|
||||
def test_tool_reminders_roundtrip(self, backend):
|
||||
backend.register_workstream("s1")
|
||||
# Build a minimal valid history: assistant turn with one
|
||||
# tool_call followed by the tool result that carries the
|
||||
# tool-channel reminder. Without the assistant turn the
|
||||
# tool row would be orphaned and stripped by the repair pass.
|
||||
tc_json = json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": "{}"},
|
||||
}
|
||||
]
|
||||
)
|
||||
backend.save_message("s1", "user", "go")
|
||||
backend.save_message("s1", "assistant", None, tool_calls=tc_json)
|
||||
payload = [{"type": "tool_error", "text": "command failed"}]
|
||||
backend.save_message(
|
||||
"s1",
|
||||
"tool",
|
||||
"boom",
|
||||
tool_call_id="c1",
|
||||
reminders=json.dumps(payload, separators=(",", ":")),
|
||||
)
|
||||
msgs = backend.load_messages("s1")
|
||||
# Find the tool message and assert reminders survived load.
|
||||
tool_msgs = [m for m in msgs if m.get("role") == "tool"]
|
||||
assert len(tool_msgs) == 1
|
||||
assert tool_msgs[0].get("_reminders") == payload
|
||||
|
||||
def test_nul_bytes_stripped_from_source_and_reminders(self, backend):
|
||||
"""NUL bytes must be stripped at the storage layer.
|
||||
|
||||
Producers (``sanitize_payload`` on the watch dispatch path,
|
||||
constants for non-watch nudges) already strip NUL today so
|
||||
nothing in production reaches this clamp — but the layer is
|
||||
the tripwire if a future producer forgets, mirroring how
|
||||
``content`` and ``provider_data`` are sanitized. PostgreSQL
|
||||
TEXT columns reject NUL outright, so the sanitization is also
|
||||
a hard correctness invariant on that backend.
|
||||
|
||||
``json.dumps`` already escapes NUL inside string values to
|
||||
``\\u0000`` so a real NUL byte can't enter ``_reminders`` via
|
||||
the normal encode path — the test feeds a raw NUL directly to
|
||||
cover the bypass case (a future producer that hand-builds the
|
||||
column string).
|
||||
"""
|
||||
backend.register_workstream("s1")
|
||||
backend.save_message(
|
||||
"s1",
|
||||
"user",
|
||||
"",
|
||||
source="system_nudge\x00",
|
||||
reminders='[{"type":"watch_triggered","text":"ok\x00bad"}]',
|
||||
)
|
||||
msgs = backend.load_messages("s1")
|
||||
assert msgs[0].get("_source") == "system_nudge"
|
||||
assert msgs[0].get("_reminders") == [{"type": "watch_triggered", "text": "okbad"}]
|
||||
|
||||
def test_malformed_reminders_json_does_not_crash_load(self, backend):
|
||||
"""A garbage string in the column must not abort the whole
|
||||
load — mirrors the ``provider_data`` JSON-decode-suppress
|
||||
pattern. Concretely: write a row with valid columns BUT a
|
||||
corrupted ``_reminders`` value via raw SQL, then verify the
|
||||
load returns the message with no ``_reminders`` key (rather
|
||||
than raising or surfacing the garbage).
|
||||
"""
|
||||
backend.register_workstream("s1")
|
||||
msg_id = backend.save_message("s1", "user", "hello")
|
||||
with backend._engine.connect() as conn:
|
||||
conn.execute(
|
||||
sa.update(conversations)
|
||||
.where(conversations.c.id == msg_id)
|
||||
.values(_reminders="this is not json {{")
|
||||
)
|
||||
conn.commit()
|
||||
msgs = backend.load_messages("s1")
|
||||
assert len(msgs) == 1
|
||||
# Garbage suppressed silently — key absent, content intact.
|
||||
assert "_reminders" not in msgs[0]
|
||||
assert msgs[0]["content"] == "hello"
|
||||
@@ -939,6 +939,64 @@ class TestTouchWorkstream:
|
||||
backend.touch_workstream("nonexistent") # must not raise
|
||||
|
||||
|
||||
# -- MCP OAuth columns ---------------------------------------------------------
|
||||
|
||||
|
||||
class TestMcpServerOauthColumns:
|
||||
def test_mcp_servers_oauth_columns_round_trip(self, backend: Any) -> None:
|
||||
"""An oauth_user row round-trips through create -> get with all
|
||||
seven OAuth text columns intact."""
|
||||
sid = "oauth-srv-1"
|
||||
backend.create_mcp_server(
|
||||
server_id=sid,
|
||||
name="oauth-srv",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="cli_abc123",
|
||||
oauth_scopes="openid profile",
|
||||
oauth_audience="https://mcp.example.com",
|
||||
oauth_registration_mode="preregistered",
|
||||
oauth_authorization_server_url="https://auth.example.com",
|
||||
oauth_as_issuer_cached="https://auth.example.com",
|
||||
)
|
||||
s = backend.get_mcp_server(sid)
|
||||
assert s is not None
|
||||
assert s["auth_type"] == "oauth_user"
|
||||
assert s["oauth_client_id"] == "cli_abc123"
|
||||
assert s["oauth_scopes"] == "openid profile"
|
||||
assert s["oauth_audience"] == "https://mcp.example.com"
|
||||
assert s["oauth_registration_mode"] == "preregistered"
|
||||
assert s["oauth_authorization_server_url"] == "https://auth.example.com"
|
||||
assert s["oauth_as_issuer_cached"] == "https://auth.example.com"
|
||||
# Phase 2 leaves the ciphertext slot NULL even when other oauth
|
||||
# fields are populated; Phase 3 wires the encryption write path.
|
||||
assert s["oauth_client_secret_ct"] is None
|
||||
|
||||
def test_update_auth_type_static_to_oauth(self, backend: Any) -> None:
|
||||
sid = "oauth-srv-2"
|
||||
backend.create_mcp_server(
|
||||
server_id=sid,
|
||||
name="static-then-oauth",
|
||||
transport="streamable-http",
|
||||
url="https://mcp.example.com/sse",
|
||||
)
|
||||
assert backend.get_mcp_server(sid)["auth_type"] == "static"
|
||||
|
||||
ok = backend.update_mcp_server(
|
||||
sid,
|
||||
auth_type="oauth_user",
|
||||
oauth_client_id="cli_after",
|
||||
oauth_audience="https://mcp.example.com",
|
||||
)
|
||||
assert ok is True
|
||||
s = backend.get_mcp_server(sid)
|
||||
assert s is not None
|
||||
assert s["auth_type"] == "oauth_user"
|
||||
assert s["oauth_client_id"] == "cli_after"
|
||||
assert s["oauth_audience"] == "https://mcp.example.com"
|
||||
|
||||
|
||||
# -- Lifecycle -----------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from turnstone.core.tool_advisory import (
|
||||
GuardAdvisory,
|
||||
MetacognitiveAdvisory,
|
||||
UserInterjection,
|
||||
escape_wrapper_tags,
|
||||
parse_priority,
|
||||
render_system_reminder,
|
||||
wrap_tool_result,
|
||||
@@ -224,6 +225,63 @@ class TestMetacognitiveAdvisory:
|
||||
assert "don't repeat tool calls" in result
|
||||
|
||||
|
||||
class TestEscapeWrapperTags:
|
||||
"""``escape_wrapper_tags`` must round-trip through
|
||||
``_entity_decode_wrapper_tags`` for any input — not just text that
|
||||
happens to contain only wrapper tags.
|
||||
"""
|
||||
|
||||
def test_short_circuit_passes_through_plain_text(self) -> None:
|
||||
"""No ``<`` and no ``&`` — ``escape_wrapper_tags`` must avoid
|
||||
the four ``replace`` chains. Common case for most tool outputs;
|
||||
the short-circuit keeps wrap_tool_result's overhead near zero."""
|
||||
text = "plain text without any markup"
|
||||
assert escape_wrapper_tags(text) == text
|
||||
|
||||
def test_escape_wrapper_tags_round_trips_preexisting_entities(self) -> None:
|
||||
"""Asymmetry guard — a tool output that happens to contain the
|
||||
literal string ``<tool_output>`` (e.g. documentation
|
||||
describing the wrapper format) must round-trip identically.
|
||||
Without escaping ``&`` first, encode→decode would produce the
|
||||
bare ``<tool_output>`` tag, fabricating an envelope the wrapper
|
||||
layer never produced."""
|
||||
from turnstone.core.history_decoration import _entity_decode_wrapper_tags
|
||||
|
||||
text = "I describe XML tags like <tool_output> in my docs."
|
||||
encoded = escape_wrapper_tags(text)
|
||||
# Sanity: the original literal got escaped to a sentinel form
|
||||
# that can't collide with our wrapper-tag escapes.
|
||||
assert "&lt;tool_output&gt;" in encoded
|
||||
assert "<tool_output>" not in encoded
|
||||
# Round-trip back to the literal source.
|
||||
assert _entity_decode_wrapper_tags(encoded) == text
|
||||
|
||||
def test_escape_wrapper_tags_round_trips_real_wrapper_tag(self) -> None:
|
||||
"""A literal ``<tool_output>`` in source text round-trips back
|
||||
correctly — encoding produces ``<tool_output>`` (no
|
||||
``&`` prefix because there was no pre-existing entity), and
|
||||
decoding restores the literal."""
|
||||
from turnstone.core.history_decoration import _entity_decode_wrapper_tags
|
||||
|
||||
text = "Here is a literal <tool_output> tag in my doc."
|
||||
encoded = escape_wrapper_tags(text)
|
||||
assert "<tool_output>" not in encoded
|
||||
assert "<tool_output>" in encoded
|
||||
assert _entity_decode_wrapper_tags(encoded) == text
|
||||
|
||||
def test_escape_wrapper_tags_round_trips_mixed_content(self) -> None:
|
||||
"""Mixed: literal wrapper tags AND pre-existing entity
|
||||
references — both round-trip."""
|
||||
from turnstone.core.history_decoration import _entity_decode_wrapper_tags
|
||||
|
||||
text = (
|
||||
"Mixed: literal <tool_output> next to escaped <system-reminder> "
|
||||
"and a stray & on its own."
|
||||
)
|
||||
encoded = escape_wrapper_tags(text)
|
||||
assert _entity_decode_wrapper_tags(encoded) == text
|
||||
|
||||
|
||||
class TestRenderSystemReminder:
|
||||
"""render_system_reminder builds a standalone <system-reminder> envelope."""
|
||||
|
||||
|
||||
+82
-5
@@ -9,6 +9,7 @@ import pytest
|
||||
|
||||
from turnstone.core.watch import (
|
||||
WatchRunner,
|
||||
build_watch_reminder,
|
||||
evaluate_condition,
|
||||
format_interval,
|
||||
format_watch_message,
|
||||
@@ -301,6 +302,62 @@ class TestFormatWatchMessage:
|
||||
assert "max polls" in msg.lower()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_watch_reminder
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildWatchReminder:
|
||||
"""The structured-reminder builder lifts ``format_watch_message``'s
|
||||
args into a dict the dispatch closure can pass to
|
||||
``WatchRunner._dispatch_result``. ``text`` matches the formatter's
|
||||
output verbatim (so compaction / channel adapters / wire splice
|
||||
keep their behaviour), and the optional fields ride alongside for
|
||||
the frontend's ``.msg.watch-result`` card.
|
||||
"""
|
||||
|
||||
def test_emits_text_body_and_fields(self):
|
||||
kwargs = dict(
|
||||
name="pr-review",
|
||||
command="gh pr view --json state",
|
||||
output='{"state": "MERGED"}',
|
||||
poll_count=5,
|
||||
max_polls=100,
|
||||
elapsed_secs=1500,
|
||||
stop_on='data["state"] == "MERGED"',
|
||||
is_final=True,
|
||||
reason='condition met: data["state"] == "MERGED"',
|
||||
)
|
||||
reminder = build_watch_reminder(**kwargs)
|
||||
# Round-trip with format_watch_message — text is the same body
|
||||
# the wire splice + channel adapters have always seen.
|
||||
assert reminder["text"] == format_watch_message(**kwargs)
|
||||
# Optional fields ride alongside.
|
||||
assert reminder["type"] == "watch_triggered"
|
||||
assert reminder["watch_name"] == "pr-review"
|
||||
assert reminder["command"] == "gh pr view --json state"
|
||||
assert reminder["poll_count"] == 5
|
||||
assert reminder["max_polls"] == 100
|
||||
assert reminder["is_final"] is True
|
||||
|
||||
def test_non_final_carries_is_final_false(self):
|
||||
reminder = build_watch_reminder(
|
||||
name="deploy",
|
||||
command="curl -s http://localhost/health",
|
||||
output="ok",
|
||||
poll_count=3,
|
||||
max_polls=50,
|
||||
elapsed_secs=90,
|
||||
stop_on=None,
|
||||
is_final=False,
|
||||
reason="",
|
||||
)
|
||||
assert reminder["is_final"] is False
|
||||
assert reminder["poll_count"] == 3
|
||||
# No "auto-cancelled" body for non-final fires.
|
||||
assert "auto-cancelled" not in reminder["text"].lower()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WatchRunner
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -450,13 +507,17 @@ class TestWatchRunner:
|
||||
runner.set_dispatch_fn("ws-1", fn1)
|
||||
runner.set_dispatch_fn("ws-2", fn2)
|
||||
|
||||
runner._dispatch_result("ws-1", "msg1")
|
||||
fn1.assert_called_once_with("msg1")
|
||||
# ``_dispatch_result`` takes a structured reminder dict, not a
|
||||
# bare string.
|
||||
reminder1 = {"type": "watch_triggered", "text": "msg1"}
|
||||
runner._dispatch_result("ws-1", reminder1, "watch-a")
|
||||
fn1.assert_called_once_with(reminder1, "watch-a")
|
||||
fn2.assert_not_called()
|
||||
|
||||
runner.remove_dispatch_fn("ws-1")
|
||||
# After removal, dispatch should try restore_fn
|
||||
runner._dispatch_result("ws-1", "msg2")
|
||||
reminder2 = {"type": "watch_triggered", "text": "msg2"}
|
||||
runner._dispatch_result("ws-1", reminder2, "watch-b")
|
||||
fn1.assert_called_once() # still just the one call
|
||||
|
||||
def test_restore_fn_called_for_evicted(self):
|
||||
@@ -464,9 +525,25 @@ class TestWatchRunner:
|
||||
restore_fn = MagicMock(return_value=restored_fn)
|
||||
runner = self._make_runner(restore_fn=restore_fn)
|
||||
|
||||
runner._dispatch_result("ws-evicted", "hello")
|
||||
reminder = {"type": "watch_triggered", "text": "hello"}
|
||||
runner._dispatch_result("ws-evicted", reminder, "watch-x")
|
||||
restore_fn.assert_called_once_with("ws-evicted")
|
||||
restored_fn.assert_called_once_with("hello")
|
||||
restored_fn.assert_called_once_with(reminder, "watch-x")
|
||||
|
||||
def test_get_dispatch_fn_returns_registered_fn(self):
|
||||
"""``get_dispatch_fn`` is the public accessor used by the
|
||||
server-side restore path to retrieve the per-ws closure that
|
||||
``set_watch_runner`` constructed during workstream rehydrate.
|
||||
"""
|
||||
runner = self._make_runner()
|
||||
fn = MagicMock()
|
||||
runner.set_dispatch_fn("ws-1", fn)
|
||||
assert runner.get_dispatch_fn("ws-1") is fn
|
||||
# Unknown ws → None.
|
||||
assert runner.get_dispatch_fn("ws-missing") is None
|
||||
# After removal → None.
|
||||
runner.remove_dispatch_fn("ws-1")
|
||||
assert runner.get_dispatch_fn("ws-1") is None
|
||||
|
||||
def test_run_command_success(self):
|
||||
runner = self._make_runner()
|
||||
|
||||
+402
-199
@@ -1,206 +1,409 @@
|
||||
"""Tests for _make_watch_dispatch error/cancel handling and concurrency guards."""
|
||||
"""Tests for the watch dispatch closure built inside ``set_watch_runner``.
|
||||
|
||||
The closure routes watch results onto the per-session :class:`NudgeQueue`
|
||||
under the unified pull-model surface. Each test focuses on one
|
||||
assertion: enqueue shape, sanitisation, soft-cap drop-oldest,
|
||||
``valid_until`` predicate, and concurrent-enqueue safety.
|
||||
|
||||
Tests in this file replace the pre-switchover suite that pinned the
|
||||
``_make_watch_dispatch`` worker-spawn / ``_watch_pending`` machinery —
|
||||
the contracts those tests pinned no longer exist. See
|
||||
``tests/test_watch.py`` for the still-relevant ``WatchRunner``
|
||||
mechanics tests, and ``tests/test_watch_integration.py`` for the
|
||||
boundary-crossing integration test covering the chat-loop drain.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from turnstone.core.session import GenerationCancelled
|
||||
from turnstone.core.workstream import Workstream
|
||||
from turnstone.server import _make_watch_dispatch
|
||||
import pytest
|
||||
|
||||
from tests._helpers import patch_session_storage
|
||||
from turnstone.core.session import _WATCH_QUEUE_SOFT_CAP, ChatSession
|
||||
|
||||
class _StubSession:
|
||||
"""Minimal ChatSession stand-in with controllable send() behaviour."""
|
||||
|
||||
def __init__(self, *, side_effect=None):
|
||||
self._watch_pending: queue.Queue = queue.Queue(maxsize=20)
|
||||
self._side_effect = side_effect
|
||||
class _NullUI:
|
||||
"""UI adapter that discards all output — local to this test module
|
||||
to avoid a cross-test-file import (mirrors the pattern in
|
||||
test_session.py / test_rewind_retry.py).
|
||||
"""
|
||||
|
||||
def send(self, msg: str) -> None:
|
||||
if self._side_effect is not None:
|
||||
raise self._side_effect
|
||||
|
||||
|
||||
class _RecordingUI:
|
||||
"""Track calls made by the dispatch error handlers."""
|
||||
|
||||
def __init__(self):
|
||||
self.errors: list[str] = []
|
||||
self.state_changes: list[str] = []
|
||||
self.stream_end_calls: int = 0
|
||||
|
||||
# -- SessionUI protocol stubs used by the dispatch code --
|
||||
|
||||
def on_error(self, message: str) -> None:
|
||||
self.errors.append(message)
|
||||
|
||||
def on_state_change(self, state: str) -> None:
|
||||
self.state_changes.append(state)
|
||||
|
||||
def on_stream_end(self) -> None:
|
||||
self.stream_end_calls += 1
|
||||
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _wait_for_worker(ws: Workstream, timeout: float = 2.0) -> None:
|
||||
"""Block until the worker thread started by dispatch() finishes."""
|
||||
t = ws.worker_thread
|
||||
if t is not None:
|
||||
t.join(timeout)
|
||||
assert not t.is_alive(), "worker thread did not finish in time"
|
||||
|
||||
|
||||
# ── GenerationCancelled path ────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_cancelled_emits_stream_end_and_idle():
|
||||
session = _StubSession(side_effect=GenerationCancelled())
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("hello")
|
||||
_wait_for_worker(ws)
|
||||
|
||||
assert ui.stream_end_calls == 1
|
||||
assert ui.state_changes == ["idle"]
|
||||
assert ui.errors == []
|
||||
|
||||
|
||||
# ── Generic exception path ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_exception_emits_stream_end_and_error():
|
||||
session = _StubSession(side_effect=RuntimeError("boom"))
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("hello")
|
||||
_wait_for_worker(ws)
|
||||
|
||||
assert ui.stream_end_calls == 1
|
||||
assert ui.state_changes == ["error"]
|
||||
assert len(ui.errors) == 1
|
||||
assert "boom" in ui.errors[0]
|
||||
|
||||
|
||||
# ── Worker-thread identity guard ────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_abandoned_thread_emits_no_events():
|
||||
"""After force-cancel sets worker_thread=None, the old thread must not
|
||||
emit stream_end or state changes."""
|
||||
barrier = threading.Event()
|
||||
|
||||
class _BlockingSession(_StubSession):
|
||||
def send(self, msg: str) -> None:
|
||||
barrier.wait(timeout=5)
|
||||
raise RuntimeError("late error")
|
||||
|
||||
session = _BlockingSession()
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("hello")
|
||||
|
||||
# Simulate force-cancel: clear the worker_thread reference.
|
||||
ws.worker_thread = None
|
||||
barrier.set()
|
||||
|
||||
# Wait for the thread to actually complete (it's still running).
|
||||
time.sleep(0.3)
|
||||
|
||||
assert ui.stream_end_calls == 0
|
||||
assert ui.state_changes == []
|
||||
assert ui.errors == []
|
||||
|
||||
|
||||
# ── Path A: busy workstream enqueue ─────────────────────────────────────────
|
||||
|
||||
|
||||
def test_busy_workstream_enqueues_message():
|
||||
"""When the workstream already has a live worker, dispatch enqueues."""
|
||||
session = _StubSession()
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
# Simulate a live worker — session_worker.send gates on
|
||||
# ``_worker_running``, not ``Thread.is_alive``.
|
||||
ws._worker_running = True
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("queued msg")
|
||||
|
||||
item = session._watch_pending.get_nowait()
|
||||
assert item == {"message": "queued msg"}
|
||||
|
||||
|
||||
def test_busy_workstream_drops_on_full_queue():
|
||||
"""When the pending queue is full, dispatch drops the message."""
|
||||
session = _StubSession()
|
||||
# Fill the queue to capacity.
|
||||
for i in range(20):
|
||||
session._watch_pending.put_nowait({"message": f"msg{i}"})
|
||||
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
ws._worker_running = True # simulate a live worker
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
# Should not block or raise — just log a warning and drop.
|
||||
dispatch("overflow msg")
|
||||
|
||||
assert session._watch_pending.full()
|
||||
|
||||
|
||||
# ── Lock guard ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_dispatch_holds_lock_during_thread_start():
|
||||
"""Dispatch acquires ws._lock before checking/starting the worker."""
|
||||
session = _StubSession()
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
acquire_count = 0
|
||||
inner = ws._lock
|
||||
|
||||
class _CountingLock:
|
||||
def __enter__(self):
|
||||
nonlocal acquire_count
|
||||
acquire_count += 1
|
||||
return inner.__enter__()
|
||||
|
||||
def __exit__(self, *args):
|
||||
return inner.__exit__(*args)
|
||||
|
||||
ws._lock = _CountingLock() # type: ignore[assignment]
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("hello")
|
||||
_wait_for_worker(ws)
|
||||
|
||||
assert acquire_count >= 1
|
||||
|
||||
|
||||
# ── Happy path ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_successful_send_no_error_events():
|
||||
"""Normal send() completion should not trigger error/cancel events."""
|
||||
session = _StubSession() # send() does nothing (success)
|
||||
ws = Workstream()
|
||||
ui = _RecordingUI()
|
||||
|
||||
dispatch = _make_watch_dispatch(ws, session, ui)
|
||||
dispatch("hello")
|
||||
_wait_for_worker(ws)
|
||||
|
||||
assert ui.stream_end_calls == 0
|
||||
assert ui.state_changes == []
|
||||
assert ui.errors == []
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
# Catch-all: any UI hook the chat loop calls becomes a no-op.
|
||||
return MagicMock()
|
||||
|
||||
|
||||
def _make_session_for_dispatch(**kwargs: Any) -> ChatSession:
|
||||
"""ChatSession built with the same minimal harness used elsewhere
|
||||
in the test suite, scoped down to what the dispatch closure needs.
|
||||
"""
|
||||
client = MagicMock()
|
||||
defaults = dict(
|
||||
client=client,
|
||||
model="test-model",
|
||||
ui=_NullUI(),
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
tool_timeout=30,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return ChatSession(**defaults)
|
||||
|
||||
|
||||
def _register_runner(session: ChatSession) -> tuple[Any, Any]:
|
||||
"""Attach a minimal stub ``WatchRunner`` to *session* and return the
|
||||
``(runner, dispatch_fn)`` pair captured by ``set_dispatch_fn``.
|
||||
"""
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
class _StubRunner:
|
||||
def set_dispatch_fn(self, ws_id: str, fn: Any) -> None:
|
||||
captured["fn"] = fn
|
||||
|
||||
runner = _StubRunner()
|
||||
session.set_watch_runner(runner)
|
||||
return runner, captured["fn"]
|
||||
|
||||
|
||||
def _reminder(text: str, **extra: Any) -> dict[str, Any]:
|
||||
"""Build a structured ``watch_triggered`` reminder dict for tests.
|
||||
|
||||
Mirrors the shape produced by :func:`turnstone.core.watch.build_watch_reminder`
|
||||
— ``text`` is the formatted body, optional fields ride alongside.
|
||||
Tests that don't care about the optional fields can call with
|
||||
``text`` only.
|
||||
"""
|
||||
out: dict[str, Any] = {"type": "watch_triggered", "text": text}
|
||||
out.update(extra)
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enqueue shape
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEnqueueShape:
|
||||
"""``set_watch_runner``'s closure produces a single
|
||||
``("watch_triggered", text, "any")`` entry per fire.
|
||||
"""
|
||||
|
||||
def test_dispatch_enqueues_watch_triggered_with_any_channel(self, tmp_db):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
dispatch(_reminder("watch fired body"), "watch-1")
|
||||
|
||||
# One entry, "watch_triggered" type, on "any" channel.
|
||||
assert len(session._nudge_queue) == 1
|
||||
assert session._nudge_queue.pending(channel="any") == [
|
||||
("watch_triggered", "watch fired body")
|
||||
]
|
||||
# NOT on "user" or "tool" channels.
|
||||
assert session._nudge_queue.pending(channel="user") == []
|
||||
assert session._nudge_queue.pending(channel="tool") == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sanitisation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSanitisation:
|
||||
"""``sanitize_payload`` runs producer-side over the formatted message
|
||||
before it ever reaches the queue. The wire-boundary
|
||||
``escape_wrapper_tags`` only protects ``<system-reminder>`` /
|
||||
``<tool_output>`` envelopes; this layer covers everything else.
|
||||
"""
|
||||
|
||||
def test_dispatch_sanitizes_payload_before_enqueue(self, tmp_db):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# Build a payload with: BEL (\x07), zero-width space (U+200B),
|
||||
# bidi RTL override (U+202E), and angle-bracket tag breakers.
|
||||
raw = "before\x07middleaftermore<thinking>tail"
|
||||
dispatch(_reminder(raw), "watch-1")
|
||||
|
||||
pending = session._nudge_queue.pending(channel="any")
|
||||
assert len(pending) == 1
|
||||
sanitized = pending[0][1]
|
||||
# Control / steering chars become spaces; angle brackets vanish.
|
||||
assert "\x07" not in sanitized
|
||||
assert "" not in sanitized
|
||||
assert "" not in sanitized
|
||||
assert "<" not in sanitized
|
||||
assert ">" not in sanitized
|
||||
# Real content survives.
|
||||
assert "before" in sanitized
|
||||
assert "thinking" in sanitized
|
||||
|
||||
def test_dispatch_preserves_newlines_for_multiline_output(self, tmp_db):
|
||||
"""Multi-line shell output must keep its layout — TAB / LF / CR
|
||||
are intentionally preserved by ``sanitize_payload`` (R8).
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
dispatch(_reminder("line1\nline2\n\tindented\n\rline3"), "watch-1")
|
||||
|
||||
pending = session._nudge_queue.pending(channel="any")
|
||||
assert len(pending) == 1
|
||||
text = pending[0][1]
|
||||
# Lines stay separated; tab kept.
|
||||
assert "\n" in text
|
||||
assert "\t" in text
|
||||
|
||||
def test_dispatch_drops_empty_after_sanitization(self, tmp_db):
|
||||
"""A payload that's all control chars sanitises to "" — no enqueue."""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# All-control + DEL + zero-width — strips to empty.
|
||||
dispatch(_reminder("\x07\x0b\x7f"), "watch-1")
|
||||
|
||||
assert len(session._nudge_queue) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Soft cap
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSoftCap:
|
||||
"""When ``"watch_triggered"`` saturates at :data:`_WATCH_QUEUE_SOFT_CAP`,
|
||||
the closure drops the OLDEST entry of that type and enqueues the new
|
||||
one — so the queue stays ≤ cap with the most recent watch outputs.
|
||||
"""
|
||||
|
||||
def test_dispatch_drop_oldest_at_soft_cap(self, tmp_db, caplog):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# Pre-fill at the cap. Each entry has a unique body so we can
|
||||
# tell which one(s) survived a drop.
|
||||
for i in range(_WATCH_QUEUE_SOFT_CAP):
|
||||
dispatch(_reminder(f"body-{i}"), "watch-1")
|
||||
assert len(session._nudge_queue) == _WATCH_QUEUE_SOFT_CAP
|
||||
|
||||
with caplog.at_level("WARNING"):
|
||||
dispatch(_reminder("overflow"), "watch-1")
|
||||
|
||||
# Total stays at cap (one dropped, one added).
|
||||
assert len(session._nudge_queue) == _WATCH_QUEUE_SOFT_CAP
|
||||
bodies = [text for _t, text in session._nudge_queue.pending(channel="any")]
|
||||
# Oldest ("body-0") gone; newest ("overflow") present.
|
||||
assert "body-0" not in bodies
|
||||
assert "overflow" in bodies
|
||||
# Warning logged.
|
||||
assert any("watch_dispatch.queue_full" in r.message for r in caplog.records), (
|
||||
"expected a watch_dispatch.queue_full warning record"
|
||||
)
|
||||
|
||||
def test_dispatch_soft_cap_does_not_evict_other_types(self, tmp_db):
|
||||
"""A watch saturation drop must only target watch-typed entries.
|
||||
Other producers (idle_children, advisories) have their own
|
||||
rate limiters and must not be collateral damage.
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# Mix in a few non-watch entries on the same queue.
|
||||
session._nudge_queue.enqueue("idle_children", "ic-1", "any")
|
||||
session._nudge_queue.enqueue("idle_children", "ic-2", "any")
|
||||
|
||||
# Saturate watches up to cap (queue holds cap+2 total).
|
||||
for i in range(_WATCH_QUEUE_SOFT_CAP):
|
||||
dispatch(_reminder(f"body-{i}"), "watch-1")
|
||||
# One more triggers drop-oldest of a "watch_triggered" entry.
|
||||
dispatch(_reminder("overflow"), "watch-1")
|
||||
|
||||
# Both idle_children entries survived — no collateral eviction.
|
||||
idle_bodies = [
|
||||
text
|
||||
for nt, text in session._nudge_queue.pending(channel="any")
|
||||
if nt == "idle_children"
|
||||
]
|
||||
assert idle_bodies == ["ic-1", "ic-2"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# valid_until predicate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestValidUntil:
|
||||
"""The ``valid_until`` predicate captured at dispatch time re-checks
|
||||
the watch's ``active`` flag at drain time, so a cancelled watch's
|
||||
last splat doesn't ride out a future wake.
|
||||
"""
|
||||
|
||||
def test_valid_until_drops_when_watch_inactive(self, tmp_db, monkeypatch):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# Storage stub returns False at drain time.
|
||||
is_active_calls = patch_session_storage(monkeypatch, active=False)
|
||||
|
||||
dispatch(_reminder("body"), "watch-1")
|
||||
# Drain fires the predicate; entry should NOT be delivered.
|
||||
out = session._nudge_queue.drain({"any"})
|
||||
assert out == []
|
||||
# Predicate ran once with the dispatched watch_id.
|
||||
assert is_active_calls == ["watch-1"]
|
||||
|
||||
def test_valid_until_drops_when_storage_raises(self, tmp_db, monkeypatch):
|
||||
"""The closure's broad-except in the predicate translates a
|
||||
storage-layer exception to ``False`` so the drain doesn't
|
||||
propagate; the predicate captured ``watch_id`` correctly
|
||||
(otherwise storage wouldn't even be touched).
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
patch_session_storage(monkeypatch, raise_on_is_active=True)
|
||||
|
||||
dispatch(_reminder("body"), "watch-bound-id")
|
||||
out = session._nudge_queue.drain({"any"})
|
||||
assert out == []
|
||||
|
||||
def test_valid_until_delivers_when_watch_active(self, tmp_db, monkeypatch):
|
||||
"""Happy-path counter-test for the predicate above: the entry
|
||||
DOES drain when the watch is still active.
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
patch_session_storage(monkeypatch, active=True)
|
||||
|
||||
dispatch(_reminder("body"), "watch-1")
|
||||
out = session._nudge_queue.drain({"any"})
|
||||
assert len(out) == 1
|
||||
assert out[0][0] == "watch_triggered"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Concurrency
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConcurrency:
|
||||
"""Two threads each fire 100 dispatches against the same session;
|
||||
the soft-cap read-then-mutate window stays bounded and the queue
|
||||
settles in a consistent state.
|
||||
|
||||
Per the plan's risk register R2: in production only one daemon
|
||||
thread (``WatchRunner``'s ``_run``) ever calls a session's dispatch
|
||||
fn, so the 3-acquisition non-atomicity is harmless. This test
|
||||
pins lock-correctness anyway against the broader race window.
|
||||
"""
|
||||
|
||||
def test_dispatch_concurrent_enqueues_thread_safe(self, tmp_db, monkeypatch):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
# Bypass the storage-touching valid_until predicate: count cap
|
||||
# behaviour, not storage round-trips.
|
||||
patch_session_storage(monkeypatch, active=True)
|
||||
|
||||
per_thread = 100
|
||||
labels = ("a", "b")
|
||||
|
||||
def fire(label: str) -> None:
|
||||
for i in range(per_thread):
|
||||
dispatch(_reminder(f"{label}-{i}"), f"watch-{label}")
|
||||
|
||||
threads = [threading.Thread(target=fire, args=(label,), daemon=True) for label in labels]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=5.0)
|
||||
for t in threads:
|
||||
assert not t.is_alive(), "dispatch thread did not finish in time"
|
||||
|
||||
# The non-atomic count-then-drop window admits at most one "slip"
|
||||
# per concurrent thread above the cap (each thread can observe a
|
||||
# sub-cap count and append before another thread's drop runs).
|
||||
depth = len(session._nudge_queue)
|
||||
assert depth <= len(threads) * per_thread
|
||||
assert depth <= _WATCH_QUEUE_SOFT_CAP + len(threads)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Empty-input / multi-call invariants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("payload", ["", " ", "\x07\x0b"])
|
||||
def test_dispatch_no_op_for_empty_payloads(tmp_db, payload: str):
|
||||
"""Whitespace-only / pure-control payloads sanitise to empty and
|
||||
do not produce a queue entry — silent drop.
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
dispatch(_reminder(payload), "watch-1")
|
||||
|
||||
assert len(session._nudge_queue) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Metadata propagation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMetadataPropagation:
|
||||
"""The dispatch closure pulls optional fields out of the structured
|
||||
``reminder`` dict and attaches them to the queue entry's
|
||||
``metadata``. Drain seams later merge ``metadata`` into the
|
||||
rendered reminder dict so the frontend can display a structured
|
||||
``.msg.watch-result`` card.
|
||||
"""
|
||||
|
||||
def test_dispatch_attaches_watch_metadata_on_enqueue(self, tmp_db):
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
reminder = _reminder(
|
||||
"$ ls\nfile.txt",
|
||||
watch_name="my-watch",
|
||||
command="ls",
|
||||
poll_count=2,
|
||||
max_polls=100,
|
||||
is_final=False,
|
||||
)
|
||||
dispatch(reminder, "watch-1")
|
||||
|
||||
# Snapshot via ``pending_with_metadata`` to inspect the full
|
||||
# entry shape. Exactly one entry, with the optional fields
|
||||
# carried verbatim onto ``metadata``.
|
||||
snapshot = session._nudge_queue.pending_with_metadata(channel="any")
|
||||
assert len(snapshot) == 1
|
||||
nt, _text, meta = snapshot[0]
|
||||
assert nt == "watch_triggered"
|
||||
assert meta == {
|
||||
"watch_name": "my-watch",
|
||||
"command": "ls",
|
||||
"poll_count": 2,
|
||||
"max_polls": 100,
|
||||
"is_final": False,
|
||||
}
|
||||
|
||||
def test_dispatch_omits_metadata_when_optional_fields_missing(self, tmp_db):
|
||||
"""A bare ``{type, text}`` reminder produces an entry with no
|
||||
metadata — the closure builds an empty dict, sees nothing to
|
||||
carry, and falls through to ``metadata=None``.
|
||||
"""
|
||||
session = _make_session_for_dispatch()
|
||||
_runner, dispatch = _register_runner(session)
|
||||
|
||||
dispatch(_reminder("just a body"), "watch-1")
|
||||
|
||||
snapshot = session._nudge_queue.pending_with_metadata(channel="any")
|
||||
assert len(snapshot) == 1
|
||||
_nt, _text, meta = snapshot[0]
|
||||
assert meta is None
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
"""Boundary-crossing integration test for the watch switchover pipeline.
|
||||
|
||||
Drives a real :class:`ChatSession` + a real :class:`WatchRunner` (with
|
||||
its daemon thread skipped — we call ``_dispatch_result`` directly to
|
||||
avoid the timer dependency) end-to-end through the chat-loop drain
|
||||
seam. The only stub is the LLM provider (patched
|
||||
``_create_stream_with_retry``); every other layer is production code:
|
||||
|
||||
* ``WatchRunner._dispatch_result`` releasing the dispatch lock before
|
||||
fan-out
|
||||
* the closure built inside ``ChatSession.set_watch_runner`` —
|
||||
``sanitize_payload`` + soft-cap check + ``valid_until`` predicate +
|
||||
``NudgeQueue.enqueue("watch_triggered", ..., "any", ...)``
|
||||
* ``ChatSession.send`` chat loop short-circuiting metacog detection
|
||||
* ``_attach_pending_user_reminders`` draining ``USER_DRAIN`` (which
|
||||
matches ``"any"``)
|
||||
* ``_apply_reminders_for_provider`` splicing the rendered envelope onto
|
||||
the user message before the wire boundary
|
||||
|
||||
Per ``feedback_tests_through_boundaries.md``: direct injection tests
|
||||
that bypass these boundaries silently mask wiring bugs. This test is
|
||||
the structural integration gate for the watch switchover.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from tests._helpers import patch_session_storage
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.watch import WatchRunner
|
||||
|
||||
|
||||
class _NullUI:
|
||||
"""UI adapter that no-ops every chat-loop hook the test triggers."""
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return MagicMock()
|
||||
|
||||
|
||||
def _make_session() -> ChatSession:
|
||||
"""Real ChatSession with the same minimal setup the unit-test suite
|
||||
uses; no LLM calls happen until a chat-loop method is exercised
|
||||
(and even then the LLM provider is patched).
|
||||
"""
|
||||
return ChatSession(
|
||||
client=MagicMock(),
|
||||
model="test-model",
|
||||
ui=_NullUI(),
|
||||
instructions=None,
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
tool_timeout=30,
|
||||
)
|
||||
|
||||
|
||||
def test_watch_fires_then_user_send_drains_envelope(tmp_db, monkeypatch):
|
||||
"""Pin the cross-PR concern that watch text reaches the model via
|
||||
the unified ``<system-reminder>`` envelope path:
|
||||
|
||||
1. WatchRunner.dispatch fires watch text against the session's
|
||||
registered closure (synchronously — no daemon thread).
|
||||
2. NudgeQueue holds one ``"watch_triggered"`` entry on ``"any"``.
|
||||
3. session.send("ok") runs the chat loop with a stubbed LLM.
|
||||
4. The drain seam drains the watch entry; the wire payload's user
|
||||
message has the watch text spliced into a ``<system-reminder>``
|
||||
envelope.
|
||||
"""
|
||||
session = _make_session()
|
||||
|
||||
# Bypass the storage-touching predicate — we want to assert the
|
||||
# envelope splice, not exercise a fresh sqlite watch row.
|
||||
patch_session_storage(monkeypatch, active=True)
|
||||
|
||||
# Real WatchRunner; we don't ``start()`` the daemon thread (that
|
||||
# would race with the test's deterministic order). Direct call
|
||||
# to ``_dispatch_result`` exercises the same dispatch path the
|
||||
# daemon would invoke. Runner-side ``storage`` is unused on this
|
||||
# path (only the polling loop touches it); a MagicMock placeholder
|
||||
# keeps the constructor signature happy.
|
||||
runner = WatchRunner(storage=MagicMock(), node_id="test-node")
|
||||
session.set_watch_runner(runner)
|
||||
|
||||
# 1. Fire a watch result synchronously.
|
||||
runner._dispatch_result(
|
||||
session._ws_id,
|
||||
{"type": "watch_triggered", "text": "watch payload body"},
|
||||
"watch-1",
|
||||
)
|
||||
|
||||
# 2. The queue holds one entry on the "any" channel.
|
||||
assert len(session._nudge_queue) == 1
|
||||
pending = session._nudge_queue.pending(channel="any")
|
||||
assert pending == [("watch_triggered", "watch payload body")]
|
||||
|
||||
# 3. Run the chat loop with the LLM patched. We don't care about
|
||||
# the assistant turn's content; only the wire payload sent to the
|
||||
# provider matters.
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=iter([])),
|
||||
patch.object(
|
||||
session,
|
||||
"_stream_response",
|
||||
return_value={"role": "assistant", "content": "ok"},
|
||||
),
|
||||
patch.object(session, "_update_token_table"),
|
||||
patch.object(session, "_print_status_line"),
|
||||
patch.object(session, "_visible_memory_count", return_value=0),
|
||||
patch("turnstone.core.session.save_message"),
|
||||
):
|
||||
session._title_generated = True # suppress orthogonal title side-thread
|
||||
session.send("ok")
|
||||
|
||||
# 4. Queue fully drained by the user-message attach seam.
|
||||
assert len(session._nudge_queue) == 0
|
||||
|
||||
# The user message that drove the assistant turn has the watch
|
||||
# text in its ``_reminders`` side-channel — the production
|
||||
# ``_apply_reminders_for_provider`` splice consumes that to wrap
|
||||
# the content in ``<system-reminder>`` at the wire boundary.
|
||||
user_msgs = [m for m in session.messages if m.get("role") == "user"]
|
||||
assert user_msgs, "expected a user message in history"
|
||||
last_user = user_msgs[-1]
|
||||
reminders = last_user.get("_reminders") or []
|
||||
assert any(
|
||||
r.get("type") == "watch_triggered" and "watch payload body" in r.get("text", "")
|
||||
for r in reminders
|
||||
), f"expected watch_triggered reminder on user message; got {reminders!r}"
|
||||
|
||||
|
||||
def test_three_back_to_back_watch_fires_drain_into_one_turn(tmp_db, monkeypatch):
|
||||
"""Behavioural delta from plan section 3.4 / risk register R3.
|
||||
|
||||
N back-to-back watch fires used to produce N successive
|
||||
``send()`` turns (each a separate model invocation, capped at
|
||||
``_MAX_WATCH_CHAIN = 5``). After the switchover, the N entries
|
||||
drain into ONE envelope splice on the next drain seam — one
|
||||
assistant turn responding to all N watch results. Pinning this
|
||||
behavioural delta protects against accidental regression to
|
||||
the old per-fire-turn shape.
|
||||
"""
|
||||
session = _make_session()
|
||||
|
||||
patch_session_storage(monkeypatch, active=True)
|
||||
|
||||
runner = WatchRunner(storage=MagicMock(), node_id="test-node")
|
||||
session.set_watch_runner(runner)
|
||||
|
||||
runner._dispatch_result(
|
||||
session._ws_id, {"type": "watch_triggered", "text": "fire one"}, "watch-1"
|
||||
)
|
||||
runner._dispatch_result(
|
||||
session._ws_id, {"type": "watch_triggered", "text": "fire two"}, "watch-1"
|
||||
)
|
||||
runner._dispatch_result(
|
||||
session._ws_id, {"type": "watch_triggered", "text": "fire three"}, "watch-1"
|
||||
)
|
||||
assert len(session._nudge_queue) == 3
|
||||
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=iter([])),
|
||||
patch.object(
|
||||
session,
|
||||
"_stream_response",
|
||||
return_value={"role": "assistant", "content": "got it"},
|
||||
),
|
||||
patch.object(session, "_update_token_table"),
|
||||
patch.object(session, "_print_status_line"),
|
||||
patch.object(session, "_visible_memory_count", return_value=0),
|
||||
patch("turnstone.core.session.save_message"),
|
||||
):
|
||||
session._title_generated = True
|
||||
session.send("user")
|
||||
|
||||
# All three drained into the single user-message attach.
|
||||
user_msgs = [m for m in session.messages if m.get("role") == "user"]
|
||||
last_user = user_msgs[-1]
|
||||
reminders = last_user.get("_reminders") or []
|
||||
watch_reminders = [r for r in reminders if r.get("type") == "watch_triggered"]
|
||||
assert len(watch_reminders) == 3
|
||||
bodies = [r.get("text", "") for r in watch_reminders]
|
||||
assert any("fire one" in b for b in bodies)
|
||||
assert any("fire two" in b for b in bodies)
|
||||
assert any("fire three" in b for b in bodies)
|
||||
# And there's exactly ONE assistant turn (not three).
|
||||
assistant_turns = [m for m in session.messages if m.get("role") == "assistant"]
|
||||
assert len(assistant_turns) == 1
|
||||
|
||||
|
||||
def test_watch_dispatch_through_restore_fn_lands_on_rehydrated_session(tmp_db, monkeypatch):
|
||||
"""Cover the production ``_watch_restore_fn`` closure surface.
|
||||
|
||||
Path under test:
|
||||
WatchRunner._dispatch_result(ws_id, msg, watch_id)
|
||||
no dispatch fn registered (original session evicted)
|
||||
restore_fn(ws_id) constructs a fresh ChatSession,
|
||||
calls session.resume(ws_id) to adopt the original ws_id,
|
||||
re-registers the dispatch closure via session.set_watch_runner,
|
||||
returns runner.get_dispatch_fn(session._ws_id)
|
||||
runner invokes the returned fn with (msg, watch_id)
|
||||
watch payload lands on the rehydrated session's NudgeQueue
|
||||
|
||||
Construction inside ``server.py``'s ``_watch_restore_fn`` is the new
|
||||
contract surface introduced by the switchover; this test pins that
|
||||
contract so a future refactor of the closure (e.g. swapping
|
||||
``manager.create + session.resume`` for ``manager.open``) doesn't
|
||||
silently break the watch-restore pipeline.
|
||||
"""
|
||||
from turnstone.core import session as session_mod
|
||||
|
||||
patch_session_storage(monkeypatch, active=True)
|
||||
|
||||
# Stage 1 — build the original session and persist a message so
|
||||
# ``session.resume`` finds the ws_id in storage.
|
||||
original = _make_session()
|
||||
original_ws_id = original._ws_id
|
||||
# Persist a stub user message so ``load_messages(original_ws_id)``
|
||||
# returns something non-empty (resume short-circuits on empty).
|
||||
session_mod.save_message(original_ws_id, "user", "kickoff message")
|
||||
|
||||
# Stage 2 — runner with NO dispatch fn registered (simulates the
|
||||
# original session being evicted between watch fire and dispatch).
|
||||
# The restore_fn captures *which* fresh ChatSession got built so the
|
||||
# test can assert the queue landed on it (not on the original).
|
||||
rehydrated_holder: dict[str, ChatSession] = {}
|
||||
|
||||
def _restore_fn(ws_id: str) -> Any:
|
||||
"""Mirror the production ``_watch_restore_fn`` closure shape:
|
||||
construct a fresh session, resume the persisted ws_id (so the
|
||||
new session adopts the original ws_id), wire the dispatch
|
||||
closure, return the dispatch fn.
|
||||
"""
|
||||
new_session = _make_session()
|
||||
ok = new_session.resume(ws_id)
|
||||
assert ok, "resume should succeed against a non-empty message log"
|
||||
new_session.set_watch_runner(runner)
|
||||
rehydrated_holder["session"] = new_session
|
||||
return runner.get_dispatch_fn(new_session._ws_id)
|
||||
|
||||
runner = WatchRunner(
|
||||
storage=MagicMock(),
|
||||
node_id="test-node",
|
||||
restore_fn=_restore_fn,
|
||||
)
|
||||
|
||||
# Sanity: no dispatch fn registered yet for the original ws_id.
|
||||
assert runner.get_dispatch_fn(original_ws_id) is None
|
||||
|
||||
# Stage 3 — fire a watch result. ``_dispatch_result`` should fall
|
||||
# through to the restore branch. The dispatch surface takes a
|
||||
# structured reminder dict.
|
||||
runner._dispatch_result(
|
||||
original_ws_id,
|
||||
{"type": "watch_triggered", "text": "post-restore body"},
|
||||
"watch-1",
|
||||
)
|
||||
|
||||
# The restore fn ran exactly once and produced a fresh session that
|
||||
# adopted the original ws_id.
|
||||
assert "session" in rehydrated_holder, "restore_fn was not invoked"
|
||||
rehydrated = rehydrated_holder["session"]
|
||||
assert rehydrated is not original
|
||||
assert rehydrated._ws_id == original_ws_id
|
||||
|
||||
# The watch payload landed on the rehydrated session's queue, not on
|
||||
# the (now-evicted) original session's queue.
|
||||
assert len(rehydrated._nudge_queue) == 1
|
||||
assert rehydrated._nudge_queue.pending(channel="any") == [
|
||||
("watch_triggered", "post-restore body")
|
||||
]
|
||||
# Original session's queue stays empty — the dispatch did NOT
|
||||
# accidentally route back to it.
|
||||
assert len(original._nudge_queue) == 0
|
||||
@@ -72,6 +72,20 @@ class TestWatchCRUD:
|
||||
assert db.delete_watch("nope") is False
|
||||
|
||||
|
||||
class TestIsWatchActive:
|
||||
def test_active_row_returns_true(self, db):
|
||||
db.create_watch(**_make_watch_kwargs())
|
||||
assert db.is_watch_active("watch_001") is True
|
||||
|
||||
def test_inactive_row_returns_false(self, db):
|
||||
db.create_watch(**_make_watch_kwargs())
|
||||
db.update_watch("watch_001", active=False)
|
||||
assert db.is_watch_active("watch_001") is False
|
||||
|
||||
def test_missing_row_returns_false(self, db):
|
||||
assert db.is_watch_active("nope") is False
|
||||
|
||||
|
||||
class TestWatchListQueries:
|
||||
def test_list_for_ws(self, db):
|
||||
db.create_watch(**_make_watch_kwargs(watch_id="w1", ws_id="ws-1", name="a"))
|
||||
|
||||
@@ -173,3 +173,50 @@ class TestResolveClient:
|
||||
def test_unknown_backend_returns_none(self):
|
||||
client = resolve_web_search_client("typo_backend", tavily_key="key")
|
||||
assert client is None
|
||||
|
||||
def test_resolve_web_search_client_rejects_oauth_user_backend(self):
|
||||
"""A web_search backend pointing at an ``auth_type=oauth_user``
|
||||
MCP server MUST be rejected at boot — per-node web_search
|
||||
cannot carry per-user tokens, so resolving the backend would
|
||||
guarantee a 401-on-call instead of a clean disablement.
|
||||
|
||||
Phase 7 invariant 8 corollary: pool tools are user-scoped;
|
||||
every entry point that lacks per-user identity (web_search
|
||||
boot resolver, eval harness, CLI default) MUST refuse them
|
||||
rather than silently produce a broken client.
|
||||
|
||||
Verified by reverting the ``server_auth_type(...) == 'oauth_user'``
|
||||
guard in ``resolve_web_search_client``: the resolver returns
|
||||
an ``MCPSearchClient`` whose ``call_tool_sync`` would surface
|
||||
a 401 / consent_required structured error on every search.
|
||||
"""
|
||||
mcp = MagicMock()
|
||||
mcp.is_mcp_tool.return_value = True # name resolves
|
||||
mcp.server_auth_type.return_value = "oauth_user"
|
||||
client = resolve_web_search_client(
|
||||
"mcp:oauth-search:search", tavily_key=None, mcp_client=mcp
|
||||
)
|
||||
assert client is None, (
|
||||
"oauth_user-backed web_search backend resolved to a non-None client; "
|
||||
"boot-time guard missing or regressed."
|
||||
)
|
||||
# Per-turn callers must read from the in-memory cache, never
|
||||
# the SQL helper — perf regression guard.
|
||||
mcp.server_auth_type.assert_called_with("oauth-search")
|
||||
assert not mcp._lookup_server_row.called, (
|
||||
"resolver issued a SQL roundtrip via _lookup_server_row; "
|
||||
"per-turn web_search backend resolution must use the "
|
||||
"in-memory server_auth_type accessor."
|
||||
)
|
||||
|
||||
def test_resolve_web_search_client_accepts_static_backend(self):
|
||||
"""Static-path (``auth_type=none`` or ``static``) MCP backends
|
||||
still resolve cleanly — the new guard ONLY rejects oauth_user.
|
||||
"""
|
||||
mcp = MagicMock()
|
||||
mcp.is_mcp_tool.return_value = True
|
||||
mcp.server_auth_type.return_value = None
|
||||
client = resolve_web_search_client(
|
||||
"mcp:static-search:search", tavily_key=None, mcp_client=mcp
|
||||
)
|
||||
assert isinstance(client, MCPSearchClient)
|
||||
|
||||
@@ -898,6 +898,97 @@ class TestHistoryInteractive:
|
||||
# Above-cap → clamps to 500 (response is still 200; we have 4 rows).
|
||||
assert client.get(base, params={"limit": 999}).status_code == 200
|
||||
|
||||
def test_returns_partial_trailing_turn_during_tool_execution(self, _inject_storage):
|
||||
"""The ``/history`` REST endpoint is a *display* read and must
|
||||
surface partial state. When the operator refreshes the page
|
||||
mid-tool-execution — assistant ``tool_calls`` saved, only some
|
||||
results saved — the trailing turn must come back on the wire so
|
||||
the UI can render what the operator was watching live.
|
||||
|
||||
Storage's ``load_messages`` defaults to a repair pass that
|
||||
strips this exact shape (correct for ``session.resume``, wrong
|
||||
for display). ``make_history_handler`` must opt out via
|
||||
``repair=False``; flipping that flag back on breaks this test.
|
||||
"""
|
||||
import json
|
||||
|
||||
ws_id = "ws-mid-exec"
|
||||
_inject_storage.register_workstream(ws_id, kind="interactive", user_id="test-user")
|
||||
_inject_storage.save_message(ws_id, "user", "kick off")
|
||||
tc_json = json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
||||
},
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
||||
},
|
||||
]
|
||||
)
|
||||
_inject_storage.save_message(ws_id, "assistant", "Working", tool_calls=tc_json)
|
||||
_inject_storage.save_message(ws_id, "tool", "file.txt", tool_call_id="call_1")
|
||||
# call_2 result not yet persisted — operator refreshes here.
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = ws_id
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.get.return_value = mock_ws
|
||||
client = _build_history_app(mock_mgr, _inject_storage)
|
||||
|
||||
r = client.get(f"/v1/api/workstreams/{ws_id}/history")
|
||||
assert r.status_code == 200
|
||||
roles = [m.get("role") for m in r.json()["messages"]]
|
||||
# All three rows survive — the trailing assistant + partial
|
||||
# tool result are what the operator was watching live. The
|
||||
# default-repair shape would have been just ``["user"]``.
|
||||
assert roles == ["user", "assistant", "tool"]
|
||||
|
||||
# Confirm the default-repair path collapses this to just the
|
||||
# user message — locks in the regression contract.
|
||||
with_repair = _inject_storage.load_messages(ws_id, repair=True)
|
||||
assert [m.get("role") for m in with_repair] == ["user"]
|
||||
|
||||
def test_history_does_not_synthesize_orphan_results(self, _inject_storage):
|
||||
"""``repair=False`` via ``/history`` must NOT splice synthetic
|
||||
``"Tool execution was cancelled."`` rows for mid-conversation
|
||||
orphaned tool_calls — the operator never saw those rows, and
|
||||
showing them would invent UI content that doesn't reflect
|
||||
persisted state.
|
||||
"""
|
||||
import json
|
||||
|
||||
ws_id = "ws-orphan-mid"
|
||||
_inject_storage.register_workstream(ws_id, kind="interactive", user_id="test-user")
|
||||
_inject_storage.save_message(ws_id, "user", "first")
|
||||
tc_json = json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
||||
},
|
||||
]
|
||||
)
|
||||
_inject_storage.save_message(ws_id, "assistant", "Working", tool_calls=tc_json)
|
||||
# Cancel landed before any tool result — next turn happens.
|
||||
_inject_storage.save_message(ws_id, "user", "second")
|
||||
_inject_storage.save_message(ws_id, "assistant", "ok")
|
||||
mock_ws = MagicMock()
|
||||
mock_ws.id = ws_id
|
||||
mock_mgr = MagicMock()
|
||||
mock_mgr.get.return_value = mock_ws
|
||||
client = _build_history_app(mock_mgr, _inject_storage)
|
||||
|
||||
r = client.get(f"/v1/api/workstreams/{ws_id}/history")
|
||||
assert r.status_code == 200
|
||||
roles = [m.get("role") for m in r.json()["messages"]]
|
||||
# No synthetic tool row spliced after the orphaned tool_calls.
|
||||
assert roles == ["user", "assistant", "user", "assistant"]
|
||||
assert all(m.get("role") != "tool" for m in r.json()["messages"])
|
||||
|
||||
|
||||
class TestBuildHistoryReminderPropagation:
|
||||
"""``_build_history`` must surface the ``_reminders`` side-channel on
|
||||
@@ -1031,6 +1122,224 @@ class TestBuildHistoryReminderPropagation:
|
||||
assert history[0]["content"] == content
|
||||
|
||||
|
||||
class TestBuildHistoryAdvisoryRoundTrip:
|
||||
"""``_build_history`` must round-trip the persisted
|
||||
``<tool_output>`` envelope (Seam 1 queued-message splice) to
|
||||
cleaned content + a wire-shape ``advisories`` array.
|
||||
|
||||
Production realism note: ``session.messages`` never carries an
|
||||
``advisories`` key — only ``decorate_history_messages`` mutates
|
||||
dicts to add it for the REST ``/history`` path, and the SSE replay
|
||||
surface bypasses that decoration entirely. The earlier
|
||||
``TestBuildHistoryAdvisoryPropagation`` class pre-populated
|
||||
``advisories`` directly on the session messages, which tested a
|
||||
passthrough that doesn't exist in production — the SSE replay code
|
||||
path silently dropped queued messages despite the green tests.
|
||||
These round-trip tests exercise the production shape (wrapped
|
||||
envelope on the tool row's ``content``) so a regression in the
|
||||
inline ``extract_advisories_from_tool_envelope`` call inside
|
||||
``_build_history`` surfaces here.
|
||||
"""
|
||||
|
||||
def _session_with_messages(self, messages: list[dict]) -> MagicMock:
|
||||
session = MagicMock()
|
||||
session.messages = messages
|
||||
return session
|
||||
|
||||
def test_build_history_round_trips_envelope_to_advisories(self):
|
||||
"""The production-realistic shape: a tool row whose ``content``
|
||||
is the wrapped ``<tool_output>`` envelope (no ``advisories``
|
||||
key set — that's the bug-1 footprint). ``_build_history``
|
||||
must extract the advisory back out and ship it on the wire as
|
||||
cleaned content + ``advisories``.
|
||||
|
||||
Reverting the inline ``extract_advisories_from_tool_envelope``
|
||||
call in ``server._build_history``'s tool-message branch breaks
|
||||
this test.
|
||||
"""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
from turnstone.server import _build_history
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"tool body",
|
||||
[UserInterjection(message="check logs", priority="notice")],
|
||||
)
|
||||
session = self._session_with_messages(
|
||||
[
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_a",
|
||||
"content": wrapped,
|
||||
}
|
||||
]
|
||||
)
|
||||
history = _build_history(session)
|
||||
# Cleaned content rides on the wire — envelope stripped.
|
||||
assert history[0]["content"] == "tool body"
|
||||
# Advisory survives as a wire-shape entry the JS can render
|
||||
# as a user bubble after the tool block.
|
||||
assert history[0]["advisories"] == [
|
||||
{"type": "user_interjection", "text": "check logs", "priority": "notice"}
|
||||
]
|
||||
|
||||
def test_build_history_round_trips_important_priority(self):
|
||||
"""The ``important`` priority preamble round-trips — pin both
|
||||
the priority detection in the parser and the projection through
|
||||
to the wire shape."""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
from turnstone.server import _build_history
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"out",
|
||||
[UserInterjection(message="urgent", priority="important")],
|
||||
)
|
||||
session = self._session_with_messages(
|
||||
[{"role": "tool", "tool_call_id": "call_a", "content": wrapped}]
|
||||
)
|
||||
history = _build_history(session)
|
||||
assert history[0]["content"] == "out"
|
||||
assert history[0]["advisories"] == [
|
||||
{"type": "user_interjection", "text": "urgent", "priority": "important"}
|
||||
]
|
||||
|
||||
def test_build_history_no_envelope_passes_through_unchanged(self):
|
||||
"""Plain tool content (no ``<tool_output>`` prefix) — no
|
||||
advisories field, content unchanged."""
|
||||
from turnstone.server import _build_history
|
||||
|
||||
session = self._session_with_messages(
|
||||
[{"role": "tool", "tool_call_id": "call_a", "content": "plain output"}]
|
||||
)
|
||||
history = _build_history(session)
|
||||
assert history[0]["content"] == "plain output"
|
||||
assert "advisories" not in history[0]
|
||||
|
||||
def test_build_history_round_trip_through_full_decoration_chain(self):
|
||||
"""End-to-end pin: persist a wrapped envelope into ``messages``,
|
||||
run the full decoration chain (``decorate_history_messages``
|
||||
followed by ``_build_history``), assert the wire shape carries
|
||||
the advisory. This pins the contract every component in the
|
||||
chain participates in — REST ``/history`` callers go through
|
||||
``decorate_history_messages``, and SSE replay goes through
|
||||
``_build_history`` — both must produce the same wire shape.
|
||||
"""
|
||||
from turnstone.core.history_decoration import decorate_history_messages
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
from turnstone.server import _build_history
|
||||
|
||||
wrapped = wrap_tool_result(
|
||||
"raw",
|
||||
[UserInterjection(message="hi", priority="notice")],
|
||||
)
|
||||
# Decorate first — REST /history shape.
|
||||
rest_messages: list[dict] = [{"role": "tool", "tool_call_id": "call_a", "content": wrapped}]
|
||||
decorate_history_messages(rest_messages, {}, {})
|
||||
# And separately drive _build_history with a fresh undecorated
|
||||
# message — SSE replay shape.
|
||||
session = self._session_with_messages(
|
||||
[{"role": "tool", "tool_call_id": "call_a", "content": wrapped}]
|
||||
)
|
||||
sse_history = _build_history(session)
|
||||
# Both surfaces produce the same advisory + cleaned content.
|
||||
assert rest_messages[0]["content"] == "raw"
|
||||
assert rest_messages[0]["advisories"] == [
|
||||
{"type": "user_interjection", "text": "hi", "priority": "notice"}
|
||||
]
|
||||
assert sse_history[0]["content"] == "raw"
|
||||
assert sse_history[0]["advisories"] == [
|
||||
{"type": "user_interjection", "text": "hi", "priority": "notice"}
|
||||
]
|
||||
|
||||
def test_build_history_extracts_advisories_from_list_content_text_part(self):
|
||||
"""List-typed tool output (image / structured MCP results)
|
||||
with a Seam 1 splice carries the wrap envelope as a separate
|
||||
text part (``session.py``'s tool-result loop appends
|
||||
``{"type": "text", "text": wrap_tool_result("", advisories)}``
|
||||
when ``output`` is a list). ``_build_history`` must walk the
|
||||
list parts, extract advisories from any wrap-envelope text
|
||||
part, and DROP that text part from the projected list — the
|
||||
cleaned inner content is empty by construction, and leaving
|
||||
the part would cause the JS replay to render the literal
|
||||
envelope text as a chunk inside the tool block AND fail to
|
||||
render the queued message as a user bubble.
|
||||
|
||||
Removing the list-content branch in ``_build_history``'s tool-
|
||||
message advisory extraction breaks this test.
|
||||
"""
|
||||
from turnstone.core.tool_advisory import UserInterjection, wrap_tool_result
|
||||
from turnstone.server import _build_history
|
||||
|
||||
wrap_text = wrap_tool_result(
|
||||
"",
|
||||
[UserInterjection(message="inspect histogram", priority="notice")],
|
||||
)
|
||||
list_content = [
|
||||
{"type": "text", "text": "the chart shows X"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,xxx"}},
|
||||
{"type": "text", "text": wrap_text},
|
||||
]
|
||||
session = self._session_with_messages(
|
||||
[{"role": "tool", "tool_call_id": "call_a", "content": list_content}]
|
||||
)
|
||||
history = _build_history(session)
|
||||
# Wire-shape content keeps the original text + image parts but
|
||||
# has the wrap text-part dropped.
|
||||
wire_content = history[0]["content"]
|
||||
assert isinstance(wire_content, list)
|
||||
assert len(wire_content) == 2
|
||||
assert wire_content[0] == {"type": "text", "text": "the chart shows X"}
|
||||
assert wire_content[1] == {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,xxx"},
|
||||
}
|
||||
# Advisory rides on the wire so JS replay renders the user
|
||||
# bubble after the tool block — same contract as the string-
|
||||
# content path.
|
||||
assert history[0]["advisories"] == [
|
||||
{
|
||||
"type": "user_interjection",
|
||||
"text": "inspect histogram",
|
||||
"priority": "notice",
|
||||
}
|
||||
]
|
||||
|
||||
def test_build_history_keeps_legitimate_envelope_text_part_with_body(self):
|
||||
"""A tool that legitimately produces output containing a
|
||||
well-formed ``<tool_output>`` envelope as a text part (e.g.
|
||||
documentation viewer, code analyzer demoing the wrapper, an
|
||||
echo tool) must NOT have that part dropped on replay. The
|
||||
list-content drop heuristic must require both an empty cleaned
|
||||
inner body AND at least one extracted advisory — the
|
||||
signature of the injected ``wrap_tool_result("", advisories)``
|
||||
carrier. A legitimate tool envelope has non-empty inner body
|
||||
OR no advisories, and stays in the projected list verbatim.
|
||||
|
||||
Removing the ``not cleaned_text and advisories_from_part``
|
||||
guard breaks this test (the legitimate envelope gets dropped
|
||||
from the wire content)."""
|
||||
from turnstone.server import _build_history
|
||||
|
||||
legit_envelope_text = (
|
||||
"<tool_output>\nThis is what a tool_output envelope looks like.\n</tool_output>"
|
||||
)
|
||||
list_content = [
|
||||
{"type": "text", "text": "doc preview:"},
|
||||
{"type": "text", "text": legit_envelope_text},
|
||||
]
|
||||
session = self._session_with_messages(
|
||||
[{"role": "tool", "tool_call_id": "call_a", "content": list_content}]
|
||||
)
|
||||
history = _build_history(session)
|
||||
# All parts survive — none dropped.
|
||||
wire_content = history[0]["content"]
|
||||
assert isinstance(wire_content, list)
|
||||
assert len(wire_content) == 2
|
||||
assert wire_content[1]["text"] == legit_envelope_text
|
||||
# No advisories surfaced (no system-reminder blocks were
|
||||
# extracted from the legitimate envelope).
|
||||
assert "advisories" not in history[0]
|
||||
|
||||
|
||||
class TestDetailInteractive:
|
||||
"""Interactive parity for the lifted ``GET /v1/api/workstreams/{ws_id}``.
|
||||
|
||||
|
||||
@@ -117,7 +117,6 @@
|
||||
|
||||
[mcp]
|
||||
# config_path = "" # Path to MCP servers config file (JSON)
|
||||
# refresh_interval = 14400 # Refresh interval in seconds (default: 4h)
|
||||
|
||||
# --- Server (node, console) ---
|
||||
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
"""turnstone - Multi-node AI orchestration platform with tool use, agent routing, and cluster simulation."""
|
||||
|
||||
__version__ = "1.5.7"
|
||||
__version__ = "1.5.9"
|
||||
|
||||
@@ -674,6 +674,16 @@ class McpServerInfo(BaseModel):
|
||||
registry_name: str | None = None
|
||||
registry_version: str = ""
|
||||
registry_meta: str = "{}"
|
||||
auth_type: str = "static"
|
||||
oauth_client_id: str | None = None
|
||||
oauth_scopes: str | None = None
|
||||
oauth_audience: str | None = None
|
||||
oauth_registration_mode: str | None = None
|
||||
oauth_authorization_server_url: str | None = None
|
||||
oauth_as_issuer_cached: str | None = None
|
||||
# Fernet ciphertext; never decrypted on the read path. Responses
|
||||
# carry the masked ``"***"`` sentinel via ``_mask_mcp_secrets``.
|
||||
oauth_client_secret_ct: str | None = None
|
||||
created: str
|
||||
updated: str
|
||||
|
||||
@@ -704,6 +714,16 @@ class CreateMcpServerRequest(BaseModel):
|
||||
env: dict[str, str] = Field(default_factory=dict)
|
||||
auto_approve: bool = False
|
||||
enabled: bool = True
|
||||
# OAuth-MCP: one of 'none' | 'static' | 'oauth_user'.
|
||||
# ``oauth_client_secret`` is plaintext input; never persisted,
|
||||
# redacted in audit log.
|
||||
auth_type: str = "static"
|
||||
oauth_client_id: str | None = None
|
||||
oauth_client_secret: str | None = None
|
||||
oauth_scopes: str | None = None
|
||||
oauth_audience: str | None = None
|
||||
oauth_registration_mode: str | None = None
|
||||
oauth_authorization_server_url: str | None = None
|
||||
|
||||
|
||||
class UpdateMcpServerRequest(BaseModel):
|
||||
@@ -716,6 +736,13 @@ class UpdateMcpServerRequest(BaseModel):
|
||||
env: dict[str, str] | None = None
|
||||
auto_approve: bool | None = None
|
||||
enabled: bool | None = None
|
||||
auth_type: str | None = None
|
||||
oauth_client_id: str | None = None
|
||||
oauth_client_secret: str | None = None
|
||||
oauth_scopes: str | None = None
|
||||
oauth_audience: str | None = None
|
||||
oauth_registration_mode: str | None = None
|
||||
oauth_authorization_server_url: str | None = None
|
||||
|
||||
|
||||
class ListMcpServersResponse(BaseModel):
|
||||
@@ -765,6 +792,34 @@ class SkillDiscoverResponse(BaseModel):
|
||||
skills: list[SkillDiscoverListing]
|
||||
|
||||
|
||||
class ParseSkillRequest(BaseModel):
|
||||
raw: str = Field(
|
||||
min_length=1,
|
||||
max_length=32_768,
|
||||
description=(
|
||||
"Raw SKILL.md text — YAML frontmatter delimited by ``---`` "
|
||||
"followed by the markdown body. Capped at 32 KiB to match "
|
||||
"``admin_create_skill``'s ``content`` ceiling and to bound "
|
||||
"the synchronous YAML parser's worst-case CPU cost. The "
|
||||
"handler reuses the Python parser at "
|
||||
"``turnstone.core.skill_parser`` so admin UIs and external "
|
||||
"import paths agree on field extraction."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ParseSkillResponse(BaseModel):
|
||||
name: str
|
||||
description: str
|
||||
content: str
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
author: str = ""
|
||||
version: str = "1.0.0"
|
||||
allowed_tools: list[str] = Field(default_factory=list)
|
||||
license: str = ""
|
||||
compatibility: str = ""
|
||||
|
||||
|
||||
class SkillInstallRequest(BaseModel):
|
||||
source: str # "skills.sh" or "github"
|
||||
skill_id: str = "" # for skills.sh
|
||||
|
||||
@@ -77,6 +77,8 @@ from turnstone.api.console_schemas import (
|
||||
NodeMetadataResponse,
|
||||
OrgInfo,
|
||||
OutputAssessmentInfo,
|
||||
ParseSkillRequest,
|
||||
ParseSkillResponse,
|
||||
RegistryInstallRequest,
|
||||
RegistrySearchResponse,
|
||||
RoleInfo,
|
||||
@@ -552,6 +554,15 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [
|
||||
error_codes=[400, 404, 409, 502],
|
||||
tags=["Admin"],
|
||||
),
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/skills/parse",
|
||||
"POST",
|
||||
"Parse a SKILL.md document and return its frontmatter fields and body",
|
||||
request_model=ParseSkillRequest,
|
||||
response_model=ParseSkillResponse,
|
||||
error_codes=[400, 413],
|
||||
tags=["Admin"],
|
||||
),
|
||||
# --- Governance: Skills ---
|
||||
EndpointSpec(
|
||||
"/v1/api/admin/skills",
|
||||
@@ -1594,6 +1605,8 @@ _ALL_MODELS: list[type[BaseModel]] = [
|
||||
SkillInstallResponse,
|
||||
SkillInfo,
|
||||
SkillVersionInfo,
|
||||
ParseSkillRequest,
|
||||
ParseSkillResponse,
|
||||
CreateSkillRequest,
|
||||
UpdateSkillRequest,
|
||||
ListSkillsResponse,
|
||||
|
||||
+5
-13
@@ -312,7 +312,7 @@ class TerminalUI(SessionUI):
|
||||
sys.stdout.write(f"{RED}{message}{RESET}\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
def _print_reminder(self, reminders: list[dict[str, str]]) -> None:
|
||||
def _print_reminder(self, reminders: list[dict[str, Any]]) -> None:
|
||||
"""Render a metacognitive reminder list as ``[metacognition · type] text``
|
||||
lines in the terminal — the CLI's equivalent of the web UI's
|
||||
yellow themed bubble. Used by both ``on_user_reminder`` and
|
||||
@@ -326,10 +326,12 @@ class TerminalUI(SessionUI):
|
||||
sys.stdout.write(f"{YELLOW}[{label}]{RESET} {text}\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
def on_user_reminder(self, reminders: list[dict[str, str]]) -> None:
|
||||
def on_user_reminder(self, reminders: list[dict[str, Any]], source: str | None = None) -> None:
|
||||
# ``source`` ignored — the CLI doesn't render a wake marker
|
||||
# (terminal output is anchored by sequence, not anchor element).
|
||||
self._print_reminder(reminders)
|
||||
|
||||
def on_tool_reminder(self, reminders: list[dict[str, str]], tool_call_id: str) -> None:
|
||||
def on_tool_reminder(self, reminders: list[dict[str, Any]], tool_call_id: str) -> None:
|
||||
# tool_call_id ignored — the CLI anchors by output sequence
|
||||
# (the line lands directly after the tool result that
|
||||
# triggered the batch's reminder).
|
||||
@@ -1015,15 +1017,6 @@ def main() -> None:
|
||||
help="Path to MCP server config file (standard mcpServers JSON format)",
|
||||
)
|
||||
|
||||
from turnstone.core.config import nonneg_float
|
||||
|
||||
parser.add_argument(
|
||||
"--mcp-refresh-interval",
|
||||
type=nonneg_float,
|
||||
default=14400,
|
||||
metavar="SECONDS",
|
||||
help="Periodic MCP tool refresh interval for servers without push notifications (default: 14400 = 4h, 0 to disable)",
|
||||
)
|
||||
judge_group = parser.add_argument_group("Judge options")
|
||||
judge_group.add_argument(
|
||||
"--judge",
|
||||
@@ -1136,7 +1129,6 @@ def main() -> None:
|
||||
|
||||
mcp_client = create_mcp_client(
|
||||
getattr(args, "mcp_config", None),
|
||||
refresh_interval=getattr(args, "mcp_refresh_interval", 14400),
|
||||
storage=_get_storage(),
|
||||
)
|
||||
|
||||
|
||||
@@ -1488,12 +1488,17 @@ class CoordinatorClient:
|
||||
is_own_child = full.get("parent_ws_id") == self._coord_ws_id
|
||||
if not (is_self or is_own_child):
|
||||
return miss
|
||||
# load_messages returns the full history in chronological order
|
||||
# (no limit param in the Protocol) — slice the tail here. Defensive
|
||||
# load_messages returns the full history in chronological order.
|
||||
# We slice the tail in Python because the SQL tail-N is
|
||||
# approximate across conversation boundaries. Defensive
|
||||
# try/except: storage errors should not break inspect.
|
||||
messages: list[Any] = []
|
||||
try:
|
||||
all_msgs = self._storage.load_messages(ws_id)
|
||||
# repair=False — inspect is a display read (admin viewing a
|
||||
# child's history in the tree UI). The LLM-context repair
|
||||
# pass would strip trailing partial turns the operator is
|
||||
# watching.
|
||||
all_msgs = self._storage.load_messages(ws_id, repair=False)
|
||||
if message_limit and message_limit > 0:
|
||||
messages = all_msgs[-message_limit:]
|
||||
else:
|
||||
@@ -1725,7 +1730,11 @@ def _last_assistant_text(storage: Any, ws_id: str) -> str | None:
|
||||
in just to surface its final turn.
|
||||
"""
|
||||
try:
|
||||
rows = storage.load_messages(ws_id, limit=_WAIT_MESSAGE_TAIL_LIMIT)
|
||||
# repair=False — this reads the tail for display ("waiting on" bubble).
|
||||
# The repair pass would strip a trailing partial assistant turn,
|
||||
# making us return the penultimate assistant message instead of the
|
||||
# one the operator is watching.
|
||||
rows = storage.load_messages(ws_id, limit=_WAIT_MESSAGE_TAIL_LIMIT, repair=False)
|
||||
except Exception:
|
||||
log.debug("coord_client.wait.load_messages_failed ws=%s", ws_id, exc_info=True)
|
||||
return None
|
||||
|
||||
@@ -0,0 +1,308 @@
|
||||
"""Observer that nudges idle coordinators with active children.
|
||||
|
||||
Subscribes to a coordinator-side :class:`SessionManager`'s state events.
|
||||
When a coord transitions to :class:`WorkstreamState.IDLE` while still
|
||||
having active interactive children, enqueues an ``idle_children`` nudge
|
||||
on the coord's :class:`NudgeQueue`. The
|
||||
:class:`turnstone.core.idle_nudge_watcher.IdleNudgeWatcher` (registered
|
||||
*after* this observer in the lifespan, so subscriber-order has the
|
||||
observer fire first on the same IDLE event) then peeks the queue and
|
||||
dispatches the wake send.
|
||||
|
||||
Gates (in order):
|
||||
|
||||
1. **Coordinator-only.** Skip non-coord workstreams. Watcher is
|
||||
kind-agnostic; this observer is the kind-aware piece.
|
||||
2. **Skip if last assistant turn used ``wait_for_workstream``.** The
|
||||
coord is already using the right tool — don't pile on with a nudge
|
||||
suggesting the same tool.
|
||||
3. **Per-(ws_id, nudge_type) hard cap** (default 3). Resets when the
|
||||
ws leaves IDLE for any non-wake-driven reason (tracked by
|
||||
:class:`ChatSession._wake_source_tag` — see below).
|
||||
4. **Active children query.** ``storage.list_workstreams`` filtered to
|
||||
interactive kind under the coord's user, with state in
|
||||
:data:`_ACTIVE_CHILD_STATES`. Empty result → no nudge.
|
||||
5. **Cooldown** via :func:`should_nudge` — default 300s per nudge type.
|
||||
|
||||
The enqueue passes ``valid_until`` that re-queries active children at
|
||||
drain time; if every child finished while the queue waited, the entry
|
||||
drops without delivering a stale snapshot.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import threading
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.metacognition import (
|
||||
_cooldown_allows,
|
||||
format_idle_children_nudge,
|
||||
should_nudge,
|
||||
)
|
||||
from turnstone.core.workstream import WorkstreamKind, WorkstreamState
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from turnstone.core.session import ChatSession
|
||||
from turnstone.core.session_manager import SessionManager
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
from turnstone.core.workstream import Workstream
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# Active = the model can act on the child (it's still working,
|
||||
# streaming, or waiting on user attention). Excludes "idle" (the
|
||||
# child is now waiting and can't be unblocked by the coord), "closed"
|
||||
# (gone), "deleted" (gone), and "error" (the model can't unblock an
|
||||
# errored child without operator intervention; cooldown handles repeat
|
||||
# fires for stuck-error children).
|
||||
_ACTIVE_CHILD_STATES: frozenset[str] = frozenset(
|
||||
{
|
||||
WorkstreamState.THINKING.value,
|
||||
WorkstreamState.RUNNING.value,
|
||||
WorkstreamState.ATTENTION.value,
|
||||
}
|
||||
)
|
||||
|
||||
# Hard cap on per-session ``idle_children`` fires. Even with the
|
||||
# cooldown and wait-tool skip gate, a coord that ignores every nudge
|
||||
# shouldn't be hammered indefinitely. Resets when the ws leaves IDLE
|
||||
# for a non-wake reason (real user input).
|
||||
_HARD_CAP_PER_SESSION = 3
|
||||
|
||||
# Soft cap on the snapshot query. Higher than ``WAIT_MAX_WS_IDS`` so
|
||||
# the SQL ``LIMIT`` (applied before the Python state filter) doesn't
|
||||
# clip genuinely-active children whose ``updated`` timestamp is older
|
||||
# than recently-closed siblings. Realistic coord histories are far
|
||||
# smaller than this; if a coord ever exceeds it, the formatter still
|
||||
# truncates to ``WAIT_MAX_WS_IDS`` for the model-facing suggestion.
|
||||
_ACTIVE_CHILDREN_QUERY_LIMIT = 200
|
||||
|
||||
|
||||
class CoordinatorIdleObserver:
|
||||
"""Subscribe to a coord SessionManager's IDLE events and enqueue
|
||||
``idle_children`` nudges when active children remain.
|
||||
"""
|
||||
|
||||
def __init__(self, manager: SessionManager, storage: StorageBackend) -> None:
|
||||
self._manager = manager
|
||||
self._storage = storage
|
||||
self._callback: Callable[[str, WorkstreamState], None] | None = None
|
||||
# Per-ws fire counts keyed by ``ws_id`` → ``{nudge_type: count}``.
|
||||
# Two-level dict makes the "any caps for this ws?" check at
|
||||
# leave-IDLE an O(1) ``ws_id in self._fire_counts`` lookup
|
||||
# instead of an O(N_caps) scan over a flat tuple-keyed map.
|
||||
# Lock protects against race with the leave-IDLE reset path
|
||||
# running on a different thread (state events fire on the
|
||||
# calling thread of ``set_state`` — currently always the
|
||||
# worker thread that did the transition, but the lock keeps
|
||||
# the contract robust).
|
||||
self._fire_counts: dict[str, dict[str, int]] = {}
|
||||
self._fire_counts_lock = threading.Lock()
|
||||
|
||||
def start(self) -> None:
|
||||
"""Idempotent — registering twice is a no-op."""
|
||||
if self._callback is not None:
|
||||
return
|
||||
|
||||
def _on_state(ws_id: str, state: WorkstreamState) -> None:
|
||||
if state is not WorkstreamState.IDLE:
|
||||
# Reset hard-cap when leaving IDLE for a *real* reason
|
||||
# (not a wake-driven exit). Skip the manager-lock /
|
||||
# session-attribute walk entirely when no caps are
|
||||
# accumulated for this ws — the common case for the
|
||||
# vast majority of state transitions.
|
||||
with self._fire_counts_lock:
|
||||
has_caps = ws_id in self._fire_counts
|
||||
if not has_caps:
|
||||
return
|
||||
ws = self._manager.get(ws_id)
|
||||
if ws is None or ws.session is None:
|
||||
return
|
||||
# ``_wake_source_tag`` is set on the session iff a
|
||||
# wake send is in flight; if set, leaving IDLE is the
|
||||
# wake's own IDLE→THINKING→RUNNING transition and the
|
||||
# cap should NOT reset. If unset, the user / a real
|
||||
# producer drove the coord forward and the cap should
|
||||
# clear so the next genuine idle bracket is fresh.
|
||||
if not ws.session._wake_source_tag:
|
||||
self._reset_caps_for(ws_id)
|
||||
return
|
||||
|
||||
# state == IDLE branch.
|
||||
try:
|
||||
self._maybe_enqueue(ws_id)
|
||||
except Exception:
|
||||
log.exception("coord_idle_observer.maybe_enqueue_failed ws=%s", ws_id[:8])
|
||||
|
||||
self._callback = _on_state
|
||||
self._manager.subscribe_to_state(_on_state)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Unsubscribe; idempotent."""
|
||||
cb = self._callback
|
||||
if cb is None:
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
self._manager.unsubscribe_from_state(cb)
|
||||
self._callback = None
|
||||
|
||||
def _maybe_enqueue(self, ws_id: str) -> None:
|
||||
ws = self._manager.get(ws_id)
|
||||
if ws is None or ws.session is None:
|
||||
return
|
||||
if ws.kind is not WorkstreamKind.COORDINATOR:
|
||||
return
|
||||
session = ws.session
|
||||
|
||||
# Gate ordering matters: cheap checks first, expensive checks
|
||||
# last. The cooldown peek + per-session hard cap are
|
||||
# microsecond-cheap dict lookups; ``_last_assistant_used_wait``
|
||||
# walks ``session.messages`` reversed; ``_active_children`` and
|
||||
# ``_visible_memory_count`` round-trip to storage. With a 300s
|
||||
# cooldown most idle events will short-circuit at the peek.
|
||||
cooldown_secs = getattr(session._mem_cfg, "nudge_cooldown", 300)
|
||||
if not _cooldown_allows(
|
||||
"idle_children", session._metacog_state, cooldown_secs=cooldown_secs
|
||||
):
|
||||
return
|
||||
|
||||
# Gate: per-session hard cap on idle_children fires.
|
||||
with self._fire_counts_lock:
|
||||
ws_caps = self._fire_counts.get(ws_id, {})
|
||||
if ws_caps.get("idle_children", 0) >= _HARD_CAP_PER_SESSION:
|
||||
return
|
||||
|
||||
# Gate: skip if the coord's last assistant turn already used
|
||||
# ``wait_for_workstream``. Don't nudge toward a tool the
|
||||
# model is already using.
|
||||
if self._last_assistant_used_wait(session):
|
||||
return
|
||||
|
||||
# Gate: query active children. Empty → nothing to nudge about.
|
||||
active = self._active_children(ws)
|
||||
if not active:
|
||||
return
|
||||
|
||||
# ``should_nudge`` re-checks cooldown AND records the timestamp
|
||||
# on success (the peek above only checks; record happens here).
|
||||
# Also enforces the message-count > 1 / memory-count > 0 sanity
|
||||
# gates we couldn't apply at the cheap-peek stage.
|
||||
if not should_nudge(
|
||||
"idle_children",
|
||||
session._metacog_state,
|
||||
message_count=len(session.messages),
|
||||
memory_count=session._visible_memory_count(),
|
||||
cooldown_secs=cooldown_secs,
|
||||
):
|
||||
return
|
||||
|
||||
text = format_idle_children_nudge(active)
|
||||
if not text: # belt-and-braces: formatter empty-input guard
|
||||
return
|
||||
|
||||
# Bind ws.id + user_id by closure so the predicate captures the
|
||||
# workstream identity (not the live ``ws`` reference, which
|
||||
# could mutate). The predicate runs at drain time outside the
|
||||
# queue lock. Use ``count_workstreams_by_state`` rather than
|
||||
# ``list_workstreams`` since the predicate only needs a
|
||||
# boolean — saves a row fetch on the chat-loop user-attach
|
||||
# path.
|
||||
bound_ws_id = ws.id
|
||||
bound_user_id = ws.user_id
|
||||
|
||||
def _still_has_active_children() -> bool:
|
||||
try:
|
||||
counts = self._storage.count_workstreams_by_state(
|
||||
parent_ws_id=bound_ws_id,
|
||||
user_id=bound_user_id,
|
||||
)
|
||||
except Exception:
|
||||
log.debug(
|
||||
"coord_idle_observer.predicate_count_failed ws=%s",
|
||||
bound_ws_id[:8],
|
||||
exc_info=True,
|
||||
)
|
||||
return False
|
||||
return any(counts.get(s, 0) > 0 for s in _ACTIVE_CHILD_STATES)
|
||||
|
||||
session._nudge_queue.enqueue(
|
||||
"idle_children",
|
||||
text,
|
||||
"any",
|
||||
valid_until=_still_has_active_children,
|
||||
)
|
||||
with self._fire_counts_lock:
|
||||
ws_caps = self._fire_counts.setdefault(ws_id, {})
|
||||
ws_caps["idle_children"] = ws_caps.get("idle_children", 0) + 1
|
||||
|
||||
log.info(
|
||||
"coord_idle_observer.enqueued ws=%s active_children=%d",
|
||||
ws_id[:8],
|
||||
len(active),
|
||||
)
|
||||
|
||||
def _last_assistant_used_wait(self, session: ChatSession) -> bool:
|
||||
"""Walk back to the most recent assistant turn; if it issued a
|
||||
``wait_for_workstream`` tool call, return ``True``.
|
||||
"""
|
||||
for msg in reversed(session.messages):
|
||||
if msg.get("role") != "assistant":
|
||||
continue
|
||||
for tc in msg.get("tool_calls") or []:
|
||||
fn = tc.get("function", {}) or {}
|
||||
if fn.get("name") == "wait_for_workstream":
|
||||
return True
|
||||
return False # found the most recent assistant turn — done
|
||||
return False
|
||||
|
||||
def _active_children(self, ws: Workstream) -> list[dict[str, str]]:
|
||||
"""Query storage for the coord's interactive children whose state
|
||||
is in :data:`_ACTIVE_CHILD_STATES`. Returns row-mapping shape.
|
||||
|
||||
``list_workstreams`` orders by ``updated DESC`` and applies its
|
||||
``LIMIT`` in SQL before any state filter, so a coord with many
|
||||
recently-closed children could clip out genuinely-active rows
|
||||
whose ``updated`` timestamp is older. We bump the limit well
|
||||
above ``NUDGE_IDLE_CHILDREN_WAIT_CAP`` to absorb that —
|
||||
realistic coord histories are far smaller than the bumped
|
||||
limit. Pushing the state filter into SQL would be the
|
||||
structural fix, but that requires a storage-protocol change;
|
||||
flagged as a follow-up.
|
||||
"""
|
||||
try:
|
||||
rows = self._storage.list_workstreams(
|
||||
limit=_ACTIVE_CHILDREN_QUERY_LIMIT,
|
||||
parent_ws_id=ws.id,
|
||||
kind=WorkstreamKind.INTERACTIVE,
|
||||
user_id=ws.user_id,
|
||||
)
|
||||
except Exception:
|
||||
log.debug("coord_idle_observer.list_failed ws=%s", ws.id[:8], exc_info=True)
|
||||
return []
|
||||
|
||||
out: list[dict[str, str]] = []
|
||||
for row in rows:
|
||||
mapping = getattr(row, "_mapping", row)
|
||||
state = mapping["state"]
|
||||
if state not in _ACTIVE_CHILD_STATES:
|
||||
continue
|
||||
out.append(
|
||||
{
|
||||
"ws_id": mapping["ws_id"],
|
||||
"name": mapping["name"] or "",
|
||||
"state": state,
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
def _reset_caps_for(self, ws_id: str) -> None:
|
||||
"""Drop every nudge-type cap counter for ``ws_id`` on a real
|
||||
(non-wake) leave-IDLE event — the next genuine idle bracket
|
||||
starts fresh. O(1) with the per-ws nested-dict layout.
|
||||
"""
|
||||
with self._fire_counts_lock:
|
||||
self._fire_counts.pop(ws_id, None)
|
||||
+1080
-166
File diff suppressed because it is too large
Load Diff
@@ -511,6 +511,49 @@ function _toggleOidcPanel(userId, username, rowEl) {
|
||||
});
|
||||
}
|
||||
|
||||
function _buildOidcRow(oid, userId, username) {
|
||||
var shortIssuer = _issuerShortName(oid.issuer || "");
|
||||
var shortSubject =
|
||||
(oid.subject || "").length > 12
|
||||
? (oid.subject || "").slice(0, 12) + "\u2026"
|
||||
: oid.subject || "";
|
||||
var lastLogin = oid.last_login ? _relativeTime(oid.last_login) : "never";
|
||||
return (
|
||||
'<div class="oidc-identity-row">' +
|
||||
'<span class="oidc-identity-issuer"><span class="scope-badge">' +
|
||||
escapeHtml(shortIssuer) +
|
||||
"</span></span>" +
|
||||
'<span class="oidc-identity-subject" title="' +
|
||||
escapeHtml(oid.subject || "") +
|
||||
'">' +
|
||||
escapeHtml(shortSubject) +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-email" title="' +
|
||||
escapeHtml(oid.email || "") +
|
||||
'">' +
|
||||
escapeHtml(oid.email || "\u2014") +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-time">' +
|
||||
escapeHtml(lastLogin) +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-actions">' +
|
||||
'<button class="admin-btn-danger" aria-label="Unlink ' +
|
||||
escapeHtml(shortIssuer) +
|
||||
" identity " +
|
||||
escapeHtml(shortSubject) +
|
||||
'" data-oidc-issuer="' +
|
||||
escapeHtml(oid.issuer || "") +
|
||||
'" data-oidc-subject="' +
|
||||
escapeHtml(oid.subject || "") +
|
||||
'" data-oidc-username="' +
|
||||
escapeHtml(username) +
|
||||
'" data-oidc-user-id="' +
|
||||
escapeHtml(userId) +
|
||||
'">unlink</button>' +
|
||||
"</span></div>"
|
||||
);
|
||||
}
|
||||
|
||||
function _renderOidcDetail(panel, identities, userId, username) {
|
||||
var body = panel.querySelector(".oidc-detail-body");
|
||||
if (!body) return;
|
||||
@@ -522,46 +565,7 @@ function _renderOidcDetail(panel, identities, userId, username) {
|
||||
}
|
||||
var html = "";
|
||||
for (var i = 0; i < identities.length; i++) {
|
||||
var oid = identities[i];
|
||||
var shortIssuer = _issuerShortName(oid.issuer || "");
|
||||
var shortSubject =
|
||||
(oid.subject || "").length > 12
|
||||
? (oid.subject || "").slice(0, 12) + "\u2026"
|
||||
: oid.subject || "";
|
||||
var lastLogin = oid.last_login ? _relativeTime(oid.last_login) : "never";
|
||||
html +=
|
||||
'<div class="oidc-identity-row">' +
|
||||
'<span class="oidc-identity-issuer"><span class="scope-badge">' +
|
||||
escapeHtml(shortIssuer) +
|
||||
"</span></span>" +
|
||||
'<span class="oidc-identity-subject" title="' +
|
||||
escapeHtml(oid.subject || "") +
|
||||
'">' +
|
||||
escapeHtml(shortSubject) +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-email" title="' +
|
||||
escapeHtml(oid.email || "") +
|
||||
'">' +
|
||||
escapeHtml(oid.email || "\u2014") +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-time">' +
|
||||
escapeHtml(lastLogin) +
|
||||
"</span>" +
|
||||
'<span class="oidc-identity-actions">' +
|
||||
'<button class="admin-btn-danger" aria-label="Unlink ' +
|
||||
escapeHtml(shortIssuer) +
|
||||
" identity " +
|
||||
escapeHtml(shortSubject) +
|
||||
'" data-oidc-issuer="' +
|
||||
escapeHtml(oid.issuer || "") +
|
||||
'" data-oidc-subject="' +
|
||||
escapeHtml(oid.subject || "") +
|
||||
'" data-oidc-username="' +
|
||||
escapeHtml(username) +
|
||||
'" data-oidc-user-id="' +
|
||||
escapeHtml(userId) +
|
||||
'">unlink</button>' +
|
||||
"</span></div>";
|
||||
html += _buildOidcRow(identities[i], userId, username);
|
||||
}
|
||||
body.innerHTML = html;
|
||||
// Update panel height for animation
|
||||
@@ -3312,9 +3316,23 @@ function _renderMcpServers(items) {
|
||||
var detailAttr = isConfig
|
||||
? 'data-mcp-detail-name="' + escapeHtml(s.name) + '"'
|
||||
: 'data-mcp-detail="' + escapeHtml(s.server_id) + '"';
|
||||
var actionBtns =
|
||||
'<button class="admin-btn-action" data-mcp-refresh="' +
|
||||
escapeHtml(s.name) +
|
||||
'">refresh</button>' +
|
||||
'<button class="admin-btn-action" data-mcp-reconnect="' +
|
||||
escapeHtml(s.name) +
|
||||
'">reconnect</button>';
|
||||
if (s.auth_type === "oauth_user") {
|
||||
actionBtns +=
|
||||
'<button class="admin-btn-action" data-mcp-oauth-connect="' +
|
||||
escapeHtml(s.name) +
|
||||
'">connect</button>';
|
||||
}
|
||||
var actions = isConfig
|
||||
? ""
|
||||
: '<button class="admin-btn-action" data-mcp-edit="' +
|
||||
? actionBtns
|
||||
: actionBtns +
|
||||
'<button class="admin-btn-action" data-mcp-edit="' +
|
||||
escapeHtml(s.server_id) +
|
||||
'">edit</button>' +
|
||||
'<button class="admin-btn-danger" data-mcp-delete="' +
|
||||
@@ -3379,6 +3397,59 @@ function _renderMcpServers(items) {
|
||||
showEditMcpModal(this.getAttribute("data-mcp-edit"));
|
||||
});
|
||||
});
|
||||
el.querySelectorAll("[data-mcp-refresh]").forEach(function (btn) {
|
||||
btn.addEventListener("click", function () {
|
||||
var name = this.getAttribute("data-mcp-refresh");
|
||||
authFetch(
|
||||
"/v1/api/admin/mcp-servers/" + encodeURIComponent(name) + "/refresh",
|
||||
{ method: "POST" },
|
||||
)
|
||||
.then(function (r) {
|
||||
if (!r.ok) throw new Error();
|
||||
return r.json();
|
||||
})
|
||||
.then(function () {
|
||||
showToast("Refreshed " + name);
|
||||
loadAdminMcp();
|
||||
})
|
||||
.catch(function () {
|
||||
showToast("Failed to refresh " + name);
|
||||
});
|
||||
});
|
||||
});
|
||||
el.querySelectorAll("[data-mcp-reconnect]").forEach(function (btn) {
|
||||
btn.addEventListener("click", function () {
|
||||
var name = this.getAttribute("data-mcp-reconnect");
|
||||
authFetch(
|
||||
"/v1/api/admin/mcp-servers/" + encodeURIComponent(name) + "/reconnect",
|
||||
{ method: "POST" },
|
||||
)
|
||||
.then(function (r) {
|
||||
if (!r.ok) throw new Error();
|
||||
return r.json();
|
||||
})
|
||||
.then(function () {
|
||||
showToast("Reconnected " + name);
|
||||
loadAdminMcp();
|
||||
})
|
||||
.catch(function () {
|
||||
showToast("Failed to reconnect " + name);
|
||||
});
|
||||
});
|
||||
});
|
||||
el.querySelectorAll("[data-mcp-oauth-connect]").forEach(function (btn) {
|
||||
btn.addEventListener("click", function () {
|
||||
var name = this.getAttribute("data-mcp-oauth-connect");
|
||||
// Open the OAuth /start endpoint in a new window so the redirect
|
||||
// chain (AS → callback → return_url) doesn't displace the admin UI.
|
||||
var url =
|
||||
"/v1/api/mcp/oauth/start?server=" +
|
||||
encodeURIComponent(name) +
|
||||
"&return_url=" +
|
||||
encodeURIComponent(window.location.href);
|
||||
window.open(url, "_blank", "noopener");
|
||||
});
|
||||
});
|
||||
el.querySelectorAll("[data-mcp-delete]").forEach(function (btn) {
|
||||
btn.addEventListener("click", function () {
|
||||
var sid = this.getAttribute("data-mcp-delete");
|
||||
@@ -3413,6 +3484,46 @@ function toggleMcpTransport() {
|
||||
v === "stdio" ? "" : "none";
|
||||
document.getElementById("mcp-http-fields").style.display =
|
||||
v === "streamable-http" ? "" : "none";
|
||||
// Re-evaluate auth-field visibility because the headers row lives
|
||||
// inside mcp-http-fields and gets toggled there.
|
||||
toggleMcpAuthFields();
|
||||
}
|
||||
|
||||
function _selectedMcpAuthType() {
|
||||
var radios = document.getElementsByName("mcp-auth-type");
|
||||
for (var i = 0; i < radios.length; i++) {
|
||||
if (radios[i].checked) return radios[i].value;
|
||||
}
|
||||
return "static";
|
||||
}
|
||||
|
||||
function toggleMcpAuthFields() {
|
||||
var authType = _selectedMcpAuthType();
|
||||
var oauthDiv = document.getElementById("mcp-oauth-fields");
|
||||
if (oauthDiv) {
|
||||
oauthDiv.style.display = authType === "oauth_user" ? "" : "none";
|
||||
}
|
||||
// The "Headers" textarea (inside mcp-http-fields) is only meaningful
|
||||
// for static auth; hide it for 'none' / 'oauth_user' so operators
|
||||
// don't accidentally configure stale credentials.
|
||||
var headersInput = document.getElementById("mcp-headers");
|
||||
if (headersInput) {
|
||||
var headersLabel = document.querySelector('label[for="mcp-headers"]');
|
||||
var show = authType === "static";
|
||||
headersInput.style.display = show ? "" : "none";
|
||||
if (headersLabel) headersLabel.style.display = show ? "" : "none";
|
||||
}
|
||||
}
|
||||
|
||||
function _wireMcpAudienceAutofill() {
|
||||
// Idempotent — only attach the listener once per page lifetime.
|
||||
var urlInput = document.getElementById("mcp-url");
|
||||
if (!urlInput || urlInput.dataset.audAutofill === "1") return;
|
||||
urlInput.dataset.audAutofill = "1";
|
||||
urlInput.addEventListener("blur", function () {
|
||||
var aud = document.getElementById("mcp-oauth-audience");
|
||||
if (aud && !aud.value.trim()) aud.value = urlInput.value.trim();
|
||||
});
|
||||
}
|
||||
|
||||
function showCreateMcpModal() {
|
||||
@@ -3431,8 +3542,20 @@ function showCreateMcpModal() {
|
||||
document.getElementById("mcp-headers").value = "";
|
||||
document.getElementById("mcp-auto-approve").checked = false;
|
||||
document.getElementById("mcp-enabled").checked = true;
|
||||
// Reset auth radios + OAuth subfields to the 'static' default.
|
||||
document.getElementById("mcp-auth-static").checked = true;
|
||||
document.getElementById("mcp-auth-none").checked = false;
|
||||
document.getElementById("mcp-auth-oauth").checked = false;
|
||||
document.getElementById("mcp-oauth-as-url").value = "";
|
||||
document.getElementById("mcp-oauth-registration").value = "preregistered";
|
||||
document.getElementById("mcp-oauth-client-id").value = "";
|
||||
document.getElementById("mcp-oauth-client-secret").value = "";
|
||||
document.getElementById("mcp-oauth-scopes").value = "";
|
||||
document.getElementById("mcp-oauth-audience").value = "";
|
||||
document.getElementById("mcp-create-error").style.display = "none";
|
||||
toggleMcpTransport();
|
||||
toggleMcpAuthFields();
|
||||
_wireMcpAudienceAutofill();
|
||||
document.getElementById("mcp-name").focus();
|
||||
_mcpCreateTrap = _installTrap("mcp-create-overlay", "mcp-create-box");
|
||||
}
|
||||
@@ -3483,7 +3606,25 @@ function showEditMcpModal(serverId) {
|
||||
document.getElementById("mcp-auto-approve").checked =
|
||||
s.auto_approve || false;
|
||||
document.getElementById("mcp-enabled").checked = s.enabled !== false;
|
||||
var authType = s.auth_type || "static";
|
||||
document.getElementById("mcp-auth-none").checked = authType === "none";
|
||||
document.getElementById("mcp-auth-static").checked =
|
||||
authType === "static";
|
||||
document.getElementById("mcp-auth-oauth").checked =
|
||||
authType === "oauth_user";
|
||||
document.getElementById("mcp-oauth-as-url").value =
|
||||
s.oauth_authorization_server_url || "";
|
||||
document.getElementById("mcp-oauth-registration").value =
|
||||
s.oauth_registration_mode || "preregistered";
|
||||
document.getElementById("mcp-oauth-client-id").value =
|
||||
s.oauth_client_id || "";
|
||||
// Secret field always blank — write-only, never read back.
|
||||
document.getElementById("mcp-oauth-client-secret").value = "";
|
||||
document.getElementById("mcp-oauth-scopes").value = s.oauth_scopes || "";
|
||||
document.getElementById("mcp-oauth-audience").value =
|
||||
s.oauth_audience || "";
|
||||
toggleMcpTransport();
|
||||
toggleMcpAuthFields();
|
||||
})
|
||||
.catch(function () {
|
||||
showToast("Failed to load server details");
|
||||
@@ -3505,11 +3646,13 @@ function _parseMcpForm() {
|
||||
return { error: "Name must match [a-zA-Z0-9._-]+" };
|
||||
if (name.indexOf("__") >= 0) return { error: "Name must not contain '__'" };
|
||||
|
||||
var authType = _selectedMcpAuthType();
|
||||
var payload = {
|
||||
name: name,
|
||||
transport: transport,
|
||||
auto_approve: document.getElementById("mcp-auto-approve").checked,
|
||||
enabled: document.getElementById("mcp-enabled").checked,
|
||||
auth_type: authType,
|
||||
};
|
||||
|
||||
if (transport === "stdio") {
|
||||
@@ -3535,19 +3678,46 @@ function _parseMcpForm() {
|
||||
payload.env = envObj;
|
||||
} else {
|
||||
payload.url = document.getElementById("mcp-url").value.trim();
|
||||
var hdrText = document.getElementById("mcp-headers").value.trim();
|
||||
var hdrObj = {};
|
||||
if (hdrText) {
|
||||
hdrText.split("\n").forEach(function (line) {
|
||||
var colon = line.indexOf(":");
|
||||
if (colon > 0)
|
||||
hdrObj[line.substring(0, colon).trim()] = line
|
||||
.substring(colon + 1)
|
||||
.trim();
|
||||
});
|
||||
if (authType === "static") {
|
||||
var hdrText = document.getElementById("mcp-headers").value.trim();
|
||||
var hdrObj = {};
|
||||
if (hdrText) {
|
||||
hdrText.split("\n").forEach(function (line) {
|
||||
var colon = line.indexOf(":");
|
||||
if (colon > 0)
|
||||
hdrObj[line.substring(0, colon).trim()] = line
|
||||
.substring(colon + 1)
|
||||
.trim();
|
||||
});
|
||||
}
|
||||
payload.headers = hdrObj;
|
||||
} else {
|
||||
// 'none' / 'oauth_user' — clear server-side static headers state.
|
||||
payload.headers = {};
|
||||
}
|
||||
payload.headers = hdrObj;
|
||||
}
|
||||
|
||||
if (authType === "oauth_user") {
|
||||
payload.oauth_authorization_server_url = document
|
||||
.getElementById("mcp-oauth-as-url")
|
||||
.value.trim();
|
||||
payload.oauth_registration_mode = document.getElementById(
|
||||
"mcp-oauth-registration",
|
||||
).value;
|
||||
payload.oauth_client_id = document
|
||||
.getElementById("mcp-oauth-client-id")
|
||||
.value.trim();
|
||||
payload.oauth_scopes = document
|
||||
.getElementById("mcp-oauth-scopes")
|
||||
.value.trim();
|
||||
payload.oauth_audience = document
|
||||
.getElementById("mcp-oauth-audience")
|
||||
.value.trim();
|
||||
var secret = document.getElementById("mcp-oauth-client-secret").value;
|
||||
// Submit only when the operator typed a value; redacted in audit log.
|
||||
if (secret) payload.oauth_client_secret = secret;
|
||||
}
|
||||
|
||||
return payload;
|
||||
}
|
||||
|
||||
|
||||
@@ -439,25 +439,6 @@
|
||||
font-size: 10px;
|
||||
}
|
||||
|
||||
/* Storage-truncation indicator — same convention as the interactive
|
||||
UI's `.tool-output-truncated` pill (transparent bg, dim border,
|
||||
small font) so the operator reads the affordance the same way on
|
||||
both surfaces. Sibling node next to .coord-tool-row-result rather
|
||||
than text-in-content so a future "best-effort JSON repair" pass
|
||||
on the result body doesn't have to strip a marker string. */
|
||||
.coord-tool-truncated {
|
||||
display: inline-block;
|
||||
margin-top: 4px;
|
||||
margin-left: 6px;
|
||||
padding: 1px 6px;
|
||||
font-size: 10px;
|
||||
font-family: var(--font-mono);
|
||||
color: var(--ink-3);
|
||||
background: transparent;
|
||||
border: 1px solid var(--ink-3);
|
||||
border-radius: 3px;
|
||||
}
|
||||
|
||||
/* memory/recall calls are background metadata — the audit trail is
|
||||
useful but they crowd the tree on workstreams with heavy memory
|
||||
usage. Dim the row by default; full opacity on hover so they
|
||||
|
||||
@@ -368,34 +368,89 @@
|
||||
return el;
|
||||
}
|
||||
|
||||
// Build a structured ``.msg.watch-result`` card for a
|
||||
// ``watch_triggered`` reminder — full-width treatment with
|
||||
// command preview header + shell output body + poll counter footer.
|
||||
// First-pass functional rendering; bespoke design polish lives in a
|
||||
// future workstream. All text goes through ``textContent`` so shell
|
||||
// output containing angle brackets / scripts / steering bytes
|
||||
// renders inertly.
|
||||
function buildWatchResultBubble(r) {
|
||||
const el = document.createElement("div");
|
||||
el.className = "msg watch-result";
|
||||
el.setAttribute("role", "article");
|
||||
el.setAttribute("data-ts-role", "watch");
|
||||
el.setAttribute("aria-label", "watch");
|
||||
const header = document.createElement("div");
|
||||
header.className = "msg-watch-header";
|
||||
header.textContent =
|
||||
"watch" + (r.watch_name ? " · " + String(r.watch_name) : "");
|
||||
el.appendChild(header);
|
||||
if (r.command) {
|
||||
const cmd = document.createElement("div");
|
||||
cmd.className = "msg-watch-cmd";
|
||||
cmd.textContent = "$ " + String(r.command);
|
||||
el.appendChild(cmd);
|
||||
}
|
||||
const body = document.createElement("pre");
|
||||
body.className = "msg-watch-body";
|
||||
body.textContent = r.text || "";
|
||||
el.appendChild(body);
|
||||
if (r.poll_count != null && r.max_polls != null) {
|
||||
const footer = document.createElement("div");
|
||||
footer.className = "msg-watch-footer";
|
||||
const finalSuffix = r.is_final ? " · final" : "";
|
||||
footer.textContent =
|
||||
"poll " +
|
||||
String(r.poll_count) +
|
||||
"/" +
|
||||
String(r.max_polls) +
|
||||
finalSuffix;
|
||||
el.appendChild(footer);
|
||||
}
|
||||
return el;
|
||||
}
|
||||
|
||||
// Default ``.msg.user-reminder`` bubble — yellow themed advisory used
|
||||
// for every metacog nudge other than ``watch_triggered``.
|
||||
function buildDefaultReminderBubble(r) {
|
||||
const el = document.createElement("div");
|
||||
el.className = "msg user-reminder";
|
||||
el.setAttribute("role", "article");
|
||||
el.setAttribute("data-ts-role", "metacognition");
|
||||
el.setAttribute("aria-label", "metacognition");
|
||||
const body = document.createElement("div");
|
||||
body.className = "msg-body";
|
||||
const labelEl = document.createElement("span");
|
||||
labelEl.className = "msg-user-reminder-label";
|
||||
labelEl.textContent =
|
||||
"metacognition" + (r.type ? " · " + String(r.type) : "");
|
||||
const textEl = document.createElement("span");
|
||||
textEl.className = "msg-user-reminder-text";
|
||||
textEl.textContent = r.text || "";
|
||||
body.appendChild(labelEl);
|
||||
body.appendChild(textEl);
|
||||
el.appendChild(body);
|
||||
return el;
|
||||
}
|
||||
|
||||
// Metacognitive reminder bubble (user-channel correction / denial /
|
||||
// resume / start / completion AND tool-channel tool_error / repeat).
|
||||
// Mirrors Pane.prototype.addUserReminder / addToolReminder in the
|
||||
// interactive UI — yellow themed bubble slotted directly below the
|
||||
// message it advises. ``anchor`` is the DOM element to anchor below;
|
||||
// when null, append at the bottom of messagesEl.
|
||||
// message it advises. ``watch_triggered`` reminders branch off into
|
||||
// the structured ``.msg.watch-result`` card. ``anchor`` is the DOM
|
||||
// element to anchor below; when null, append at the bottom of
|
||||
// messagesEl.
|
||||
function appendReminderBubble(reminders, anchor) {
|
||||
if (!Array.isArray(reminders) || !reminders.length) return;
|
||||
let cursor = anchor;
|
||||
for (let i = 0; i < reminders.length; i++) {
|
||||
const r = reminders[i] || {};
|
||||
const el = document.createElement("div");
|
||||
el.className = "msg user-reminder";
|
||||
el.setAttribute("role", "article");
|
||||
el.setAttribute("data-ts-role", "metacognition");
|
||||
el.setAttribute("aria-label", "metacognition");
|
||||
const body = document.createElement("div");
|
||||
body.className = "msg-body";
|
||||
const labelEl = document.createElement("span");
|
||||
labelEl.className = "msg-user-reminder-label";
|
||||
labelEl.textContent =
|
||||
"metacognition" + (r.type ? " · " + String(r.type) : "");
|
||||
const textEl = document.createElement("span");
|
||||
textEl.className = "msg-user-reminder-text";
|
||||
textEl.textContent = r.text || "";
|
||||
body.appendChild(labelEl);
|
||||
body.appendChild(textEl);
|
||||
el.appendChild(body);
|
||||
const el =
|
||||
r.type === "watch_triggered"
|
||||
? buildWatchResultBubble(r)
|
||||
: buildDefaultReminderBubble(r);
|
||||
if (cursor) {
|
||||
cursor.insertAdjacentElement("afterend", el);
|
||||
cursor = el;
|
||||
@@ -406,12 +461,40 @@
|
||||
_scheduleScroll();
|
||||
}
|
||||
|
||||
// Thin ``.msg.user.system-nudge`` marker rendered as the anchor for
|
||||
// wake-driven reminder bubbles. Replaces the previously-invisible
|
||||
// synthetic empty user turn with a visible-but-subtle DOM element so
|
||||
// the bubble below it lands in the right place even when the wake
|
||||
// fires long after the user's last real message.
|
||||
function appendSystemNudgeMarker() {
|
||||
const el = document.createElement("div");
|
||||
el.className = "msg user system-nudge";
|
||||
el.setAttribute("data-source", "system_nudge");
|
||||
el.setAttribute("aria-label", "system nudge");
|
||||
el.textContent = "system nudge";
|
||||
messagesEl.appendChild(el);
|
||||
return el;
|
||||
}
|
||||
|
||||
// Live SSE for user-channel reminders — anchors below the most
|
||||
// recent user message. On a non-originating tab there may be no
|
||||
// user message rendered yet; we append and the next /history reload
|
||||
// corrects. (Same caveat as the interactive UI; tracked there.)
|
||||
function appendUserReminderLive(reminders) {
|
||||
const userMsgs = messagesEl.querySelectorAll(".msg.user");
|
||||
//
|
||||
// ``source`` widens the live SSE event to carry the wake's
|
||||
// ``"system_nudge"`` tag so the marker renders on every connected
|
||||
// tab — without this, only the originating tab (which sees the
|
||||
// synthesised empty user turn live) would render the wake bubble in
|
||||
// the right place.
|
||||
function appendUserReminderLive(reminders, source) {
|
||||
if (source === "system_nudge") {
|
||||
const marker = appendSystemNudgeMarker();
|
||||
appendReminderBubble(reminders, marker);
|
||||
return;
|
||||
}
|
||||
const userMsgs = messagesEl.querySelectorAll(
|
||||
".msg.user:not(.system-nudge)",
|
||||
);
|
||||
const anchor = userMsgs.length ? userMsgs[userMsgs.length - 1] : null;
|
||||
appendReminderBubble(reminders, anchor);
|
||||
}
|
||||
@@ -866,10 +949,6 @@
|
||||
if (!row) return;
|
||||
const existing = row.querySelector(".coord-tool-row-result");
|
||||
if (existing) existing.remove();
|
||||
// Re-fires (cancel + rerun, error + retry) clear any prior
|
||||
// truncation pill so it doesn't stack on the new result.
|
||||
const existingTrunc = row.querySelector(".coord-tool-truncated");
|
||||
if (existingTrunc) existingTrunc.remove();
|
||||
if (isError) {
|
||||
row.classList.add("error");
|
||||
// Lift the row's error onto the enclosing batch so the left
|
||||
@@ -925,19 +1004,6 @@
|
||||
body.textContent = pretty;
|
||||
block.appendChild(body);
|
||||
row.appendChild(block);
|
||||
// Storage-truncation indicator — sibling pill (not text inside
|
||||
// the result body) so renderers / parsers / copy-as-text paths
|
||||
// see the unmodified output. Same convention as interactive's
|
||||
// .tool-output-truncated; styled by .coord-tool-truncated in
|
||||
// coordinator.css.
|
||||
if (opts && opts.truncated) {
|
||||
const pill = document.createElement("span");
|
||||
pill.className = "coord-tool-truncated";
|
||||
pill.textContent = "… truncated in storage";
|
||||
pill.title =
|
||||
"Full tool output was sent to the model live; only the first 10000 characters are persisted to the conversation row.";
|
||||
row.appendChild(pill);
|
||||
}
|
||||
}
|
||||
|
||||
function _makeActionButton(label, role, kbdHint, ariaLabel) {
|
||||
@@ -2028,9 +2094,11 @@
|
||||
case "user_reminder":
|
||||
// Metacognitive user-channel nudge — render below the most
|
||||
// recent user message as a yellow themed bubble. Same shape
|
||||
// as the interactive UI's case.
|
||||
// as the interactive UI's case. When ``source === "system_nudge"``
|
||||
// (wake-driven), render the thin .msg.user.system-nudge
|
||||
// marker first so the bubble anchors below it.
|
||||
if (Array.isArray(ev.reminders) && ev.reminders.length) {
|
||||
appendUserReminderLive(ev.reminders);
|
||||
appendUserReminderLive(ev.reminders, ev.source || "");
|
||||
}
|
||||
break;
|
||||
case "tool_reminder":
|
||||
@@ -3839,112 +3907,108 @@
|
||||
callOutcomes.set(m.tool_call_id, outcome);
|
||||
});
|
||||
|
||||
// Render an assistant turn's tool_calls as a single batch
|
||||
// construct. Synthesises one batch per assistant turn so a
|
||||
// parallel fan-out (tool_calls.length ≥ 2) reads as one cohesive
|
||||
// dispatch, matching how live SSE renders the same flow via
|
||||
// approve_request / tool_info. Resolved when every call_id has
|
||||
// a matching tool result; otherwise --running (see the
|
||||
// resolvedCallIds rationale above). SSE upgrades --running in
|
||||
// place when it knows more.
|
||||
function renderAssistantToolBatch(m) {
|
||||
const items = m.tool_calls.map((tc) => {
|
||||
const fn = (tc && tc.function) || {};
|
||||
const name = String(fn.name || "tool");
|
||||
const callId = String((tc && tc.id) || "");
|
||||
const argsRaw = String(fn.arguments || "");
|
||||
let parsedArgs = null;
|
||||
try {
|
||||
parsedArgs = JSON.parse(argsRaw || "{}");
|
||||
} catch (_) {
|
||||
/* malformed — fall back to raw string in preview */
|
||||
}
|
||||
if (callId) toolNameByCallId.set(callId, name);
|
||||
const item = synthesizeHistoricalToolCall(
|
||||
name,
|
||||
callId,
|
||||
parsedArgs,
|
||||
argsRaw,
|
||||
);
|
||||
// Server attaches the persisted intent_verdict to each
|
||||
// tc on /history (newest-wins per call_id; LLM upgrade
|
||||
// beats heuristic when both exist). Stamp on the item
|
||||
// under the field name the render path already consumes
|
||||
// (judge_verdict for LLM tier, heuristic_verdict
|
||||
// otherwise) so the verdict pill paints on history rows
|
||||
// without a render-path fork. Also seed the
|
||||
// judgeVerdicts cache so a later live SSE event for the
|
||||
// same call_id reads "already painted" and skips the
|
||||
// rebuild.
|
||||
if (tc && tc.verdict) {
|
||||
if (tc.verdict.tier === "llm") {
|
||||
item.judge_verdict = tc.verdict;
|
||||
} else {
|
||||
item.heuristic_verdict = tc.verdict;
|
||||
}
|
||||
if (callId) _cacheJudgeVerdict(callId, tc.verdict);
|
||||
}
|
||||
// Output-guard finding — surface as the same
|
||||
// "[output guard] ..." chat line the live handler emits
|
||||
// (case "output_warning" above). Stamp on the item so
|
||||
// the post-batch loop below can read + emit; rendering
|
||||
// anchored next to the call gives the operator the same
|
||||
// adjacency they'd see live.
|
||||
if (tc && tc.output_assessment) {
|
||||
item.output_assessment = tc.output_assessment;
|
||||
}
|
||||
// needs_approval is unknown at replay time (the
|
||||
// assistant.tool_calls history payload doesn't persist
|
||||
// the bit). Leave it unset; the upgrade-in-place path
|
||||
// refreshes per-row state via _refreshRowStatus from the
|
||||
// authoritative SSE item when approve_request /
|
||||
// tool_info actually arrives, so we never tag the wrong
|
||||
// row as needing approval.
|
||||
return item;
|
||||
});
|
||||
// Classify the batch as a whole:
|
||||
// - any call_id without an outcome at all → orphan,
|
||||
// render as --running (SSE will upgrade in place)
|
||||
// - any call_id outcome === "denied" → resolved-denied
|
||||
// - else → resolved-approved (a runtime error doesn't
|
||||
// change the approval verdict; the per-row .error class
|
||||
// comes from the tool_result branch below)
|
||||
const outcomes = items.map((it) =>
|
||||
it.call_id ? callOutcomes.get(it.call_id) : "ok",
|
||||
);
|
||||
const allResolved = outcomes.every((o) => o !== undefined);
|
||||
if (!allResolved) {
|
||||
appendToolBatch(items, { running: true });
|
||||
} else if (outcomes.some((o) => o === "denied")) {
|
||||
appendToolBatch(items, { resolved: { approved: false } });
|
||||
} else {
|
||||
appendToolBatch(items, { resolved: { approved: true } });
|
||||
}
|
||||
// Output-guard findings — render each one as a chip
|
||||
// anchored to the .coord-tool-row that tripped the guard
|
||||
// rather than a generic "[output guard]" chat line.
|
||||
// Anchored placement preserves per-call adjacency on
|
||||
// multi-tool batches (live + replay) and the chip's
|
||||
// severity styling makes the visual weight match the
|
||||
// verdict pill on the same row.
|
||||
for (let oi = 0; oi < items.length; oi++) {
|
||||
const oa = items[oi].output_assessment;
|
||||
if (!oa || !oa.risk_level || oa.risk_level === "none") continue;
|
||||
const cid = items[oi].call_id || "";
|
||||
if (!cid) continue;
|
||||
const entry = toolRows.get(cid);
|
||||
if (!entry || !entry.row) continue;
|
||||
_attachOutputWarningChip(entry.row, oa);
|
||||
}
|
||||
}
|
||||
|
||||
(hist.messages || []).forEach((m) => {
|
||||
const role = m.role || "tool";
|
||||
|
||||
// Assistant tool_calls — synthesize one batch construct per
|
||||
// assistant turn so a parallel fan-out (tool_calls.length ≥ 2)
|
||||
// reads as one cohesive dispatch, matching how live SSE
|
||||
// renders the same flow via approve_request / tool_info.
|
||||
// Resolved when every call_id has a matching tool result;
|
||||
// otherwise --running (see the resolvedCallIds rationale
|
||||
// above). SSE upgrades --running in place when it knows
|
||||
// more.
|
||||
if (
|
||||
role === "assistant" &&
|
||||
Array.isArray(m.tool_calls) &&
|
||||
m.tool_calls.length
|
||||
) {
|
||||
const items = m.tool_calls.map((tc) => {
|
||||
const fn = (tc && tc.function) || {};
|
||||
const name = String(fn.name || "tool");
|
||||
const callId = String((tc && tc.id) || "");
|
||||
const argsRaw = String(fn.arguments || "");
|
||||
let parsedArgs = null;
|
||||
try {
|
||||
parsedArgs = JSON.parse(argsRaw || "{}");
|
||||
} catch (_) {
|
||||
/* malformed — fall back to raw string in preview */
|
||||
}
|
||||
if (callId) toolNameByCallId.set(callId, name);
|
||||
const item = synthesizeHistoricalToolCall(
|
||||
name,
|
||||
callId,
|
||||
parsedArgs,
|
||||
argsRaw,
|
||||
);
|
||||
// Server attaches the persisted intent_verdict to each
|
||||
// tc on /history (newest-wins per call_id; LLM upgrade
|
||||
// beats heuristic when both exist). Stamp on the item
|
||||
// under the field name the render path already consumes
|
||||
// (judge_verdict for LLM tier, heuristic_verdict
|
||||
// otherwise) so the verdict pill paints on history rows
|
||||
// without a render-path fork. Also seed the
|
||||
// judgeVerdicts cache so a later live SSE event for the
|
||||
// same call_id reads "already painted" and skips the
|
||||
// rebuild.
|
||||
if (tc && tc.verdict) {
|
||||
if (tc.verdict.tier === "llm") {
|
||||
item.judge_verdict = tc.verdict;
|
||||
} else {
|
||||
item.heuristic_verdict = tc.verdict;
|
||||
}
|
||||
if (callId) _cacheJudgeVerdict(callId, tc.verdict);
|
||||
}
|
||||
// Output-guard finding — surface as the same
|
||||
// "[output guard] ..." chat line the live handler emits
|
||||
// (case "output_warning" above). Stamp on the item so
|
||||
// the post-batch loop below can read + emit; rendering
|
||||
// anchored next to the call gives the operator the same
|
||||
// adjacency they'd see live.
|
||||
if (tc && tc.output_assessment) {
|
||||
item.output_assessment = tc.output_assessment;
|
||||
}
|
||||
// needs_approval is unknown at replay time (the
|
||||
// assistant.tool_calls history payload doesn't persist
|
||||
// the bit). Leave it unset; the upgrade-in-place path
|
||||
// refreshes per-row state via _refreshRowStatus from the
|
||||
// authoritative SSE item when approve_request /
|
||||
// tool_info actually arrives, so we never tag the wrong
|
||||
// row as needing approval.
|
||||
return item;
|
||||
});
|
||||
// Classify the batch as a whole:
|
||||
// - any call_id without an outcome at all → orphan,
|
||||
// render as --running (SSE will upgrade in place)
|
||||
// - any call_id outcome === "denied" → resolved-denied
|
||||
// - else → resolved-approved (a runtime error doesn't
|
||||
// change the approval verdict; the per-row .error class
|
||||
// comes from the tool_result branch below)
|
||||
const outcomes = items.map((it) =>
|
||||
it.call_id ? callOutcomes.get(it.call_id) : "ok",
|
||||
);
|
||||
const allResolved = outcomes.every((o) => o !== undefined);
|
||||
if (!allResolved) {
|
||||
appendToolBatch(items, { running: true });
|
||||
} else if (outcomes.some((o) => o === "denied")) {
|
||||
appendToolBatch(items, { resolved: { approved: false } });
|
||||
} else {
|
||||
appendToolBatch(items, { resolved: { approved: true } });
|
||||
}
|
||||
// Output-guard findings — render each one as a chip
|
||||
// anchored to the .coord-tool-row that tripped the guard
|
||||
// rather than a generic "[output guard]" chat line.
|
||||
// Anchored placement preserves per-call adjacency on
|
||||
// multi-tool batches (live + replay) and the chip's
|
||||
// severity styling makes the visual weight match the
|
||||
// verdict pill on the same row.
|
||||
for (let oi = 0; oi < items.length; oi++) {
|
||||
const oa = items[oi].output_assessment;
|
||||
if (!oa || !oa.risk_level || oa.risk_level === "none") continue;
|
||||
const cid = items[oi].call_id || "";
|
||||
if (!cid) continue;
|
||||
const entry = toolRows.get(cid);
|
||||
if (!entry || !entry.row) continue;
|
||||
_attachOutputWarningChip(entry.row, oa);
|
||||
}
|
||||
}
|
||||
|
||||
// User messages with attachments arrive as multipart list
|
||||
// content (text + image_url/document parts) and may carry an
|
||||
// ``_attachments_meta`` side-channel with display metadata
|
||||
@@ -4007,42 +4071,65 @@
|
||||
const toolName =
|
||||
(callId && toolNameByCallId.get(callId)) || m.tool_name || "tool";
|
||||
const isError = callOutcomes.get(callId) === "error";
|
||||
// Storage truncation surfaces as a sibling pill next to
|
||||
// the result (see _appendResultToRow's opts.truncated
|
||||
// branch) rather than as text inside the result body — a
|
||||
// future "best-effort JSON repair" pass would otherwise
|
||||
// need to strip a marker string before parsing.
|
||||
appendToolResult(toolName, callId, content || "", isError, {
|
||||
truncated: !!m.truncated,
|
||||
});
|
||||
appendToolResult(toolName, callId, content || "", isError);
|
||||
// Tool-channel metacog reminders ride the same _reminders
|
||||
// side-channel as the user channel; surface as a themed
|
||||
// bubble below the .coord-tool-batch construct.
|
||||
if (Array.isArray(m.reminders) && m.reminders.length) {
|
||||
appendToolReminderLive(m.reminders, callId);
|
||||
}
|
||||
// Queued user messages spliced into the last tool-result
|
||||
// envelope of a batch (Seam 1) replay as proper user bubbles
|
||||
// after the tool block. ``decorate_history_messages``
|
||||
// extracts the user_interjection advisory from the persisted
|
||||
// envelope and the wire layer projects it onto
|
||||
// ``m.advisories``; rendering through
|
||||
// ``appendUserMessageWithAttachments`` matches the live shape
|
||||
// a Seam 2/3 message would produce. The walk/filter is
|
||||
// shared via ``replayAdvisoriesAfterTool`` in
|
||||
// ``shared/utils.js`` so coord and interactive can never drift
|
||||
// on advisory-shape filtering.
|
||||
replayAdvisoriesAfterTool(m.advisories, function (text) {
|
||||
appendUserMessageWithAttachments(text, [], { label: "user" });
|
||||
});
|
||||
} else if (role === "assistant") {
|
||||
// Empty content with tool_calls only means the assistant
|
||||
// turn was just tool dispatch — the synthesized tool-call
|
||||
// rows above already cover it; skip the empty bubble.
|
||||
if (!content) return;
|
||||
// Run assistant content through the markdown pipeline
|
||||
// (renderMarkdown + post-render hljs / mermaid / KaTeX) so a
|
||||
// reconnect / page-reload renders the same way a live stream
|
||||
// does. appendText would only escape and dump the raw text —
|
||||
// markdown tables, code fences, math, and links would all
|
||||
// render as literal characters.
|
||||
const el = appendMsg(role, "", { label: role });
|
||||
const body = el.querySelector(".msg-body");
|
||||
if (body && typeof streamingRenderFinalize === "function") {
|
||||
try {
|
||||
streamingRenderFinalize(body, content);
|
||||
} catch (e) {
|
||||
console.warn("coordinator history render failed", e);
|
||||
// Render content BEFORE the tool batch so DOM order matches
|
||||
// chronological order (the model emits text first, then
|
||||
// dispatches tools). Whitespace-only content (e.g. "\n\n"
|
||||
// from a reasoning-parser model that strips <think>…</think>
|
||||
// and leaves only trailing newlines before the tool call) is
|
||||
// treated as empty — without the .trim() guard it would
|
||||
// render a visible-but-empty .msg.assistant card on replay,
|
||||
// which the live stream never showed (the live path didn't
|
||||
// accumulate the trailing whitespace as a visible bubble).
|
||||
if (content && content.trim()) {
|
||||
// Run assistant content through the markdown pipeline
|
||||
// (renderMarkdown + post-render hljs / mermaid / KaTeX) so
|
||||
// a reconnect / page-reload renders the same way a live
|
||||
// stream does. appendText would only escape and dump the
|
||||
// raw text — markdown tables, code fences, math, and links
|
||||
// would all render as literal characters.
|
||||
const el = appendMsg(role, "", { label: role });
|
||||
const body = el.querySelector(".msg-body");
|
||||
if (body && typeof streamingRenderFinalize === "function") {
|
||||
try {
|
||||
streamingRenderFinalize(body, content);
|
||||
} catch (e) {
|
||||
console.warn("coordinator history render failed", e);
|
||||
body.textContent = content;
|
||||
}
|
||||
} else if (body) {
|
||||
body.textContent = content;
|
||||
}
|
||||
} else if (body) {
|
||||
body.textContent = content;
|
||||
}
|
||||
// Tool batch comes after the content card so the DOM matches
|
||||
// the chronological order the model emitted (text → dispatch).
|
||||
// Hoisting this out of the role-agnostic top of the loop —
|
||||
// the prior shape rendered tool_calls before the assistant
|
||||
// text that announced them, putting parallel batches
|
||||
// visually above their narrating message on rehydrate.
|
||||
if (Array.isArray(m.tool_calls) && m.tool_calls.length) {
|
||||
renderAssistantToolBatch(m);
|
||||
}
|
||||
} else {
|
||||
// user / reasoning / system / other roles render as plain
|
||||
@@ -4053,6 +4140,17 @@
|
||||
// text when the message carried attachments — even when the
|
||||
// text portion is empty (image-only sends).
|
||||
if (role === "user") {
|
||||
const isSystemNudge = m.source === "system_nudge";
|
||||
if (isSystemNudge) {
|
||||
// Wake-driven empty user turn: render the thin marker
|
||||
// (replaces the previously-skipped synthetic empty
|
||||
// bubble) and anchor reminder bubbles below it.
|
||||
const marker = appendSystemNudgeMarker();
|
||||
if (Array.isArray(m.reminders) && m.reminders.length) {
|
||||
appendReminderBubble(m.reminders, marker);
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (!content && userAttachments.length === 0) return;
|
||||
appendUserMessageWithAttachments(content, userAttachments, {
|
||||
label: role,
|
||||
|
||||
@@ -839,6 +839,179 @@ function _renderGovSkills(items) {
|
||||
});
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SKILL.md paste auto-fill — sniff frontmatter on paste, hit the parse
|
||||
// endpoint, and populate the form so users don't have to retype name /
|
||||
// description / tags / etc. when importing an Anthropic-style skill.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// Trigger on any paste whose first non-whitespace bytes look like an opening
|
||||
// YAML frontmatter delimiter — restrictive enough to ignore normal markdown
|
||||
// pastes, permissive enough to catch CRLF and trailing-space variants.
|
||||
var _SKILL_FRONTMATTER_RE = /^---\s*\r?\n/;
|
||||
|
||||
var _SKILL_FIELD_MAP = {
|
||||
name: "ctm-name",
|
||||
description: "skill-description",
|
||||
tags: "skill-tags",
|
||||
author: "skill-author",
|
||||
version: "skill-version",
|
||||
license: "skill-license",
|
||||
compatibility: "skill-compatibility",
|
||||
allowed_tools: "csk-allowed-tools",
|
||||
};
|
||||
|
||||
// Inflight paste-parse fetch — referenced from hideCreateTemplateModal so a
|
||||
// modal close cancels the request, and from _handleSkillContentPaste so a
|
||||
// fresh paste supersedes the previous one. Acts as a generation token: any
|
||||
// callback that observes _ctmPasteController != its captured controller knows
|
||||
// the modal moved on and must not touch the DOM.
|
||||
var _ctmPasteController = null;
|
||||
|
||||
// Returns "filled" if we set the value, "skipped" if the field was already
|
||||
// non-empty (we don't clobber user input), or "absent" if we couldn't find or
|
||||
// match the option. Tracking this lets the caller report what actually
|
||||
// happened so the user knows whether their pre-typed values survived.
|
||||
function _setSkillFormField(id, value) {
|
||||
var el = document.getElementById(id);
|
||||
if (!el) return "absent";
|
||||
if (el.value && String(el.value).trim()) return "skipped";
|
||||
if (el.tagName === "SELECT") {
|
||||
// License is a fixed option list — only set the value if it matches an
|
||||
// option. Custom licenses fall through to the default "— not specified —"
|
||||
// and the user can edit manually.
|
||||
for (var i = 0; i < el.options.length; i++) {
|
||||
if (el.options[i].value === value) {
|
||||
el.value = value;
|
||||
return "filled";
|
||||
}
|
||||
}
|
||||
return "absent";
|
||||
}
|
||||
el.value = value;
|
||||
return "filled";
|
||||
}
|
||||
|
||||
function _applyParsedSkill(parsed, contentTextarea, fieldMap) {
|
||||
// The textarea is the explicit paste target — replacing its full content
|
||||
// matches the user's mental model ("I pasted a SKILL.md, the body should
|
||||
// become the content"). Side metadata fields use the non-destructive
|
||||
// _setSkillFormField rule below so half-typed values aren't lost.
|
||||
contentTextarea.value = parsed.content || "";
|
||||
contentTextarea.dispatchEvent(new Event("input", { bubbles: true }));
|
||||
|
||||
var filled = 0;
|
||||
var skipped = 0;
|
||||
function _apply(id, value) {
|
||||
var outcome = _setSkillFormField(id, value);
|
||||
if (outcome === "filled") filled++;
|
||||
else if (outcome === "skipped") skipped++;
|
||||
}
|
||||
|
||||
if (parsed.name) _apply(fieldMap.name, parsed.name);
|
||||
if (parsed.description) _apply(fieldMap.description, parsed.description);
|
||||
if (parsed.tags && parsed.tags.length)
|
||||
_apply(fieldMap.tags, parsed.tags.join(", "));
|
||||
if (parsed.author) _apply(fieldMap.author, parsed.author);
|
||||
if (parsed.version) _apply(fieldMap.version, parsed.version);
|
||||
if (parsed.license) _apply(fieldMap.license, parsed.license);
|
||||
if (parsed.compatibility)
|
||||
_apply(fieldMap.compatibility, parsed.compatibility);
|
||||
if (parsed.allowed_tools && parsed.allowed_tools.length)
|
||||
_apply(fieldMap.allowed_tools, parsed.allowed_tools.join(", "));
|
||||
|
||||
return { filled: filled, skipped: skipped };
|
||||
}
|
||||
|
||||
function _setSkillPasteHintBusy(busy) {
|
||||
var hint = document.getElementById("ctm-paste-hint");
|
||||
if (!hint) return;
|
||||
var rest = hint.querySelector(".skill-paste-hint-rest");
|
||||
var busyEl = hint.querySelector(".skill-paste-hint-busy");
|
||||
if (rest) rest.style.display = busy ? "none" : "";
|
||||
if (busyEl) busyEl.style.display = busy ? "" : "none";
|
||||
}
|
||||
|
||||
function _handleSkillContentPaste(event, fieldMap) {
|
||||
var clipboard = event.clipboardData || window.clipboardData;
|
||||
if (!clipboard) return;
|
||||
var text = clipboard.getData("text/plain");
|
||||
if (!text || !_SKILL_FRONTMATTER_RE.test(text)) return;
|
||||
|
||||
event.preventDefault();
|
||||
var textarea = event.target;
|
||||
|
||||
// Cancel any prior paste fetch — a fresh paste supersedes whatever was in
|
||||
// flight. The previous handler's callbacks see _ctmPasteController !=
|
||||
// their captured controller and bail before touching the DOM.
|
||||
if (_ctmPasteController) _ctmPasteController.abort();
|
||||
var controller = new AbortController();
|
||||
_ctmPasteController = controller;
|
||||
|
||||
// Optimistic paint — drop the raw text into the textarea immediately so the
|
||||
// user sees their paste landed, then disable the field and flip the hint
|
||||
// line into a "Parsing..." state. On a slow network the round-trip can
|
||||
// stretch past 400ms; without a visible state the user thinks nothing
|
||||
// happened and re-pastes (or hits Create with empty fields).
|
||||
textarea.value = text;
|
||||
textarea.disabled = true;
|
||||
textarea.setAttribute("aria-busy", "true");
|
||||
textarea.dispatchEvent(new Event("input", { bubbles: true }));
|
||||
_setSkillPasteHintBusy(true);
|
||||
|
||||
function _isCurrent() {
|
||||
return _ctmPasteController === controller;
|
||||
}
|
||||
|
||||
authFetch("/v1/api/admin/skills/parse", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ raw: text }),
|
||||
signal: controller.signal,
|
||||
})
|
||||
.then(function (r) {
|
||||
return r.json().then(function (d) {
|
||||
return { ok: r.ok, data: d };
|
||||
});
|
||||
})
|
||||
.then(function (res) {
|
||||
if (!_isCurrent()) return;
|
||||
if (res.ok) {
|
||||
var counts = _applyParsedSkill(res.data, textarea, fieldMap);
|
||||
var msg = "Populated from SKILL.md";
|
||||
if (counts.skipped) {
|
||||
msg += " (" + counts.filled + " set, " + counts.skipped + " kept)";
|
||||
}
|
||||
showToast(msg);
|
||||
} else {
|
||||
// Frontmatter looked plausible but the parser rejected it (missing
|
||||
// name, malformed YAML beyond the retry, etc). The raw text is
|
||||
// already in the textarea from the optimistic paint above so the
|
||||
// user can fix the YAML in place.
|
||||
showToast(
|
||||
"Couldn't parse SKILL.md: " +
|
||||
((res.data && res.data.error) || "unknown error"),
|
||||
"error",
|
||||
);
|
||||
}
|
||||
})
|
||||
.catch(function (err) {
|
||||
if (!_isCurrent()) return;
|
||||
// AbortError fires when the modal closed or a fresher paste superseded
|
||||
// this one — silent, the new lifecycle owns the UI.
|
||||
if (err && err.name === "AbortError") return;
|
||||
showToast("Network error — pasted as plain text", "error");
|
||||
})
|
||||
.finally(function () {
|
||||
if (!_isCurrent()) return;
|
||||
_ctmPasteController = null;
|
||||
textarea.disabled = false;
|
||||
textarea.removeAttribute("aria-busy");
|
||||
_setSkillPasteHintBusy(false);
|
||||
textarea.focus();
|
||||
});
|
||||
}
|
||||
|
||||
function _detectTemplateVars(content) {
|
||||
var matches = content.match(/\{\{(\w+)\}\}/g) || [];
|
||||
var seen = {};
|
||||
@@ -874,11 +1047,15 @@ function showCreateTemplateModal() {
|
||||
document.getElementById("skill-license").value = "";
|
||||
document.getElementById("skill-compatibility").value = "";
|
||||
document.getElementById("skill-activation").value = "named";
|
||||
document.getElementById("ctm-content").value = "";
|
||||
var ctmContent = document.getElementById("ctm-content");
|
||||
ctmContent.value = "";
|
||||
document.getElementById("ctm-variables").textContent = "(none)";
|
||||
document.getElementById("ctm-content").oninput = function () {
|
||||
ctmContent.oninput = function () {
|
||||
_updateVarsDisplay("ctm-content", "ctm-variables");
|
||||
};
|
||||
ctmContent.onpaste = function (event) {
|
||||
_handleSkillContentPaste(event, _SKILL_FIELD_MAP);
|
||||
};
|
||||
document.getElementById("ctm-default").checked = false;
|
||||
// Session config fields
|
||||
document.getElementById("csk-model").value = "";
|
||||
@@ -907,6 +1084,23 @@ function showCreateTemplateModal() {
|
||||
}
|
||||
|
||||
function hideCreateTemplateModal() {
|
||||
// Cancel any inflight paste-parse so a late response can't reach into a
|
||||
// closed (or freshly reopened) modal and clobber state. AbortController
|
||||
// also short-circuits the .then chain — see _handleSkillContentPaste.
|
||||
// After abort, the handler's .catch/.finally bail via _isCurrent() before
|
||||
// resetting the textarea, so we proactively restore the paste-induced
|
||||
// visible state here. Otherwise reopening would land on a disabled
|
||||
// textarea stuck on "Parsing…".
|
||||
if (_ctmPasteController) {
|
||||
_ctmPasteController.abort();
|
||||
_ctmPasteController = null;
|
||||
var ctmContent = document.getElementById("ctm-content");
|
||||
if (ctmContent) {
|
||||
ctmContent.disabled = false;
|
||||
ctmContent.removeAttribute("aria-busy");
|
||||
}
|
||||
_setSkillPasteHintBusy(false);
|
||||
}
|
||||
document.getElementById("create-template-overlay").style.display = "none";
|
||||
_ctmTrapHandler = _removeTrap(_ctmTrapHandler);
|
||||
if (_ctmTriggerEl && _ctmTriggerEl.focus) {
|
||||
|
||||
@@ -3094,17 +3094,33 @@
|
||||
<h3 class="skill-spec-heading">
|
||||
Skill Content
|
||||
<span class="label-hint"
|
||||
>system message — {{model}}, {{ws_id}},
|
||||
{{node_id}}</span
|
||||
>available: {{model}}, {{ws_id}}, {{node_id}}</span
|
||||
>
|
||||
</h3>
|
||||
<div
|
||||
class="skill-paste-hint"
|
||||
id="ctm-paste-hint"
|
||||
aria-live="polite"
|
||||
>
|
||||
<span class="skill-paste-hint-rest">
|
||||
Tip: paste a SKILL.md (with <code>---</code> frontmatter) to
|
||||
auto-fill the form.
|
||||
</span>
|
||||
<span
|
||||
class="skill-paste-hint-busy"
|
||||
style="display: none"
|
||||
>
|
||||
Parsing SKILL.md…
|
||||
</span>
|
||||
</div>
|
||||
<textarea
|
||||
id="ctm-content"
|
||||
class="skill-content-area"
|
||||
aria-describedby="ctm-paste-hint"
|
||||
placeholder="You are a code reviewer using {{model}}..."
|
||||
></textarea>
|
||||
<div class="skill-vars-row">
|
||||
<span class="skill-vars-label">Variables</span>
|
||||
<span class="skill-vars-label">Used</span>
|
||||
<div
|
||||
id="ctm-variables"
|
||||
class="skill-vars-display label-hint"
|
||||
@@ -3379,12 +3395,12 @@
|
||||
<h3 class="skill-spec-heading">
|
||||
Skill Content
|
||||
<span class="label-hint"
|
||||
>{{model}}, {{ws_id}}, {{node_id}}</span
|
||||
>available: {{model}}, {{ws_id}}, {{node_id}}</span
|
||||
>
|
||||
</h3>
|
||||
<textarea id="etm-content" class="skill-content-area"></textarea>
|
||||
<div class="skill-vars-row">
|
||||
<span class="skill-vars-label">Variables</span>
|
||||
<span class="skill-vars-label">Used</span>
|
||||
<div
|
||||
id="etm-variables"
|
||||
class="skill-vars-display label-hint"
|
||||
@@ -3661,6 +3677,114 @@
|
||||
placeholder="Authorization: Bearer ..."
|
||||
></textarea>
|
||||
</div>
|
||||
<fieldset
|
||||
id="mcp-auth-section"
|
||||
style="
|
||||
border: 1px solid var(--border);
|
||||
padding: 10px 12px;
|
||||
margin-top: 12px;
|
||||
"
|
||||
>
|
||||
<legend style="font-size: 12px; padding: 0 6px">
|
||||
Multitenant Authorization
|
||||
</legend>
|
||||
<label
|
||||
style="display: block; margin: 4px 0; font-weight: 400; font-size: 13px"
|
||||
>
|
||||
<input
|
||||
type="radio"
|
||||
name="mcp-auth-type"
|
||||
id="mcp-auth-none"
|
||||
value="none"
|
||||
onchange="toggleMcpAuthFields()"
|
||||
style="margin-right: 6px"
|
||||
/>No authorization
|
||||
</label>
|
||||
<label
|
||||
style="display: block; margin: 4px 0; font-weight: 400; font-size: 13px"
|
||||
>
|
||||
<input
|
||||
type="radio"
|
||||
name="mcp-auth-type"
|
||||
id="mcp-auth-static"
|
||||
value="static"
|
||||
onchange="toggleMcpAuthFields()"
|
||||
checked
|
||||
style="margin-right: 6px"
|
||||
/>Static headers (single shared identity)
|
||||
</label>
|
||||
<label
|
||||
style="display: block; margin: 4px 0; font-weight: 400; font-size: 13px"
|
||||
>
|
||||
<input
|
||||
type="radio"
|
||||
name="mcp-auth-type"
|
||||
id="mcp-auth-oauth"
|
||||
value="oauth_user"
|
||||
onchange="toggleMcpAuthFields()"
|
||||
style="margin-right: 6px"
|
||||
/>Per-user OAuth 2.1 (recommended)
|
||||
</label>
|
||||
<div id="mcp-oauth-fields" style="display: none; margin-top: 8px">
|
||||
<label for="mcp-oauth-as-url"
|
||||
>Authorization Server URL
|
||||
<span style="font-weight: 400; text-transform: none"
|
||||
>(optional — override discovery when MCP server URL is not the
|
||||
OAuth issuer)</span
|
||||
></label
|
||||
>
|
||||
<input
|
||||
type="text"
|
||||
id="mcp-oauth-as-url"
|
||||
placeholder="https://auth.example.com"
|
||||
/>
|
||||
<label for="mcp-oauth-registration">Client Registration</label>
|
||||
<select id="mcp-oauth-registration">
|
||||
<option value="preregistered">preregistered</option>
|
||||
<option value="dcr">dcr (Dynamic Client Registration)</option>
|
||||
</select>
|
||||
<label for="mcp-oauth-client-id">Client ID</label>
|
||||
<input
|
||||
type="text"
|
||||
id="mcp-oauth-client-id"
|
||||
placeholder="(operator-issued client_id)"
|
||||
/>
|
||||
<label for="mcp-oauth-client-secret"
|
||||
>Client Secret
|
||||
<span style="font-weight: 400; text-transform: none"
|
||||
>(write-only, never displayed)</span
|
||||
></label
|
||||
>
|
||||
<input
|
||||
type="password"
|
||||
id="mcp-oauth-client-secret"
|
||||
placeholder="***"
|
||||
autocomplete="off"
|
||||
/>
|
||||
<label for="mcp-oauth-scopes"
|
||||
>Scopes
|
||||
<span style="font-weight: 400; text-transform: none"
|
||||
>(space-separated)</span
|
||||
></label
|
||||
>
|
||||
<input
|
||||
type="text"
|
||||
id="mcp-oauth-scopes"
|
||||
placeholder="openid profile"
|
||||
/>
|
||||
<label for="mcp-oauth-audience"
|
||||
>Audience
|
||||
<span style="font-weight: 400; text-transform: none"
|
||||
>(auto-populated from URL)</span
|
||||
></label
|
||||
>
|
||||
<input
|
||||
type="text"
|
||||
id="mcp-oauth-audience"
|
||||
placeholder="https://mcp.example.com"
|
||||
/>
|
||||
</div>
|
||||
</fieldset>
|
||||
<div style="display: flex; gap: 20px; margin-top: 14px">
|
||||
<label style="margin: 0; font-size: 12px; color: var(--fg-dim)"
|
||||
><input
|
||||
|
||||
@@ -908,11 +908,14 @@
|
||||
}
|
||||
|
||||
/* ==========================================================================
|
||||
Toast override — position above cluster status bar
|
||||
Toast override — position above cluster status bar AND above admin modal
|
||||
overlays (which sit at z-index 600). Without this, toasts fired while a
|
||||
modal is open — e.g. paste-to-fill on the Create Skill modal — render
|
||||
behind the dimmed backdrop and never reach the user.
|
||||
========================================================================== */
|
||||
#toast {
|
||||
bottom: 56px;
|
||||
z-index: 200;
|
||||
z-index: 700;
|
||||
color: var(--fg-bright);
|
||||
border-color: var(--border-strong);
|
||||
}
|
||||
@@ -1909,6 +1912,23 @@ h3.skill-spec-heading {
|
||||
opacity: 1;
|
||||
}
|
||||
|
||||
/* Tip line above the Skill Content textarea announcing the paste-to-fill
|
||||
affordance. Sized to match the dim hint on the heading rather than the
|
||||
default body text, so it doesn't out-shout the rest of the modal. */
|
||||
.skill-paste-hint {
|
||||
font-size: 11px;
|
||||
color: var(--fg-dim);
|
||||
margin: -2px 0 6px;
|
||||
line-height: 1.5;
|
||||
}
|
||||
.skill-paste-hint code {
|
||||
font-size: 10.5px;
|
||||
padding: 0 4px;
|
||||
background: var(--code-bg);
|
||||
border-radius: 2px;
|
||||
color: var(--fg);
|
||||
}
|
||||
|
||||
.skill-spec-section-content {
|
||||
flex: 1;
|
||||
display: flex;
|
||||
|
||||
@@ -34,6 +34,24 @@ Action-name conventions (non-exhaustive — grep
|
||||
``prompt_policy``, ``setting``, ``token``,
|
||||
``conversation``, ``memory``, ``org``.
|
||||
|
||||
mcp_server.oauth.* OAuth-MCP delegated authorization events
|
||||
(``mcp_server.oauth.client_secret_set`` from
|
||||
admin handlers when the operator stores or
|
||||
clears a per-server OAuth client secret;
|
||||
``mcp_server.oauth.token_decrypt_failure`` from
|
||||
``MCPTokenStore.get_user_token`` when no
|
||||
installed key can decrypt a stored token;
|
||||
``mcp_server.oauth.consent_started`` /
|
||||
``.consent_completed`` / ``.consent_failed`` for
|
||||
the per-user authorization-flow handlers;
|
||||
``mcp_server.oauth.token_refreshed`` /
|
||||
``.token_revoked`` from
|
||||
``get_user_access_token`` when the
|
||||
refresh-grant exchange runs;
|
||||
``mcp_server.oauth.dcr_registered`` when a
|
||||
client_id was provisioned dynamically against
|
||||
an AS that exposes ``registration_endpoint``).
|
||||
|
||||
When adding a new namespace, prefer extending an existing prefix over
|
||||
inventing a synonym (e.g. ``mcp_server.refresh`` rather than
|
||||
``mcp.refresh`` — ``mcp_server.*`` is already the established prefix).
|
||||
|
||||
+136
-57
@@ -15,6 +15,7 @@ always accessible without authentication.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
@@ -35,6 +36,17 @@ if TYPE_CHECKING:
|
||||
from turnstone.core.oidc import OIDCConfig
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.oidc import (
|
||||
OIDC_STATE_TTL_SECONDS,
|
||||
OIDCError,
|
||||
OIDCKeyNotFoundError,
|
||||
build_authorize_url,
|
||||
exchange_code,
|
||||
fetch_jwks,
|
||||
generate_pkce_verifier,
|
||||
provision_oidc_user,
|
||||
validate_id_token,
|
||||
)
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
@@ -521,6 +533,13 @@ def required_scope(method: str, path: str) -> str:
|
||||
if method == "POST" and normalized in APPROVE_PATHS:
|
||||
return "approve"
|
||||
|
||||
# Path-keyed internal admin actions: /api/_internal/mcp-{refresh,reconnect}/{name}
|
||||
if method == "POST" and (
|
||||
normalized.startswith("/api/_internal/mcp-refresh/")
|
||||
or normalized.startswith("/api/_internal/mcp-reconnect/")
|
||||
):
|
||||
return "approve"
|
||||
|
||||
# Write endpoints
|
||||
if method == "POST" and normalized in WRITE_PATHS:
|
||||
return "write"
|
||||
@@ -582,6 +601,10 @@ def required_scope(method: str, path: str) -> str:
|
||||
if proxied:
|
||||
if proxied in APPROVE_PATHS:
|
||||
return "approve"
|
||||
if proxied.startswith("/api/_internal/mcp-refresh/") or proxied.startswith(
|
||||
"/api/_internal/mcp-reconnect/"
|
||||
):
|
||||
return "approve"
|
||||
if proxied in WRITE_PATHS:
|
||||
return "write"
|
||||
# Parametric workstream sub-resource mutations
|
||||
@@ -1122,8 +1145,7 @@ async def handle_auth_status(request: Request) -> Response:
|
||||
has_users = False
|
||||
if storage is not None:
|
||||
try:
|
||||
users = storage.list_users()
|
||||
has_users = len(users) > 0
|
||||
has_users = await asyncio.to_thread(storage.count_users) > 0
|
||||
except Exception:
|
||||
log.warning("Failed to check user existence for auth status", exc_info=True)
|
||||
|
||||
@@ -1385,17 +1407,16 @@ async def handle_auth_refresh(request: Request, audience: str) -> Response:
|
||||
return response
|
||||
|
||||
|
||||
def _build_oidc_redirect_uri(request: Request, oidc_config: OIDCConfig) -> str:
|
||||
"""Build the OIDC callback redirect URI.
|
||||
def _build_oidc_redirect_uri(oidc_config: OIDCConfig) -> str:
|
||||
"""Build the OIDC callback redirect URI from the pinned ``redirect_base``.
|
||||
|
||||
Uses ``redirect_base`` from OIDC config when set (recommended for
|
||||
reverse-proxy deployments), otherwise falls back to the request Host header.
|
||||
``initialize_oidc_state`` refuses to enable OIDC unless ``redirect_base``
|
||||
is set, so any caller reaching this point may assume it is non-empty.
|
||||
A previous Host-header fallback was removed because a permissive front
|
||||
proxy could let an attacker spoof ``Host`` and mint an authorize URL
|
||||
pointing to an attacker-controlled callback origin.
|
||||
"""
|
||||
if oidc_config.redirect_base:
|
||||
return f"{oidc_config.redirect_base}/v1/api/auth/oidc/callback"
|
||||
scheme = "https" if is_secure_request(dict(request.headers), request.url.scheme) else "http"
|
||||
host = request.headers.get("host", "localhost")
|
||||
return f"{scheme}://{host}/v1/api/auth/oidc/callback"
|
||||
return f"{oidc_config.redirect_base}/v1/api/auth/oidc/callback"
|
||||
|
||||
|
||||
async def handle_oidc_authorize(request: Request, audience: str) -> Response:
|
||||
@@ -1421,31 +1442,73 @@ async def handle_oidc_authorize(request: Request, audience: str) -> Response:
|
||||
|
||||
# Require setup to be complete before allowing OIDC login
|
||||
try:
|
||||
users = storage.list_users()
|
||||
users_count = await asyncio.to_thread(storage.count_users)
|
||||
except Exception:
|
||||
return JSONResponse({"error": "Storage unavailable"}, status_code=503)
|
||||
if not users:
|
||||
if users_count == 0:
|
||||
return JSONResponse(
|
||||
{"error": "Initial setup required before OIDC login"},
|
||||
status_code=403,
|
||||
)
|
||||
|
||||
from turnstone.core.oidc import build_authorize_url, generate_pkce_pair
|
||||
|
||||
state = secrets.token_urlsafe(32)
|
||||
nonce = secrets.token_urlsafe(32)
|
||||
code_verifier, _code_challenge = generate_pkce_pair()
|
||||
code_verifier = generate_pkce_verifier()
|
||||
|
||||
# Store pending state in database
|
||||
storage.create_oidc_pending_state(state, nonce, code_verifier, audience)
|
||||
await asyncio.to_thread(
|
||||
storage.create_oidc_pending_state, state, nonce, code_verifier, audience
|
||||
)
|
||||
|
||||
# Build redirect URI (pinned by TURNSTONE_OIDC_REDIRECT_BASE when set)
|
||||
redirect_uri = _build_oidc_redirect_uri(request, oidc_config)
|
||||
# Build redirect URI (pinned by TURNSTONE_OIDC_REDIRECT_BASE)
|
||||
redirect_uri = _build_oidc_redirect_uri(oidc_config)
|
||||
|
||||
url = build_authorize_url(oidc_config, redirect_uri, state, nonce, code_verifier)
|
||||
return RedirectResponse(url, status_code=302)
|
||||
|
||||
|
||||
_OIDC_STATE_CLEANUP_INTERVAL_S = 60.0
|
||||
|
||||
|
||||
def _resolve_kid(jwks: dict[str, Any] | None, kid: str | None) -> bool:
|
||||
"""Return True iff *jwks* contains the supplied *kid* (or has a single key when kid is None)."""
|
||||
if jwks is None:
|
||||
return False
|
||||
keys = jwks.get("keys", [])
|
||||
if not isinstance(keys, list):
|
||||
return False
|
||||
if kid is None:
|
||||
return len(keys) == 1
|
||||
return any(isinstance(k, dict) and k.get("kid") == kid for k in keys)
|
||||
|
||||
|
||||
async def _refetch_jwks_locked(
|
||||
request: Request, jwks_uri: str, kid: str | None
|
||||
) -> dict[str, Any] | None:
|
||||
"""Acquire the per-app JWKS refetch lock, re-check the cache, and refetch on miss.
|
||||
|
||||
Returns the JWKS dict on success or ``None`` if the fetch failed. When
|
||||
another concurrent caller already refreshed the cache to include *kid*
|
||||
we return the existing snapshot without issuing another network call.
|
||||
"""
|
||||
lock = getattr(request.app.state, "jwks_refetch_lock", None)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
request.app.state.jwks_refetch_lock = lock
|
||||
http_client = getattr(request.app.state, "oidc_http_client", None)
|
||||
async with lock:
|
||||
cached: dict[str, Any] | None = getattr(request.app.state, "jwks_data", None)
|
||||
if _resolve_kid(cached, kid):
|
||||
return cached
|
||||
try:
|
||||
fresh = await fetch_jwks(jwks_uri, client=http_client)
|
||||
except Exception:
|
||||
log.warning("JWKS fetch failed from %s", jwks_uri, exc_info=True)
|
||||
return None
|
||||
request.app.state.jwks_data = fresh
|
||||
return fresh
|
||||
|
||||
|
||||
async def handle_oidc_callback(request: Request, audience: str) -> Response:
|
||||
"""Shared ``GET /api/auth/oidc/callback`` handler — exchange code, provision user, issue JWT."""
|
||||
from starlette.responses import JSONResponse, RedirectResponse
|
||||
@@ -1468,11 +1531,17 @@ async def handle_oidc_callback(request: Request, audience: str) -> Response:
|
||||
if not ip_ok:
|
||||
return RedirectResponse("/?oidc_error=Too+many+login+attempts", status_code=302)
|
||||
|
||||
# Lazy cleanup of expired pending states
|
||||
try:
|
||||
storage.cleanup_expired_oidc_states(300)
|
||||
except Exception:
|
||||
log.debug("OIDC state cleanup failed", exc_info=True)
|
||||
# Lazy cleanup of expired pending states — gated to once per
|
||||
# _OIDC_STATE_CLEANUP_INTERVAL_S so a high-rate callback path
|
||||
# doesn't fire a full DELETE per login.
|
||||
last_cleanup = getattr(request.app.state, "oidc_last_cleanup_monotonic", 0.0)
|
||||
now_mono = time.monotonic()
|
||||
if now_mono - last_cleanup > _OIDC_STATE_CLEANUP_INTERVAL_S:
|
||||
request.app.state.oidc_last_cleanup_monotonic = now_mono
|
||||
try:
|
||||
await asyncio.to_thread(storage.cleanup_expired_oidc_states, OIDC_STATE_TTL_SECONDS)
|
||||
except Exception:
|
||||
log.debug("OIDC state cleanup failed", exc_info=True)
|
||||
|
||||
def _record_oidc_failure() -> None:
|
||||
if login_limiter is not None:
|
||||
@@ -1487,68 +1556,75 @@ async def handle_oidc_callback(request: Request, audience: str) -> Response:
|
||||
|
||||
# Validate state
|
||||
state = request.query_params.get("state", "")
|
||||
pending = storage.pop_oidc_pending_state(state, max_age_seconds=300)
|
||||
pending = await asyncio.to_thread(storage.pop_oidc_pending_state, state, OIDC_STATE_TTL_SECONDS)
|
||||
if not pending:
|
||||
_record_oidc_failure()
|
||||
return RedirectResponse("/?oidc_error=Login+session+expired", status_code=302)
|
||||
|
||||
# Build redirect URI (must match what was sent in authorize)
|
||||
redirect_uri = _build_oidc_redirect_uri(request, oidc_config)
|
||||
redirect_uri = _build_oidc_redirect_uri(oidc_config)
|
||||
|
||||
try:
|
||||
from turnstone.core.oidc import (
|
||||
OIDCError,
|
||||
exchange_code,
|
||||
fetch_jwks,
|
||||
provision_oidc_user,
|
||||
validate_id_token,
|
||||
)
|
||||
|
||||
# Exchange code for tokens
|
||||
code = request.query_params.get("code", "")
|
||||
tokens = await exchange_code(oidc_config, code, redirect_uri, pending["code_verifier"])
|
||||
http_client = getattr(request.app.state, "oidc_http_client", None)
|
||||
tokens = await exchange_code(
|
||||
oidc_config,
|
||||
code,
|
||||
redirect_uri,
|
||||
pending["code_verifier"],
|
||||
client=http_client,
|
||||
)
|
||||
id_token = tokens.get("id_token")
|
||||
if not isinstance(id_token, str) or not id_token:
|
||||
raise OIDCError("Token endpoint response missing id_token")
|
||||
|
||||
# Validate ID token against cached JWKS keys (no I/O).
|
||||
# On unknown kid, refresh JWKS once (async) for key rotation.
|
||||
jwks_data: dict[str, Any] | None = getattr(request.app.state, "jwks_data", None)
|
||||
if jwks_data is None and oidc_config.jwks_uri:
|
||||
# Lazy fetch: JWKS may have failed at startup but IdP recovered
|
||||
try:
|
||||
jwks_data = await fetch_jwks(oidc_config.jwks_uri)
|
||||
request.app.state.jwks_data = jwks_data
|
||||
except OIDCError:
|
||||
log.warning("JWKS fetch failed from %s", oidc_config.jwks_uri, exc_info=True)
|
||||
# Lazy fetch: JWKS may have failed at startup but IdP recovered.
|
||||
# Coalesced via the per-app refetch lock so concurrent callbacks
|
||||
# don't fan out N parallel JWKS GETs.
|
||||
jwks_data = await _refetch_jwks_locked(request, oidc_config.jwks_uri, kid=None)
|
||||
if jwks_data is None:
|
||||
return RedirectResponse("/?oidc_error=OIDC+temporarily+unavailable", status_code=302)
|
||||
|
||||
try:
|
||||
id_claims = validate_id_token(
|
||||
tokens["id_token"],
|
||||
id_token,
|
||||
jwks_data,
|
||||
oidc_config,
|
||||
pending["nonce"],
|
||||
)
|
||||
except OIDCError as first_err:
|
||||
if "not found in JWKS" not in str(first_err):
|
||||
raise
|
||||
# Key rotation: re-fetch JWKS and retry once.
|
||||
except OIDCKeyNotFoundError:
|
||||
# Key rotation: re-fetch JWKS once (coalesced) and retry.
|
||||
import jwt as _jwt
|
||||
|
||||
try:
|
||||
_kid = _jwt.get_unverified_header(id_token).get("kid")
|
||||
except Exception:
|
||||
_kid = None
|
||||
log.info("JWKS key not found — refreshing for possible key rotation")
|
||||
jwks_data = await fetch_jwks(oidc_config.jwks_uri)
|
||||
request.app.state.jwks_data = jwks_data
|
||||
refreshed = await _refetch_jwks_locked(request, oidc_config.jwks_uri, kid=_kid)
|
||||
if refreshed is None:
|
||||
raise
|
||||
jwks_data = refreshed
|
||||
id_claims = validate_id_token(
|
||||
tokens["id_token"],
|
||||
id_token,
|
||||
jwks_data,
|
||||
oidc_config,
|
||||
pending["nonce"],
|
||||
)
|
||||
|
||||
# Verify setup is complete
|
||||
users = storage.list_users()
|
||||
if not users:
|
||||
users_count = await asyncio.to_thread(storage.count_users)
|
||||
if users_count == 0:
|
||||
return RedirectResponse("/?oidc_error=Initial+setup+required", status_code=302)
|
||||
|
||||
# Provision or match user
|
||||
user = provision_oidc_user(storage, oidc_config, id_claims)
|
||||
# Provision or match user (chains apply_role_mapping +
|
||||
# potentially a write to user_roles — wrap as a unit).
|
||||
user = await asyncio.to_thread(provision_oidc_user, storage, oidc_config, id_claims)
|
||||
|
||||
except OIDCError as exc:
|
||||
log.warning("OIDC callback failed: %s", exc)
|
||||
@@ -1560,13 +1636,16 @@ async def handle_oidc_callback(request: Request, audience: str) -> Response:
|
||||
return RedirectResponse("/?oidc_error=Authentication+failed", status_code=302)
|
||||
|
||||
# Load permissions and issue Turnstone JWT
|
||||
perms = _load_user_permissions(storage, user["user_id"])
|
||||
perms = await asyncio.to_thread(_load_user_permissions, storage, user["user_id"])
|
||||
scopes = _permissions_to_scopes(perms)
|
||||
jwt_token = ""
|
||||
if jwt_secret:
|
||||
# Use the audience stored during authorize (not the handler param)
|
||||
# to bind the JWT to the service that initiated the flow
|
||||
jwt_audience = pending.get("audience", audience)
|
||||
# Bind the JWT to the audience stored at /authorize so it cannot be
|
||||
# silently re-targeted at the callback's handler-supplied audience.
|
||||
# ``pending["audience"]`` is always present (NOT NULL TEXT column,
|
||||
# filled by create_oidc_pending_state); falling back to *audience*
|
||||
# only matters if the column ever holds an empty string.
|
||||
jwt_audience = pending.get("audience") or audience
|
||||
jwt_token = create_jwt(
|
||||
user_id=user["user_id"],
|
||||
scopes=scopes,
|
||||
|
||||
@@ -124,7 +124,6 @@ _CONFIG_MAP: dict[str, dict[str, str]] = {
|
||||
},
|
||||
"mcp": {
|
||||
"config_path": "mcp_config",
|
||||
"refresh_interval": "mcp_refresh_interval",
|
||||
},
|
||||
"ratelimit": {
|
||||
"enabled": "ratelimit_enabled",
|
||||
|
||||
@@ -22,24 +22,14 @@ import json
|
||||
from typing import Any
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.tool_advisory import (
|
||||
_USER_INTERJECTION_BODY_MARKER,
|
||||
_USER_INTERJECTION_IMPORTANT_PREAMBLE,
|
||||
_USER_INTERJECTION_NOTICE_PREAMBLE,
|
||||
)
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# Tool results are clamped at this length per row at storage time
|
||||
# (see ``session.py``'s ``store_text = raw_output[:TOOL_RESULT_STORAGE_CAP]``).
|
||||
# Keeping the constant here lets the truncation flag detection in
|
||||
# ``decorate_history_messages`` stay in sync without a magic number
|
||||
# duplicated across server.py / session.py.
|
||||
#
|
||||
# Raised from 2000 → 10000 because a 2000-char clip routinely cut
|
||||
# the body of a single grep / file read mid-line, leaving the
|
||||
# historical record useless for retrospective debugging. FTS5
|
||||
# index + row size grow proportionally; the per-tool upper bound is
|
||||
# still bounded upstream by ``_truncate_output``'s context-budget
|
||||
# clamp (so a single huge result can't blow past the live context
|
||||
# window).
|
||||
TOOL_RESULT_STORAGE_CAP = 10000
|
||||
|
||||
|
||||
def load_verdict_indexes(
|
||||
ws_id: str,
|
||||
@@ -174,6 +164,110 @@ def decorate_tool_call(
|
||||
tc["output_assessment"] = assessment
|
||||
|
||||
|
||||
def _entity_decode_wrapper_tags(text: str) -> str:
|
||||
"""Reverse :func:`tool_advisory.escape_wrapper_tags` on extraction.
|
||||
|
||||
The wrap layer escapes the four wrapper-tag forms to HTML entities
|
||||
so embedded user / advisory text cannot fabricate or close an
|
||||
envelope. When the replay decorator pulls advisories back out of
|
||||
the persisted envelope, the inner text needs to be returned to its
|
||||
literal form for UI rendering.
|
||||
|
||||
Decodes ``&`` last so a tool output that contains the literal
|
||||
string ``<tool_output>`` round-trips identically to its
|
||||
source: encode produces ``&lt;tool_output&gt;`` (no
|
||||
collision with wrapper-tag escapes), decode walks the wrapper
|
||||
escapes first, then strips the ``&`` sentinel back to ``&``.
|
||||
The short-circuit on ``"&" not in text`` covers the common case
|
||||
where no escaped entities are present.
|
||||
"""
|
||||
if "&" not in text:
|
||||
return text
|
||||
return (
|
||||
text.replace("</tool_output>", "</tool_output>")
|
||||
.replace("<tool_output>", "<tool_output>")
|
||||
.replace("<system-reminder>", "<system-reminder>")
|
||||
.replace("</system-reminder>", "</system-reminder>")
|
||||
.replace("&", "&")
|
||||
)
|
||||
|
||||
|
||||
def _classify_advisory(render_text: str) -> dict[str, str] | None:
|
||||
"""Map a ``<system-reminder>`` body back to a wire-shape advisory.
|
||||
|
||||
Returns a dict with ``type`` / ``text`` / optional ``priority`` for
|
||||
advisory shapes the UI knows how to render, or ``None`` to suppress
|
||||
the advisory entirely (output-guard findings already render via the
|
||||
``output_assessment`` audit-table decoration; doubling them would
|
||||
paint two warning bubbles). Unknown advisory shapes fall through
|
||||
to ``None`` rather than rendering an opaque envelope blob.
|
||||
"""
|
||||
if render_text.startswith("Output guard:"):
|
||||
return None
|
||||
if _USER_INTERJECTION_BODY_MARKER in render_text:
|
||||
# UserInterjection is the only producer that uses this marker.
|
||||
# The preamble disambiguates priority: "important" gets the
|
||||
# MUST-address framing, "notice" gets the incorporate-if-relevant
|
||||
# framing. The body sits after the marker. Preamble + marker
|
||||
# constants are imported from ``tool_advisory`` so the parser
|
||||
# and producer can never drift on wording.
|
||||
if render_text.startswith(_USER_INTERJECTION_IMPORTANT_PREAMBLE):
|
||||
priority = "important"
|
||||
elif render_text.startswith(_USER_INTERJECTION_NOTICE_PREAMBLE):
|
||||
priority = "notice"
|
||||
else:
|
||||
# Marker present but preamble drifted — still render as a
|
||||
# notice rather than dropping the user's text.
|
||||
priority = "notice"
|
||||
body = render_text.split(_USER_INTERJECTION_BODY_MARKER, 1)[1]
|
||||
# Suppress empty/whitespace-only advisories — ``queue_message``
|
||||
# accepts any non-None text including ``""`` / ``" "``, and a
|
||||
# blank body would paint a featureless empty user bubble on
|
||||
# replay. Dropping at the classifier keeps the wire-shape
|
||||
# contract uniform (no empty advisories ever ride the wire).
|
||||
if not body.strip():
|
||||
return None
|
||||
return {"type": "user_interjection", "text": body, "priority": priority}
|
||||
return None
|
||||
|
||||
|
||||
def extract_advisories_from_tool_envelope(
|
||||
content: str,
|
||||
) -> tuple[str, list[dict[str, str]]] | None:
|
||||
"""Strip a ``<tool_output>`` envelope and return ``(clean, advisories)``.
|
||||
|
||||
Returns ``None`` when *content* doesn't look like a wrapped tool
|
||||
result — caller should leave the message unchanged. When the
|
||||
envelope parses but no advisories survive classification (e.g.
|
||||
only an output_guard advisory rode along), returns the cleaned
|
||||
output with an empty advisories list — the caller still needs to
|
||||
strip the envelope from the rendered content.
|
||||
"""
|
||||
if not content.startswith("<tool_output>\n"):
|
||||
return None
|
||||
close = content.find("\n</tool_output>")
|
||||
if close == -1:
|
||||
return None
|
||||
inner = content[len("<tool_output>\n") : close]
|
||||
rest = content[close + len("\n</tool_output>") :]
|
||||
advisories: list[dict[str, str]] = []
|
||||
cursor = 0
|
||||
while True:
|
||||
open_idx = rest.find("<system-reminder>\n", cursor)
|
||||
if open_idx == -1:
|
||||
break
|
||||
close_idx = rest.find("\n</system-reminder>", open_idx)
|
||||
if close_idx == -1:
|
||||
break
|
||||
body = rest[open_idx + len("<system-reminder>\n") : close_idx]
|
||||
decoded = _entity_decode_wrapper_tags(body)
|
||||
classified = _classify_advisory(decoded)
|
||||
if classified is not None:
|
||||
advisories.append(classified)
|
||||
cursor = close_idx + len("\n</system-reminder>")
|
||||
return _entity_decode_wrapper_tags(inner), advisories
|
||||
|
||||
|
||||
def decorate_history_messages(
|
||||
messages: list[dict[str, Any]],
|
||||
verdicts_by_call_id: dict[str, dict[str, Any]],
|
||||
@@ -184,8 +278,12 @@ def decorate_history_messages(
|
||||
Used by the ``/history`` REST endpoint after ``load_messages``
|
||||
returns. For each assistant message with ``tool_calls``, runs
|
||||
:func:`decorate_tool_call` on every entry. For each tool message
|
||||
whose content hits the storage cap, sets ``truncated: True`` so
|
||||
the client can render the "… truncated in storage" pill.
|
||||
whose ``content`` carries a ``<tool_output>`` envelope (queued
|
||||
user message spliced via :class:`UserInterjection` during a tool
|
||||
batch), strips the envelope, restores literal wrapper tags inside
|
||||
the body, and surfaces the extracted advisories on
|
||||
``msg["advisories"]`` so the wire layer can replay them as user
|
||||
bubbles after the tool result.
|
||||
|
||||
Pure transform — no I/O. Async callers should pre-load the
|
||||
indexes via :func:`load_verdict_indexes` (in ``to_thread``) and
|
||||
@@ -201,5 +299,18 @@ def decorate_history_messages(
|
||||
decorate_tool_call(tc, verdicts_by_call_id, assessments_by_call_id)
|
||||
elif role == "tool":
|
||||
content = msg.get("content")
|
||||
if isinstance(content, str) and len(content) >= TOOL_RESULT_STORAGE_CAP:
|
||||
msg["truncated"] = True
|
||||
if not isinstance(content, str):
|
||||
continue
|
||||
try:
|
||||
extracted = extract_advisories_from_tool_envelope(content)
|
||||
except Exception:
|
||||
# Defensive — on any unexpected parse failure leave the
|
||||
# message untouched rather than crashing the replay.
|
||||
log.debug("advisory extraction failed; leaving content intact", exc_info=True)
|
||||
continue
|
||||
if extracted is None:
|
||||
continue
|
||||
cleaned, advisories = extracted
|
||||
msg["content"] = cleaned
|
||||
if advisories:
|
||||
msg["advisories"] = advisories
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Idle wake-trigger for the metacog NudgeQueue pipeline.
|
||||
|
||||
Hosts :class:`IdleNudgeWatcher` plus the
|
||||
:func:`install_idle_nudge_watcher` / :func:`shutdown_idle_nudge_watchers`
|
||||
lifespan helpers. Pulled out of :mod:`turnstone.core.metacognition`
|
||||
because the watcher is subscriber-lifecycle / runtime-orchestration
|
||||
code with different concerns from the static nudge-text templates and
|
||||
detection heuristics that live in metacognition; mixing them grew the
|
||||
metacog module past its single-responsibility line.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from turnstone.core import session_worker
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.nudge_queue import USER_DRAIN
|
||||
from turnstone.core.workstream import WorkstreamState
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from turnstone.core.session_manager import SessionManager
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
|
||||
class IdleNudgeWatcher:
|
||||
"""Convert a workstream IDLE transition into a wake send when the
|
||||
session has queued nudges.
|
||||
|
||||
Subscribes to :meth:`SessionManager.subscribe_to_state` and listens
|
||||
for ``WorkstreamState.IDLE``. If the workstream's
|
||||
:class:`NudgeQueue` has any drainable entry for the wake's drain
|
||||
filter (``USER_DRAIN`` — channels ``"user"`` or ``"any"``),
|
||||
dispatches via ``session_worker.send`` with a no-op ``enqueue``
|
||||
callback. Tool-only entries don't fire the wake — they belong to
|
||||
the next tool-result seam, not a synthetic empty user turn —
|
||||
otherwise every IDLE event with a queued tool advisory would spawn
|
||||
a wake daemon that immediately no-ops at
|
||||
``deliver_wake_nudge_from_queue``'s drain guard.
|
||||
|
||||
**Race semantics.** ``session_worker.send`` decides atomically
|
||||
under ``ws._lock`` whether a worker thread already owns the
|
||||
workstream. Three outcomes:
|
||||
|
||||
* No worker → spawn a new daemon that calls
|
||||
:meth:`ChatSession.deliver_wake_nudge_from_queue` (the wake
|
||||
drains its own queue and runs the synthetic empty-user turn).
|
||||
* Worker running → call our ``enqueue`` lambda, which is a no-op.
|
||||
The wake is silently dropped; the queued nudge stays in
|
||||
``NudgeQueue`` and the in-flight worker picks it up at its next
|
||||
user-message-attach or tool-result seam (whichever fires first
|
||||
for the entry's channel). This is the load-bearing fallback —
|
||||
we never spawn a competing worker.
|
||||
* Workstream gone (``ws is None``) or session not built
|
||||
(``ws.session is None``) → bail.
|
||||
|
||||
**Subscription order matters.** When a workstream-kind-specific
|
||||
observer (e.g. ``CoordinatorIdleObserver``) needs to *enqueue* a
|
||||
nudge on the same IDLE event before this watcher *peeks* the
|
||||
queue, the observer must register first so that
|
||||
``SessionManager.set_state``'s subscriber loop fires it earlier in
|
||||
the same synchronous fan-out.
|
||||
|
||||
**Kind-agnostic.** Fires for any workstream regardless of
|
||||
:class:`WorkstreamKind`. Producers decide what to enqueue.
|
||||
"""
|
||||
|
||||
def __init__(self, manager: SessionManager) -> None:
|
||||
self._manager = manager
|
||||
self._callback: Callable[[str, WorkstreamState], None] | None = None
|
||||
|
||||
def start(self) -> None:
|
||||
"""Idempotent — registering twice is a no-op."""
|
||||
if self._callback is not None:
|
||||
return
|
||||
|
||||
def _on_state(ws_id: str, state: WorkstreamState) -> None:
|
||||
if state is not WorkstreamState.IDLE:
|
||||
return
|
||||
ws = self._manager.get(ws_id)
|
||||
if ws is None or ws.session is None:
|
||||
return
|
||||
session = ws.session
|
||||
if not session._nudge_queue.has_pending(USER_DRAIN):
|
||||
return
|
||||
session_worker.send(
|
||||
ws,
|
||||
enqueue=lambda: None,
|
||||
run=session.deliver_wake_nudge_from_queue,
|
||||
thread_name=f"wake-nudge-{ws.id[:8]}",
|
||||
)
|
||||
|
||||
self._callback = _on_state
|
||||
self._manager.subscribe_to_state(_on_state)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Unsubscribe; idempotent."""
|
||||
cb = self._callback
|
||||
if cb is None:
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
self._manager.unsubscribe_from_state(cb)
|
||||
self._callback = None
|
||||
|
||||
|
||||
_APP_STATE_ATTR = "_idle_nudge_watchers"
|
||||
|
||||
|
||||
def install_idle_nudge_watcher(app: Any, manager: SessionManager) -> IdleNudgeWatcher:
|
||||
"""Construct + start an :class:`IdleNudgeWatcher` and register it
|
||||
for lifespan teardown via :func:`shutdown_idle_nudge_watchers`.
|
||||
|
||||
Multiple watchers may be installed against different
|
||||
:class:`SessionManager` instances on the same ``app`` (e.g. the
|
||||
interactive manager + the coord manager on a multi-kind host).
|
||||
All of them get torn down by a single
|
||||
:func:`shutdown_idle_nudge_watchers` call.
|
||||
|
||||
Returns the watcher so the caller can run additional setup
|
||||
against the same manager — but the typical site doesn't need
|
||||
the return value.
|
||||
"""
|
||||
watcher = IdleNudgeWatcher(manager)
|
||||
watcher.start()
|
||||
watchers: list[IdleNudgeWatcher] = getattr(app.state, _APP_STATE_ATTR, [])
|
||||
if not watchers:
|
||||
# First watcher on this app — initialise the list. Avoids
|
||||
# mutating a default arg or sharing the list across apps.
|
||||
setattr(app.state, _APP_STATE_ATTR, watchers)
|
||||
watchers.append(watcher)
|
||||
return watcher
|
||||
|
||||
|
||||
def shutdown_idle_nudge_watchers(app: Any) -> None:
|
||||
"""Shut down every watcher installed via
|
||||
:func:`install_idle_nudge_watcher`. No-op if none.
|
||||
"""
|
||||
watchers: list[IdleNudgeWatcher] = getattr(app.state, _APP_STATE_ATTR, [])
|
||||
for watcher in watchers:
|
||||
watcher.shutdown()
|
||||
watchers.clear()
|
||||
+3682
-376
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,579 @@
|
||||
"""Token-at-rest encryption for OAuth-MCP.
|
||||
|
||||
Uses cryptography.fernet (AES-128-CBC + HMAC-SHA256, 256-bit total key
|
||||
material, encrypt-then-MAC). Single-key chosen for v1; rotation supported
|
||||
via cryptography.fernet.MultiFernet.
|
||||
|
||||
Operator note: when rotating keys, place the NEW key first in
|
||||
``mcp_token_encryption_keys``. MultiFernet writes with the first key and
|
||||
tries each in order on read. Old keys can be retired once all rows are
|
||||
re-encrypted by a future operator-driven migration.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, TypedDict
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken, MultiFernet
|
||||
|
||||
from turnstone.core.audit import record_audit
|
||||
from turnstone.core.log import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from turnstone.core.storage._protocol import StorageBackend
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
# Constants
|
||||
_KEY_BYTES = 32 # Fernet requires 32 bytes
|
||||
_KEY_FINGERPRINT_BYTES = 8 # short hex prefix for audit/error fields
|
||||
|
||||
# Operator-facing hint for malformed/missing keys.
|
||||
_KEY_GEN_HINT = (
|
||||
"regenerate with: python -c 'from cryptography.fernet import Fernet; "
|
||||
"print(Fernet.generate_key().decode())'"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MCPCryptoError(Exception):
|
||||
"""Base class for MCP token-at-rest encryption errors."""
|
||||
|
||||
|
||||
class MCPTokenDecryptError(MCPCryptoError):
|
||||
"""No installed key can decrypt the ciphertext.
|
||||
|
||||
Maps to RFC's ``mcp_token_undecryptable_key_unknown`` error class.
|
||||
Carries ``key_fingerprints_attempted: tuple[str, ...]`` for audit.
|
||||
|
||||
Critical: callers MUST NOT auto-delete the row on this error.
|
||||
The row is still valid; this node just doesn't have the right key.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, *, key_fingerprints_attempted: tuple[str, ...]) -> None:
|
||||
super().__init__(message)
|
||||
self.key_fingerprints_attempted = key_fingerprints_attempted
|
||||
|
||||
|
||||
class MCPTokenKeyConfigError(MCPCryptoError):
|
||||
"""Key material in config.toml is malformed or missing."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Plaintext shape returned by ``MCPTokenStore.get_user_token``
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MCPUserTokenPlain(TypedDict):
|
||||
"""Plaintext shape returned by ``MCPTokenStore.get_user_token``.
|
||||
|
||||
Mirrors ``MCPUserToken`` (storage row shape) minus the ``_ct`` suffix
|
||||
on token columns and with plaintext bytes-decoded values.
|
||||
"""
|
||||
|
||||
user_id: str
|
||||
server_name: str
|
||||
access_token: str
|
||||
refresh_token: str | None
|
||||
expires_at: str | None
|
||||
scopes: str | None
|
||||
as_issuer: str
|
||||
audience: str
|
||||
created: str
|
||||
last_refreshed: str | None
|
||||
|
||||
|
||||
class MCPUserTokenMetadata(TypedDict):
|
||||
"""Non-secret subset of ``MCPUserToken`` for the settings UI.
|
||||
|
||||
Token ciphertext is intentionally absent: a list view never needs
|
||||
the access/refresh secrets, and decrypt happens only at MCP-call
|
||||
time.
|
||||
"""
|
||||
|
||||
user_id: str
|
||||
server_name: str
|
||||
expires_at: str | None
|
||||
scopes: str | None
|
||||
as_issuer: str
|
||||
audience: str
|
||||
created: str
|
||||
last_refreshed: str | None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config dataclass + loader
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True, repr=False)
|
||||
class MCPTokenCipherConfig:
|
||||
"""Validated key material loaded from config.toml.
|
||||
|
||||
``keys`` are raw 32-byte secrets; the first is the encryption key,
|
||||
all are tried in order on read. The cipher wrapper re-encodes them
|
||||
via ``base64.urlsafe_b64encode`` for ``Fernet(...)`` at construction
|
||||
time.
|
||||
|
||||
``__repr__`` is overridden to redact the raw key bytes — the default
|
||||
dataclass repr would emit them verbatim into logs / tracebacks.
|
||||
"""
|
||||
|
||||
keys: tuple[bytes, ...]
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"MCPTokenCipherConfig(keys=<{len(self.keys)} key(s) redacted>)"
|
||||
|
||||
|
||||
def _validate_key(raw: str, *, label: str) -> bytes:
|
||||
"""Decode + validate a single base64 url-safe key. Raises ``MCPTokenKeyConfigError``."""
|
||||
if not isinstance(raw, str) or not raw.strip():
|
||||
raise MCPTokenKeyConfigError(f"{label}: key is empty or not a string. {_KEY_GEN_HINT}")
|
||||
try:
|
||||
decoded = base64.urlsafe_b64decode(raw.encode("ascii"))
|
||||
except Exception as exc:
|
||||
raise MCPTokenKeyConfigError(
|
||||
f"{label}: not valid base64 url-safe ({exc}). {_KEY_GEN_HINT}"
|
||||
) from exc
|
||||
if len(decoded) != _KEY_BYTES:
|
||||
raise MCPTokenKeyConfigError(
|
||||
f"{label}: decoded key must be exactly {_KEY_BYTES} bytes, "
|
||||
f"got {len(decoded)}. {_KEY_GEN_HINT}"
|
||||
)
|
||||
return decoded
|
||||
|
||||
|
||||
def _key_fingerprint(key: bytes) -> str:
|
||||
"""Stable, non-reversible 8-hex prefix of SHA-256(key)."""
|
||||
digest = hashlib.sha256(key).hexdigest()
|
||||
return digest[: _KEY_FINGERPRINT_BYTES * 2]
|
||||
|
||||
|
||||
def load_mcp_token_cipher_config() -> MCPTokenCipherConfig | None:
|
||||
"""Read ``[security] mcp_token_encryption_keys`` (plural) or
|
||||
``mcp_token_encryption_key`` (singular) from config.toml.
|
||||
|
||||
Plural takes precedence when both are present. Returns ``None`` when
|
||||
neither key is configured (caller decides whether that's fatal).
|
||||
Raises ``MCPTokenKeyConfigError`` on malformed key material.
|
||||
"""
|
||||
from turnstone.core.config import load_config
|
||||
|
||||
sec_cfg = load_config("security")
|
||||
raw_list_value = sec_cfg.get("mcp_token_encryption_keys")
|
||||
raw_single_value = sec_cfg.get("mcp_token_encryption_key")
|
||||
|
||||
raw_keys: list[str]
|
||||
if isinstance(raw_list_value, list) and raw_list_value:
|
||||
raw_keys = []
|
||||
for idx, item in enumerate(raw_list_value):
|
||||
if not isinstance(item, str):
|
||||
raise MCPTokenKeyConfigError(
|
||||
f"mcp_token_encryption_keys[{idx}]: must be a string. {_KEY_GEN_HINT}"
|
||||
)
|
||||
raw_keys.append(item)
|
||||
elif raw_list_value is not None and not isinstance(raw_list_value, list):
|
||||
raise MCPTokenKeyConfigError(
|
||||
f"mcp_token_encryption_keys: must be a list of base64 url-safe strings. {_KEY_GEN_HINT}"
|
||||
)
|
||||
elif isinstance(raw_single_value, str) and raw_single_value.strip():
|
||||
raw_keys = [raw_single_value]
|
||||
else:
|
||||
return None
|
||||
|
||||
decoded_keys: list[bytes] = []
|
||||
for idx, raw in enumerate(raw_keys):
|
||||
label = (
|
||||
f"mcp_token_encryption_keys[{idx}]" if len(raw_keys) > 1 else "mcp_token_encryption_key"
|
||||
)
|
||||
decoded_keys.append(_validate_key(raw, label=label))
|
||||
return MCPTokenCipherConfig(keys=tuple(decoded_keys))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cipher wrapper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MCPTokenCipher:
|
||||
"""Encrypt/decrypt with one or more Fernet keys.
|
||||
|
||||
First key in ``cfg.keys`` is the encryption key. All keys are tried
|
||||
(in declared order) for decryption. On total decryption failure,
|
||||
raises ``MCPTokenDecryptError`` with the fingerprints attempted.
|
||||
"""
|
||||
|
||||
def __init__(self, cfg: MCPTokenCipherConfig) -> None:
|
||||
if not cfg.keys:
|
||||
raise MCPTokenKeyConfigError(
|
||||
f"MCPTokenCipher requires at least one key. {_KEY_GEN_HINT}"
|
||||
)
|
||||
self._cfg = cfg
|
||||
self._fingerprints = tuple(_key_fingerprint(k) for k in cfg.keys)
|
||||
# Re-encode raw bytes to the base64-url-safe form Fernet expects.
|
||||
fernets = [Fernet(base64.urlsafe_b64encode(k)) for k in cfg.keys]
|
||||
self._encrypter = fernets[0]
|
||||
self._multi = MultiFernet(fernets)
|
||||
|
||||
def encrypt(self, plaintext: bytes) -> bytes:
|
||||
"""Encrypt ``plaintext`` with the active (first) key."""
|
||||
return self._encrypter.encrypt(plaintext)
|
||||
|
||||
def decrypt(self, ciphertext: bytes) -> bytes:
|
||||
"""Try every installed key in declared order.
|
||||
|
||||
Raises ``MCPTokenDecryptError`` carrying the fingerprints
|
||||
attempted when all fail.
|
||||
"""
|
||||
try:
|
||||
return self._multi.decrypt(ciphertext)
|
||||
except InvalidToken as exc:
|
||||
raise MCPTokenDecryptError(
|
||||
"no installed key can decrypt the ciphertext",
|
||||
key_fingerprints_attempted=self._fingerprints,
|
||||
) from exc
|
||||
|
||||
@property
|
||||
def key_fingerprints(self) -> tuple[str, ...]:
|
||||
"""Stable fingerprints of installed keys, in declared order.
|
||||
|
||||
Useful for audit events and operator-facing error messages.
|
||||
"""
|
||||
return self._fingerprints
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Token store: ciphertext-aware CRUD layered on the storage protocol
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MCPTokenStore:
|
||||
"""Encrypt/decrypt OAuth tokens at the storage boundary.
|
||||
|
||||
Wraps a :class:`StorageBackend`'s ciphertext-only token CRUD with a
|
||||
plaintext-facing API. ``audit_storage`` + ``node_id`` are optional;
|
||||
when both are set, decrypt failures are recorded as audit events
|
||||
under ``mcp_server.oauth.token_decrypt_failure`` with the
|
||||
fingerprints attempted.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
storage: StorageBackend,
|
||||
cipher: MCPTokenCipher,
|
||||
*,
|
||||
node_id: str = "",
|
||||
audit_storage: StorageBackend | None = None,
|
||||
) -> None:
|
||||
self._storage = storage
|
||||
self._cipher = cipher
|
||||
self._node_id = node_id
|
||||
self._audit_storage = audit_storage
|
||||
|
||||
@property
|
||||
def cipher(self) -> MCPTokenCipher:
|
||||
"""The underlying cipher (exposed for callers that need to encrypt
|
||||
non-token blobs, e.g., the MCP-server admin form's
|
||||
``oauth_client_secret`` plaintext input)."""
|
||||
return self._cipher
|
||||
|
||||
def create_user_token(
|
||||
self,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
*,
|
||||
access_token: str,
|
||||
refresh_token: str | None,
|
||||
expires_at: str | None,
|
||||
scopes: str | None,
|
||||
as_issuer: str,
|
||||
audience: str,
|
||||
) -> None:
|
||||
"""Encrypt the access (and optional refresh) token and persist."""
|
||||
access_ct = self._cipher.encrypt(access_token.encode("utf-8"))
|
||||
refresh_ct = self._cipher.encrypt(refresh_token.encode("utf-8")) if refresh_token else None
|
||||
self._storage.create_mcp_user_token(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token_ct=access_ct,
|
||||
refresh_token_ct=refresh_ct,
|
||||
expires_at=expires_at,
|
||||
scopes=scopes,
|
||||
as_issuer=as_issuer,
|
||||
audience=audience,
|
||||
)
|
||||
|
||||
def get_user_token(self, user_id: str, server_name: str) -> MCPUserTokenPlain | None:
|
||||
"""Returns plaintext dict or None.
|
||||
|
||||
Raises ``MCPTokenDecryptError`` on key mismatch — caller MUST NOT
|
||||
auto-delete the row. If ``audit_storage`` + ``node_id`` are
|
||||
configured, emits ``mcp_server.oauth.token_decrypt_failure``
|
||||
audit event.
|
||||
"""
|
||||
row = self._storage.get_mcp_user_token(user_id, server_name)
|
||||
if row is None:
|
||||
return None
|
||||
try:
|
||||
access_pt = self._cipher.decrypt(row["access_token_ct"]).decode("utf-8")
|
||||
refresh_pt: str | None
|
||||
if row["refresh_token_ct"] is not None:
|
||||
refresh_pt = self._cipher.decrypt(row["refresh_token_ct"]).decode("utf-8")
|
||||
else:
|
||||
refresh_pt = None
|
||||
except MCPTokenDecryptError as exc:
|
||||
self._audit_decrypt_failure(server_name, exc.key_fingerprints_attempted)
|
||||
raise
|
||||
return MCPUserTokenPlain(
|
||||
user_id=row["user_id"],
|
||||
server_name=row["server_name"],
|
||||
access_token=access_pt,
|
||||
refresh_token=refresh_pt,
|
||||
expires_at=row["expires_at"],
|
||||
scopes=row["scopes"],
|
||||
as_issuer=row["as_issuer"],
|
||||
audience=row["audience"],
|
||||
created=row["created"],
|
||||
last_refreshed=row["last_refreshed"],
|
||||
)
|
||||
|
||||
def update_user_token_after_refresh(
|
||||
self,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
*,
|
||||
access_token: str,
|
||||
refresh_token: str | None,
|
||||
expires_at: str | None,
|
||||
) -> bool:
|
||||
"""Atomic write of new tokens after a refresh-grant exchange.
|
||||
|
||||
Returns True when a row was updated.
|
||||
|
||||
``refresh_token=None`` CLEARS the column — it does NOT preserve
|
||||
the existing value. Per RFC 6749 §6, an authorization server MAY
|
||||
omit ``refresh_token`` from the refresh response; in that case
|
||||
the OAuth-flow caller MUST pre-resolve whether to keep the
|
||||
existing refresh token or drop it before invoking this method.
|
||||
This API has no "leave unchanged" sentinel.
|
||||
"""
|
||||
access_ct = self._cipher.encrypt(access_token.encode("utf-8"))
|
||||
refresh_ct = self._cipher.encrypt(refresh_token.encode("utf-8")) if refresh_token else None
|
||||
return self._storage.update_mcp_user_token_after_refresh(
|
||||
user_id,
|
||||
server_name,
|
||||
access_token_ct=access_ct,
|
||||
refresh_token_ct=refresh_ct,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
def delete_user_token(self, user_id: str, server_name: str) -> bool:
|
||||
"""Delete the user-token row. Returns True if existed."""
|
||||
return self._storage.delete_mcp_user_token(user_id, server_name)
|
||||
|
||||
def list_user_token_metadata(self, user_id: str) -> list[MCPUserTokenMetadata]:
|
||||
"""Return non-secret metadata for every token row owned by ``user_id``.
|
||||
|
||||
Storage layer projects the metadata columns at the SQL boundary
|
||||
(``list_mcp_user_token_metadata_by_user``) so ciphertext blobs
|
||||
never cross the wire on this list-view path. Rows arrive in
|
||||
``created`` ASC order. Decrypt is intentionally skipped — the
|
||||
list view has no need for the secret material.
|
||||
"""
|
||||
rows = self._storage.list_mcp_user_token_metadata_by_user(user_id)
|
||||
return [
|
||||
MCPUserTokenMetadata(
|
||||
user_id=row["user_id"],
|
||||
server_name=row["server_name"],
|
||||
expires_at=row["expires_at"],
|
||||
scopes=row["scopes"],
|
||||
as_issuer=row["as_issuer"],
|
||||
audience=row["audience"],
|
||||
created=row["created"],
|
||||
last_refreshed=row["last_refreshed"],
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def set_oauth_client_secret(self, server_id: str, plaintext_secret: str | None) -> bool:
|
||||
"""Encrypt plaintext and persist via the dedicated storage writer.
|
||||
|
||||
Pass ``None`` to clear the column. Empty string is encrypted
|
||||
normally (Fernet accepts empty plaintext); callers that treat
|
||||
empty as "clear" must convert to ``None`` at their API boundary
|
||||
first — the admin form does this before invoking the helper.
|
||||
|
||||
Returns ``False`` when ``server_id`` does not exist.
|
||||
"""
|
||||
if plaintext_secret is None:
|
||||
return self._storage.set_mcp_oauth_client_secret_ct(server_id, None)
|
||||
secret_ct = self._cipher.encrypt(plaintext_secret.encode("utf-8"))
|
||||
return self._storage.set_mcp_oauth_client_secret_ct(server_id, secret_ct)
|
||||
|
||||
def get_oauth_client_secret(self, server_id: str) -> str | None:
|
||||
"""Decrypt and return the per-server OAuth client secret, or None.
|
||||
|
||||
Returns ``None`` when the row is missing or the column is NULL.
|
||||
Raises :class:`MCPTokenDecryptError` on key mismatch — the caller
|
||||
decides whether to treat that as a missing-secret case (e.g. log +
|
||||
prompt re-consent) or surface as a configuration failure.
|
||||
"""
|
||||
secret_ct = self._storage.get_mcp_oauth_client_secret_ct(server_id)
|
||||
if secret_ct is None:
|
||||
return None
|
||||
return self._cipher.decrypt(secret_ct).decode("utf-8")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _audit_decrypt_failure(self, server_name: str, fingerprints: tuple[str, ...]) -> None:
|
||||
"""Best-effort audit emit on decrypt failure (no-op when unconfigured).
|
||||
|
||||
Uses ``server_id`` (PK UUID) as ``resource_id`` so admin-driven
|
||||
server renames don't break event correlation. Falls back to
|
||||
``server_name`` when the lookup misses.
|
||||
"""
|
||||
if self._audit_storage is None:
|
||||
return
|
||||
resource_id = server_name
|
||||
try:
|
||||
row = self._audit_storage.get_mcp_server_by_name(server_name)
|
||||
except Exception:
|
||||
row = None
|
||||
if row is not None:
|
||||
resource_id = str(row.get("server_id") or server_name)
|
||||
try:
|
||||
record_audit(
|
||||
self._audit_storage,
|
||||
user_id="",
|
||||
action="mcp_server.oauth.token_decrypt_failure",
|
||||
resource_type="mcp_server",
|
||||
resource_id=resource_id,
|
||||
detail={
|
||||
"server_name": server_name,
|
||||
"key_fingerprints_attempted": list(fingerprints),
|
||||
"node_id": self._node_id,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
log.warning(
|
||||
"mcp_server.oauth.audit_emit_failed",
|
||||
action="token_decrypt_failure",
|
||||
server_name=server_name,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lifespan integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def initialize_mcp_crypto_state(app_state: object, *, node_id: str = "") -> None:
|
||||
"""Validate Fernet key config + install :class:`MCPTokenCipher` /
|
||||
:class:`MCPTokenStore` on ``app_state``.
|
||||
|
||||
Called from the server / console lifespan after OIDC initialization.
|
||||
|
||||
Behavior:
|
||||
|
||||
1. ``load_mcp_token_cipher_config()`` — wrapped in try/except. Raises
|
||||
:class:`SystemExit(1)` on :class:`MCPTokenKeyConfigError` after
|
||||
logging.
|
||||
2. Counts ``mcp_servers`` rows with ``auth_type='oauth_user'``. If
|
||||
any exist AND no key is configured, raises ``SystemExit(1)``.
|
||||
3. On success, sets ``app_state.mcp_token_cipher`` and
|
||||
``app_state.mcp_token_store`` (both possibly ``None`` when no
|
||||
key + no oauth_user rows).
|
||||
|
||||
The helper is shared by ``turnstone/server.py:_lifespan`` and
|
||||
``turnstone/console/server.py:_lifespan``. A separate
|
||||
:func:`close_mcp_crypto_state` mirrors :func:`close_oidc_state` for
|
||||
parity even though the cipher itself owns no resources.
|
||||
"""
|
||||
from turnstone.core.storage import get_storage
|
||||
|
||||
try:
|
||||
cipher_cfg = load_mcp_token_cipher_config()
|
||||
except MCPTokenKeyConfigError as exc:
|
||||
log.error("mcp_server.oauth.key_config_invalid: %s", exc)
|
||||
raise SystemExit(1) from exc
|
||||
|
||||
storage = get_storage()
|
||||
oauth_user_count = sum(
|
||||
1 for row in storage.list_mcp_servers() if row.get("auth_type") == "oauth_user"
|
||||
)
|
||||
|
||||
if oauth_user_count > 0 and cipher_cfg is None:
|
||||
log.error(
|
||||
"mcp.oauth: %d server(s) configured with auth_type='oauth_user' but no "
|
||||
"[security] mcp_token_encryption_keys (rotation list) or "
|
||||
"mcp_token_encryption_key (single) in config.toml. Generate a key with: "
|
||||
"python -c 'from cryptography.fernet import Fernet; "
|
||||
"print(Fernet.generate_key().decode())' "
|
||||
"and add it to your config.toml.",
|
||||
oauth_user_count,
|
||||
)
|
||||
raise SystemExit(1)
|
||||
|
||||
if cipher_cfg is None:
|
||||
# No oauth_user rows + no key configured: zero new code paths
|
||||
# exercised; install None sentinels so callers can fast-path.
|
||||
app_state.mcp_token_cipher = None # type: ignore[attr-defined]
|
||||
app_state.mcp_token_store = None # type: ignore[attr-defined]
|
||||
log.debug("mcp_server.oauth.disabled (no key configured, no oauth_user rows)")
|
||||
return
|
||||
|
||||
cipher = MCPTokenCipher(cipher_cfg)
|
||||
app_state.mcp_token_cipher = cipher # type: ignore[attr-defined]
|
||||
app_state.mcp_token_store = MCPTokenStore( # type: ignore[attr-defined]
|
||||
storage,
|
||||
cipher,
|
||||
node_id=node_id,
|
||||
audit_storage=storage,
|
||||
)
|
||||
log.info(
|
||||
"mcp_server.oauth.cipher_installed",
|
||||
keys=len(cipher.key_fingerprints),
|
||||
active_fp=cipher.key_fingerprints[0],
|
||||
)
|
||||
|
||||
|
||||
def close_mcp_crypto_state(app_state: object) -> None:
|
||||
"""Drop references to the cipher / token store on shutdown.
|
||||
|
||||
Mirrors :func:`turnstone.core.oidc.close_oidc_state` for parity.
|
||||
The cipher itself owns no network resources, so this is a simple
|
||||
attribute clear.
|
||||
"""
|
||||
if hasattr(app_state, "mcp_token_store"):
|
||||
app_state.mcp_token_store = None
|
||||
if hasattr(app_state, "mcp_token_cipher"):
|
||||
app_state.mcp_token_cipher = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Re-exports for callers that don't need the storage backend
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
__all__ = [
|
||||
"MCPCryptoError",
|
||||
"MCPTokenCipher",
|
||||
"MCPTokenCipherConfig",
|
||||
"MCPTokenDecryptError",
|
||||
"MCPTokenKeyConfigError",
|
||||
"MCPTokenStore",
|
||||
"MCPUserTokenMetadata",
|
||||
"MCPUserTokenPlain",
|
||||
"close_mcp_crypto_state",
|
||||
"initialize_mcp_crypto_state",
|
||||
"load_mcp_token_cipher_config",
|
||||
]
|
||||
@@ -0,0 +1,255 @@
|
||||
"""HTTP header parsing helpers shared by the MCP client and OAuth modules.
|
||||
|
||||
Both ``mcp_client`` and ``mcp_oauth`` need to extract structured values from
|
||||
``WWW-Authenticate: Bearer ...`` headers — ``mcp_client`` to classify
|
||||
401/403 responses for the user-pool dispatcher, and ``mcp_oauth`` to pull
|
||||
the ``resource_metadata`` URL out of a discovery challenge. This module
|
||||
hosts the shared primitives so both modules can call them without
|
||||
duplicating fragile substring scanners.
|
||||
|
||||
The earlier hand-rolled scanners (``_parse_www_authenticate_scope`` /
|
||||
``_parse_www_authenticate_error`` in ``mcp_client``) used
|
||||
``header.lower().find(needle, i)`` to locate parameter names. That made
|
||||
them vulnerable to:
|
||||
|
||||
* matching ``scope`` inside ``xscope`` or ``ascope``,
|
||||
* matching the literal text ``scope=...`` embedded inside the quoted
|
||||
``realm`` value of a preceding ``auth-param``,
|
||||
* O(N**2) behaviour on pathological input (each ``find`` rescans the prefix).
|
||||
|
||||
This module replaces those with a single tokenizer that walks the RFC 7235
|
||||
``challenge → auth-param`` grammar once, tracks quoted-string state, and
|
||||
returns a normalised ``{key.lower(): value}`` dict. The thin extraction
|
||||
wrappers (``parse_www_authenticate_scope`` / ``parse_www_authenticate_error``)
|
||||
preserve the original return shapes so call sites only need to swap the
|
||||
import.
|
||||
|
||||
Also hosts ``MAX_INSUFFICIENT_SCOPE_REPORTED`` — the shared defensive cap
|
||||
on scope-list lengths consumed by both the WWW-Authenticate parser
|
||||
(``mcp_client``) and the ``/v1/api/mcp/oauth/start?scopes=`` step-up
|
||||
handler (``mcp_oauth``). Living here avoids one module importing a
|
||||
private name from the other.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def _parse_quoted_string(text: str, start: int) -> tuple[str, int] | None:
|
||||
"""Parse an RFC 7230 ``quoted-string`` starting at ``text[start]``.
|
||||
|
||||
Returns ``(value, end_index)`` where ``end_index`` is the index just
|
||||
past the closing quote, or ``None`` if the input is malformed (no
|
||||
opening quote, unterminated string).
|
||||
|
||||
Handles ``\\"`` and ``\\\\`` escapes per RFC 7230 section 3.2.6 — the
|
||||
prior naive ``([^"]+)`` regex truncated the URL at the first
|
||||
unescaped quote and silently dropped backslash escapes from the
|
||||
value.
|
||||
"""
|
||||
if start >= len(text) or text[start] != '"':
|
||||
return None
|
||||
out: list[str] = []
|
||||
i = start + 1
|
||||
while i < len(text):
|
||||
ch = text[i]
|
||||
if ch == "\\" and i + 1 < len(text):
|
||||
out.append(text[i + 1])
|
||||
i += 2
|
||||
continue
|
||||
if ch == '"':
|
||||
return "".join(out), i + 1
|
||||
out.append(ch)
|
||||
i += 1
|
||||
return None
|
||||
|
||||
|
||||
# Maximum header length we'll attempt to parse. Real ASes emit a handful
|
||||
# of short auth-params; anything past this is either malformed or
|
||||
# adversarial. Returning ``{}`` (rather than raising) keeps callers' error
|
||||
# paths uniform with "unparseable header → no signal".
|
||||
_MAX_HEADER_LEN = 4096
|
||||
|
||||
_TOKEN_DELIMS = frozenset('()<>@,;:\\"/[]?={} \t')
|
||||
|
||||
|
||||
def _is_token_char(ch: str) -> bool:
|
||||
"""RFC 7230 token character: visible ASCII minus the delimiter set."""
|
||||
return ch.isascii() and ch.isprintable() and ch not in _TOKEN_DELIMS
|
||||
|
||||
|
||||
def _looks_like_bearer_challenge_start(header: str, i: int) -> bool:
|
||||
"""Peek at ``header[i:]`` for the start of a fresh ``Bearer`` challenge.
|
||||
|
||||
Returns True when the slice begins with the case-insensitive token
|
||||
``Bearer`` followed by whitespace — the RFC 7235 marker for a new
|
||||
``challenge`` after a separator comma. This is the cue the bearer
|
||||
tokenizer uses to stop parsing rather than fold a second challenge's
|
||||
auth-params into the first challenge's dict.
|
||||
"""
|
||||
n = len(header)
|
||||
if i + 6 > n:
|
||||
return False
|
||||
if header[i : i + 6].lower() != "bearer":
|
||||
return False
|
||||
after = i + 6
|
||||
# ``Bearer`` must be followed by whitespace to qualify as a scheme
|
||||
# boundary; ``Bearer-like-token`` is just a regular token.
|
||||
return after < n and header[after] in " \t"
|
||||
|
||||
|
||||
def parse_www_authenticate_bearer(header: str) -> dict[str, str]:
|
||||
"""Extract ``auth-param``s from a ``WWW-Authenticate: Bearer ...`` header.
|
||||
|
||||
Walks the RFC 7235 challenge grammar once, returning a dict of
|
||||
``{lowercased-key: value}`` pairs. Quoted-strings are unquoted (with
|
||||
backslash escapes resolved). Unknown / malformed input returns an
|
||||
empty dict — never raises.
|
||||
|
||||
Only ``Bearer`` challenges are recognised. The function ignores any
|
||||
leading whitespace before the scheme. When a second ``Bearer``
|
||||
challenge appears after a separator comma — as it would when
|
||||
httpx joins repeated ``WWW-Authenticate`` headers via
|
||||
``response.headers.get(...)`` — the tokenizer stops at the
|
||||
challenge boundary rather than folding the second challenge's
|
||||
auth-params into the first challenge's dict. This is the
|
||||
parser-side defence-in-depth mirror of the
|
||||
``response.headers.get_list(...)[0]`` guard in the dispatcher's
|
||||
capturing httpx factory; either layer alone neutralises the
|
||||
multi-header injection vector but both run together so a
|
||||
regression in one cannot silently re-open it.
|
||||
A ``realm`` value that contains the literal text ``scope=fake`` is
|
||||
correctly attributed to ``realm`` because the tokenizer respects
|
||||
quoted-string boundaries.
|
||||
"""
|
||||
if not header or len(header) > _MAX_HEADER_LEN:
|
||||
return {}
|
||||
|
||||
n = len(header)
|
||||
i = 0
|
||||
|
||||
# Skip leading whitespace then the ``Bearer`` scheme token.
|
||||
while i < n and header[i] in " \t":
|
||||
i += 1
|
||||
scheme_start = i
|
||||
while i < n and _is_token_char(header[i]):
|
||||
i += 1
|
||||
scheme = header[scheme_start:i]
|
||||
if scheme.lower() != "bearer":
|
||||
return {}
|
||||
# Require at least one space between scheme and first auth-param.
|
||||
if i >= n or header[i] not in " \t":
|
||||
return {}
|
||||
|
||||
out: dict[str, str] = {}
|
||||
while i < n:
|
||||
# Skip whitespace and stray commas between params.
|
||||
while i < n and header[i] in " \t,":
|
||||
i += 1
|
||||
if i >= n:
|
||||
break
|
||||
# If a fresh ``Bearer`` challenge starts here, the upstream is
|
||||
# multi-challenge — stop before reading any of its auth-params.
|
||||
if _looks_like_bearer_challenge_start(header, i):
|
||||
break
|
||||
# Read the param key (a token).
|
||||
key_start = i
|
||||
while i < n and _is_token_char(header[i]):
|
||||
i += 1
|
||||
if i == key_start:
|
||||
# Not a valid token start — skip one char to make forward
|
||||
# progress and continue. This bounds total cost to O(N).
|
||||
i += 1
|
||||
continue
|
||||
key = header[key_start:i].lower()
|
||||
# Optional whitespace, then ``=``.
|
||||
while i < n and header[i] in " \t":
|
||||
i += 1
|
||||
if i >= n or header[i] != "=":
|
||||
# Param without a value — skip.
|
||||
continue
|
||||
i += 1
|
||||
while i < n and header[i] in " \t":
|
||||
i += 1
|
||||
if i >= n:
|
||||
break
|
||||
# Value: either a quoted-string or a token.
|
||||
if header[i] == '"':
|
||||
parsed = _parse_quoted_string(header, i)
|
||||
if parsed is None:
|
||||
# Unterminated quoted-string — treat the rest of the
|
||||
# header as garbage and stop. Returning what we already
|
||||
# have is safer than guessing where the value ends.
|
||||
break
|
||||
value, i = parsed
|
||||
out.setdefault(key, value)
|
||||
else:
|
||||
val_start = i
|
||||
while i < n and header[i] not in ", \t":
|
||||
i += 1
|
||||
value = header[val_start:i]
|
||||
out.setdefault(key, value)
|
||||
return out
|
||||
|
||||
|
||||
# Defensive cap on the number of scopes reported in
|
||||
# ``mcp_insufficient_scope`` audit/error payloads and accepted from the
|
||||
# ``/v1/api/mcp/oauth/start?scopes=`` step-up query param. Real ASes
|
||||
# return single-digit scope counts; the cap stops a malicious upstream
|
||||
# (or buggy client) from bloating either surface via a thousand-token
|
||||
# scope list. Lives here so the WWW-Authenticate parser (consumer:
|
||||
# ``mcp_client``) and the ``/start`` handler (consumer: ``mcp_oauth``)
|
||||
# share a single source of truth without one importing a private name
|
||||
# from the other.
|
||||
MAX_INSUFFICIENT_SCOPE_REPORTED = 32
|
||||
|
||||
|
||||
def is_valid_scope_token(token: str) -> bool:
|
||||
"""Return True iff ``token`` is a valid RFC 6749 §3.3 ``scope-token``.
|
||||
|
||||
The grammar restricts scope tokens to visible ASCII (``0x21..0x7E``)
|
||||
excluding ``"`` (``0x22``) and ``\\`` (``0x5C``). The empty string
|
||||
is rejected — a zero-length token has no semantic meaning in the
|
||||
space-separated scope list.
|
||||
|
||||
Used by the WWW-Authenticate parser to filter AS-supplied scope
|
||||
sets, and by ``/v1/api/mcp/oauth/start`` to reject caller-supplied
|
||||
scope query params that could smuggle CR/LF/tab/control bytes
|
||||
through the AS round-trip into downstream log or notification
|
||||
paths.
|
||||
"""
|
||||
if not token:
|
||||
return False
|
||||
return all(0x21 <= ord(c) <= 0x7E and c not in ('"', "\\") for c in token)
|
||||
|
||||
|
||||
def parse_www_authenticate_scope(header: str) -> tuple[str, ...]:
|
||||
"""Return the ``scope=...`` value as a tuple of individual scopes.
|
||||
|
||||
Splits on a single space per RFC 6749 section 3.3 (``scope-token``
|
||||
sequence). Returns ``()`` when the header is malformed or carries no
|
||||
``scope`` parameter.
|
||||
|
||||
Each token is validated against the RFC 6749 §3.3 ``scope-token``
|
||||
grammar (visible ASCII ``0x21..0x7E`` excluding ``"`` and ``\\``)
|
||||
via :func:`is_valid_scope_token` so that a malicious or buggy AS
|
||||
cannot smuggle CR/LF/tab/control bytes through a future log or
|
||||
notification path. Today scopes are JSON-encoded everywhere
|
||||
downstream so no concrete exploit exists, but the validation is
|
||||
cheap and forecloses regressions in structured-error rendering.
|
||||
"""
|
||||
params = parse_www_authenticate_bearer(header)
|
||||
value = params.get("scope")
|
||||
if not value:
|
||||
return ()
|
||||
return tuple(s for s in value.split(" ") if is_valid_scope_token(s))
|
||||
|
||||
|
||||
def parse_www_authenticate_error(header: str) -> str | None:
|
||||
"""Return the ``error=...`` value or ``None`` when absent.
|
||||
|
||||
The tokenizer naturally distinguishes ``error`` from
|
||||
``error_description`` / ``error_uri`` because ``_`` is not a valid
|
||||
token-character delimiter — they parse as separate keys.
|
||||
"""
|
||||
params = parse_www_authenticate_bearer(header)
|
||||
return params.get("error") or None
|
||||
File diff suppressed because it is too large
Load Diff
@@ -41,11 +41,18 @@ def save_message(
|
||||
tool_call_id: str | None = None,
|
||||
provider_data: str | None = None,
|
||||
tool_calls: str | None = None,
|
||||
source: str | None = None,
|
||||
reminders: str | None = None,
|
||||
) -> int:
|
||||
"""Log a message to the conversations table.
|
||||
|
||||
Returns the inserted row id, or ``0`` on failure (preserving the
|
||||
module's no-raise contract).
|
||||
|
||||
``source`` / ``reminders`` mirror the in-memory ``_source`` /
|
||||
``_reminders`` side-channels (``reminders`` JSON-encoded). Both
|
||||
default to ``None`` for the common case where no metacog payload
|
||||
rides the row.
|
||||
"""
|
||||
try:
|
||||
return get_storage().save_message(
|
||||
@@ -56,6 +63,8 @@ def save_message(
|
||||
tool_call_id,
|
||||
provider_data,
|
||||
tool_calls=tool_calls,
|
||||
source=source,
|
||||
reminders=reminders,
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to save message for ws=%s role=%s", ws_id, role, exc_info=True)
|
||||
@@ -70,10 +79,10 @@ def save_messages_bulk(rows: list[dict[str, Any]]) -> None:
|
||||
log.warning("Failed to bulk-save %d messages", len(rows), exc_info=True)
|
||||
|
||||
|
||||
def load_messages(ws_id: str) -> list[dict[str, Any]]:
|
||||
def load_messages(ws_id: str, *, repair: bool = True) -> list[dict[str, Any]]:
|
||||
"""Load messages for a workstream and reconstruct OpenAI message format."""
|
||||
try:
|
||||
return get_storage().load_messages(ws_id)
|
||||
return get_storage().load_messages(ws_id, repair=repair)
|
||||
except Exception:
|
||||
log.warning("Failed to load messages for ws=%s", ws_id, exc_info=True)
|
||||
return []
|
||||
|
||||
@@ -1,4 +1,14 @@
|
||||
"""Metacognitive prompting — situational nudges for proactive memory use."""
|
||||
"""Metacognitive prompting — situational nudges for proactive memory use.
|
||||
|
||||
Static nudge text templates (``NUDGE_*``), detection heuristics
|
||||
(``detect_correction``, ``detect_completion``), the :class:`RepeatDetector`
|
||||
streak counter, and the cooldown-aware :func:`should_nudge` /
|
||||
:func:`format_nudge` / :func:`format_idle_children_nudge` helpers.
|
||||
|
||||
The wake-trigger lifecycle (``IdleNudgeWatcher`` plus the
|
||||
``install_idle_nudge_watcher`` / ``shutdown_idle_nudge_watchers``
|
||||
lifespan helpers) lives in :mod:`turnstone.core.idle_nudge_watcher`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -108,8 +118,159 @@ _NUDGE_MAP: dict[str, str] = {
|
||||
"start": NUDGE_START,
|
||||
"tool_error": NUDGE_TOOL_ERROR,
|
||||
"repeat": NUDGE_REPEAT,
|
||||
# idle_children and watch_triggered carry no static body — the
|
||||
# per-fire text comes from a producer (``format_idle_children_nudge``
|
||||
# for the former, ``format_watch_message`` + ``sanitize_payload``
|
||||
# in the watch dispatch closure for the latter). Empty string
|
||||
# here keeps :func:`format_nudge` round-tripping honestly while
|
||||
# still letting :func:`should_nudge` and ``_NUDGE_MAP``-as-registry
|
||||
# consumers recognise the type.
|
||||
"idle_children": "",
|
||||
"watch_triggered": "",
|
||||
}
|
||||
|
||||
|
||||
# Display cap for the ``idle_children`` body — list at most this many
|
||||
# children inline, append "...and N more" overflow line beyond that.
|
||||
NUDGE_IDLE_CHILDREN_DISPLAY_CAP = 6
|
||||
|
||||
# Suggested ``wait_for_workstream(ws_ids=[...])`` cap — matches
|
||||
# ``WAIT_MAX_WS_IDS`` in :mod:`turnstone.core.coordinator_client` so the
|
||||
# emitted suggestion is callable as-is.
|
||||
NUDGE_IDLE_CHILDREN_WAIT_CAP = 32
|
||||
|
||||
NUDGE_IDLE_CHILDREN_HEADER = (
|
||||
"You went idle but still have active child workstreams. Either "
|
||||
"continue the user's work or block on the listed children "
|
||||
"explicitly:"
|
||||
)
|
||||
|
||||
|
||||
# ASCII control chars + Unicode steering vectors (bidi-override,
|
||||
# zero-width, line/paragraph separators, BOM, tag chars). Treated
|
||||
# uniformly as control chars and replaced with a space; angle-bracket
|
||||
# tag-breakers are stripped separately below. Defense-in-depth today
|
||||
# (self-injection within one user's tenant — children inherit parent
|
||||
# ``user_id`` and watch commands are user-supplied), but becomes
|
||||
# load-bearing the moment a producer ingests payloads from a different
|
||||
# trust boundary (a future watch trigger consuming external webhook
|
||||
# bodies, etc).
|
||||
#
|
||||
# Two classes, picked at the call site by the caller's structural
|
||||
# requirements:
|
||||
# * :data:`_NAME_CONTROL_CHARS` — STRICT: also strips TAB/LF/CR.
|
||||
# Used by :func:`sanitize_name` for single-line user-controlled
|
||||
# fields (workstream ``name`` rendered as bullet items by
|
||||
# :func:`format_idle_children_nudge` — a name with ``\n`` in it
|
||||
# would otherwise break the bullet's one-line structure and let
|
||||
# a malicious child name forge sibling rows).
|
||||
# * :data:`_PAYLOAD_CONTROL_CHARS` — PERMISSIVE: preserves TAB/LF/CR.
|
||||
# Used by :func:`sanitize_payload` for multi-line payloads where
|
||||
# line layout is part of the signal (watch shell output —
|
||||
# stripping LF/CR would collapse multi-line output to one line).
|
||||
_CONTROL_CHARS_TAIL = (
|
||||
r"\u200b-\u200f" # zero-width / LRM / RLM
|
||||
r"\u202a-\u202e" # bidi overrides
|
||||
r"\u2066-\u2069" # bidi isolates
|
||||
r"\u2028\u2029" # line / paragraph separator
|
||||
r"\ufeff" # BOM
|
||||
r"]"
|
||||
r"|[\U000e0000-\U000e007f]" # Unicode tag chars (separate range above BMP)
|
||||
)
|
||||
_NAME_CONTROL_CHARS = re.compile(
|
||||
r"[\x00-\x1f\x7f" + _CONTROL_CHARS_TAIL # ASCII control (incl. \t\n\r) + DEL
|
||||
)
|
||||
_PAYLOAD_CONTROL_CHARS = re.compile(
|
||||
r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f" + _CONTROL_CHARS_TAIL # ASCII control (skip \t\n\r) + DEL
|
||||
)
|
||||
_PAYLOAD_TAG_BREAKERS = re.compile(r"[<>]")
|
||||
|
||||
|
||||
def sanitize_name(text: str) -> str:
|
||||
"""Strict sanitiser for single-line user-controlled name fields.
|
||||
|
||||
Strips ASCII control chars **including** TAB/LF/CR plus Unicode
|
||||
steering vectors and angle-bracket tag breakers. Use for fields
|
||||
rendered as a single bullet item / label where embedded newlines
|
||||
would break the surrounding structure (workstream ``name`` in
|
||||
:func:`format_idle_children_nudge`, where a ``\\n`` in the name
|
||||
would otherwise forge a fake sibling bullet).
|
||||
"""
|
||||
if not text:
|
||||
return ""
|
||||
cleaned = _NAME_CONTROL_CHARS.sub(" ", text)
|
||||
cleaned = _PAYLOAD_TAG_BREAKERS.sub("", cleaned)
|
||||
return cleaned.strip()
|
||||
|
||||
|
||||
def sanitize_payload(text: str) -> str:
|
||||
"""Permissive sanitiser for multi-line user-controlled nudge payloads.
|
||||
|
||||
Used by ``format_watch_message`` output rendered into the
|
||||
``watch_triggered`` nudge body.
|
||||
|
||||
The wire-boundary :func:`escape_wrapper_tags` only protects the
|
||||
``<system-reminder>`` and ``<tool_output>`` envelopes; other
|
||||
angle-bracketed markers (``</thinking>``, ``<answer>``,
|
||||
``<artifact>``, …) and Unicode steering vectors (RTL override,
|
||||
zero-width chars, tag chars) can still steer some models. Strip
|
||||
both classes before interpolation — self-injection only today
|
||||
(watch commands are user-supplied), but the cost is one ``re.sub``
|
||||
per payload.
|
||||
|
||||
TAB / LF / CR are preserved (see ``_PAYLOAD_CONTROL_CHARS``) so
|
||||
multi-line shell output in watch payloads keeps its line structure.
|
||||
For single-line name fields where newlines would break surrounding
|
||||
structure, use :func:`sanitize_name` instead.
|
||||
"""
|
||||
if not text:
|
||||
return ""
|
||||
cleaned = _PAYLOAD_CONTROL_CHARS.sub(" ", text)
|
||||
cleaned = _PAYLOAD_TAG_BREAKERS.sub("", cleaned)
|
||||
return cleaned.strip()
|
||||
|
||||
|
||||
def format_idle_children_nudge(children: list[dict[str, str]]) -> str:
|
||||
"""Render the ``idle_children`` reminder body.
|
||||
|
||||
*children* is a list of dicts with ``ws_id``, ``name``, ``state``
|
||||
keys — the row-mapping shape coordinator-side storage exposes.
|
||||
Returns raw text *without* the ``<system-reminder>`` envelope; the
|
||||
side-channel :func:`_apply_reminders_for_provider` splice wraps it
|
||||
at the wire boundary.
|
||||
|
||||
User-controlled ``name`` strings get sanitized via
|
||||
:func:`sanitize_name` before interpolation so a workstream
|
||||
named ``</thinking>...`` can't steer the model's reasoning
|
||||
channels through the rendered body, and an embedded ``\\n`` in
|
||||
a name can't forge a fake sibling bullet.
|
||||
|
||||
Display caps at :data:`NUDGE_IDLE_CHILDREN_DISPLAY_CAP` with an
|
||||
overflow line; the trailing ``wait_for_workstream`` suggestion's
|
||||
``ws_ids`` list caps at :data:`NUDGE_IDLE_CHILDREN_WAIT_CAP`.
|
||||
Empty input returns the empty string so callers can short-circuit
|
||||
on ``if not text: return``.
|
||||
"""
|
||||
if not children:
|
||||
return ""
|
||||
lines = [NUDGE_IDLE_CHILDREN_HEADER, ""]
|
||||
shown = children[:NUDGE_IDLE_CHILDREN_DISPLAY_CAP]
|
||||
for c in shown:
|
||||
ws_id = c.get("ws_id", "")
|
||||
name = sanitize_name(c.get("name", "")) or "(unnamed)"
|
||||
state = c.get("state", "?")
|
||||
lines.append(f" - {ws_id[:8]} ({state}): {name}")
|
||||
overflow = len(children) - len(shown)
|
||||
if overflow > 0:
|
||||
lines.append(f" ...and {overflow} more")
|
||||
lines.append("")
|
||||
wait_ids = [c.get("ws_id", "") for c in children[:NUDGE_IDLE_CHILDREN_WAIT_CAP]]
|
||||
lines.append(
|
||||
f'To block on them: wait_for_workstream(ws_ids={wait_ids!r}, mode="any", timeout=120).'
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Detection heuristics — strong/weak tiers
|
||||
#
|
||||
@@ -194,6 +355,26 @@ def detect_completion(message: str) -> bool:
|
||||
return any(p.search(message) for p in _WEAK_COMPLETION)
|
||||
|
||||
|
||||
def _cooldown_allows(
|
||||
nudge_type: str,
|
||||
state: dict[str, float],
|
||||
*,
|
||||
cooldown_secs: int = _COOLDOWN_SECS,
|
||||
) -> bool:
|
||||
"""Read-only cooldown peek — does NOT record a fire timestamp.
|
||||
|
||||
Use this as a cheap pre-gate before expensive work (storage queries,
|
||||
message walks). The follow-up :func:`should_nudge` call re-checks
|
||||
cooldown AND records the timestamp atomically. A producer that
|
||||
races between this peek and ``should_nudge`` would just lose the
|
||||
fire to the other producer — benign.
|
||||
"""
|
||||
last = state.get(nudge_type)
|
||||
if last is None:
|
||||
return True
|
||||
return time.monotonic() - last >= cooldown_secs
|
||||
|
||||
|
||||
def should_nudge(
|
||||
nudge_type: str,
|
||||
state: dict[str, float],
|
||||
|
||||
@@ -0,0 +1,288 @@
|
||||
"""Thread-safe FIFO queue for metacognitive nudges with channel filtering.
|
||||
|
||||
Replaces the dual ``_pending_user_advisories`` / ``_pending_tool_advisories``
|
||||
list pair with a single channel-tagged queue per session. Producers
|
||||
(`_queue_user_advisory`, `_queue_tool_advisory`,
|
||||
`CoordinatorIdleObserver`, the future watch dispatcher) all enqueue onto
|
||||
the same queue with an explicit ``channel``; consumers
|
||||
(`_attach_pending_user_reminders` on user-message attach,
|
||||
`_collect_advisories` on tool-result wrap,
|
||||
`IdleNudgeWatcher` on workstream-IDLE) drain by channel filter.
|
||||
|
||||
Channels:
|
||||
* ``"user"`` — only drains at user-turn seams.
|
||||
* ``"tool"`` — only drains at tool-result seams.
|
||||
* ``"any"`` — drains at whichever seam fires first (used for
|
||||
wake-trigger-driven nudges that should not be pinned to a
|
||||
specific drain seam).
|
||||
|
||||
Drain preserves FIFO order; non-matching entries stay queued. Each
|
||||
entry can carry an optional ``valid_until`` predicate that drain
|
||||
evaluates outside the queue lock; entries whose predicate returns
|
||||
``False`` (or raises) are silently dropped without delivery — used by
|
||||
producers whose payload becomes stale if the underlying state changes
|
||||
between enqueue and drain (e.g. ``idle_children`` re-checks the active
|
||||
child set, dropping the nudge if every child finished while the queue
|
||||
sat). Operations are atomic under an internal :class:`threading.Lock`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from collections import deque
|
||||
from typing import TYPE_CHECKING, Any, Literal, NamedTuple
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
Channel = Literal["user", "tool", "any"]
|
||||
_VALID_CHANNELS: frozenset[str] = frozenset({"user", "tool", "any"})
|
||||
|
||||
# Module-level filter constants — most callers want one of these and
|
||||
# pre-allocating spares us a frozenset construction at every drain seam.
|
||||
USER_DRAIN: frozenset[str] = frozenset({"user", "any"})
|
||||
TOOL_DRAIN: frozenset[str] = frozenset({"tool", "any"})
|
||||
|
||||
|
||||
class _Entry(NamedTuple):
|
||||
nudge_type: str
|
||||
text: str
|
||||
channel: Channel
|
||||
valid_until: Callable[[], bool] | None = None
|
||||
# Producer-supplied optional fields that ride alongside ``text`` when
|
||||
# drained — used by ``watch_triggered`` to carry ``watch_name`` /
|
||||
# ``command`` / ``poll_count`` / ``max_polls`` / ``is_final`` into the
|
||||
# rendered reminder dict so the frontend can render a structured
|
||||
# ``.msg.watch-result`` card instead of a plain advisory bubble.
|
||||
# Other producers leave it ``None`` and consumers see only
|
||||
# ``{type, text}``. Atomicity guarantee: text + metadata land on the
|
||||
# same enqueue call, so a concurrent drain can't observe text without
|
||||
# the matching metadata.
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class NudgeQueue:
|
||||
"""Single-session FIFO queue with channel-tagged entries."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._items: deque[_Entry] = deque()
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def enqueue(
|
||||
self,
|
||||
nudge_type: str,
|
||||
text: str,
|
||||
channel: Channel,
|
||||
*,
|
||||
valid_until: Callable[[], bool] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Append a nudge. ``channel`` MUST be in :data:`_VALID_CHANNELS`.
|
||||
|
||||
``channel`` is required so producers ingesting untrusted text
|
||||
(future child workstream names, watch payloads) must pick a
|
||||
seam consciously rather than silently routing to whichever
|
||||
seam drains first via an implicit default.
|
||||
|
||||
If ``valid_until`` is provided, drain re-evaluates it before
|
||||
delivering the entry; a falsy result drops the entry silently
|
||||
(the producer's signal that the snapshot it enqueued is now
|
||||
stale). The predicate is called outside the queue lock so
|
||||
producers can do non-trivial work (e.g. re-querying storage
|
||||
for active children) without blocking other producers.
|
||||
|
||||
``metadata`` carries optional producer-specific fields that
|
||||
drain returns alongside ``(nudge_type, text)``. ``watch_triggered``
|
||||
uses it for ``watch_name`` / ``command`` / ``poll_count`` /
|
||||
``max_polls`` / ``is_final`` so the frontend can render a
|
||||
structured card; other producers leave it ``None``.
|
||||
"""
|
||||
if channel not in _VALID_CHANNELS:
|
||||
raise ValueError(f"channel={channel!r}; expected one of {sorted(_VALID_CHANNELS)}")
|
||||
with self._lock:
|
||||
self._items.append(_Entry(nudge_type, text, channel, valid_until, metadata))
|
||||
|
||||
def drain(
|
||||
self, channels: frozenset[str] | set[str]
|
||||
) -> list[tuple[str, str, dict[str, Any] | None]]:
|
||||
"""Drain entries whose channel is in ``channels``.
|
||||
|
||||
Entries with non-matching channels stay in the queue, in order.
|
||||
Returns ``(nudge_type, text, metadata)`` tuples in insertion
|
||||
order; ``metadata`` is the producer-supplied dict (or ``None``
|
||||
when unset).
|
||||
|
||||
Entries with a ``valid_until`` predicate get re-checked outside
|
||||
the queue lock; falsy / raising predicates drop the entry
|
||||
without delivering it. Already-removed-from-queue either way —
|
||||
dropped entries don't ride a future drain.
|
||||
"""
|
||||
with self._lock:
|
||||
if not self._items:
|
||||
return []
|
||||
# Fast path: every entry matches → swap deque rather than
|
||||
# walk + partition + per-entry append. This is the common
|
||||
# case in practice since the chat loop's drain seams use
|
||||
# ``USER_DRAIN`` / ``TOOL_DRAIN`` (channel + "any") and
|
||||
# most queues hold only one channel's entries at a time.
|
||||
if all(entry.channel in channels for entry in self._items):
|
||||
candidates: list[_Entry] = list(self._items)
|
||||
self._items = deque()
|
||||
else:
|
||||
kept: deque[_Entry] = deque()
|
||||
candidates = []
|
||||
for entry in self._items:
|
||||
if entry.channel in channels:
|
||||
candidates.append(entry)
|
||||
else:
|
||||
kept.append(entry)
|
||||
self._items = kept
|
||||
# Predicates evaluate outside the lock — they may do storage
|
||||
# I/O or other work that shouldn't block other producers /
|
||||
# the drain consumer's other queues.
|
||||
out: list[tuple[str, str, dict[str, Any] | None]] = []
|
||||
for entry in candidates:
|
||||
if entry.valid_until is None:
|
||||
out.append((entry.nudge_type, entry.text, entry.metadata))
|
||||
continue
|
||||
try:
|
||||
if entry.valid_until():
|
||||
out.append((entry.nudge_type, entry.text, entry.metadata))
|
||||
except Exception:
|
||||
# Predicate raising is treated as "no longer valid" —
|
||||
# drop silently rather than letting one bad predicate
|
||||
# poison the whole drain batch.
|
||||
pass
|
||||
return out
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Current depth. Used by the future IdleNudgeWatcher gate."""
|
||||
with self._lock:
|
||||
return len(self._items)
|
||||
|
||||
def clear(self) -> int:
|
||||
"""Drop every entry; return the count cleared. Used in cancel paths."""
|
||||
with self._lock:
|
||||
n = len(self._items)
|
||||
self._items.clear()
|
||||
return n
|
||||
|
||||
def count_by_type(self, nudge_type: str, channel: Channel | None = None) -> int:
|
||||
"""Return the number of queued entries matching ``nudge_type``.
|
||||
|
||||
With ``channel=None`` counts across all channels; with a specific
|
||||
channel filters to that channel only. Walks ``_items`` once
|
||||
under the lock without materialising tuples — cheaper than
|
||||
``len(pending(channel=...))`` for callers that only need the
|
||||
count (e.g. the watch dispatcher's soft-cap pre-check).
|
||||
Producer-side soft caps that pair this with
|
||||
:meth:`drop_oldest_by_type` should pass the same ``channel`` to
|
||||
both halves so the count snapshot and the drop walk over the
|
||||
same entry set.
|
||||
"""
|
||||
with self._lock:
|
||||
if channel is None:
|
||||
return sum(1 for e in self._items if e.nudge_type == nudge_type)
|
||||
return sum(
|
||||
1 for e in self._items if e.nudge_type == nudge_type and e.channel == channel
|
||||
)
|
||||
|
||||
def drop_oldest_by_type(self, nudge_type: str, channel: Channel | None = None) -> bool:
|
||||
"""Remove the earliest-enqueued entry whose type matches ``nudge_type``.
|
||||
|
||||
With ``channel=None`` searches across all channels; with a specific
|
||||
channel filters to that channel only. Returns ``True`` if an
|
||||
entry was dropped, ``False`` if no matching entry was found.
|
||||
The call itself is atomic under the queue lock; producers that
|
||||
need an atomic count-and-drop pair (no interleave with concurrent
|
||||
drains) should use :meth:`cap_at_or_drop_oldest` instead.
|
||||
"""
|
||||
with self._lock:
|
||||
for i, entry in enumerate(self._items):
|
||||
if entry.nudge_type != nudge_type:
|
||||
continue
|
||||
if channel is not None and entry.channel != channel:
|
||||
continue
|
||||
del self._items[i]
|
||||
return True
|
||||
return False
|
||||
|
||||
def cap_at_or_drop_oldest(
|
||||
self,
|
||||
nudge_type: str,
|
||||
max_depth: int,
|
||||
channel: Channel | None = None,
|
||||
) -> bool:
|
||||
"""If queued ``nudge_type`` entries reach ``max_depth``, drop the
|
||||
earliest matching entry — under a single lock acquisition so a
|
||||
concurrent drain can't slip between the count and the drop.
|
||||
|
||||
Returns ``True`` iff a drop happened. Producers with a per-type
|
||||
soft cap call this on the enqueue path; ``max_depth <= 0`` is a
|
||||
defensive no-op returning ``False``.
|
||||
"""
|
||||
if max_depth <= 0:
|
||||
return False
|
||||
with self._lock:
|
||||
oldest_index = -1
|
||||
count = 0
|
||||
for i, entry in enumerate(self._items):
|
||||
if entry.nudge_type != nudge_type:
|
||||
continue
|
||||
if channel is not None and entry.channel != channel:
|
||||
continue
|
||||
if oldest_index == -1:
|
||||
oldest_index = i
|
||||
count += 1
|
||||
if count >= max_depth:
|
||||
del self._items[oldest_index]
|
||||
return True
|
||||
return False
|
||||
|
||||
def pending(self, channel: Channel | None = None) -> list[tuple[str, str]]:
|
||||
"""Non-mutating snapshot for tests / introspection.
|
||||
|
||||
With ``channel=None`` returns every queued entry as
|
||||
``(nudge_type, text)`` tuples in insertion order; with a
|
||||
specific channel filters to that channel only. Production
|
||||
code that wants to *consume* entries should call :meth:`drain`
|
||||
instead — pending entries are by definition unconsumed and
|
||||
will redraw at the next matching seam.
|
||||
|
||||
``metadata`` is intentionally NOT projected here — tests that
|
||||
need to assert producer-specific fields call
|
||||
:meth:`pending_with_metadata` (or :meth:`drain` directly).
|
||||
"""
|
||||
with self._lock:
|
||||
if channel is None:
|
||||
return [(e.nudge_type, e.text) for e in self._items]
|
||||
return [(e.nudge_type, e.text) for e in self._items if e.channel == channel]
|
||||
|
||||
def pending_with_metadata(
|
||||
self, channel: Channel | None = None
|
||||
) -> list[tuple[str, str, dict[str, Any] | None]]:
|
||||
"""Non-mutating snapshot including each entry's ``metadata``.
|
||||
|
||||
Used by tests / introspection paths that need to assert
|
||||
producer-specific optional fields (e.g. the watch dispatcher's
|
||||
``watch_name`` / ``command`` / ``poll_count`` payload). Production
|
||||
consumers should still call :meth:`drain`.
|
||||
"""
|
||||
with self._lock:
|
||||
if channel is None:
|
||||
return [(e.nudge_type, e.text, e.metadata) for e in self._items]
|
||||
return [(e.nudge_type, e.text, e.metadata) for e in self._items if e.channel == channel]
|
||||
|
||||
def has_pending(self, channels: frozenset[str] | set[str]) -> bool:
|
||||
"""Short-circuiting existence check.
|
||||
|
||||
Returns ``True`` as soon as a queued entry's channel matches
|
||||
``channels``. Cheaper than :meth:`pending` for callers that
|
||||
only need a boolean — used by
|
||||
:class:`turnstone.core.idle_nudge_watcher.IdleNudgeWatcher` to
|
||||
gate wake dispatch on whether ``USER_DRAIN`` would actually
|
||||
deliver anything before paying the worker-thread spawn. No
|
||||
list allocation, lock released on first match.
|
||||
"""
|
||||
with self._lock:
|
||||
return any(e.channel in channels for e in self._items)
|
||||
@@ -0,0 +1,224 @@
|
||||
"""Shared SSRF and same-origin validation for OAuth/OIDC endpoint URLs.
|
||||
|
||||
Extracted from :mod:`turnstone.core.oidc` so the per-(user, server) MCP
|
||||
OAuth flow (see :mod:`turnstone.core.mcp_oauth`) can reuse the exact same
|
||||
guards without depending on the OIDC module.
|
||||
|
||||
The canonical exception is :class:`OAuthSSRFError`. The OIDC module wraps
|
||||
calls to these helpers and re-raises ``OIDCError`` so its public API is
|
||||
unchanged. The MCP OAuth module catches :class:`OAuthSSRFError` directly.
|
||||
|
||||
DNS-rebinding limitation: this module resolves the hostname during
|
||||
validation, but the subsequent ``httpx`` call resolves again. A hostname
|
||||
the operator points at could in principle rebind between the two resolves
|
||||
to expose an internal address. Callers must ensure the AS / IdP hostname
|
||||
is operator-controlled — the SSRF guard prevents private-IP responses for
|
||||
hostnames the operator points at, but does not prevent rebinding by a
|
||||
hostile DNS authority. Pinning a single resolution into the ``httpx``
|
||||
transport is a future hardening step.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import socket
|
||||
import urllib.parse
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Trusted-host allowlist for well-known multi-origin IdPs / authorization
|
||||
# servers whose discovery documents legitimately reference endpoints on
|
||||
# hostnames distinct from the issuer hostname. eTLD+1 matching does not
|
||||
# work here (e.g. google.com vs googleapis.com), so an explicit allow-map
|
||||
# is the only safe option.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
KNOWN_TRUSTED_OAUTH_ENDPOINT_HOSTS: dict[str, frozenset[str]] = {
|
||||
"accounts.google.com": frozenset(
|
||||
{
|
||||
"accounts.google.com",
|
||||
"oauth2.googleapis.com",
|
||||
"www.googleapis.com",
|
||||
"openidconnect.googleapis.com",
|
||||
}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class OAuthSSRFError(Exception):
|
||||
"""Raised when an SSRF/same-origin validation fails.
|
||||
|
||||
OIDC callers wrap this and re-raise as ``OIDCError`` to preserve the
|
||||
existing public API.
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def is_localhost(hostname: str) -> bool:
|
||||
"""Return True if *hostname* refers to the loopback interface."""
|
||||
return hostname in ("localhost", "127.0.0.1", "::1") or hostname.endswith(".localhost")
|
||||
|
||||
|
||||
def sanitize_log_text(text: str, limit: int = 200) -> str:
|
||||
"""Escape control characters and truncate untrusted text for log/audit inclusion.
|
||||
|
||||
Untrusted bytes (e.g. an AS error body, ``error_description`` from a
|
||||
callback redirect) embedded in log lines or exception messages must not
|
||||
be able to forge fake log records via CR/LF or hide content via NULs /
|
||||
other control characters. ``unicode_escape`` renders these as visible
|
||||
``\\r``, ``\\n``, ``\\x00`` etc., and *limit* caps the *rendered* length.
|
||||
|
||||
Shared with the OIDC module — its private ``_sanitize_log_text`` is a
|
||||
legacy alias that forwards here.
|
||||
"""
|
||||
if not text:
|
||||
return ""
|
||||
return text.encode("unicode_escape").decode("ascii")[:limit]
|
||||
|
||||
|
||||
def effective_port(parsed: urllib.parse.ParseResult) -> int | None:
|
||||
"""Return the explicit port if set, else the scheme default."""
|
||||
if parsed.port is not None:
|
||||
return parsed.port
|
||||
return {"http": 80, "https": 443}.get(parsed.scheme)
|
||||
|
||||
|
||||
def validate_url_no_ssrf(url: str, *, allow_http: bool) -> urllib.parse.ParseResult:
|
||||
"""Run the scheme/userinfo/SSRF checks shared by issuer and discovered URLs.
|
||||
|
||||
Returns the parsed URL on success. Raises :class:`OAuthSSRFError` on
|
||||
failure. The ``allow_http`` flag is the only knob: when ``True``,
|
||||
``http://`` is accepted *if* the hostname is also a localhost form;
|
||||
when ``False``, only ``https://`` is accepted.
|
||||
"""
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
raise OAuthSSRFError(f"endpoint URL has no hostname: {url}")
|
||||
|
||||
if parsed.username or parsed.password:
|
||||
raise OAuthSSRFError("endpoint URL must not contain embedded credentials (userinfo)")
|
||||
|
||||
if parsed.scheme != "https":
|
||||
if allow_http and parsed.scheme == "http" and is_localhost(hostname):
|
||||
pass
|
||||
else:
|
||||
raise OAuthSSRFError(f"endpoint URL must use HTTPS (got {parsed.scheme}://): {url}")
|
||||
|
||||
try:
|
||||
addr_infos = socket.getaddrinfo(hostname, None, proto=socket.IPPROTO_TCP)
|
||||
except socket.gaierror as exc:
|
||||
raise OAuthSSRFError(f"endpoint hostname cannot be resolved: {hostname}") from exc
|
||||
|
||||
for _family, _type, _proto, _canonname, sockaddr in addr_infos:
|
||||
try:
|
||||
addr = ipaddress.ip_address(sockaddr[0])
|
||||
except ValueError as exc:
|
||||
raise OAuthSSRFError(
|
||||
f"endpoint hostname resolved to invalid IP {sockaddr[0]!r}: {hostname}"
|
||||
) from exc
|
||||
if not addr.is_global and not is_localhost(hostname):
|
||||
raise OAuthSSRFError(f"endpoint URL resolves to non-public address ({addr}): {url}")
|
||||
|
||||
return parsed
|
||||
|
||||
|
||||
def validate_discovered_endpoint(
|
||||
url: str,
|
||||
issuer_parsed: urllib.parse.ParseResult,
|
||||
*,
|
||||
allow_http: bool,
|
||||
trusted_endpoint_hosts: frozenset[str],
|
||||
) -> None:
|
||||
"""Validate an endpoint URL pulled from an OIDC/OAuth discovery document.
|
||||
|
||||
Applies :func:`validate_url_no_ssrf` plus the same-origin / trusted-host
|
||||
constraint: the endpoint host must equal the issuer host, be in the
|
||||
well-known trust map, or be in the operator-supplied
|
||||
``trusted_endpoint_hosts``. Effective port (with scheme defaults
|
||||
applied) and scheme must match the issuer.
|
||||
|
||||
Raises :class:`OAuthSSRFError` on validation failure.
|
||||
"""
|
||||
parsed = validate_url_no_ssrf(url, allow_http=allow_http)
|
||||
|
||||
issuer_hostname = (issuer_parsed.hostname or "").lower()
|
||||
endpoint_hostname = (parsed.hostname or "").lower()
|
||||
|
||||
if parsed.scheme != issuer_parsed.scheme:
|
||||
raise OAuthSSRFError(
|
||||
f"discovered endpoint scheme ({parsed.scheme}) "
|
||||
f"does not match issuer ({issuer_parsed.scheme}): {url}"
|
||||
)
|
||||
|
||||
known_trusted = KNOWN_TRUSTED_OAUTH_ENDPOINT_HOSTS.get(issuer_hostname, frozenset())
|
||||
host_allowed = (
|
||||
endpoint_hostname == issuer_hostname
|
||||
or endpoint_hostname in known_trusted
|
||||
or endpoint_hostname in trusted_endpoint_hosts
|
||||
)
|
||||
if not host_allowed:
|
||||
raise OAuthSSRFError(
|
||||
f"discovered endpoint host ({endpoint_hostname}) "
|
||||
f"does not match issuer ({issuer_hostname}) and is not trusted: {url}"
|
||||
)
|
||||
|
||||
endpoint_port = effective_port(parsed)
|
||||
issuer_port = effective_port(issuer_parsed)
|
||||
if endpoint_port != issuer_port:
|
||||
raise OAuthSSRFError(
|
||||
f"discovered endpoint port ({endpoint_port}) "
|
||||
f"does not match issuer ({issuer_port}): {url}"
|
||||
)
|
||||
|
||||
|
||||
async def validate_url_no_ssrf_async(url: str, *, allow_http: bool) -> urllib.parse.ParseResult:
|
||||
"""Async variant of :func:`validate_url_no_ssrf` for hot-path callers.
|
||||
|
||||
The synchronous variant calls ``socket.getaddrinfo``, which blocks
|
||||
the event loop. Async OAuth flows (notably
|
||||
:mod:`turnstone.core.mcp_oauth`) wrap their validation calls in
|
||||
:func:`asyncio.to_thread` to keep the loop responsive. This wrapper
|
||||
centralises that wrapping so callers don't repeat the idiom.
|
||||
"""
|
||||
return await asyncio.to_thread(validate_url_no_ssrf, url, allow_http=allow_http)
|
||||
|
||||
|
||||
async def validate_discovered_endpoint_async(
|
||||
url: str,
|
||||
issuer_parsed: urllib.parse.ParseResult,
|
||||
*,
|
||||
allow_http: bool,
|
||||
trusted_endpoint_hosts: frozenset[str],
|
||||
) -> None:
|
||||
"""Async variant of :func:`validate_discovered_endpoint`."""
|
||||
await asyncio.to_thread(
|
||||
validate_discovered_endpoint,
|
||||
url,
|
||||
issuer_parsed,
|
||||
allow_http=allow_http,
|
||||
trusted_endpoint_hosts=trusted_endpoint_hosts,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"KNOWN_TRUSTED_OAUTH_ENDPOINT_HOSTS",
|
||||
"OAuthSSRFError",
|
||||
"effective_port",
|
||||
"is_localhost",
|
||||
"sanitize_log_text",
|
||||
"validate_discovered_endpoint",
|
||||
"validate_discovered_endpoint_async",
|
||||
"validate_url_no_ssrf",
|
||||
"validate_url_no_ssrf_async",
|
||||
]
|
||||
+431
-128
@@ -7,22 +7,34 @@ event loop.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import dataclasses
|
||||
import hashlib
|
||||
import ipaddress
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import socket
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Mapping
|
||||
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.oauth_ssrf import (
|
||||
OAuthSSRFError,
|
||||
is_localhost,
|
||||
)
|
||||
from turnstone.core.oauth_ssrf import (
|
||||
validate_discovered_endpoint as _ssrf_validate_discovered_endpoint,
|
||||
)
|
||||
from turnstone.core.oauth_ssrf import (
|
||||
validate_url_no_ssrf as _ssrf_validate_url_no_ssrf,
|
||||
)
|
||||
|
||||
log = get_logger(__name__)
|
||||
|
||||
@@ -30,6 +42,11 @@ log = get_logger(__name__)
|
||||
# Not a valid bcrypt hash -- verify_password() always rejects it.
|
||||
OIDC_PASSWORD_SENTINEL = "!oidc"
|
||||
|
||||
# Lifetime of an OIDC authorization-flow pending-state row. Bounds the window
|
||||
# between /authorize and /callback; longer than typical IdP latency, shorter
|
||||
# than a stale browser tab.
|
||||
OIDC_STATE_TTL_SECONDS = 300
|
||||
|
||||
# Sanitisation pattern: only keep safe username characters.
|
||||
_USERNAME_SAFE_RE = re.compile(r"[^a-zA-Z0-9._-]")
|
||||
|
||||
@@ -49,9 +66,8 @@ _ALLOWED_ID_TOKEN_ALGS = [
|
||||
"PS512",
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Exception
|
||||
# Exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@@ -59,6 +75,22 @@ class OIDCError(Exception):
|
||||
"""Raised when an OIDC operation fails."""
|
||||
|
||||
|
||||
class OIDCKeyNotFoundError(OIDCError):
|
||||
"""Raised when an ID token's signing key is absent from the cached JWKS.
|
||||
|
||||
Distinguishing this from generic OIDCError lets the callback retry once
|
||||
after re-fetching JWKS (key rotation), without depending on substring
|
||||
matching of the error message.
|
||||
"""
|
||||
|
||||
|
||||
def _sanitize_log_text(s: str, limit: int) -> str:
|
||||
"""Legacy alias for the shared :func:`oauth_ssrf.sanitize_log_text`."""
|
||||
from turnstone.core.oauth_ssrf import sanitize_log_text
|
||||
|
||||
return sanitize_log_text(s, limit)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configuration
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -66,7 +98,20 @@ class OIDCError(Exception):
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OIDCConfig:
|
||||
"""OIDC provider configuration -- immutable after startup."""
|
||||
"""OIDC provider configuration -- immutable after startup.
|
||||
|
||||
The dataclass has a two-phase lifecycle:
|
||||
|
||||
Startup-config fields (set by :func:`load_oidc_config`):
|
||||
``enabled``, ``issuer``, ``client_id``, ``client_secret``, ``scopes``,
|
||||
``provider_name``, ``role_claim``, ``role_map``, ``password_enabled``,
|
||||
``redirect_base``, ``trusted_endpoint_hosts``.
|
||||
|
||||
Discovery-derived fields (set by :func:`discover_oidc`; empty before
|
||||
discovery completes):
|
||||
``authorization_endpoint``, ``token_endpoint``, ``userinfo_endpoint``,
|
||||
``jwks_uri``.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
issuer: str = ""
|
||||
@@ -78,6 +123,7 @@ class OIDCConfig:
|
||||
role_map: dict[str, str] = field(default_factory=dict)
|
||||
password_enabled: bool = True
|
||||
redirect_base: str = ""
|
||||
trusted_endpoint_hosts: tuple[str, ...] = ()
|
||||
# Discovered from .well-known/openid-configuration
|
||||
authorization_endpoint: str = ""
|
||||
token_endpoint: str = ""
|
||||
@@ -98,6 +144,32 @@ def _parse_role_map(raw: str) -> dict[str, str]:
|
||||
return result
|
||||
|
||||
|
||||
def _parse_trusted_endpoint_hosts(raw: str) -> tuple[str, ...]:
|
||||
"""Parse a comma-separated host list into a normalised tuple."""
|
||||
hosts: list[str] = []
|
||||
for entry in raw.split(","):
|
||||
host = entry.strip().lower()
|
||||
if host:
|
||||
hosts.append(host)
|
||||
return tuple(hosts)
|
||||
|
||||
|
||||
def _env_or_cfg_str(env_name: str, cfg: Mapping[str, Any], key: str, default: str = "") -> str:
|
||||
"""Resolve a string field: env var (stripped, non-empty) wins, else config, else default."""
|
||||
val = os.environ.get(env_name, "").strip()
|
||||
if not val:
|
||||
val = str(cfg.get(key, default)).strip()
|
||||
return val
|
||||
|
||||
|
||||
def _env_or_cfg_bool(env_name: str, cfg: Mapping[str, Any], key: str, default: bool) -> bool:
|
||||
"""Resolve a boolean field: env var (stripped) wins when set, else config, else default."""
|
||||
raw = os.environ.get(env_name, "").strip().lower()
|
||||
if raw:
|
||||
return raw in ("true", "1", "yes")
|
||||
return bool(cfg.get(key, default))
|
||||
|
||||
|
||||
def load_oidc_config() -> OIDCConfig:
|
||||
"""Build :class:`OIDCConfig` from env vars with config.toml fallback.
|
||||
|
||||
@@ -108,30 +180,15 @@ def load_oidc_config() -> OIDCConfig:
|
||||
|
||||
cfg = load_config("oidc")
|
||||
|
||||
# Start with config.toml values, then override with env vars.
|
||||
issuer = os.environ.get("TURNSTONE_OIDC_ISSUER", "").strip()
|
||||
if not issuer:
|
||||
issuer = str(cfg.get("issuer", "")).strip()
|
||||
|
||||
client_id = os.environ.get("TURNSTONE_OIDC_CLIENT_ID", "").strip()
|
||||
if not client_id:
|
||||
client_id = str(cfg.get("client_id", "")).strip()
|
||||
|
||||
client_secret = os.environ.get("TURNSTONE_OIDC_CLIENT_SECRET", "").strip()
|
||||
if not client_secret:
|
||||
client_secret = str(cfg.get("client_secret", "")).strip()
|
||||
|
||||
scopes = os.environ.get("TURNSTONE_OIDC_SCOPES", "").strip()
|
||||
if not scopes:
|
||||
scopes = str(cfg.get("scopes", "openid email profile")).strip()
|
||||
|
||||
provider_name = os.environ.get("TURNSTONE_OIDC_PROVIDER_NAME", "").strip()
|
||||
if not provider_name:
|
||||
provider_name = str(cfg.get("provider_name", "SSO")).strip()
|
||||
|
||||
role_claim = os.environ.get("TURNSTONE_OIDC_ROLE_CLAIM", "").strip()
|
||||
if not role_claim:
|
||||
role_claim = str(cfg.get("role_claim", "")).strip()
|
||||
issuer = _env_or_cfg_str("TURNSTONE_OIDC_ISSUER", cfg, "issuer")
|
||||
client_id = _env_or_cfg_str("TURNSTONE_OIDC_CLIENT_ID", cfg, "client_id")
|
||||
client_secret = _env_or_cfg_str("TURNSTONE_OIDC_CLIENT_SECRET", cfg, "client_secret")
|
||||
scopes = _env_or_cfg_str("TURNSTONE_OIDC_SCOPES", cfg, "scopes", "openid email profile")
|
||||
provider_name = _env_or_cfg_str("TURNSTONE_OIDC_PROVIDER_NAME", cfg, "provider_name", "SSO")
|
||||
role_claim = _env_or_cfg_str("TURNSTONE_OIDC_ROLE_CLAIM", cfg, "role_claim")
|
||||
password_enabled = _env_or_cfg_bool(
|
||||
"TURNSTONE_OIDC_PASSWORD_ENABLED", cfg, "password_enabled", True
|
||||
)
|
||||
|
||||
# Role map: env var is "admin:builtin-admin,eng:builtin-operator"
|
||||
role_map_raw = os.environ.get("TURNSTONE_OIDC_ROLE_MAP", "").strip()
|
||||
@@ -141,16 +198,21 @@ def load_oidc_config() -> OIDCConfig:
|
||||
cfg_role_map = cfg.get("role_map", {})
|
||||
role_map = dict(cfg_role_map) if isinstance(cfg_role_map, dict) else {}
|
||||
|
||||
password_raw = os.environ.get("TURNSTONE_OIDC_PASSWORD_ENABLED", "").strip().lower()
|
||||
if password_raw:
|
||||
password_enabled = password_raw in ("true", "1", "yes")
|
||||
trusted_hosts_raw = os.environ.get("TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS", "").strip()
|
||||
if trusted_hosts_raw:
|
||||
trusted_endpoint_hosts = _parse_trusted_endpoint_hosts(trusted_hosts_raw)
|
||||
else:
|
||||
password_enabled = bool(cfg.get("password_enabled", True))
|
||||
cfg_trusted = cfg.get("trusted_endpoint_hosts", "")
|
||||
if isinstance(cfg_trusted, list):
|
||||
trusted_endpoint_hosts = _parse_trusted_endpoint_hosts(
|
||||
",".join(str(h) for h in cfg_trusted)
|
||||
)
|
||||
else:
|
||||
trusted_endpoint_hosts = _parse_trusted_endpoint_hosts(str(cfg_trusted))
|
||||
|
||||
redirect_base = os.environ.get("TURNSTONE_OIDC_REDIRECT_BASE", "").strip()
|
||||
if not redirect_base:
|
||||
redirect_base = str(cfg.get("redirect_base", "")).strip()
|
||||
redirect_base = redirect_base.rstrip("/")
|
||||
redirect_base = _env_or_cfg_str("TURNSTONE_OIDC_REDIRECT_BASE", cfg, "redirect_base").rstrip(
|
||||
"/"
|
||||
)
|
||||
if redirect_base:
|
||||
parsed = urllib.parse.urlparse(redirect_base)
|
||||
if parsed.scheme not in ("https", "http"):
|
||||
@@ -212,17 +274,25 @@ def load_oidc_config() -> OIDCConfig:
|
||||
role_map=role_map,
|
||||
password_enabled=password_enabled,
|
||||
redirect_base=redirect_base,
|
||||
trusted_endpoint_hosts=trusted_endpoint_hosts,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSRF validation
|
||||
#
|
||||
# The actual checks live in :mod:`turnstone.core.oauth_ssrf` so the MCP OAuth
|
||||
# flow can reuse them. The wrappers below preserve OIDC's public API by
|
||||
# converting :class:`OAuthSSRFError` to :class:`OIDCError`.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _is_localhost(hostname: str) -> bool:
|
||||
"""Return True if *hostname* refers to the loopback interface."""
|
||||
return hostname in ("localhost", "127.0.0.1", "::1") or hostname.endswith(".localhost")
|
||||
def _validate_url_no_ssrf(url: str, *, allow_http: bool) -> urllib.parse.ParseResult:
|
||||
"""OIDC-flavoured wrapper around :func:`oauth_ssrf.validate_url_no_ssrf`."""
|
||||
try:
|
||||
return _ssrf_validate_url_no_ssrf(url, allow_http=allow_http)
|
||||
except OAuthSSRFError as exc:
|
||||
raise OIDCError(str(exc)) from exc
|
||||
|
||||
|
||||
def validate_issuer_url(url: str) -> None:
|
||||
@@ -235,39 +305,44 @@ def validate_issuer_url(url: str) -> None:
|
||||
|
||||
Raises :class:`OIDCError` on validation failure.
|
||||
"""
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
_validate_url_no_ssrf(url, allow_http=True)
|
||||
|
||||
# Require a hostname.
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
raise OIDCError(f"OIDC issuer URL has no hostname: {url}")
|
||||
|
||||
# Reject embedded credentials — redact userinfo from error message.
|
||||
if parsed.username or parsed.password:
|
||||
raise OIDCError("OIDC issuer URL must not contain embedded credentials (userinfo)")
|
||||
def validate_discovered_endpoint(
|
||||
url: str,
|
||||
issuer_parsed: urllib.parse.ParseResult,
|
||||
*,
|
||||
allow_http: bool,
|
||||
trusted_endpoint_hosts: frozenset[str],
|
||||
) -> None:
|
||||
"""Validate an endpoint pulled from an IdP discovery document.
|
||||
|
||||
# Require HTTPS (allow HTTP only for localhost development).
|
||||
if parsed.scheme != "https":
|
||||
if parsed.scheme == "http" and _is_localhost(hostname):
|
||||
pass # Allow http://localhost for dev
|
||||
else:
|
||||
raise OIDCError(f"OIDC issuer URL must use HTTPS (got {parsed.scheme}://): {url}")
|
||||
Applies the same scheme/userinfo/SSRF rules as :func:`validate_issuer_url`,
|
||||
then constrains the host: by default the endpoint must share the issuer's
|
||||
hostname. Strict equality is intentional — a hostile or compromised IdP
|
||||
must not be able to redirect ``token_endpoint`` to a third-party host where
|
||||
``client_secret`` would leak. Multi-origin IdPs (e.g. Google) are
|
||||
accommodated via :data:`turnstone.core.oauth_ssrf.KNOWN_TRUSTED_OAUTH_ENDPOINT_HOSTS`
|
||||
plus an operator-configurable ``trusted_endpoint_hosts`` list.
|
||||
|
||||
# Resolve hostname and reject non-globally-routable addresses.
|
||||
The scheme must match the issuer's scheme, and the *effective* port (with
|
||||
scheme defaults applied) must match — so ``https://host`` and
|
||||
``https://host:443`` are treated as identical.
|
||||
|
||||
``allow_http`` should track whether the *issuer* URL was localhost, so the
|
||||
whole flow is allowed to be HTTP only in dev mode.
|
||||
|
||||
Raises :class:`OIDCError` on validation failure.
|
||||
"""
|
||||
try:
|
||||
addr_infos = socket.getaddrinfo(hostname, None, proto=socket.IPPROTO_TCP)
|
||||
except socket.gaierror as exc:
|
||||
raise OIDCError(f"OIDC issuer hostname cannot be resolved: {hostname}") from exc
|
||||
|
||||
for _family, _type, _proto, _canonname, sockaddr in addr_infos:
|
||||
try:
|
||||
addr = ipaddress.ip_address(sockaddr[0])
|
||||
except ValueError as exc:
|
||||
raise OIDCError(
|
||||
f"OIDC issuer hostname resolved to invalid IP {sockaddr[0]!r}: {hostname}"
|
||||
) from exc
|
||||
if not addr.is_global and not _is_localhost(hostname):
|
||||
raise OIDCError(f"OIDC issuer URL resolves to non-public address ({addr}): {url}")
|
||||
_ssrf_validate_discovered_endpoint(
|
||||
url,
|
||||
issuer_parsed,
|
||||
allow_http=allow_http,
|
||||
trusted_endpoint_hosts=trusted_endpoint_hosts,
|
||||
)
|
||||
except OAuthSSRFError as exc:
|
||||
raise OIDCError(str(exc)) from exc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -275,28 +350,50 @@ def validate_issuer_url(url: str) -> None:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def discover_oidc(config: OIDCConfig) -> OIDCConfig:
|
||||
async def discover_oidc(
|
||||
config: OIDCConfig,
|
||||
*,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> OIDCConfig:
|
||||
"""Fetch OIDC discovery document and return updated config with endpoints.
|
||||
|
||||
On failure, logs a warning and returns config with ``enabled=False``.
|
||||
|
||||
A long-lived ``client`` may be supplied to amortise TLS / connection
|
||||
setup across calls; when ``None`` a transient client is used (the
|
||||
legacy shape, kept so tests don't need lifecycle management).
|
||||
"""
|
||||
if not config.issuer:
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
try:
|
||||
validate_issuer_url(config.issuer)
|
||||
issuer_parsed = _validate_url_no_ssrf(config.issuer, allow_http=True)
|
||||
except OIDCError as exc:
|
||||
log.warning("OIDC issuer URL rejected: %s", exc)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
url = config.issuer.rstrip("/") + "/.well-known/openid-configuration"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.get(url)
|
||||
if client is not None:
|
||||
resp = await client.get(url, timeout=10.0)
|
||||
resp.raise_for_status()
|
||||
doc = resp.json()
|
||||
except Exception as exc:
|
||||
log.warning("OIDC discovery failed for %s: %s", config.issuer, exc)
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as transient:
|
||||
resp = await transient.get(url)
|
||||
resp.raise_for_status()
|
||||
doc = resp.json()
|
||||
except (httpx.HTTPError, ValueError, KeyError) as exc:
|
||||
# ValueError covers json.JSONDecodeError (subclass).
|
||||
log.warning("OIDC discovery failed for %s: %s", config.issuer, exc, exc_info=True)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
if not isinstance(doc, dict):
|
||||
log.warning(
|
||||
"OIDC discovery document for %s is not a JSON object (got %s)",
|
||||
config.issuer,
|
||||
type(doc).__name__,
|
||||
)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
authorization_endpoint = str(doc.get("authorization_endpoint", ""))
|
||||
@@ -311,6 +408,41 @@ async def discover_oidc(config: OIDCConfig) -> OIDCConfig:
|
||||
)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
allow_http = is_localhost(issuer_parsed.hostname or "")
|
||||
trusted_hosts = frozenset(h.lower() for h in config.trusted_endpoint_hosts)
|
||||
required = (
|
||||
("authorization_endpoint", authorization_endpoint),
|
||||
("token_endpoint", token_endpoint),
|
||||
("jwks_uri", jwks_uri),
|
||||
)
|
||||
for name, endpoint_url in required:
|
||||
try:
|
||||
validate_discovered_endpoint(
|
||||
endpoint_url,
|
||||
issuer_parsed,
|
||||
allow_http=allow_http,
|
||||
trusted_endpoint_hosts=trusted_hosts,
|
||||
)
|
||||
except OIDCError as exc:
|
||||
log.warning("OIDC discovered %s rejected (url=%s): %s", name, endpoint_url, exc)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
if userinfo_endpoint:
|
||||
try:
|
||||
validate_discovered_endpoint(
|
||||
userinfo_endpoint,
|
||||
issuer_parsed,
|
||||
allow_http=allow_http,
|
||||
trusted_endpoint_hosts=trusted_hosts,
|
||||
)
|
||||
except OIDCError as exc:
|
||||
log.warning(
|
||||
"OIDC discovered userinfo_endpoint rejected (url=%s): %s",
|
||||
userinfo_endpoint,
|
||||
exc,
|
||||
)
|
||||
return dataclasses.replace(config, enabled=False)
|
||||
|
||||
log.info("OIDC discovery complete: %s", config.issuer)
|
||||
return dataclasses.replace(
|
||||
config,
|
||||
@@ -326,38 +458,154 @@ async def discover_oidc(config: OIDCConfig) -> OIDCConfig:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def fetch_jwks(jwks_uri: str) -> dict[str, Any]:
|
||||
async def fetch_jwks(
|
||||
jwks_uri: str,
|
||||
*,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Fetch the JWKS key set from the IdP.
|
||||
|
||||
Returns the parsed JSON document (``{"keys": [...]}``) . Called during
|
||||
startup discovery and on-demand when an unknown ``kid`` is encountered
|
||||
(key rotation). Uses ``httpx.AsyncClient`` — never blocks the event loop.
|
||||
|
||||
Raises :class:`OIDCError` on network failures or malformed responses.
|
||||
A long-lived ``client`` may be supplied to share connection pooling;
|
||||
when ``None`` a transient client is used.
|
||||
|
||||
Raises :class:`OIDCError` on HTTP error or malformed JSON.
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.get(jwks_uri)
|
||||
if client is not None:
|
||||
resp = await client.get(jwks_uri, timeout=10.0)
|
||||
resp.raise_for_status()
|
||||
result: dict[str, Any] = resp.json()
|
||||
except Exception as exc:
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as transient:
|
||||
resp = await transient.get(jwks_uri)
|
||||
resp.raise_for_status()
|
||||
result = resp.json()
|
||||
except (httpx.HTTPError, ValueError) as exc:
|
||||
# ValueError covers json.JSONDecodeError (subclass).
|
||||
raise OIDCError(f"JWKS fetch failed: {exc}") from exc
|
||||
if not isinstance(result, dict):
|
||||
raise OIDCError("JWKS document is not a JSON object")
|
||||
if not isinstance(result.get("keys"), list):
|
||||
raise OIDCError("JWKS document missing 'keys' array")
|
||||
return result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lifespan integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def initialize_oidc_state(app_state: Any) -> None:
|
||||
"""Run OIDC discovery + JWKS prefetch and stash results on ``app_state``.
|
||||
|
||||
Reads ``app_state.oidc_config`` (already set to a non-discovered config
|
||||
by the lifespan), runs :func:`discover_oidc` and :func:`fetch_jwks`, and
|
||||
writes back the populated ``oidc_config`` plus ``jwks_data``.
|
||||
|
||||
Post-conditions on ``app_state``:
|
||||
|
||||
- On disable (config flag off, discovery failure, discovery-returned-disabled,
|
||||
missing ``redirect_base``): ``oidc_config.enabled is False``,
|
||||
``jwks_data is None``, ``oidc_http_client is None``.
|
||||
- On JWKS-prefetch failure with otherwise-valid config: ``oidc_config.enabled``
|
||||
stays ``True`` and ``jwks_data is None`` so the callback's lazy-fetch
|
||||
retry path can recover from a transient IdP failure at startup. The
|
||||
long-lived ``oidc_http_client`` stays open for that retry.
|
||||
- On full success: ``oidc_config`` populated with discovered endpoints,
|
||||
``jwks_data`` populated, ``oidc_http_client`` open for the runtime
|
||||
callback path. Pair with :func:`close_oidc_state` in lifespan teardown.
|
||||
"""
|
||||
cfg: OIDCConfig = app_state.oidc_config
|
||||
app_state.jwks_refetch_lock = asyncio.Lock()
|
||||
if not cfg.enabled:
|
||||
app_state.jwks_data = None
|
||||
app_state.oidc_http_client = None
|
||||
return
|
||||
|
||||
# Use a transient client for discovery so the long-lived client is only
|
||||
# installed once we know OIDC will actually be enabled. The disable
|
||||
# branches below would otherwise leak sockets until shutdown.
|
||||
async with httpx.AsyncClient(timeout=10.0) as transient_client:
|
||||
# Discovery is operator-controlled config; any unexpected failure
|
||||
# must disable OIDC rather than escape and bring down the service.
|
||||
try:
|
||||
cfg = await discover_oidc(cfg, client=transient_client)
|
||||
except Exception:
|
||||
log.warning("OIDC discovery failed -- OIDC login disabled", exc_info=True)
|
||||
app_state.oidc_config = dataclasses.replace(cfg, enabled=False)
|
||||
app_state.jwks_data = None
|
||||
app_state.oidc_http_client = None
|
||||
return
|
||||
|
||||
if not cfg.enabled:
|
||||
app_state.oidc_config = cfg
|
||||
app_state.jwks_data = None
|
||||
app_state.oidc_http_client = None
|
||||
return
|
||||
|
||||
if not cfg.redirect_base:
|
||||
log.error(
|
||||
"OIDC enabled but TURNSTONE_OIDC_REDIRECT_BASE is unset. "
|
||||
"This is required to prevent Host-header-derived redirect_uri spoofing. "
|
||||
"Set it to your service's externally-visible URL "
|
||||
"(e.g. https://idp.example.com). OIDC will be disabled."
|
||||
)
|
||||
app_state.oidc_config = dataclasses.replace(cfg, enabled=False)
|
||||
app_state.jwks_data = None
|
||||
app_state.oidc_http_client = None
|
||||
return
|
||||
|
||||
http_client = httpx.AsyncClient(timeout=10.0)
|
||||
app_state.oidc_http_client = http_client
|
||||
|
||||
try:
|
||||
jwks_data = await fetch_jwks(cfg.jwks_uri, client=http_client)
|
||||
except OIDCError:
|
||||
# Keep enabled=True so the callback's lazy-fetch retry path can
|
||||
# recover if the IdP transiently failed during startup. The
|
||||
# http_client stays open for that retry.
|
||||
log.warning("OIDC JWKS prefetch failed -- will retry on first login", exc_info=True)
|
||||
app_state.oidc_config = cfg
|
||||
app_state.jwks_data = None
|
||||
return
|
||||
|
||||
app_state.oidc_config = cfg
|
||||
app_state.jwks_data = jwks_data
|
||||
log.info("OIDC enabled: %s (%s)", cfg.provider_name, cfg.issuer)
|
||||
|
||||
|
||||
async def close_oidc_state(app_state: Any) -> None:
|
||||
"""Close the long-lived OIDC HTTP client installed by :func:`initialize_oidc_state`.
|
||||
|
||||
Safe to call when OIDC was never enabled — does nothing.
|
||||
"""
|
||||
client = getattr(app_state, "oidc_http_client", None)
|
||||
if client is not None:
|
||||
try:
|
||||
await client.aclose()
|
||||
except Exception:
|
||||
log.debug("OIDC http client close failed", exc_info=True)
|
||||
app_state.oidc_http_client = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PKCE helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_pkce_pair() -> tuple[str, str]:
|
||||
"""Generate a PKCE code_verifier and code_challenge pair."""
|
||||
code_verifier = secrets.token_urlsafe(48)
|
||||
digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
|
||||
code_challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
|
||||
return code_verifier, code_challenge
|
||||
def generate_pkce_verifier() -> str:
|
||||
"""Generate a PKCE code_verifier.
|
||||
|
||||
The matching code_challenge is recomputed from the verifier inside
|
||||
:func:`build_authorize_url`, so callers that only need the verifier
|
||||
(the value stored in pending state for the callback's token exchange)
|
||||
don't have to discard a separately-returned challenge.
|
||||
"""
|
||||
return secrets.token_urlsafe(48)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -399,9 +647,14 @@ async def exchange_code(
|
||||
code: str,
|
||||
redirect_uri: str,
|
||||
code_verifier: str,
|
||||
*,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Exchange authorization code for tokens at the token endpoint.
|
||||
|
||||
A long-lived ``client`` may be supplied; when ``None`` a transient
|
||||
client is used.
|
||||
|
||||
Raises :class:`OIDCError` on non-200 response.
|
||||
"""
|
||||
data = {
|
||||
@@ -413,16 +666,24 @@ async def exchange_code(
|
||||
"code_verifier": code_verifier,
|
||||
}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.post(config.token_endpoint, data=data)
|
||||
if client is not None:
|
||||
resp = await client.post(config.token_endpoint, data=data, timeout=10.0)
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as transient:
|
||||
resp = await transient.post(config.token_endpoint, data=data)
|
||||
except Exception as exc:
|
||||
raise OIDCError(f"Token exchange request failed: {exc}") from exc
|
||||
|
||||
if resp.status_code != 200:
|
||||
raise OIDCError(f"Token endpoint returned {resp.status_code}: {resp.text[:500]}")
|
||||
raise OIDCError(
|
||||
f"Token endpoint returned {resp.status_code}: {_sanitize_log_text(resp.text, 500)}"
|
||||
)
|
||||
|
||||
result: dict[str, Any] = resp.json()
|
||||
return result
|
||||
result = resp.json()
|
||||
if not isinstance(result, dict):
|
||||
raise OIDCError("Token endpoint returned non-dict body")
|
||||
typed: dict[str, Any] = result
|
||||
return typed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -479,7 +740,7 @@ def validate_id_token(
|
||||
raise OIDCError(f"Failed to parse signing key: {exc}") from exc
|
||||
|
||||
if signing_key is None:
|
||||
raise OIDCError(f"Signing key '{kid}' not found in JWKS")
|
||||
raise OIDCKeyNotFoundError(f"Signing key '{kid}' not found in JWKS")
|
||||
|
||||
try:
|
||||
claims: dict[str, Any] = jwt.decode(
|
||||
@@ -503,6 +764,41 @@ def validate_id_token(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ensure_default_role(
|
||||
storage: Any,
|
||||
user_id: str,
|
||||
desired_role_ids: set[str] | None = None,
|
||||
) -> None:
|
||||
"""Self-heal safety-net: assign builtin-viewer if the user has zero roles.
|
||||
|
||||
Runs after :func:`apply_role_mapping` on both the new-user and
|
||||
existing-identity paths so a user stranded by a transient failure
|
||||
during initial role mapping (e.g. a DB blip after ``create_oidc_user``
|
||||
committed) recovers on next login. ``assigned_by="oidc-default"``
|
||||
deliberately differs from ``"oidc"`` so claim-driven revocation in
|
||||
``apply_role_mapping`` leaves it alone.
|
||||
|
||||
The optional ``desired_role_ids`` is a hint: when the caller already
|
||||
knows claim-driven mapping populated at least one role, we skip the
|
||||
``list_user_roles`` query.
|
||||
|
||||
No-op when builtin-viewer is unavailable (admin removed it from the
|
||||
role table) or the user already has at least one role.
|
||||
|
||||
Note: if an admin manually strips all roles from an OIDC user, this
|
||||
helper will re-grant viewer on the next login. The documented way to
|
||||
deny an OIDC user access is to unlink their OIDC identity, not to
|
||||
strip roles.
|
||||
"""
|
||||
if desired_role_ids:
|
||||
return
|
||||
if storage.get_role("builtin-viewer") is None:
|
||||
return
|
||||
if storage.list_user_roles(user_id):
|
||||
return
|
||||
storage.assign_role(user_id, "builtin-viewer", "oidc-default")
|
||||
|
||||
|
||||
def provision_oidc_user(
|
||||
storage: Any,
|
||||
config: OIDCConfig,
|
||||
@@ -512,10 +808,13 @@ def provision_oidc_user(
|
||||
|
||||
Looks up an existing OIDC identity by (issuer, sub). If found,
|
||||
updates ``last_login`` and applies role mapping. Otherwise creates
|
||||
a new user and OIDC identity record.
|
||||
a new user and OIDC identity record atomically.
|
||||
|
||||
Raises :class:`OIDCError` if user creation fails.
|
||||
Raises :class:`OIDCError` if user creation fails or if a concurrent
|
||||
callback wins the username / identity race.
|
||||
"""
|
||||
from turnstone.core.storage import StorageConflictError
|
||||
|
||||
issuer = config.issuer
|
||||
sub = str(claims["sub"])
|
||||
email = str(claims.get("email", ""))
|
||||
@@ -526,25 +825,32 @@ def provision_oidc_user(
|
||||
if identity is not None:
|
||||
user_id = identity["user_id"]
|
||||
storage.update_oidc_identity_login(issuer, sub)
|
||||
apply_role_mapping(storage, user_id, claims, config)
|
||||
desired_role_ids = apply_role_mapping(storage, user_id, claims, config)
|
||||
_ensure_default_role(storage, user_id, desired_role_ids)
|
||||
user: dict[str, str] | None = storage.get_user(user_id)
|
||||
if user is None:
|
||||
raise OIDCError(f"OIDC identity references missing user: {user_id}")
|
||||
return user
|
||||
|
||||
# New user -- derive username
|
||||
# New user -- derive username, then create user + identity atomically.
|
||||
username = _derive_username(storage, claims)
|
||||
user_id = uuid.uuid4().hex
|
||||
|
||||
storage.create_user(user_id, username, display_name, OIDC_PASSWORD_SENTINEL)
|
||||
storage.create_oidc_identity(issuer, sub, user_id, email)
|
||||
apply_role_mapping(storage, user_id, claims, config)
|
||||
try:
|
||||
storage.create_oidc_user(
|
||||
user_id,
|
||||
username,
|
||||
display_name,
|
||||
OIDC_PASSWORD_SENTINEL,
|
||||
issuer,
|
||||
sub,
|
||||
email,
|
||||
)
|
||||
except StorageConflictError as exc:
|
||||
raise OIDCError(f"OIDC provisioning failed: {exc}") from exc
|
||||
|
||||
# Ensure new OIDC users have at least a default role so they can
|
||||
# access the application. builtin-viewer grants read-only access.
|
||||
user_roles = storage.list_user_roles(user_id)
|
||||
if not user_roles and storage.get_role("builtin-viewer") is not None:
|
||||
storage.assign_role(user_id, "builtin-viewer", "oidc-default")
|
||||
desired_role_ids = apply_role_mapping(storage, user_id, claims, config)
|
||||
_ensure_default_role(storage, user_id, desired_role_ids)
|
||||
|
||||
created_user: dict[str, str] | None = storage.get_user(user_id)
|
||||
if created_user is None:
|
||||
@@ -570,14 +876,13 @@ def _derive_username(storage: Any, claims: dict[str, Any]) -> str:
|
||||
if not sanitised:
|
||||
sanitised = "user"
|
||||
|
||||
# Check validity and uniqueness.
|
||||
if is_valid_username(sanitised) and storage.get_user_by_username(sanitised) is None:
|
||||
return sanitised
|
||||
|
||||
# Deduplicate: append suffix.
|
||||
for suffix in range(2, 11):
|
||||
candidate = f"{sanitised[:60]}{suffix}"
|
||||
if is_valid_username(candidate) and storage.get_user_by_username(candidate) is None:
|
||||
# Build the full set of bounded candidates (base + 2..10 suffixes), strip
|
||||
# invalid forms, then ask storage which ones are already taken in one query.
|
||||
candidates = [sanitised, *(f"{sanitised[:60]}{n}" for n in range(2, 11))]
|
||||
valid_candidates = [c for c in candidates if is_valid_username(c)]
|
||||
existing = storage.find_existing_usernames(valid_candidates)
|
||||
for candidate in valid_candidates:
|
||||
if candidate not in existing:
|
||||
return candidate
|
||||
|
||||
# Last resort: full UUID suffix with validation + uniqueness check.
|
||||
@@ -600,18 +905,22 @@ def apply_role_mapping(
|
||||
user_id: str,
|
||||
claims: dict[str, Any],
|
||||
config: OIDCConfig,
|
||||
) -> None:
|
||||
"""Sync Turnstone roles from OIDC claims.
|
||||
) -> set[str]:
|
||||
"""Sync Turnstone roles from OIDC claims. Returns desired role id set.
|
||||
|
||||
If ``config.role_claim`` is set, reads the corresponding claim value,
|
||||
normalises it to a list, and maps each value via ``config.role_map``
|
||||
to a Turnstone role ID. Roles assigned by OIDC on previous logins
|
||||
that are no longer present in the claims are revoked (IdP demotions
|
||||
propagate). Roles assigned manually or by other sources are never
|
||||
touched.
|
||||
propagate). Roles assigned manually or by other sources (including
|
||||
the ``oidc-default`` builtin-viewer fallback) are never touched.
|
||||
|
||||
The returned ``desired_role_ids`` lets the caller decide whether to
|
||||
apply the new-user fallback role without a second ``list_user_roles``
|
||||
round-trip.
|
||||
"""
|
||||
if not config.role_claim or not config.role_map:
|
||||
return
|
||||
return set()
|
||||
|
||||
claim_value = claims.get(config.role_claim)
|
||||
|
||||
@@ -632,16 +941,10 @@ def apply_role_mapping(
|
||||
if role_id and storage.get_role(role_id) is not None:
|
||||
desired_role_ids.add(role_id)
|
||||
|
||||
# Add new roles from claims.
|
||||
for role_id in desired_role_ids:
|
||||
storage.assign_role(user_id, role_id, "oidc")
|
||||
added, removed = storage.replace_oidc_roles(user_id, desired_role_ids)
|
||||
for role_id in added:
|
||||
log.debug("Assigned role %s to user %s via OIDC claim", role_id, user_id)
|
||||
for role_id in removed:
|
||||
log.info("Revoked role %s from user %s (removed from IdP claims)", role_id, user_id)
|
||||
|
||||
# Revoke OIDC-assigned roles no longer present in claims.
|
||||
current_roles = storage.list_user_roles(user_id)
|
||||
for role in current_roles:
|
||||
if role.get("assigned_by") == "oidc" and role["role_id"] not in desired_role_ids:
|
||||
storage.unassign_role(user_id, role["role_id"])
|
||||
log.info(
|
||||
"Revoked role %s from user %s (removed from IdP claims)", role["role_id"], user_id
|
||||
)
|
||||
return desired_role_ids
|
||||
|
||||
+721
-211
File diff suppressed because it is too large
Load Diff
@@ -2269,7 +2269,10 @@ def make_history_handler(cfg: SessionEndpointConfig) -> Handler:
|
||||
messages: list[dict[str, Any]] = []
|
||||
if storage is not None:
|
||||
try:
|
||||
messages = await asyncio.to_thread(storage.load_messages, ws_id, limit=limit)
|
||||
# repair=False — display read; see reconstruct_messages docstring.
|
||||
messages = await asyncio.to_thread(
|
||||
storage.load_messages, ws_id, limit=limit, repair=False
|
||||
)
|
||||
except Exception:
|
||||
log.debug("ws.history.load_failed ws=%s", ws_id[:8], exc_info=True)
|
||||
# Audit-trail decoration — attach persisted intent_verdict and
|
||||
@@ -2287,7 +2290,14 @@ def make_history_handler(cfg: SessionEndpointConfig) -> Handler:
|
||||
)
|
||||
|
||||
indexes = await asyncio.to_thread(load_verdict_indexes, ws_id)
|
||||
decorate_history_messages(messages, indexes[0], indexes[1])
|
||||
# Pure transform but iterates every message and every
|
||||
# tool_call dict — for a long workstream the pass takes
|
||||
# tens of milliseconds and would otherwise block the
|
||||
# event loop's hot path on the request handler.
|
||||
# ``decorate_history_messages`` is thread-safe (no
|
||||
# shared mutable state beyond the per-call message
|
||||
# list) so the off-loop hop is free.
|
||||
await asyncio.to_thread(decorate_history_messages, messages, indexes[0], indexes[1])
|
||||
except Exception:
|
||||
# Operationally interesting: a persistent decoration
|
||||
# failure (missing migration, driver mismatch, schema
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user