mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-28 06:44:51 -06:00
124615cce0
Extends the Phase 7 per-(user, server) ClientSession pool to cover
RFC §3.2 (resources/read) and §3.3 (prompts/get) on the same shape
already proven for tools/call. Pool discovery is capability-gated so
servers without resources/ or prompts/ stay free of extra round-trips.
API additions / widenings (MCPClientManager):
- ``read_resource_sync(uri, *, user_id=None, timeout=120)`` —
per-user-first dispatch; falls through to the byte-identical static
path when ``user_id`` is None or the URI doesn't resolve to an
``oauth_user`` pool entry.
- ``get_prompt_sync(prefixed_name, arguments=None, *, user_id=None,
timeout=30)`` — same dispatch shape; structured-error responses
surface via ``RuntimeError`` so the agent-loop's ``except Exception``
block renders the JSON without polluting the prompt-protocol return
shape.
- ``get_resources(user_id=None)`` / ``get_prompts(user_id=None)`` —
per-user merged catalogs (admin/global call still passes None).
- ``add_{resource,prompt}_listener`` /
``remove_{resource,prompt}_listener`` — ``user_id`` keyword scopes
the listener so a pool-only catalog change for one user does not
wake another user's session.
- ``resource_count_for_user(user_id=None)`` /
``prompt_count_for_user(user_id=None)`` — method-form variants used
by ChatSession's ``read_resource`` / ``use_prompt`` tool gating; the
legacy ``resource_count`` / ``prompt_count`` properties remain
static-only for admin paths.
- ``_dispatch_pool_resource`` / ``_dispatch_pool_prompt`` async coros
— mirror ``_dispatch_pool`` for the new SDK calls; share the
carrier-race-and-cancel core via ``_dispatch_pool_with_entry_call``.
- ``_handle_auth_403`` extended with ``kind=Literal["tool",
"resource", "prompt"]`` so the per-operation ``mcp_*_forbidden``
code surfaces (kind="tool" remains the default for back-compat).
- Pool notification handler now refreshes resources / prompts on
``ResourceListChangedNotification`` / ``PromptListChangedNotification``
via ``_refresh_pool_server_resources`` / ``_refresh_pool_server_prompts``.
ChatSession (``turnstone/core/session.py``) call-site updates:
- 12 sites threaded the session-bound ``user_id`` through
``add_*_listener`` / ``remove_*_listener``, ``get_resources`` /
``get_prompts``, gating, ``read_resource_sync`` /
``get_prompt_sync``, and ``is_mcp_prompt`` so the per-user merged
catalog drives both the visible-tool set and dispatch.
- ``/mcp`` slash command now lists this user's pool resources and
prompts alongside tools (Phase 7 already scoped tools).
Scope decisions:
- Per-user-first URI ordering (decision 0.1): the dispatcher attempts
the user's pool catalog first, falling back to the static catalog
only when no pool entry resolves the URI / prefixed name. Pool-only
users never see the static catalog leak into their resolution.
- Method-form ``*_count_for_user`` (vs property) keeps the legacy
``resource_count`` / ``prompt_count`` properties intact for admin
endpoints whose contract is "static catalog size only".
- Shared ``_dispatch_pool_with_entry_call`` helper accepts an
``sdk_call: Callable[[ClientSession], Awaitable[Any]]`` closure,
keeping the entry-locked carrier-race / classification / retry
plumbing single-source instead of a 3x copy across tool / resource
/ prompt paths.
R6 (anyio uniformity): every pool-side list / read / get path uses
``async with asyncio.timeout(...)`` — ``asyncio.wait_for`` is
forbidden in those paths because it wraps the inner awaitable in a
fresh task and surfaces ``CancelledError`` from inside
``streamablehttp_client``'s anyio TaskGroup on Python 3.11
(per ``feedback_asyncio_timeout_vs_wait_for.md``).
Tests:
- ``test_mcp_pool_auth_resource_integration.py`` — 9 real-transport
resource tests (FastMCP upstream + ``BehaviorMiddleware``):
401-refresh-retry success, persistent 401 -> consent_required,
403+insufficient_scope, 403 generic -> mcp_resource_read_forbidden,
breaker-isolation under repeated auth failures, missing-token,
decrypt-failure, http:// URL guard, unknown-URI ValueError.
- ``test_mcp_pool_auth_prompt_integration.py`` — 9 mirror tests for
the prompt path; structured-error responses verified via
``RuntimeError`` payload shape.
- ``test_mcp_user_catalog.py`` — extended unit coverage for per-user
resource / prompt rebuild + collision policy + symmetric eviction.
- ``test_sessions.py::TestMCPToolGating`` — pool-only-user canary
asserts ``read_resource`` / ``use_prompt`` stay visible when the
static catalog is empty but the user has pool entries.
Round-1 review fixes (4-finder review applied, no push yet):
- bug-1: ``_exec_use_prompt`` was hardcoding ``"MCP prompt error: failed
to invoke prompt"`` — discarding the structured-error JSON that
``_dispatch_pool_prompt_sync`` raises via ``RuntimeError``. Now uses
``f"MCP prompt error: {e}"`` mirroring ``_exec_mcp_tool``; pool-prompt
consent_required / insufficient_scope / forbidden errors now reach
the LLM as intended.
- bug-2 + bug-3: resource template discovery was uncapped —
``_cap_server_resources`` covered ``res_result.resources`` but the
separate ``tmpl_result.resourceTemplates`` loop appended every
template a server returned. Added ``_MAX_RESOURCE_TEMPLATES_PER_SERVER``
(1000) + ``_cap_server_resource_templates`` helper, applied at both
the initial discovery site (``_connect_one_pool``) and the refresh
site (``_refresh_pool_server_resources``). Mirrors the existing
``_MAX_TOOLS_PER_SERVER`` / ``_MAX_PROMPTS_PER_SERVER`` defensive
ceilings.
- sec-1 + sec-2: ``emit_insufficient_scope_audit`` generalized to
``emit_oauth_failure_audit(kind, code, ...)``, called from both the
insufficient_scope branch AND the previously-silent generic 403
branch. Audit detail now records ``{"kind": kind, "code": code,
"scopes_required": [...]}`` so operators can distinguish tool-call
vs resource-read vs prompt-get 403s in audit logs and so cross-
tenant probing on the generic 403 path leaves a trail. The Phase 7
inherited gap (``mcp_tool_call_forbidden`` had the same silence) is
closed in the same refactor.
- perf-1: pool resource discovery now uses ``asyncio.gather(
list_resources, list_resource_templates)`` inside the existing
``async with asyncio.timeout(...)`` budget — disjoint catalogs, no
ordering dependency. Typical-case 2-RTT cold-connect resource block
collapses to 1-RTT. Same change applied at ``_refresh_pool_server_resources``.
- q-1: ``_rebuild_user_prompt_map`` docstring corrected RFC §3.2 →
§3.3 (resources are §3.2; prompts are §3.3).
- q-2: ``_refresh_pool_server_prompts`` docstring now carries the
R6 / mcp-loop note that the resource sibling already had — both
refresh paths now declare the asyncio.timeout invariant explicitly.
- q-5: added the ``_user_resource_map`` / DB-mismatch guard to
``read_resource_sync`` for parity with ``get_prompt_sync``. A stale
per-user map entry with no matching oauth_user row now raises a
specific ValueError instead of silently falling through to a
generic ``Unknown MCP resource``.
- q-6: ``_dispatch_pool_with_entry`` (now a single-caller wrapper
after the ``_dispatch_pool_with_entry_call`` extraction) gains a
one-line docstring explaining why the wrapper is preserved
(tool-decode localization + stack-trace identity for debugging).
- q-7: added 1 resource + 1 prompt end-to-end integration test that
drive REAL discovery + dispatch in the same connect (no
``_seed_pool_*_map`` shortcuts), mirroring the tool path's
``test_integration_pool_reuse_401_refresh_and_retry_succeeds``.
The seeded-map tests stay (faster, focused on dispatch); the new
e2e tests cover the connect-discover-dispatch composition that
caught Phase 6's carrier-on-entry bug.
Pre-push round-1 review fixes (3-finder review on the final state —
the lesson from Phase 7 round-3's q-1 regression: round-2 catches
what the round-1 apply pass missed):
- q-1 (MAJOR): the bug-1 sibling that round-1 missed —
``_exec_read_resource`` was hardcoding ``"MCP resource error: failed
to read resource"`` while ``_exec_use_prompt`` (post-bug-1) preserved
the structured-error JSON via ``f"... error: {e}"``. The round-1
apply pass patched the prompt side but not the resource side. q-5's
per-user-map / DB-mismatch ValueError was being swallowed at the
agent loop boundary, defeating the operator-diagnostic intent. Now
``_exec_read_resource`` mirrors ``_exec_mcp_tool`` and ``_exec_use_prompt``.
- q-6 (nit): defensive-cap comment block at module-level cited
"(RFC §3.2)" while covering both resource and prompt list paths;
prompts are §3.3. Now reads "(RFC §3.2 for resources, §3.3 for
prompts)" matching the convention the q-1 apply established.
- q-5 (rejected with better justification): the reviewer flagged
``_dispatch_pool_with_entry`` as a single-caller wrapper that should
be inlined. After examination — the autouse fixture
``tests/test_mcp_pool_auth_introspection.py::_install_capture_intercept``
monkeypatches this method to stash ``entry.auth_capture`` for the
fake call_tool stubs in dispatcher-asserting tests. Inlining would
redirect the patch to ``_dispatch_pool_with_entry_call`` (different
kwargs shape) and require re-validating every test that depends on
the interception. The wrapper IS load-bearing; q-6 docstring updated
to cite the test-fixture rationale instead of the thin "stack-trace
identity" claim.
Deferred to follow-up (documented rationale):
- perf-2: single-pass partition for system-message resource list
(concrete vs templates). Sub-microsecond at expected scale;
opportunistic-only.
- q-2 (pre-push): ~200 lines of fixture infrastructure
(``BehaviorMiddleware``, ``_build_server``, ``_seed_oauth_server``,
``running_loop_mgr``, etc.) duplicated across three pool-integration
test files. Real maintenance cost, but a 200-line conftest extraction
is a focused refactor that earns its own commit / PR. Tracking as
follow-up rather than balloon Phase 7b's diff further.
- q-3 / q-4 (refactor): extract shared dispatcher / scheduler
helpers to compress three near-identical 90-line bodies (round-1
q-3 was the same root cause; the pre-push q-3/q-4 reviewer
reaffirmed it concretely). Three named methods preserve readability
for the codebase's hottest correctness path; follow-up if
duplication grows further or if a per-path divergence ships.
- q-4 (round-1, distinct from pre-push q-4): split pool concerns
into ``mcp_pool.py``. Out-of-scope per finder; future refactor as
the file approaches the navigation/merge-conflict threshold.
3.13: 5590 passed (5541 baseline -> +49 net; pre-review +47, q-7
e2e tests added +2). Existing audit-detail tests updated in-place
to expect the new ``kind`` and ``code`` fields.
3.11: 5590 passed (parity gate per ``feedback_pytest_env_parity.md``).
969 lines
35 KiB
Python
969 lines
35 KiB
Python
"""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)
|
|
|
|
result = mgr.call_tool_sync(
|
|
"mcp__pool-srv__do_thing",
|
|
{},
|
|
user_id="user-1",
|
|
timeout=5,
|
|
)
|
|
payload = json.loads(result)
|
|
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
|
|
|
|
result = mgr.call_tool_sync(
|
|
"mcp__pool-srv__do_thing",
|
|
{},
|
|
user_id="user-1",
|
|
timeout=5,
|
|
)
|
|
payload = json.loads(result)
|
|
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",
|
|
)
|
|
|
|
result = mgr.call_tool_sync(
|
|
"mcp__pool-srv__do_thing",
|
|
{},
|
|
user_id="user-1",
|
|
timeout=5,
|
|
)
|
|
payload = json.loads(result)
|
|
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)
|
|
|
|
result = mgr.call_tool_sync(
|
|
"mcp__pool-srv__do_thing",
|
|
{},
|
|
user_id="user-1",
|
|
timeout=5,
|
|
)
|
|
payload = json.loads(result)
|
|
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 == {}
|