mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
feat(mcp): SDK 401/403 introspection via httpx response hook
Phase 6 of OAuth-MCP. Recovers upstream 401/403 from MCP servers via a
capturing httpx_client_factory: an async response hook records 4xx
status + WWW-Authenticate header into a per-dispatch carrier before
the SDK's post_writer swallows the underlying httpx.HTTPStatusError.
Splits _classify_failure into auth_401 (refresh-and-retry once) vs
auth_403 (parse insufficient_scope, emit mcp_insufficient_scope with
parsed scope set). The 401 retry runs on a fresh asyncio.Task via
run_coroutine_threadsafe in _dispatch_pool_sync, escaping the anyio
cancel-scope state of the prior dispatch's TaskGroup.
WWW-Authenticate parsing extracted to a new mcp_http_parsers module
with an RFC 7235 challenge tokenizer (replaces hand-rolled substring
scanners). Two-layer defense against multi-Bearer-challenge injection:
the hook uses get_list("www-authenticate")[0] to drop attacker's
second challenge, the parser truncates at challenge boundary as
belt-and-braces. Scope set capped at 32 entries before hitting the
audit row or the LLM-visible structured-error JSON.
Auth failures (401/403) never trip the per-server circuit breaker
(server-only breaker invariant). Static path remains byte-identical.
_PgRefreshLock untouched. Pool dispatch still reachable from the
agent loop only via Phase 7 catalog scoping; Phase 6 behaviour is
testable via direct call_tool_sync.
5557 tests pass. 33 tokenizer unit tests in tests/test_mcp_http_parsers
cover the RFC 7235 grammar + the scope/error wrappers + the 4 KB input
cap. 7 integration tests in tests/test_mcp_pool_auth_integration drive
real upstream 401/403 through streamablehttp_client + a FastMCP
subprocess fixture — the structural exit gate that makes
HTTPStatusError-injection-only unit tests insufficient.
This commit is contained in:
@@ -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"
|
||||
@@ -0,0 +1,718 @@
|
||||
"""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,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
assert payload["error"]["code"] == "mcp_consent_required"
|
||||
assert payload["error"]["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,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
assert payload["error"]["code"] == "mcp_insufficient_scope"
|
||||
assert payload["error"]["scopes_required"] == ["files:write", "mail:send"]
|
||||
# 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,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
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,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
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,
|
||||
):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
|
||||
payload = json.loads(result)
|
||||
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):
|
||||
result = mgr.call_tool_sync(
|
||||
"mcp__pool-srv__echo", {"payload": "x"}, user_id="user-1", timeout=15
|
||||
)
|
||||
payload = json.loads(result)
|
||||
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"
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
+34
-42
@@ -517,23 +517,29 @@ class TestClassifyFailure:
|
||||
mgr = MCPClientManager({})
|
||||
assert mgr._classify_failure(ValueError("nope")) == "other"
|
||||
|
||||
def test_http_401_classified_as_auth(self) -> None:
|
||||
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"
|
||||
assert mgr._classify_failure(exc) == "auth_401"
|
||||
|
||||
def test_http_403_classified_as_auth(self) -> None:
|
||||
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"
|
||||
assert mgr._classify_failure(exc) == "auth_403"
|
||||
|
||||
def test_http_500_not_classified_as_auth(self) -> None:
|
||||
import httpx
|
||||
@@ -577,37 +583,6 @@ class TestDispatchFailureWiring:
|
||||
|
||||
_run_on_loop(loop, _seed())
|
||||
|
||||
def test_dispatch_pool_401_does_not_trip_breaker(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
import httpx
|
||||
|
||||
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)
|
||||
|
||||
req = httpx.Request("POST", "https://mcp.example.com/sse")
|
||||
resp = httpx.Response(401, request=req)
|
||||
self._seed_connected_session(
|
||||
mgr, loop, httpx.HTTPStatusError("unauthorized", request=req, response=resp)
|
||||
)
|
||||
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
mgr.call_tool_sync(
|
||||
"mcp__pool-srv__do_thing",
|
||||
{},
|
||||
user_id="user-1",
|
||||
timeout=5,
|
||||
)
|
||||
# Auth failure must NOT tick the per-server breaker.
|
||||
assert mgr._consecutive_failures.get("pool-srv", 0) == 0
|
||||
# But the pool entry's session must be cleared so the next call
|
||||
# re-authenticates with a fresh bearer.
|
||||
entry = mgr._user_pool_entries[("user-1", "pool-srv")]
|
||||
assert entry.session is None
|
||||
|
||||
def test_dispatch_pool_transport_failure_trips_breaker(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
@@ -811,11 +786,27 @@ class TestLruInterlock:
|
||||
|
||||
|
||||
class TestConcurrentDispatch:
|
||||
def test_pool_concurrent_dispatch_to_same_user_server_is_concurrent(
|
||||
def test_pool_concurrent_dispatch_to_same_user_server_is_serialized(
|
||||
self, running_loop_mgr, storage: SQLiteBackend
|
||||
) -> None:
|
||||
"""Two tool calls on the SAME (user, server) must overlap in flight,
|
||||
not serialize on ``open_lock``."""
|
||||
"""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")
|
||||
@@ -834,7 +825,8 @@ class TestConcurrentDispatch:
|
||||
in_flight += 1
|
||||
observed_max_concurrency = max(observed_max_concurrency, in_flight)
|
||||
try:
|
||||
# Hold a moment so concurrent calls overlap.
|
||||
# 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"
|
||||
@@ -880,9 +872,9 @@ class TestConcurrentDispatch:
|
||||
|
||||
assert errors == []
|
||||
assert results == ["ok", "ok"]
|
||||
# Both dispatches were in flight at the same time — open_lock did
|
||||
# not serialize them.
|
||||
assert observed_max_concurrency == 2
|
||||
# ``open_lock`` held across ``call_tool`` — the second dispatch
|
||||
# waits for the first to release before entering call_tool.
|
||||
assert observed_max_concurrency == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
+444
-63
@@ -35,11 +35,21 @@ import mcp.types as mcp_types
|
||||
from mcp import ClientSession, McpError, StdioServerParameters
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
from mcp.shared._httpx_utils import (
|
||||
MCP_DEFAULT_SSE_READ_TIMEOUT,
|
||||
MCP_DEFAULT_TIMEOUT,
|
||||
McpHttpClientFactory,
|
||||
)
|
||||
|
||||
from turnstone.core.config import load_config
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.mcp_http_parsers import (
|
||||
parse_www_authenticate_error,
|
||||
parse_www_authenticate_scope,
|
||||
)
|
||||
from turnstone.core.mcp_oauth import (
|
||||
TokenLookupResult,
|
||||
emit_insufficient_scope_audit,
|
||||
get_user_access_token_classified,
|
||||
)
|
||||
|
||||
@@ -81,6 +91,116 @@ def _validate_oauth_user_url(url: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pool dispatch auth introspection (response-hook carrier)
|
||||
# ---------------------------------------------------------------------------
|
||||
#
|
||||
# The MCP SDK's ``streamable_http`` transport raises
|
||||
# :class:`httpx.HTTPStatusError` inside ``_handle_post_request`` and the
|
||||
# enclosing ``post_writer`` swallows it (``mcp/client/streamable_http.py``
|
||||
# logger.exception path). The dispatcher then sees only
|
||||
# ``McpError(CONNECTION_CLOSED)`` with no status / headers preserved.
|
||||
# To recover the upstream 401/403 we plug into the SDK's documented
|
||||
# extension point — ``streamablehttp_client(httpx_client_factory=...)``
|
||||
# — and pass a factory that builds the ``httpx.AsyncClient`` with a
|
||||
# response hook. The hook fires after headers arrive but BEFORE
|
||||
# ``raise_for_status()`` runs, so the carrier is populated before the
|
||||
# SDK swallow.
|
||||
#
|
||||
# Forward-compat: ``streamablehttp_client`` is ``@deprecated`` in SDK
|
||||
# 1.27 in favour of ``streamable_http_client(http_client=...)`` which
|
||||
# accepts a pre-built client. The same factory pattern translates
|
||||
# 1:1 against the new entry point when we migrate.
|
||||
|
||||
|
||||
# Defensive cap on the number of scopes we report in
|
||||
# ``mcp_insufficient_scope`` audit/error payloads. Real ASes return
|
||||
# single-digit scope counts; the cap stops a malicious upstream from
|
||||
# bloating either surface via a thousand-token ``scope=`` value.
|
||||
_MAX_INSUFFICIENT_SCOPE_REPORTED = 32
|
||||
|
||||
|
||||
@dataclass
|
||||
class _AuthCapture:
|
||||
"""Carrier populated by the response hook on 4xx upstream responses."""
|
||||
|
||||
status: int | None = None
|
||||
www_authenticate: str | None = None
|
||||
|
||||
|
||||
class _PoolDispatchRetryRequested(BaseException): # noqa: N818
|
||||
"""Module-private signal: the auth_401 branch wants the sync caller to retry.
|
||||
|
||||
Inherits :class:`BaseException` (not ``Exception``) so that nothing in
|
||||
the SDK / anyio path accidentally swallows it inside an ``except
|
||||
Exception`` block. ``_dispatch_pool_sync`` is the only caller that
|
||||
catches it and the only handler that re-issues the dispatch on a
|
||||
fresh ``asyncio.Task`` via ``run_coroutine_threadsafe``.
|
||||
"""
|
||||
|
||||
|
||||
def _make_capturing_http_factory(capture: _AuthCapture) -> McpHttpClientFactory:
|
||||
"""Return an ``httpx`` factory that records 4xx auth signals into ``capture``.
|
||||
|
||||
The hook is ``async`` because :class:`httpx.AsyncClient` invokes
|
||||
response hooks via ``await hook(response)`` — a sync function would
|
||||
return ``None`` and ``await None`` raises ``TypeError`` inside the
|
||||
SDK's :meth:`client.stream` call. Even though our work is purely
|
||||
synchronous (read ``status_code`` and a header), the contract
|
||||
requires an awaitable. The first attempt's hook would still
|
||||
populate the carrier (the body runs before the ``await``), but the
|
||||
``TypeError`` poisons the SDK's anyio TaskGroup teardown so the
|
||||
next ``streamablehttp_client(...)`` invocation surfaces a stray
|
||||
``CancelledError`` from inside its own scope. Empirically verified
|
||||
via a minimal repro: a sync hook breaks back-to-back connects in
|
||||
the same process; the async form does not.
|
||||
"""
|
||||
|
||||
async def _hook(response: httpx.Response) -> None:
|
||||
# Only record on auth-relevant statuses to keep the carrier
|
||||
# focused. ``capture`` is mutated in place; the dispatcher
|
||||
# consults it after ``call_tool`` returns/raises. No I/O, no
|
||||
# other awaits — the hook stays cancellation-safe.
|
||||
#
|
||||
# Use ``get_list(...)[0]`` rather than ``get(...)`` so a
|
||||
# malicious upstream that emits multiple ``WWW-Authenticate``
|
||||
# headers cannot inject auth-params into the parser via the
|
||||
# comma-joined value httpx returns from ``get(...)``. Repeated
|
||||
# headers join with ``, `` which the RFC 7235 tokenizer would
|
||||
# otherwise consume as a continuation of the first challenge,
|
||||
# silently folding attacker scopes into the parsed dict. We
|
||||
# discard every challenge after the first; defence-in-depth
|
||||
# mirror lives in ``parse_www_authenticate_bearer``.
|
||||
status = response.status_code
|
||||
if status in (401, 403):
|
||||
capture.status = status
|
||||
headers = response.headers.get_list("www-authenticate")
|
||||
capture.www_authenticate = headers[0] if headers else None
|
||||
|
||||
def _factory(
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: httpx.Timeout | None = None,
|
||||
auth: httpx.Auth | None = None,
|
||||
) -> httpx.AsyncClient:
|
||||
kwargs: dict[str, Any] = {
|
||||
"follow_redirects": True,
|
||||
"event_hooks": {"response": [_hook]},
|
||||
}
|
||||
if timeout is None:
|
||||
kwargs["timeout"] = httpx.Timeout(
|
||||
MCP_DEFAULT_TIMEOUT, read=MCP_DEFAULT_SSE_READ_TIMEOUT
|
||||
)
|
||||
else:
|
||||
kwargs["timeout"] = timeout
|
||||
if headers is not None:
|
||||
kwargs["headers"] = headers
|
||||
if auth is not None:
|
||||
kwargs["auth"] = auth
|
||||
return httpx.AsyncClient(**kwargs)
|
||||
|
||||
return _factory
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCP ↔ OpenAI schema conversion
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -156,9 +276,9 @@ class PoolEntryState:
|
||||
session: Any | None = None
|
||||
stack: AsyncExitStack | None = None
|
||||
streams: tuple[Any, Any] | None = None
|
||||
# Catalog state — populated when Phase 6 wires per-user discovery in;
|
||||
# left ``None`` here so 200-entry pools don't retain 600 empty list
|
||||
# objects.
|
||||
# Catalog state — populated lazily once per-user discovery wires
|
||||
# in; left ``None`` here so 200-entry pools don't retain 600 empty
|
||||
# list objects.
|
||||
tools: list[dict[str, Any]] | None = None
|
||||
resources: list[dict[str, Any]] | None = None
|
||||
prompts: list[dict[str, Any]] | None = None
|
||||
@@ -501,10 +621,27 @@ class MCPClientManager:
|
||||
errors could mask the original exception. CancelledError is caught
|
||||
explicitly because it is the primary failure mode (stray cancel from
|
||||
broken anyio scope) and is BaseException, not Exception.
|
||||
|
||||
``BaseExceptionGroup`` is also caught: anyio's TaskGroup wraps any
|
||||
unhandled task exception (e.g., the SDK's
|
||||
``HTTPStatusError`` raised inside ``post_writer``'s
|
||||
``tg.start_soon(handle_request_async)`` after a 401/403)
|
||||
into a ``BaseExceptionGroup`` on ``__aexit__``. Without this catch
|
||||
the auth-retry path's eager teardown would propagate the SDK's
|
||||
own collected fallout.
|
||||
|
||||
The 5s ``asyncio.wait_for`` is a deliberate guard against
|
||||
``aclose()`` hanging on a broken stack (e.g., a never-completing
|
||||
anyio task during teardown). The auth_401 retry runs on a
|
||||
fresh :class:`asyncio.Task` scheduled via
|
||||
:func:`asyncio.run_coroutine_threadsafe` (see
|
||||
:meth:`_dispatch_pool_sync`), so this helper is never invoked
|
||||
from inside an active cancellation context — ``wait_for``'s
|
||||
own cancellation observation does not abort the aclose early.
|
||||
"""
|
||||
try:
|
||||
await asyncio.wait_for(stack.aclose(), timeout=5)
|
||||
except (Exception, asyncio.CancelledError):
|
||||
except (Exception, asyncio.CancelledError, BaseExceptionGroup):
|
||||
log.debug("Error closing AsyncExitStack; ignoring", exc_info=True)
|
||||
|
||||
async def _safe_teardown_on_connect_failure(
|
||||
@@ -785,6 +922,8 @@ class MCPClientManager:
|
||||
key: tuple[str, str],
|
||||
cfg: dict[str, Any],
|
||||
access_token: str,
|
||||
*,
|
||||
auth_capture: _AuthCapture | None = None,
|
||||
) -> PoolEntryState:
|
||||
"""Connect a single per-(user, server) pool entry.
|
||||
|
||||
@@ -795,8 +934,17 @@ class MCPClientManager:
|
||||
* ``Authorization: Bearer {access_token}`` injected into headers
|
||||
alongside any operator-supplied static headers.
|
||||
* No catalog discovery (tools / resources / prompts) and no
|
||||
notification handler — those land in Phase 6 once per-user
|
||||
catalog state is in place.
|
||||
notification handler — those land later once per-user catalog
|
||||
state is in place.
|
||||
|
||||
When ``auth_capture`` is supplied, the underlying ``httpx``
|
||||
client is built via a factory whose response hook records 401/403
|
||||
status + ``WWW-Authenticate`` into the carrier, recovering the
|
||||
upstream auth signal that the SDK's ``post_writer`` would
|
||||
otherwise swallow. Static-path callers
|
||||
(:meth:`_connect_one`) MUST NOT pass this — the static path
|
||||
must remain byte-identical, which means the SDK's default
|
||||
``create_mcp_http_client`` factory.
|
||||
|
||||
MUST run on the mcp-loop. Caller holds ``entry.open_lock``.
|
||||
"""
|
||||
@@ -834,12 +982,16 @@ class MCPClientManager:
|
||||
headers: dict[str, str] = dict(cfg.get("headers") or {})
|
||||
headers["Authorization"] = f"Bearer {access_token}"
|
||||
|
||||
client_kwargs: dict[str, Any] = {"url": url, "headers": headers}
|
||||
if auth_capture is not None:
|
||||
client_kwargs["httpx_client_factory"] = _make_capturing_http_factory(auth_capture)
|
||||
|
||||
stack = AsyncExitStack()
|
||||
await stack.__aenter__()
|
||||
try:
|
||||
await self._tcp_probe(key, url)
|
||||
read, write, _ = await asyncio.wait_for(
|
||||
stack.enter_async_context(streamablehttp_client(url=url, headers=headers)),
|
||||
stack.enter_async_context(streamablehttp_client(**client_kwargs)),
|
||||
timeout=self._CONNECT_TIMEOUT,
|
||||
)
|
||||
entry.streams = (read, write)
|
||||
@@ -1011,8 +1163,11 @@ class MCPClientManager:
|
||||
# -- failure classification (pool dispatch) ------------------------------
|
||||
|
||||
def _classify_failure(
|
||||
self, exc: BaseException
|
||||
) -> Literal["transport", "auth", "protocol", "other"]:
|
||||
self,
|
||||
exc: BaseException,
|
||||
*,
|
||||
capture: _AuthCapture | None = None,
|
||||
) -> Literal["transport", "auth_401", "auth_403", "protocol", "other"]:
|
||||
"""Classify a dispatch-time exception for circuit-breaker gating.
|
||||
|
||||
Only ``transport`` failures trip the per-server breaker. Auth
|
||||
@@ -1020,16 +1175,24 @@ class MCPClientManager:
|
||||
breaker. Protocol errors (``McpError``) come from a healthy
|
||||
connection that rejected the request.
|
||||
|
||||
``auth`` is detected via :class:`httpx.HTTPStatusError` since the
|
||||
streamable-http transport surfaces upstream HTTP errors through
|
||||
``response.raise_for_status()``. ``McpError`` payloads do not
|
||||
carry a clean status code; bare ``McpError`` therefore stays
|
||||
``protocol``.
|
||||
Auth detection prefers ``capture.status`` (response-hook
|
||||
introspection — the SDK swallows :class:`httpx.HTTPStatusError`
|
||||
in its ``post_writer`` so the carrier is the only signal that
|
||||
reaches us in production). The ``HTTPStatusError`` fallback is
|
||||
defense-in-depth for the non-SDK refresh path
|
||||
(:func:`turnstone.core.mcp_oauth._refresh_and_persist`) where
|
||||
``httpx`` errors propagate directly.
|
||||
"""
|
||||
if capture is not None and capture.status == 401:
|
||||
return "auth_401"
|
||||
if capture is not None and capture.status == 403:
|
||||
return "auth_403"
|
||||
if isinstance(exc, httpx.HTTPStatusError):
|
||||
status = exc.response.status_code
|
||||
if status in (401, 403):
|
||||
return "auth"
|
||||
if status == 401:
|
||||
return "auth_401"
|
||||
if status == 403:
|
||||
return "auth_403"
|
||||
if isinstance(exc, McpError):
|
||||
return "protocol"
|
||||
if isinstance(exc, BrokenPipeError | ConnectionResetError | EOFError | TimeoutError):
|
||||
@@ -1046,10 +1209,9 @@ class MCPClientManager:
|
||||
|
||||
Every entry sourced from ``_static_servers`` is, by construction,
|
||||
NOT ``auth_type='oauth_user'`` — :func:`_db_servers_to_config`
|
||||
strips oauth_user rows on the way in. Phase 5 pool catalogs do
|
||||
not contribute to ``_tool_map``; pool tools become reachable
|
||||
via ``_tool_map`` only when Phase 7 (catalog scoping) lands
|
||||
per-user catalogs.
|
||||
strips oauth_user rows on the way in. Pool catalogs do not
|
||||
contribute to ``_tool_map`` today; pool tools become reachable
|
||||
via ``_tool_map`` only once per-user catalog scoping lands.
|
||||
"""
|
||||
new_tools: list[dict[str, Any]] = []
|
||||
new_map: dict[str, tuple[str, str]] = {}
|
||||
@@ -1912,14 +2074,14 @@ class MCPClientManager:
|
||||
def is_mcp_tool(self, func_name: str) -> bool:
|
||||
"""Check whether *func_name* belongs to an MCP server.
|
||||
|
||||
Phase 5 caveat: only static-path tools (``auth_type ∈ {none,
|
||||
static}``) populate ``_tool_map``; pool tools become reachable
|
||||
via this method only when Phase 7 (catalog scoping) lands
|
||||
per-user catalogs. Until then, pool dispatch is reachable from
|
||||
production callers like ``ChatSession._exec_mcp_tool`` only when
|
||||
the LLM produces a ``mcp__{server}__{tool}`` name that bypasses
|
||||
``is_mcp_tool``-style gating, or via direct ``call_tool_sync``
|
||||
with a known prefixed name (the path the new pool tests use).
|
||||
Caveat: only static-path tools (``auth_type ∈ {none, static}``)
|
||||
populate ``_tool_map``; pool tools become reachable via this
|
||||
method only once per-user catalog scoping lands. Until then,
|
||||
pool dispatch is reachable from production callers like
|
||||
``ChatSession._exec_mcp_tool`` only when the LLM produces a
|
||||
``mcp__{server}__{tool}`` name that bypasses ``is_mcp_tool``-
|
||||
style gating, or via direct ``call_tool_sync`` with a known
|
||||
prefixed name (the path the new pool tests use).
|
||||
"""
|
||||
return func_name in self._tool_map
|
||||
|
||||
@@ -2025,6 +2187,17 @@ class MCPClientManager:
|
||||
When ``user_id`` is supplied AND the resolved server's ``auth_type``
|
||||
is ``oauth_user``, dispatch goes through the per-(user, server)
|
||||
pool. Otherwise the call takes the byte-identical static path.
|
||||
|
||||
Pool-path 401/403 handling: the SDK's ``post_writer`` swallows
|
||||
``httpx.HTTPStatusError``; we recover the upstream auth signal
|
||||
via a response-hook carrier on the per-dispatch
|
||||
``httpx.AsyncClient``. A 401 triggers one refresh-and-retry
|
||||
with ``force_refresh=True``; persistent 401 emits
|
||||
``mcp_consent_required``. A 403 with
|
||||
``WWW-Authenticate: error="insufficient_scope"`` emits
|
||||
``mcp_insufficient_scope`` with the parsed scope set. Other
|
||||
403s emit ``mcp_tool_call_forbidden``. Auth failures NEVER
|
||||
trip the per-server breaker.
|
||||
"""
|
||||
mapping = self._tool_map.get(func_name)
|
||||
server_name: str | None = None
|
||||
@@ -2033,10 +2206,10 @@ class MCPClientManager:
|
||||
server_name, original_name = mapping
|
||||
|
||||
# Pool dispatch is gated on (a) caller passing user_id and
|
||||
# (b) the server row being auth_type=oauth_user. The Phase 5
|
||||
# tool-map only carries static-path entries; pool catalogs land
|
||||
# in Phase 7. Consequently, in Phase 5 a pool dispatch reaches
|
||||
# this branch only when the caller supplied func_name as
|
||||
# (b) the server row being auth_type=oauth_user. The current
|
||||
# tool-map only carries static-path entries; pool catalogs are
|
||||
# not merged in. Consequently, a pool dispatch reaches this
|
||||
# branch only when the caller supplied func_name as
|
||||
# ``mcp__{server}__{tool}`` and we resolve auth_type via storage.
|
||||
if user_id and self._app_state is not None and self._storage is not None:
|
||||
pool_target = self._resolve_pool_target(func_name, server_name, original_name)
|
||||
@@ -2098,7 +2271,7 @@ class MCPClientManager:
|
||||
"""Resolve ``(server_name, original_name, server_row)`` for pool dispatch.
|
||||
|
||||
Returns ``None`` when the call should fall through to the static
|
||||
path. Phase 5: pool catalogs do not contribute to ``_tool_map``,
|
||||
path. Pool catalogs do not contribute to ``_tool_map`` today,
|
||||
so a pool tool is invoked by passing the prefixed name
|
||||
``mcp__{server}__{tool}`` directly. ``mcp_servers.auth_type``
|
||||
confirms pool eligibility.
|
||||
@@ -2155,10 +2328,86 @@ class MCPClientManager:
|
||||
(consent required, key mismatch, etc.). ``server_row`` is the
|
||||
row already resolved by ``_resolve_pool_target`` and is reused
|
||||
verbatim by ``_dispatch_pool`` to skip a duplicate DB hop.
|
||||
|
||||
Retry-on-401: the 401 branch in :meth:`_dispatch_pool` raises
|
||||
:class:`_PoolDispatchRetryRequested` after refreshing the
|
||||
bearer. Catching the signal HERE — at the sync boundary — means
|
||||
the retry is scheduled via a fresh
|
||||
:func:`asyncio.run_coroutine_threadsafe` call, which runs the
|
||||
retry coroutine in a brand-new :class:`asyncio.Task` with no
|
||||
inherited anyio cancel-scope state from the prior connect's
|
||||
``streamablehttp_client`` TaskGroup. An in-task retry inherits
|
||||
that scope state across :meth:`asyncio.Task.uncancel` and
|
||||
``loop.create_task`` and surfaces ``CancelledError`` from inside
|
||||
the retry's own anyio scope. The retry-count ceiling is one;
|
||||
callers see consent_required if both attempts 401.
|
||||
|
||||
The ``timeout`` is a wall-clock budget across both attempts —
|
||||
the retry's ``future.result`` window is reduced by however long
|
||||
the first attempt consumed before raising
|
||||
:class:`_PoolDispatchRetryRequested`. Without this, a slow
|
||||
first attempt followed by a stuck retry could double the
|
||||
caller-observed timeout.
|
||||
"""
|
||||
assert self._loop is not None
|
||||
start = time.monotonic()
|
||||
try:
|
||||
return self._run_pool_dispatch_attempt(
|
||||
retry_count=0,
|
||||
timeout=timeout,
|
||||
original_timeout=timeout,
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
original_name=original_name,
|
||||
arguments=arguments,
|
||||
server_row=server_row,
|
||||
)
|
||||
except _PoolDispatchRetryRequested:
|
||||
# auth_401: refresh already happened on the prior task;
|
||||
# re-issue on a fresh task so the retry's anyio scope
|
||||
# state is independent of the prior connect's TaskGroup
|
||||
# teardown.
|
||||
remaining = max(1, int(timeout - (time.monotonic() - start)))
|
||||
return self._run_pool_dispatch_attempt(
|
||||
retry_count=1,
|
||||
timeout=remaining,
|
||||
original_timeout=timeout,
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
original_name=original_name,
|
||||
arguments=arguments,
|
||||
server_row=server_row,
|
||||
)
|
||||
|
||||
def _run_pool_dispatch_attempt(
|
||||
self,
|
||||
*,
|
||||
retry_count: int,
|
||||
timeout: int,
|
||||
original_timeout: int,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
original_name: str,
|
||||
arguments: dict[str, Any],
|
||||
server_row: dict[str, Any],
|
||||
) -> str:
|
||||
"""Schedule one ``_dispatch_pool`` attempt and wait for the result.
|
||||
|
||||
Split out of :meth:`_dispatch_pool_sync` so the two retry
|
||||
attempts share scheduling + timeout-bookkeeping without
|
||||
re-introducing the ``for retry_count in (0, 1)`` loop (which
|
||||
gave both attempts the full ``timeout`` and required an
|
||||
unreachable ``RuntimeError`` fallback).
|
||||
|
||||
``original_timeout`` is the wall-clock budget the caller
|
||||
requested; ``timeout`` is what's left for this specific
|
||||
attempt. The ``TimeoutError`` message reports the original so
|
||||
callers see the budget they set, not the trimmed window.
|
||||
"""
|
||||
assert self._loop is not None
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self._dispatch_pool(
|
||||
retry_count=retry_count,
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
original_name=original_name,
|
||||
@@ -2172,7 +2421,7 @@ class MCPClientManager:
|
||||
except concurrent.futures.TimeoutError:
|
||||
future.cancel()
|
||||
self._cb_record_failure(server_name)
|
||||
raise TimeoutError(f"MCP tool call timed out after {timeout}s") from None
|
||||
raise TimeoutError(f"MCP tool call timed out after {original_timeout}s") from None
|
||||
|
||||
async def _dispatch_pool(
|
||||
self,
|
||||
@@ -2182,6 +2431,7 @@ class MCPClientManager:
|
||||
original_name: str,
|
||||
arguments: dict[str, Any],
|
||||
server_row: dict[str, Any],
|
||||
retry_count: int = 0,
|
||||
) -> str:
|
||||
"""Pool-side coroutine: resolve token, connect-or-reuse, dispatch.
|
||||
|
||||
@@ -2190,6 +2440,18 @@ class MCPClientManager:
|
||||
is supplied by ``_resolve_pool_target`` so this path doesn't
|
||||
re-issue the ``mcp_servers`` lookup; it's also pre-validated to
|
||||
have ``auth_type='oauth_user'``.
|
||||
|
||||
``retry_count`` is supplied by :meth:`_dispatch_pool_sync` and
|
||||
bounds the auth_401 refresh-and-retry to one re-issue. The first
|
||||
attempt (``retry_count == 0``) refreshes the token via
|
||||
``force_refresh=True`` and raises
|
||||
:class:`_PoolDispatchRetryRequested` so the sync caller schedules
|
||||
the retry on a fresh :class:`asyncio.Task` (an in-task retry
|
||||
inherits anyio cancel-scope state from the prior connect's
|
||||
TaskGroup teardown and surfaces ``CancelledError`` from inside
|
||||
the retry's own anyio scope; the cross-task hop avoids that).
|
||||
At the ceiling (``retry_count == 1``) the auth_401 branch emits
|
||||
``mcp_consent_required`` instead.
|
||||
"""
|
||||
if self._app_state is None:
|
||||
raise RuntimeError("Pool dispatch requires set_app_state() to have been called")
|
||||
@@ -2199,8 +2461,17 @@ class MCPClientManager:
|
||||
# gate runs only AFTER we have a usable access token; that way a
|
||||
# spurious cooldown-expired probe never lands here without ending
|
||||
# in either a real success or a real transport failure.
|
||||
#
|
||||
# ``retry_count >= 1`` forces the AS round-trip: the previous
|
||||
# attempt's auth_401 branch evicted the session and signalled
|
||||
# this retry; the local cached token is the one the AS just
|
||||
# rejected, so reading it back without ``force_refresh=True``
|
||||
# would re-attempt with the same (rejected) bearer.
|
||||
lookup: TokenLookupResult = await get_user_access_token_classified(
|
||||
app_state=self._app_state, user_id=user_id, server_name=server_name
|
||||
app_state=self._app_state,
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
force_refresh=retry_count > 0,
|
||||
)
|
||||
if lookup.kind == "missing":
|
||||
return _structured_error(
|
||||
@@ -2258,6 +2529,13 @@ class MCPClientManager:
|
||||
key = (user_id, server_name)
|
||||
entry = await self._ensure_pool_entry(key)
|
||||
|
||||
# First attempt — fresh capture per dispatch so a prior
|
||||
# dispatch's 401 cannot leak into this one's classification.
|
||||
# Even though ``open_lock`` is held across ``call_tool``,
|
||||
# allocating per-call (rather than per-entry) protects against
|
||||
# future concurrent-multiplex regressions if the lock scope
|
||||
# ever shrinks.
|
||||
capture = _AuthCapture()
|
||||
try:
|
||||
result = await self._dispatch_pool_with_entry(
|
||||
entry=entry,
|
||||
@@ -2266,23 +2544,48 @@ class MCPClientManager:
|
||||
access_token=access_token,
|
||||
original_name=original_name,
|
||||
arguments=arguments,
|
||||
auth_capture=capture,
|
||||
)
|
||||
except BaseException as exc:
|
||||
classification = self._classify_failure(exc)
|
||||
if classification == "auth":
|
||||
# 401/403 from the upstream — pool-entry-only, never
|
||||
# affects the breaker. Drop the cached session so the
|
||||
# next dispatch re-authenticates with a fresh token.
|
||||
evict = self._user_pool_entries.get(key)
|
||||
if evict is not None:
|
||||
evict.session = None
|
||||
log.debug("mcp_pool.auth_failure", exc_info=exc)
|
||||
raise
|
||||
classification = self._classify_failure(exc, capture=capture)
|
||||
if classification == "auth_401":
|
||||
self._evict_session(key)
|
||||
if retry_count == 0:
|
||||
# 401 on the initial attempt: signal the sync caller
|
||||
# to re-issue on a fresh :class:`asyncio.Task`. The
|
||||
# retry's :meth:`_dispatch_pool` invocation runs the
|
||||
# token lookup with ``force_refresh=True`` (the
|
||||
# ``retry_count > 0`` branch above), guaranteeing
|
||||
# the bearer attached to the retry's connect is
|
||||
# different from the one the AS just rejected. The
|
||||
# cross-task hop avoids the in-task anyio
|
||||
# cancel-scope inheritance that surfaces a
|
||||
# ``CancelledError`` from inside the retry's own
|
||||
# ``streamablehttp_client`` TaskGroup — see
|
||||
# :meth:`_dispatch_pool_sync` for the architectural
|
||||
# rationale.
|
||||
log.debug("mcp_pool.auth_401_initial", exc_info=exc)
|
||||
raise _PoolDispatchRetryRequested from None
|
||||
# retry_count == 1 — refreshed bearer also rejected;
|
||||
# emit consent_required so the user/operator re-grants.
|
||||
log.debug("mcp_pool.auth_401_retry_failed", exc_info=exc)
|
||||
return _structured_error(
|
||||
code="mcp_consent_required",
|
||||
server=server_name,
|
||||
detail="Refreshed token still rejected. Re-consent required.",
|
||||
)
|
||||
if classification == "auth_403":
|
||||
self._evict_session(key)
|
||||
log.debug("mcp_pool.auth_403", exc_info=exc)
|
||||
return await self._handle_auth_403(
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
server_row=server_row,
|
||||
capture=capture,
|
||||
)
|
||||
if classification == "transport":
|
||||
self._cb_record_failure(server_name)
|
||||
evict = self._user_pool_entries.get(key)
|
||||
if evict is not None:
|
||||
evict.session = None
|
||||
self._evict_session(key)
|
||||
log.debug("mcp_pool.transport_failure", exc_info=exc)
|
||||
raise
|
||||
# protocol / other — don't trip the breaker.
|
||||
@@ -2291,6 +2594,63 @@ class MCPClientManager:
|
||||
self._cb_record_success(server_name)
|
||||
return result
|
||||
|
||||
def _evict_session(self, key: tuple[str, str]) -> None:
|
||||
"""Drop the cached session on a pool entry. Stack/streams left for reconnect.
|
||||
|
||||
Auth/transport branches both call this — the next connect's
|
||||
``_connect_one_pool`` tears down the stale stack lazily via the
|
||||
stale-entry guard at the top of the method. Closing eagerly
|
||||
from here is incorrect under cancellation: ``stack.aclose()``
|
||||
must run inside the same anyio scope it was entered in, which
|
||||
the next connect arranges.
|
||||
"""
|
||||
evict = self._user_pool_entries.get(key)
|
||||
if evict is not None:
|
||||
evict.session = None
|
||||
|
||||
async def _handle_auth_403(
|
||||
self,
|
||||
*,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
server_row: dict[str, Any],
|
||||
capture: _AuthCapture,
|
||||
) -> str:
|
||||
"""Map a 403 + WWW-Authenticate into a structured error.
|
||||
|
||||
``error="insufficient_scope"`` becomes ``mcp_insufficient_scope``
|
||||
with the parsed ``scope=...`` set so the dashboard renderer
|
||||
can construct an authorize URL with the union of original +
|
||||
new scopes — re-consenting with the original scopes alone
|
||||
would loop because the AS would re-issue the same insufficient
|
||||
token. Other 403s become a generic ``mcp_tool_call_forbidden``
|
||||
with no retry; the user lacks permission and a step-up
|
||||
wouldn't help.
|
||||
"""
|
||||
header = capture.www_authenticate or ""
|
||||
error_token = parse_www_authenticate_error(header)
|
||||
if error_token == "insufficient_scope":
|
||||
scopes = parse_www_authenticate_scope(header)
|
||||
scopes = scopes[:_MAX_INSUFFICIENT_SCOPE_REPORTED]
|
||||
await emit_insufficient_scope_audit(
|
||||
app_state=self._app_state,
|
||||
user_id=user_id,
|
||||
server_name=server_name,
|
||||
server_row=server_row,
|
||||
scopes=scopes,
|
||||
)
|
||||
return _structured_error(
|
||||
code="mcp_insufficient_scope",
|
||||
server=server_name,
|
||||
detail=("Tool requires elevated scopes. Re-consent flow with new scopes required."),
|
||||
scopes_required=list(scopes),
|
||||
)
|
||||
return _structured_error(
|
||||
code="mcp_tool_call_forbidden",
|
||||
server=server_name,
|
||||
detail="Tool call forbidden by upstream policy.",
|
||||
)
|
||||
|
||||
async def _dispatch_pool_with_entry(
|
||||
self,
|
||||
*,
|
||||
@@ -2300,15 +2660,23 @@ class MCPClientManager:
|
||||
access_token: str,
|
||||
original_name: str,
|
||||
arguments: dict[str, Any],
|
||||
auth_capture: _AuthCapture,
|
||||
) -> str:
|
||||
"""Acquire ``entry.open_lock`` only across connect-or-reuse, then dispatch.
|
||||
"""Hold ``entry.open_lock`` across connect-or-reuse AND ``call_tool``.
|
||||
|
||||
Releasing ``open_lock`` before ``call_tool`` lets two concurrent
|
||||
tool calls from the SAME user against the SAME server multiplex
|
||||
on a shared :class:`mcp.ClientSession` (request_id-correlated by
|
||||
the SDK). The eviction interlock now uses ``entry.in_flight``:
|
||||
eviction skips entries with in-flight calls so a long-running
|
||||
``call_tool`` can't have its session yanked mid-await.
|
||||
Lock held across ``call_tool`` because the per-dispatch
|
||||
``_AuthCapture`` is keyed off the httpx event hook; releasing
|
||||
would let a concurrent same-(user, server) dispatch overwrite
|
||||
the carrier mid-flight, attributing one caller's auth failure
|
||||
to another. Holding the lock serialises same-(user, server)
|
||||
calls — acceptable because that contention scenario is rare
|
||||
in practice (ChatSession dispatches sequentially and per-
|
||||
(user, server) parallelism is not a production requirement).
|
||||
|
||||
``entry.in_flight`` accounting is preserved for the eviction
|
||||
interlock — it's belt-and-braces here since ``open_lock.locked()``
|
||||
already signals "do not evict", but the in-flight counter
|
||||
remains the source of truth for :meth:`_close_pool_entry_if_idle`.
|
||||
"""
|
||||
async with entry.open_lock:
|
||||
entry.last_used = time.monotonic()
|
||||
@@ -2316,7 +2684,9 @@ class MCPClientManager:
|
||||
session = entry.session
|
||||
if session is None:
|
||||
# Lazy connect — also covers post-eviction recovery.
|
||||
fresh = await self._connect_one_pool(key, cfg, access_token)
|
||||
fresh = await self._connect_one_pool(
|
||||
key, cfg, access_token, auth_capture=auth_capture
|
||||
)
|
||||
session = fresh.session
|
||||
if session is None:
|
||||
raise RuntimeError(f"Pool connect for {key!r} produced no session")
|
||||
@@ -2482,13 +2852,24 @@ def _decode_tool_result(result: Any) -> str:
|
||||
return output
|
||||
|
||||
|
||||
def _structured_error(*, code: str, server: str, detail: str) -> str:
|
||||
def _structured_error(
|
||||
*,
|
||||
code: str,
|
||||
server: str,
|
||||
detail: str,
|
||||
scopes_required: list[str] | None = None,
|
||||
) -> str:
|
||||
"""Encode a pool-dispatch failure as a JSON string.
|
||||
|
||||
Returned to the agent through ``_exec_mcp_tool`` so the LLM can
|
||||
narrate "tool unavailable, consent required" rather than crashing
|
||||
the workstream. Schema mirrors RFC §6 (consent UX) plus the
|
||||
decrypt-failure code introduced in RFC §5.3.
|
||||
the workstream. Schema covers ``mcp_consent_required``,
|
||||
decrypt-failure (``mcp_token_undecryptable_key_unknown``), and the
|
||||
``mcp_insufficient_scope`` step-up shape.
|
||||
|
||||
``scopes_required`` is omitted from the payload when ``None`` —
|
||||
the dashboard renderer keys on its presence to construct an
|
||||
authorize URL with the union of original + new scopes.
|
||||
|
||||
Operator-actionable encryption-key fingerprints are intentionally
|
||||
NOT included in this payload: they are already captured server-side
|
||||
@@ -2496,14 +2877,14 @@ def _structured_error(*, code: str, server: str, detail: str) -> str:
|
||||
to the LLM (and through it to the model provider) would be
|
||||
unnecessary disclosure.
|
||||
"""
|
||||
payload: dict[str, Any] = {
|
||||
"error": {
|
||||
err: dict[str, Any] = {
|
||||
"code": code,
|
||||
"server": server,
|
||||
"detail": detail,
|
||||
}
|
||||
}
|
||||
return json.dumps(payload)
|
||||
if scopes_required is not None:
|
||||
err["scopes_required"] = scopes_required
|
||||
return json.dumps({"error": err})
|
||||
|
||||
|
||||
def _pool_cfg_from_row(row: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
"""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.
|
||||
"""
|
||||
|
||||
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
|
||||
|
||||
|
||||
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 ``\\``)
|
||||
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 s and all(0x21 <= ord(c) <= 0x7E and c not in '"\\' for c in 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
|
||||
+71
-47
@@ -26,7 +26,6 @@ import concurrent.futures
|
||||
import contextlib
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
import urllib.parse
|
||||
@@ -39,6 +38,7 @@ import httpx
|
||||
from turnstone.core.audit import record_audit
|
||||
from turnstone.core.log import get_logger
|
||||
from turnstone.core.mcp_crypto import MCPTokenDecryptError
|
||||
from turnstone.core.mcp_http_parsers import parse_www_authenticate_bearer
|
||||
from turnstone.core.oauth_ssrf import (
|
||||
OAuthSSRFError,
|
||||
sanitize_log_text,
|
||||
@@ -127,53 +127,19 @@ class ASMetadata:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_PRM_RESOURCE_METADATA_KEY_RE = re.compile(
|
||||
r"resource_metadata\s*=\s*",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
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 §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
|
||||
|
||||
|
||||
def _parse_prm_url_from_www_authenticate(header: str) -> str | None:
|
||||
"""Extract ``resource_metadata`` URL from a ``WWW-Authenticate: Bearer`` header.
|
||||
|
||||
Returns the URL string or ``None`` when the header lacks the param,
|
||||
is malformed, or terminates the quoted-string prematurely.
|
||||
is malformed, or terminates the quoted-string prematurely. Delegates
|
||||
to :func:`parse_www_authenticate_bearer` so quoted-string handling
|
||||
(RFC 7230 §3.2.6 backslash escapes) and the multi-challenge
|
||||
defence-in-depth guard live in one place.
|
||||
"""
|
||||
if not header:
|
||||
return None
|
||||
for match in _PRM_RESOURCE_METADATA_KEY_RE.finditer(header):
|
||||
parsed = _parse_quoted_string(header, match.end())
|
||||
if parsed is not None:
|
||||
return parsed[0]
|
||||
return None
|
||||
params = parse_www_authenticate_bearer(header)
|
||||
return params.get("resource_metadata") or None
|
||||
|
||||
|
||||
async def _fetch_prm_issuer(
|
||||
@@ -1145,8 +1111,8 @@ class TokenLookupResult:
|
||||
rejected" — three states that the ``Optional[str]``-returning
|
||||
:func:`get_user_access_token` collapses to ``None``. The
|
||||
distinction matters because they map to different user-facing
|
||||
errors (RFC §5.3 forbids emitting ``mcp_consent_required`` on a
|
||||
decrypt failure).
|
||||
errors — emitting ``mcp_consent_required`` on a decrypt failure
|
||||
would be wrong (the user can't fix it; only an operator can).
|
||||
"""
|
||||
|
||||
kind: Literal["token", "missing", "decrypt_failure", "refresh_failed"] = "missing"
|
||||
@@ -1177,11 +1143,12 @@ async def get_user_access_token(*, app_state: Any, user_id: str, server_name: st
|
||||
|
||||
|
||||
async def get_user_access_token_classified(
|
||||
*, app_state: Any, user_id: str, server_name: str
|
||||
*, app_state: Any, user_id: str, server_name: str, force_refresh: bool = False
|
||||
) -> TokenLookupResult:
|
||||
"""Tagged token lookup with refresh-on-expiry.
|
||||
|
||||
Walks the §1.5 / RFC §6 state machine and returns a tagged result so
|
||||
Walks the token-lookup state machine (token / missing /
|
||||
decrypt_failure / refresh_failed) and returns a tagged result so
|
||||
the dispatcher can map each failure mode to the right user-facing
|
||||
error.
|
||||
|
||||
@@ -1196,12 +1163,28 @@ async def get_user_access_token_classified(
|
||||
surviving local caller then serializes against other nodes via
|
||||
the cluster lock. The re-read inside the locked block collapses
|
||||
both contention windows.
|
||||
|
||||
``force_refresh=True`` bypasses the local freshness check and
|
||||
forces an AS round-trip — used by the dispatch path when the
|
||||
upstream AS rejected a token our cache still considered fresh
|
||||
(e.g., AS-side revocation). Concurrent ``force_refresh=True``
|
||||
callers still collapse to one round-trip via the dual-layer lock:
|
||||
the second caller sees ``last_refreshed > t_lock_request_started``
|
||||
and reuses the freshly-refreshed token.
|
||||
"""
|
||||
token_store: MCPTokenStore | None = getattr(app_state, "mcp_token_store", None)
|
||||
if token_store is None:
|
||||
log.debug("mcp_server.oauth.token_store_unconfigured")
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
# Captured BEFORE we acquire any lock so the inside-lock guard can
|
||||
# tell whether another caller refreshed under contention. Truncated
|
||||
# to seconds because ``last_refreshed`` storage has second-precision
|
||||
# ISO8601 — comparing microsecond-precision against second-precision
|
||||
# would race when the refresh and the contention occur in the same
|
||||
# wall-clock second.
|
||||
t_lock_request_started = datetime.now(UTC).replace(microsecond=0)
|
||||
|
||||
try:
|
||||
plain = await asyncio.to_thread(token_store.get_user_token, user_id, server_name)
|
||||
except MCPTokenDecryptError as exc:
|
||||
@@ -1219,7 +1202,7 @@ async def get_user_access_token_classified(
|
||||
return TokenLookupResult(kind="missing")
|
||||
|
||||
expires_at = plain.get("expires_at")
|
||||
if not _token_needs_refresh(expires_at):
|
||||
if not force_refresh and not _token_needs_refresh(expires_at):
|
||||
return TokenLookupResult(kind="token", token=plain["access_token"])
|
||||
|
||||
storage = _get_storage(app_state)
|
||||
@@ -1272,7 +1255,19 @@ async def get_user_access_token_classified(
|
||||
if plain2 is None:
|
||||
return TokenLookupResult(kind="missing")
|
||||
expires_at2 = plain2.get("expires_at")
|
||||
# Reuse the freshly-refreshed token under two conditions:
|
||||
# 1. ``force_refresh=False`` and the cached token is still fresh
|
||||
# (existing fast path).
|
||||
# 2. ``force_refresh=True`` BUT another same-key caller already
|
||||
# refreshed under contention since we started waiting for the
|
||||
# lock. ``last_refreshed`` is the storage-side replacement
|
||||
# timestamp; if it advanced past ``t_lock_request_started``,
|
||||
# we lost the race and should reuse rather than refresh again.
|
||||
if not _token_needs_refresh(expires_at2):
|
||||
if not force_refresh:
|
||||
return TokenLookupResult(kind="token", token=plain2["access_token"])
|
||||
last_refreshed = _parse_iso_to_utc(plain2.get("last_refreshed") or "")
|
||||
if last_refreshed is not None and last_refreshed >= t_lock_request_started:
|
||||
return TokenLookupResult(kind="token", token=plain2["access_token"])
|
||||
refresh_value2 = plain2.get("refresh_token")
|
||||
if not refresh_value2:
|
||||
@@ -1413,7 +1408,7 @@ async def _refresh_and_persist(
|
||||
if not isinstance(new_access, str) or not new_access:
|
||||
raise MCPOAuthRefreshFailed("refresh response missing access_token")
|
||||
|
||||
# RFC 6749 §6 — the AS MAY omit ``refresh_token`` from the refresh
|
||||
# RFC 6749 section 6 — the AS MAY omit ``refresh_token`` from the refresh
|
||||
# response. Most production ASes (Google, default Auth0, default
|
||||
# Okta) do not rotate the refresh token; clearing the column on
|
||||
# every refresh would force the user to re-consent every hour.
|
||||
@@ -1533,6 +1528,35 @@ async def _audit_event(
|
||||
log.debug("mcp_server.oauth.audit_emit_failed", action=action, exc_info=True)
|
||||
|
||||
|
||||
async def emit_insufficient_scope_audit(
|
||||
*,
|
||||
app_state: Any,
|
||||
user_id: str,
|
||||
server_name: str,
|
||||
server_row: dict[str, Any],
|
||||
scopes: tuple[str, ...],
|
||||
) -> None:
|
||||
"""Emit ``mcp_server.oauth.insufficient_scope_emitted`` audit event.
|
||||
|
||||
Best-effort: :func:`_audit_event` already swallows storage / write
|
||||
failures internally so audit emission never breaks dispatch.
|
||||
Operators tracking step-up patterns consume this via the standard
|
||||
audit log. Called by the pool dispatcher after classifying a 403
|
||||
``WWW-Authenticate: error="insufficient_scope"``.
|
||||
"""
|
||||
if app_state is None:
|
||||
return
|
||||
server_id = str(server_row.get("server_id") or "") if server_row else ""
|
||||
await _audit_event(
|
||||
app_state,
|
||||
server_id=server_id,
|
||||
user_id=user_id,
|
||||
action="mcp_server.oauth.insufficient_scope_emitted",
|
||||
server_name=server_name,
|
||||
detail={"scopes_required": list(scopes)},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HTTP handlers — /api/mcp/oauth/start and /api/mcp/oauth/callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user