diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index 1201a156c00..623f07bcac1 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -128,6 +128,7 @@ def wire_server( class OwnedHTTPServer(ThreadingHTTPServer): daemon_threads = False + request_queue_size = 128 def server_bind(self) -> None: super().server_bind() diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index 13759ce245c..6b0034b35bb 100644 --- a/tests/integration/mcp/test_mcp_access_matrix.py +++ b/tests/integration/mcp/test_mcp_access_matrix.py @@ -1,19 +1,28 @@ +import re +import textwrap import uuid +from collections.abc import Iterator +from pathlib import Path from typing import Final import pytest -from integration._support.client import Gateway +from integration._support.client import Gateway, gateway_from_environment from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, McpCaller, McpPeer, + Outcome, PeerKind, + ScriptedTool, peer_of, register_mcp, + scripted_peer, + text_result, tool_calls, ) from integration._support.mcp_grants import SUBJECTS, Subject, grant +from integration._support.process import owned_proxy CALLABLE: Final = {"add": {"a": 1, "b": 2}, "multiply": {"a": 2, "b": 3}} RESULTS: Final = {"add": "3", "multiply": "6"} @@ -122,3 +131,87 @@ def test_same_tool_name_on_two_servers_routes_by_prefix(gateway: Gateway) -> Non assert outcome.ok and outcome.text == "10", outcome.raw assert tool_calls(first.drain()) == () assert [call["body"]["params"]["name"] for call in tool_calls(second.drain())] == ["add"] + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +@pytest.mark.parametrize("subject", SUBJECTS) +def test_each_subjects_call_is_evaluated_only_against_the_catalog_its_own_listing_served( + echo_rig: Gateway, subject: Subject +) -> None: + described: Final = "Adds under grant " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + group: Final = "grp" + uuid.uuid4().hex[:8] + alias: Final = "cat" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, mcp_access_groups=[group]) + caller: Final = grant( + scenario, subject, (identity,), (identity,), access_group=group, allowed_tools={identity: ("add",)} + ) + reach: Final = McpCaller(echo_rig, caller.key, "mcp", alias, caller.headers) + assert reach.initialize().ok + cold: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE})) + listed: Final = reach.list_tools() + assert listed.ok and f"{alias}-add" in listed.tools, listed.raw + warm: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE})) + assert (cold, warm) == (_UNLISTED, described), (cold, warm) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" + + +def test_end_users_of_one_key_share_its_catalog_slot_because_the_identity_excludes_the_end_user( + echo_rig: Gateway, +) -> None: + described: Final = "Adds for end users " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "eu" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = grant(scenario, "end_user", (identity,), (identity,)) + first: Final = McpCaller(echo_rig, granted.key, "mcp", alias, granted.headers) + second: Final = McpCaller( + echo_rig, granted.key, "mcp", alias, {"x-litellm-end-user-id": "integration-" + uuid.uuid4().hex[:10]} + ) + assert second.initialize().ok + assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == _UNLISTED + listed: Final = first.list_tools() + assert listed.ok and f"{alias}-add" in listed.tools, listed.raw + assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == described, ( + "the end-user header is intentionally not part of the catalog identity: one key, one slot" + ) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_accounting_guardrails.py b/tests/integration/mcp/test_mcp_accounting_guardrails.py index 323afad40db..594950a9c13 100644 --- a/tests/integration/mcp/test_mcp_accounting_guardrails.py +++ b/tests/integration/mcp/test_mcp_accounting_guardrails.py @@ -1,11 +1,23 @@ +import json import uuid -from collections.abc import Iterator +from collections.abc import Generator, Iterator, Mapping from contextlib import contextmanager +from dataclasses import dataclass from hashlib import sha256 +from pathlib import Path from typing import Final +import httpx import pytest -from integration._support.client import Gateway, JsonValue, Scenario, eventually +import yaml +from integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, +) from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, @@ -13,10 +25,16 @@ from integration._support.mcp import ( McpCaller, McpPeer, Outcome, + ScriptedTool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, ) +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue DEFAULT_COST: Final = 0.25 ADD_COST: Final = 0.5 @@ -191,3 +209,472 @@ def test_guardrail_removal_stops_blocking_without_restart(gateway: Gateway) -> N lambda calls: len(calls) >= 1, seconds=40, ) + + +MASK_ME: Final = "mask-integration-secret" +MASKED: Final = "[MASKED]" +COUNT_MISMATCH: Final = "count-mismatch-marker" +LOOKUP_DESCRIPTION: Final = "Look up one record" +LOOKUP_SCHEMA: Final = { + "type": "object", + "properties": {"record": {"type": "string", "description": "record identifier"}}, +} +SPEND_ROW: Final = 'SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s' +_RECORDER_CODE: Final = """\ +import os + +import httpx +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger + +SINK = "{sink}/native" + + +def _record(stage, data, call_type): + logging_obj = data.get("litellm_logging_obj") + return {{ + "stage": stage, + "pid": os.getpid(), + "call_type": call_type, + "litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id, + "messages": data.get("messages"), + "mcp_tool_name": data.get("mcp_tool_name"), + "mcp_arguments": data.get("mcp_arguments"), + "mcp_tool_description": data.get("mcp_tool_description"), + "mcp_input_schema": data.get("mcp_input_schema"), + }} + + +async def _post(record): + async with httpx.AsyncClient(timeout=5) as client: + await client.post(SINK, json=record) + + +class HookRecorder(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + await _post(_record("pre", data, call_type)) + + async def async_moderation_hook(self, data, user_api_key_dict, call_type): + await _post(_record("during", data, call_type)) + + async def async_post_mcp_tool_call_hook(self, kwargs, response_obj, start_time, end_time): + await _post( + {{ + "stage": "post", + "pid": os.getpid(), + "litellm_call_id": kwargs.get("litellm_call_id"), + "tool": kwargs.get("mcp_tool_call_metadata"), + "content": [item.model_dump() for item in response_obj.mcp_tool_call_response], + }} + ) + + +class SinkGuardrail(CustomGuardrail): + def __init__(self, api_base, **kwargs): + super().__init__(**kwargs) + self.api_base = api_base + + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + payload = {{ + "pid": os.getpid(), + "input_type": input_type, + "litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id, + "texts": inputs.get("texts"), + "tools": inputs.get("tools"), + "structured_messages": inputs.get("structured_messages"), + "mcp_tool_name": request_data.get("mcp_tool_name"), + }} + async with httpx.AsyncClient(timeout=5) as client: + verdict = (await client.post(self.api_base, json=payload)).json() + return {{**inputs, "texts": verdict["texts"]}} + + +recorder = HookRecorder() +""" + + +@dataclass(frozen=True, slots=True) +class Sunk: + target: str + body: dict[str, JsonValue] + + +@dataclass(frozen=True, slots=True) +class HooksRig: + gateway: Gateway + sibling: Gateway + sink: Wire + guardrail: str + + def sunk(self) -> tuple[Sunk, ...]: + return tuple(Sunk(request.target, JSON_OBJECT.validate_json(request.body)) for request in self.sink.drain()) + + +def _guardrail_sink(request: Request) -> Reply: + if not request.target.startswith("/guardrail"): + return Reply() + texts: Final = JSON_OBJECT.validate_json(request.body).get("texts") + assert isinstance(texts, list), texts + masked: Final = [str(text).replace(MASK_ME, MASKED) for text in texts] + extra: Final = ["extra"] if any(COUNT_MISMATCH in text for text in masked) else [] + return Reply(body=json.dumps({"texts": [*masked, *extra]}).encode()) + + +@pytest.fixture(scope="module") +def hooks_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[HooksRig]: + directory: Final = tmp_path_factory.mktemp("guardrail-payloads") + guardrail: Final = "sink" + uuid.uuid4().hex[:8] + with wire_server(_guardrail_sink) as sink: + (directory / "hook_recorder.py").write_text(_RECORDER_CODE.format(sink=sink.url)) + config: Final = JSON_OBJECT.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + config["guardrails"] = [ + { + "guardrail_name": guardrail, + "litellm_params": { + "guardrail": "hook_recorder.SinkGuardrail", + "mode": ["pre_mcp_call", "post_mcp_call"], + "default_on": True, + "api_base": f"{sink.url}/guardrail", + }, + } + ] + config["litellm_settings"] = { + **object_value(config["litellm_settings"]), + "callbacks": ["hook_recorder.recorder"], + } + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, {}, config=path, workers=2) as candidate, + owned_proxy(gateway, directory, {}, config=path) as sibling, + ): + yield HooksRig(candidate, sibling, sink, guardrail) + + +def _worker(gateway: Gateway) -> int: + response: Final = gateway.client.get("/debug/memory/summary", headers={"x-litellm-api-key": gateway.key}) + assert response.status_code == 200, response.text + worker: Final = JSON_OBJECT.validate_json(response.content)["worker_pid"] + assert isinstance(worker, int), response.text + return worker + + +@contextmanager +def _pinned(gateway: Gateway) -> Generator[tuple[Gateway, int], None, None]: + limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1) + with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=limits) as client: + pinned: Final = Gateway(client, gateway.key, gateway.upstream_url) + yield pinned, _worker(pinned) + + +def _lookup_tool(result: str = "found") -> ScriptedTool: + return ScriptedTool( + "lookup", lambda _: text_result(result), description=LOOKUP_DESCRIPTION, input_schema=LOOKUP_SCHEMA + ) + + +def _generic(sunk: tuple[Sunk, ...], input_type: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(item.body for item in sunk if item.target == "/guardrail" and item.body["input_type"] == input_type) + + +def _native(sunk: tuple[Sunk, ...], stage: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(item.body for item in sunk if item.target == "/native" and item.body["stage"] == stage) + + +def _only(records: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]: + assert len(records) == 1, records + return records[0] + + +def _scan(sunk: tuple[Sunk, ...], call_id: JsonValue) -> dict[str, JsonValue]: + return _only(tuple(record for record in _generic(sunk, "request") if record["litellm_call_id"] == call_id)) + + +def _texts(content: JsonValue) -> list[JsonValue]: + assert isinstance(content, list), content + return [object_value(item)["text"] for item in content] + + +def _has_lookup(listing: Outcome) -> bool: + return any(tool.endswith("lookup") for tool in listing.tools) + + +def _spend_row(call_id: JsonValue) -> dict[str, JsonValue]: + assert isinstance(call_id, str), call_id + rows: Final = eventually(lambda: read_rows(SPEND_ROW, (call_id,)), lambda found: len(found) == 1, seconds=70) + return rows[0] + + +def _synthetic_message(name: str, arguments: Mapping[str, str]) -> list[dict[str, str]]: + return [{"role": "user", "content": f"Tool: {name}\nArguments: {dict(arguments)}"}] + + +def test_generic_sink_and_native_hooks_receive_listed_metadata_on_typed_keys_with_the_message_bytes_unchanged( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "payload" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + hooks_rig.sunk() + listed: Final = eventually(caller.list_tools, _has_lookup, seconds=30) + name: Final = next(tool for tool in listed.tools if tool.endswith("lookup")) + scans: Final = _generic(hooks_rig.sunk(), "request") + assert scans and all(scan["texts"] == [LOOKUP_DESCRIPTION, "record identifier"] for scan in scans), scans + arguments: Final = {"record": "r-1"} + outcome: Final = caller.call(name, arguments) + assert outcome.text == "found", outcome.raw + assert _worker(pinned) == worker + sunk: Final = hooks_rig.sunk() + pre: Final = _only(_native(sunk, "pre")) + call_id: Final = pre["litellm_call_id"] + generic: Final = _scan(sunk, call_id) + assert generic["texts"] == [LOOKUP_DESCRIPTION, "record identifier", "r-1"], generic + assert generic["tools"] == [ + { + "type": "function", + "function": { + "name": "lookup", + "description": LOOKUP_DESCRIPTION, + "parameters": {**LOOKUP_SCHEMA, "additionalProperties": False}, + "strict": False, + }, + } + ], generic + response: Final = _only(_generic(sunk, "response")) + assert (response["texts"], response["pid"]) == (["found"], worker), response + during: Final = _only(_native(sunk, "during")) + post: Final = _only(_native(sunk, "post")) + assert all(record["pid"] == worker for record in (pre, during, post)), sunk + assert all(record["litellm_call_id"] == call_id for record in (pre, during, post)), sunk + assert pre["call_type"] == "call_mcp_tool" and during["call_type"] == "call_mcp_tool", sunk + assert pre["messages"] == _synthetic_message("lookup", arguments), pre + assert during["messages"] == _synthetic_message("lookup", arguments), during + assert (pre["mcp_tool_name"], pre["mcp_arguments"]) == ("lookup", arguments), pre + assert pre["mcp_tool_description"] == LOOKUP_DESCRIPTION, pre + assert pre["mcp_input_schema"] == LOOKUP_SCHEMA, pre + assert (during["mcp_tool_description"], during["mcp_input_schema"]) == (None, None), during + assert _texts(post["content"]) == ["found"], post + row: Final = _spend_row(call_id) + assert row["status"] == "success", row + assert _tool_metadata(row)["name"] == "lookup", row + metadata: Final = row["metadata"] + assert isinstance(metadata, dict) and metadata["applied_guardrails"] == [hooks_rig.guardrail], metadata + + +def test_pre_call_mask_reaches_the_peer_and_post_call_mask_reaches_the_caller_on_one_call_id( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool(f"found {MASK_ME}")) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "mask" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + peer.drain() + hooks_rig.sunk() + outcome: Final = caller.call(name, {"record": MASK_ME}) + assert outcome.text == f"found {MASKED}", outcome.raw + assert _worker(pinned) == worker + reached: Final = tool_calls(peer.drain()) + assert len(reached) == 1, reached + params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"]) + assert params["arguments"] == {"record": MASKED}, params + sunk: Final = hooks_rig.sunk() + generic: Final = _scan(sunk, _only(_native(sunk, "pre"))["litellm_call_id"]) + scanned: Final = generic["texts"] + assert isinstance(scanned, list) and scanned[-1] == MASK_ME and MASKED not in scanned, generic + assert _only(_generic(sunk, "response"))["texts"] == [f"found {MASK_ME}"], sunk + during: Final = _only(_native(sunk, "during")) + assert during["mcp_arguments"] == {"record": MASKED}, during + assert during["messages"] == _synthetic_message("lookup", {"record": MASKED}), during + assert _only(_native(sunk, "post"))["litellm_call_id"] == generic["litellm_call_id"], sunk + assert _spend_row(generic["litellm_call_id"])["status"] == "success" + + +def test_call_time_description_and_schema_come_from_the_catalog_of_the_worker_that_served_the_listing( + hooks_rig: HooksRig, +) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + hooks_rig.gateway.scenario() as scenario, + _pinned(hooks_rig.sibling) as (second, second_worker), + ): + alias: Final = "local" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + other: Final = McpCaller(second, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(other.initialize, lambda outcome: outcome.ok, seconds=60).ok + with _pinned(hooks_rig.gateway) as (first, first_worker): + assert first_worker != second_worker + lister: Final = McpCaller(first, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(lister.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(lister.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + hooks_rig.sunk() + assert other.call(name, {"record": "r-2"}).text == "found" + assert _worker(second) == second_worker + elsewhere: Final = hooks_rig.sunk() + unlisted: Final = _only(_native(elsewhere, "pre")) + assert unlisted["pid"] == second_worker, unlisted + assert (unlisted["mcp_tool_description"], unlisted["mcp_input_schema"]) == (None, None), unlisted + assert _scan(elsewhere, unlisted["litellm_call_id"])["texts"] == ["r-2"], elsewhere + assert lister.call(name, {"record": "r-3"}).text == "found" + assert _worker(first) == first_worker + at_lister: Final = hooks_rig.sunk() + listed: Final = _only(_native(at_lister, "pre")) + assert listed["pid"] == first_worker, listed + assert (listed["mcp_tool_description"], listed["mcp_input_schema"]) == (LOOKUP_DESCRIPTION, LOOKUP_SCHEMA) + assert _scan(at_lister, listed["litellm_call_id"])["texts"] == [ + LOOKUP_DESCRIPTION, + "record identifier", + "r-3", + ] + assert _has_lookup(other.list_tools()) + hooks_rig.sunk() + assert other.call(name, {"record": "r-4"}).text == "found" + assert _worker(second) == second_worker + populated: Final = _only(_native(hooks_rig.sunk(), "pre")) + assert populated["pid"] == second_worker, populated + assert populated["mcp_tool_description"] == LOOKUP_DESCRIPTION, populated + + +@pytest.mark.parametrize("entry", ("mcp", "rest")) +@pytest.mark.parametrize("listed", (False, True)) +@pytest.mark.parametrize("record", ("r-1", MASK_ME)) +def test_long_descriptions_do_not_refuse_small_tpm_calls_or_change_masked_message_bytes( + hooks_rig: HooksRig, entry: EntryPoint, listed: bool, record: str +) -> None: + description: Final = "Gateway tool metadata. " * 300 + tool: Final = ScriptedTool( + "lookup", lambda _: text_result("found"), description=description, input_schema=LOOKUP_SCHEMA + ) + with ( + scripted_peer(tool) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "quota" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}, tpm_limit=64) + caller: Final = McpCaller(pinned, key, entry, headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + if listed: + listing: Final = eventually(lambda: caller.list_tools(identity), _has_lookup, seconds=30) + assert _has_lookup(listing), listing.raw + hooks_rig.sunk() + peer.drain() + arguments: Final = {"record": record} + outcome: Final = caller.call(f"{alias}-lookup", arguments, identity) + assert outcome.ok and outcome.text == "found", outcome.raw + assert _worker(pinned) == worker + reached: Final = tool_calls(peer.drain()) + assert len(reached) == 1, reached + params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"]) + masked_arguments: Final = {"record": record.replace(MASK_ME, MASKED)} + assert params["arguments"] == masked_arguments, params + sunk: Final = hooks_rig.sunk() + pre: Final = _only(tuple(hook for hook in _native(sunk, "pre") if hook["mcp_tool_name"] == "lookup")) + during: Final = _only(tuple(hook for hook in _native(sunk, "during") if hook["mcp_tool_name"] == "lookup")) + assert pre["messages"] == _synthetic_message("lookup", arguments), pre + assert during["messages"] == _synthetic_message("lookup", masked_arguments), during + assert _spend_row(pre["litellm_call_id"])["status"] == "success" + + +def _nested_schema(levels: int) -> dict[str, object]: + if levels == 0: + return {"type": "string", "description": "deepest leaf"} + return {"type": "object", "properties": {"a": _nested_schema(levels - 1)}} + + +def test_a_schema_past_the_scan_depth_is_not_published_while_one_at_the_limit_is_scanned_on_the_call( + hooks_rig: HooksRig, +) -> None: + shallow: Final = ScriptedTool("shallow", lambda _: text_result("found"), input_schema=_nested_schema(49)) + deep: Final = ScriptedTool("deep", lambda _: text_result("found"), input_schema=_nested_schema(50)) + with ( + scripted_peer(shallow, deep) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "depth" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + listing: Final = eventually(caller.list_tools, lambda outcome: len(outcome.tools) > 0, seconds=30) + assert listing.tools == (f"{alias}-shallow",), listing.raw + assert _worker(pinned) == worker + hooks_rig.sunk() + assert caller.call(f"{alias}-shallow", {"record": "r-1"}).text == "found" + scanned: Final = _only(_generic(hooks_rig.sunk(), "request")) + assert scanned["texts"] == ["deepest leaf", "r-1"], scanned + unpublished: Final = caller.call(f"{alias}-deep", {"record": "r-2"}) + assert _worker(pinned) == worker + assert unpublished.error is not None and "Tool 'deep' not found" in unpublished.raw, unpublished.raw + reached: Final = tool_calls(peer.drain()) + assert [object_value(JSON_OBJECT.validate_python(call["body"])["params"])["name"] for call in reached] == [ + "shallow" + ] + relisted: Final = _generic(hooks_rig.sunk(), "request") + assert all((scan["mcp_tool_name"], scan["litellm_call_id"]) == ("shallow", None) for scan in relisted), relisted + rows: Final = _rows(key, 3) + assert [(row["call_type"], row["status"]) for row in rows] == [ + ("list_mcp_tools", "success"), + ("call_mcp_tool", "success"), + ("call_mcp_tool", "failure"), + ], rows + assert [_tool_metadata(row)["name"] for row in rows[1:]] == ["shallow", "deep"], rows + failed: Final = object_value(JSON_OBJECT.validate_python(rows[2]["metadata"])["error_information"]) + assert failed["error_message"] == "404: Tool 'deep' not found", failed + + +def test_an_adapter_returning_the_wrong_number_of_texts_fails_closed_before_the_peer(hooks_rig: HooksRig) -> None: + with ( + scripted_peer(_lookup_tool()) as peer, + _pinned(hooks_rig.gateway) as (pinned, worker), + pinned.scenario() as scenario, + ): + alias: Final = "count" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias}) + assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok + name: Final = next( + tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool + ) + assert _worker(pinned) == worker + hooks_rig.sunk() + blocked: Final = caller.call(name, {"record": COUNT_MISMATCH}) + assert _worker(pinned) == worker + assert blocked.error is not None, blocked.raw + assert ( + "guardrail returned 4 texts for 3 MCP tool strings, so the redaction cannot be mapped back" in blocked.raw + ) + assert tool_calls(peer.drain()) == (), "the blocked call reached the peer" + sunk: Final = hooks_rig.sunk() + assert _only(_generic(sunk, "request"))["texts"] == [LOOKUP_DESCRIPTION, "record identifier", COUNT_MISMATCH] + assert _native(sunk, "post") == (), sunk + rows: Final = _rows(key, 2) + assert [(row["call_type"], row["status"]) for row in rows] == [ + ("list_mcp_tools", "success"), + ("call_mcp_tool", "failure"), + ], rows + assert _tool_metadata(rows[1])["arguments"] == {"record": COUNT_MISMATCH}, rows[1] diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index e53707a38cc..16dcaab274a 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -1,22 +1,32 @@ import base64 +import re +import textwrap import uuid +from collections.abc import Iterator +from pathlib import Path from typing import Final import pytest -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, McpCaller, McpPeer, + Outcome, + ScriptedTool, call_tool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, tool_names, ) from integration._support.oauth_server import oauth_server +from integration._support.process import owned_proxy +from pydantic import TypeAdapter ADD: Final = {"a": 2, "b": 3} STATIC_MODES: Final = ( @@ -301,3 +311,180 @@ def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_us listings: Final = _listings(peer) assert len(listings) == 1, listings assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr" + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +_PROBE_ARGUMENTS: Final = {"probe": _PROBE} +_HEADERS: Final = TypeAdapter(dict[str, str]) + + +def _store_byok_credential(scenario: Scenario, identity: str, key: str, secret: str) -> None: + stored: Final = scenario.gateway.client.post( + f"/v1/mcp/server/{identity}/user-credential", json={"credential": secret}, headers={"x-litellm-api-key": key} + ) + assert stored.status_code in (200, 201), stored.text + scenario.cleanups.callback( + scenario.gateway.client.delete, f"/v1/mcp/server/{identity}/user-credential", headers={"x-litellm-api-key": key} + ) + + +def test_rotating_the_credential_drops_the_callers_listing_until_it_lists_again(echo_rig: Gateway) -> None: + with mcp_peer() as peer, echo_rig.scenario() as scenario: + first: Final = "cred-" + uuid.uuid4().hex + second: Final = "cred-" + uuid.uuid4().hex + alias: Final = "rot" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, auth_type="bearer_token", credentials={"auth_value": first} + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(echo_rig, key, "mcp", alias) + assert caller.list_tools().ok + assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers" + rotated: Final = echo_rig.request( + "PUT", "/v1/mcp/server", {"server_id": identity, "credentials": {"auth_value": second}} + ) + assert rotated.status_code == 202, rotated.text + eventually( + lambda: _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)), lambda seen: seen == _UNLISTED + ) + peer.drain() + assert caller.list_tools().ok + assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers" + relisted: Final = _listings(peer) + assert len(relisted) == 1, relisted + assert _header(relisted[0], b"authorization") == f"Bearer {second}".encode(), relisted[0]["headers"] + + +def test_byok_callers_are_evaluated_against_their_own_listing_and_the_stored_secret_never_keys_the_slot( + echo_rig: Gateway, +) -> None: + with mcp_peer() as peer, echo_rig.scenario() as scenario: + alias: Final = "byok" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type="api_key", is_byok=True) + owner_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]}) + stranger_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]}) + owner_secret: Final = "byok-" + uuid.uuid4().hex + replacement: Final = "byok-" + uuid.uuid4().hex + _store_byok_credential(scenario, identity, owner_key, owner_secret) + _store_byok_credential(scenario, identity, stranger_key, "byok-" + uuid.uuid4().hex) + owner: Final = McpCaller(echo_rig, owner_key, "mcp", alias) + stranger: Final = McpCaller(echo_rig, stranger_key, "mcp", alias) + peer.drain() + assert owner.list_tools().ok + listings: Final = _listings(peer) + assert [_header(item, b"x-api-key") for item in listings] == [owner_secret.encode()], listings + own: Final = _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) + other: Final = _echoed_description(stranger.call(f"{alias}-add", _PROBE_ARGUMENTS)) + assert (own, other) == ("Add two integers", _UNLISTED), (own, other) + _store_byok_credential(scenario, identity, owner_key, replacement) + assert _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers", ( + "the slot is keyed by the client-supplied header, never by the stored credential" + ) + sent: Final = eventually( + lambda: (owner.call(f"{alias}-add", ADD).ok, tool_calls(peer.drain())), + lambda value: any(_header(call, b"x-api-key") == replacement.encode() for call in value[1]), + ) + assert sent[0], sent + + +def test_callers_with_different_server_scoped_auth_headers_are_evaluated_against_their_own_listings( + echo_rig: Gateway, +) -> None: + tool: Final = ScriptedTool( + "add", + lambda _: text_result("3"), + description=lambda headers: "Adds for " + headers.get("authorization", "nobody"), + ) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "scoped" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + acme_token: Final = "acme-" + uuid.uuid4().hex + globex_token: Final = "globex-" + uuid.uuid4().hex + acme: Final = McpCaller(echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {acme_token}"}) + globex: Final = McpCaller( + echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {globex_token}"} + ) + assert acme.list_tools().ok and globex.list_tools().ok + seen: Final = ( + _echoed_description(acme.call(f"{alias}-add", _PROBE_ARGUMENTS)), + _echoed_description(globex.call(f"{alias}-add", _PROBE_ARGUMENTS)), + ) + assert seen == (f"Adds for Bearer {acme_token}", f"Adds for Bearer {globex_token}"), seen + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" + + +def test_deprecated_string_x_mcp_auth_callers_on_a_user_less_key_own_separate_listings(echo_rig: Gateway) -> None: + tool: Final = ScriptedTool( + "add", + lambda _: text_result("3"), + description=lambda headers: "Adds for " + headers.get("authorization", "nobody"), + ) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "legacy" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token") + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + first_token: Final = "first-" + uuid.uuid4().hex + second_token: Final = "second-" + uuid.uuid4().hex + first: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {first_token}"}) + second: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {second_token}"}) + assert _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED + peer.drain() + assert first.list_tools().ok + listings: Final = _listings(peer) + assert len(listings) == 1, listings + listed_with: Final = _HEADERS.validate_python(listings[0]["headers"]) + assert listed_with.get("authorization") == f"Bearer {first_token}", listed_with + warm: Final = eventually( + lambda: _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)), + lambda seen: seen != _UNLISTED, + ) + assert warm == f"Adds for Bearer {first_token}", warm + assert _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED + assert second.list_tools().ok + seen: Final = eventually( + lambda: ( + _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)), + _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)), + ), + lambda pair: _UNLISTED not in pair, + ) + assert seen == (f"Adds for Bearer {first_token}", f"Adds for Bearer {second_token}"), seen + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_lifecycle.py b/tests/integration/mcp/test_mcp_lifecycle.py index fa253f03520..eed3cbf658e 100644 --- a/tests/integration/mcp/test_mcp_lifecycle.py +++ b/tests/integration/mcp/test_mcp_lifecycle.py @@ -1,7 +1,9 @@ import functools import json import uuid +from collections.abc import Mapping from contextlib import ExitStack +from hashlib import sha256 from pathlib import Path from typing import Final @@ -14,15 +16,27 @@ from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests from integration._support.mcp import ( + JsonRpc, McpCaller, Outcome, + ScriptedTool, call_tool, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, tool_names, ) from integration._support.process import owned_proxy +from pydantic import TypeAdapter + +_SPEND_NONCES: Final = ( + "SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce" + ' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s' +) +_OBJECTS: Final = TypeAdapter(Mapping[str, object]) +_STRINGS: Final = TypeAdapter(Mapping[str, str]) @pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport") @@ -463,9 +477,7 @@ def _update_tool_permissions( assert updated.status_code == 200, updated.text -def _listing_on_both( - gateway: Gateway, peer: Gateway, key: str, expected: set[str] -) -> None: +def _listing_on_both(gateway: Gateway, peer: Gateway, key: str, expected: set[str]) -> None: for worker in (gateway, peer): listing: Final = eventually( functools.partial(_granted_view, worker, key), @@ -525,3 +537,56 @@ def test_key_update_tool_permission_widen_narrow_and_clear_apply_on_both_workers _listing_on_both(gateway, peer, key, all_tools) nulled: Final = _multiply_outcome_on_both(gateway, peer, key, alias) assert [call.text for call in nulled] == ["6", "6"], [call.raw for call in nulled] + + +def _nonce_echo(params: JsonRpc) -> JsonRpc: + return text_result(_STRINGS.validate_python(params["arguments"])["nonce"]) + + +def _call_params(call: Mapping[str, object]) -> Mapping[str, object]: + return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"]) + + +def _listed(caller: McpCaller, name: str) -> None: + listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45) + assert listing.error is None, (caller.gateway.client.base_url, listing.raw) + + +def test_tool_calls_on_both_workers_stay_base_compatible_after_each_worker_lists( + gateway: Gateway, peer: Gateway +) -> None: + schema: Final = {"type": "object", "properties": {"nonce": {"type": "string"}}} + tool: Final = ScriptedTool("echo", _nonce_echo, description="Echo the nonce back", input_schema=schema) + with scripted_peer(tool) as upstream, gateway.scenario() as scenario: + alias: Final = "compat" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-echo" + first: Final = McpCaller(gateway, key, "mcp", alias) + second: Final = McpCaller(peer, key, "mcp", alias) + _listed(first, name) + _listed(second, name) + upstream.drain() + nonces: Final = (uuid.uuid4().hex, uuid.uuid4().hex) + outcomes: Final = ( + first.call(name, {"nonce": nonces[0]}), + second.call(name, {"nonce": nonces[1]}), + first.call(name, {"nonce": nonces[0]}), + ) + assert [outcome.text for outcome in outcomes] == [nonces[0], nonces[1], nonces[0]], [o.raw for o in outcomes] + params: Final = [_call_params(call) for call in tool_calls(upstream.drain())] + assert [set(entry) - {"_meta"} for entry in params] == [{"name", "arguments"}] * 3, params + assert [entry["name"] for entry in params] == ["echo"] * 3, params + assert [entry["arguments"] for entry in params] == [ + {"nonce": nonces[0]}, + {"nonce": nonces[1]}, + {"nonce": nonces[0]}, + ], params + rows: Final = eventually( + lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")), + lambda found: len(found) >= 3, + seconds=70, + ) + assert sorted((str(row["status"]), str(row["nonce"])) for row in rows) == sorted( + ("success", nonce) for nonce in (nonces[0], nonces[1], nonces[0]) + ), rows diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py index 92cf56e9899..e892617d9ab 100644 --- a/tests/integration/mcp/test_mcp_listed_tool_metadata.py +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -7,16 +7,21 @@ what metadata the gateway attached to the hook """ import json +import threading import uuid -from collections.abc import Iterator, Mapping +from collections.abc import Callable, Generator, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack, contextmanager from pathlib import Path from typing import Final +import httpx import pytest import yaml -from integration._support.client import Gateway, gateway_from_environment +from integration._support.client import Gateway, eventually, gateway_from_environment from integration._support.mcp import ( EntryPoint, + JsonRpc, McpCaller, ScriptedTool, listed_tools, @@ -24,11 +29,24 @@ from integration._support.mcp import ( register_mcp, scripted_peer, text_result, + tool_calls, ) from integration._support.process import owned_proxy +from pydantic import TypeAdapter _ECHO: Final = "catalog-echo:" _PROBE: Final = "catalog-probe" +_CALLERS_PER_SERVER: Final = 256 +_PIN_SECONDS: Final = 120 +_COLD: Final[tuple[str, JsonRpc]] = ("", {"type": "object", "properties": {}, "additionalProperties": False}) +_PID: Final = TypeAdapter(int) +_HEADERS: Final = TypeAdapter(dict[bytes, bytes]) +_SCHEMA: Final = TypeAdapter(dict[str, object]) +_LOOKUP_SCHEMA: Final = { + "type": "object", + "properties": {"probe": {"type": "string", "description": "a probe marker"}}, + "additionalProperties": False, +} _GUARDRAIL_CODE: Final = ( "def apply_guardrail(inputs, request_data, input_type):\n" ' texts = list(inputs.get("texts") or [])\n' @@ -58,9 +76,15 @@ def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: }, } ] + config["general_settings"] = {**config["general_settings"], "proxy_config_reload_interval_seconds": 1} path: Final = directory / "config.yaml" path.write_text(yaml.safe_dump(config)) - with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path) as candidate: + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate, + ExitStack() as stack, + ): + _two_workers(stack, candidate) yield candidate @@ -93,6 +117,84 @@ def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Ma return _echoed(outcome.raw) +def _worker(gateway: Gateway) -> int: + response: Final = gateway.request("GET", "/debug/memory/summary") + assert response.status_code == 200, response.text + return _PID.validate_python(response.json()["worker_pid"]) + + +@contextmanager +def _pinned(rig: Gateway) -> Generator[Gateway, None, None]: + """A single keep-alive connection, so every request on it is served by the worker that accepted it.""" + limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1) + with httpx.Client(base_url=rig.client.base_url, timeout=15, trust_env=False, limits=limits) as client: + yield Gateway(client, rig.key, rig.upstream_url) + + +def _connection(stack: ExitStack, rig: Gateway, wanted: Callable[[int], bool]) -> tuple[Gateway, int]: + """A pinned connection to a worker ``wanted`` accepts; a worker still starting up accepts nothing yet.""" + + def attempt() -> tuple[Gateway, int] | None: + with ExitStack() as candidate: + gateway: Final = candidate.enter_context(_pinned(rig)) + pid: Final = _worker(gateway) + if not wanted(pid): + return None + stack.enter_context(candidate.pop_all()) + return gateway, pid + + found: Final = eventually(attempt, lambda pair: pair is not None, seconds=_PIN_SECONDS) + assert found is not None + return found + + +def _two_workers(stack: ExitStack, rig: Gateway) -> tuple[tuple[Gateway, int], tuple[Gateway, int]]: + """Connects while the first worker is busy answering, so the idle worker wins the accept race.""" + first: Final = _connection(stack, rig, lambda _: True) + stop: Final = threading.Event() + + def keep_busy() -> None: + while not stop.is_set(): + _worker(first[0]) + + with ThreadPoolExecutor(max_workers=1) as pool: + busy: Final = pool.submit(keep_busy) + try: + other: Final = _connection(stack, rig, lambda pid: pid != first[1]) + finally: + stop.set() + busy.result() + return first, other + + +def _served_name(gateway: Gateway, key: str, identity: str, tool: str) -> str: + """The prefixed name this worker lists for ``tool`` once its registry reload carries the server.""" + listing: Final = eventually( + lambda: McpCaller(gateway, key, "rest").list_tools(server_id=identity), + lambda value: value.ok and any(full.endswith(tool) for full in value.tools), + ) + return next(full for full in listing.tools if full.endswith(tool)) + + +def _settled_probe( + gateway: Gateway, key: str, name: str, identity: str +) -> tuple[str | None, Mapping[str, object] | None]: + """The hook echo for a direct call, once this worker's registry reload carries the server.""" + outcome: Final = eventually( + lambda: McpCaller(gateway, key, "rest").call(name, {"probe": _PROBE}, server_id=identity), + lambda value: value.error is not None and _ECHO in value.raw, + ) + return _echoed(outcome.raw) + + +def _forwarded_tenants(observed: tuple[dict[str, object], ...]) -> frozenset[bytes]: + return frozenset(_HEADERS.validate_python(item["headers"]).get(b"x-tenant", b"") for item in observed) - {b""} + + +def _lookup_tool(description: str | Callable[[Mapping[str, str]], str] = "Look up one record") -> ScriptedTool: + return ScriptedTool("lookup", lambda _: text_result("found"), description=description, input_schema=_LOOKUP_SCHEMA) + + @pytest.mark.parametrize("entry", ["rest", "mcp"]) def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed( rig: Gateway, entry: EntryPoint @@ -192,3 +294,200 @@ def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the assert seen == "Fetch one [MASKED] pet", ( "the guarded key must be evaluated against its own listing, not the opted-out key's later one" ) + + +def test_direct_call_without_a_listing_hands_the_hook_no_metadata_on_either_worker(rig: Gateway) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "cold" + uuid.uuid4().hex[:8]) + observer: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, observer, identity, "lookup") + seen: Final = tuple(_settled_probe(gateway, key, name, identity) for gateway in (first, other)) + assert seen == (_COLD, _COLD), (seen, first_pid, other_pid) + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == (), "the probe is blocked at the hook, before the upstream" + + +def test_warm_metadata_is_local_to_the_worker_that_served_the_listing(rig: Gateway) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "local" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, key, identity, "lookup") + warm: Final = ("Look up one record", _LOOKUP_SCHEMA) + assert _probe(McpCaller(first, key, "rest"), name, identity) == warm, first_pid + assert _settled_probe(other, key, name, identity) == _COLD, ( + "a worker that never served this caller a listing has no catalog for it", + other_pid, + ) + assert _served_name(other, key, identity, "lookup") == name + assert _probe(McpCaller(other, key, "rest"), name, identity) == warm, other_pid + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == () + + +def test_listed_tools_without_a_description_still_hand_the_hook_the_schema_the_listing_served(rig: Gateway) -> None: + open_schema: Final[JsonRpc] = {"type": "object", "properties": {}, "additionalProperties": True} + undescribed: Final = ScriptedTool("undescribed", lambda _: text_result("ok"), input_schema=open_schema) + blank: Final = ScriptedTool("blank", lambda _: text_result("ok"), description="", input_schema=open_schema) + with scripted_peer(undescribed, blank) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + identity: Final = register_mcp(scenario, peer, "bare" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(worker, key, identity) + names: Final = tuple(next(full for full in served if full.endswith(tool)) for tool in ("undescribed", "blank")) + assert tuple((served[name].get("description") or "", served[name]["inputSchema"]) for name in names) == ( + ("", open_schema), + ("", open_schema), + ), served + seen: Final = tuple(_probe(McpCaller(worker, key, "rest"), name, identity) for name in names) + assert seen == (("", open_schema), ("", open_schema)), (seen, _COLD) + assert _worker(worker) == pid + assert tool_calls(peer.drain()) == () + + +def test_hook_receives_the_nested_schema_with_the_leaves_the_listing_masked(rig: Gateway) -> None: + schema: Final = { + "type": "object", + "required": ["filter"], + "additionalProperties": False, + "properties": { + "filter": { + "type": "object", + "properties": { + "path": {"type": "string", "description": "SECRET path"}, + "tags": {"type": "array", "items": {"type": "string", "description": "one SECRET tag"}}, + }, + } + }, + } + masked: Final = _SCHEMA.validate_python(json.loads(json.dumps(schema).replace("SECRET", "[MASKED]"))) + tool: Final = ScriptedTool( + "search", lambda _: text_result("hit"), description="Search SECRET records", input_schema=schema + ) + with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + identity: Final = register_mcp(scenario, peer, "nested" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + served: Final = listed_tools(worker, key, identity) + name: Final = next(full for full in served if full.endswith("search")) + assert (served[name]["description"], served[name]["inputSchema"]) == ("Search [MASKED] records", masked), served + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Search [MASKED] records", masked), ( + "the hook must be handed the nested schema exactly as the listing served it" + ) + assert _worker(worker) == pid + assert tool_calls(peer.drain()) == () + + +def test_a_server_definition_update_drops_the_catalog_on_every_worker_until_the_caller_lists_again( + rig: Gateway, +) -> None: + with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack: + (first, first_pid), (other, other_pid) = _two_workers(stack, rig) + with first.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "upd" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = _served_name(first, key, identity, "lookup") + assert _served_name(other, key, identity, "lookup") == name + before: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other)) + assert before == (("Look up one record", _LOOKUP_SCHEMA),) * 2, before + updated: Final = first.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}}, + ) + assert updated.status_code == 202, updated.text + assert _probe(McpCaller(first, key, "rest"), name, identity) == _COLD, ( + "the worker that applied the update must drop its catalog at once", + first_pid, + ) + assert ( + eventually(lambda: _probe(McpCaller(other, key, "rest"), name, identity), lambda seen: seen == _COLD) + == _COLD + ), other_pid + relisted: Final = tuple( + listed_tools(gateway, key, identity)[name]["description"] for gateway in (first, other) + ) + assert relisted == ("Audited lookup", "Audited lookup"), relisted + after: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other)) + assert after == (("Audited lookup", _LOOKUP_SCHEMA),) * 2, after + assert (_worker(first), _worker(other)) == (first_pid, other_pid) + assert tool_calls(peer.drain()) == () + + +def test_a_listing_in_flight_across_a_server_update_does_not_resurrect_the_old_catalog(rig: Gateway) -> None: + started: Final = threading.Event() + release: Final = threading.Event() + + def describe(_: Mapping[str, str]) -> str: + started.set() + assert release.wait(20), "the listing was never released" + return "Look up one record" + + with ( + scripted_peer(_lookup_tool(describe)) as peer, + ExitStack() as stack, + ThreadPoolExecutor(max_workers=1) as pool, + ): + worker, pid = _connection(stack, rig, lambda _: True) + sibling, _ = _connection(stack, rig, lambda candidate: candidate == pid) + with worker.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "stale" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + release.set() + name: Final = _served_name(worker, key, identity, "lookup") + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Look up one record", _LOOKUP_SCHEMA) + started.clear() + release.clear() + pending: Final = pool.submit(McpCaller(worker, key, "rest").list_tools, identity) + assert started.wait(10), "the upstream never saw the in-flight listing" + updated: Final = sibling.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}}, + ) + assert updated.status_code == 202, updated.text + release.set() + stale: Final = pending.result(timeout=20) + assert stale.ok, stale.raw + assert _probe(McpCaller(worker, key, "rest"), name, identity) == _COLD, ( + "a listing fetched before the update must not be recorded after it" + ) + assert listed_tools(worker, key, identity)[name]["description"] == "Audited lookup" + assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Audited lookup", _LOOKUP_SCHEMA) + assert (_worker(worker), _worker(sibling)) == (pid, pid) + assert tool_calls(peer.drain()) == () + + +def test_a_server_keeps_the_newest_256_caller_catalogs_and_evicts_the_oldest(rig: Gateway) -> None: + tool: Final = _lookup_tool(lambda headers: f"Lookup for tenant {headers.get('x-tenant', 'nobody')}") + with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario: + pid: Final = _worker(worker) + alias: Final = "cap" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + tenants: Final = tuple(f"t{index}" for index in range(_CALLERS_PER_SERVER + 1)) + callers: Final = { + tenant: McpCaller(worker, key, "server_mcp", alias, headers={"x-tenant": tenant}) for tenant in tenants + } + listings: Final = tuple(callers[tenant].list_tools() for tenant in tenants) + assert all(listing.ok for listing in listings), [listing.raw for listing in listings if not listing.ok] + name: Final = next(full for full in listings[0].tools if full.endswith("lookup")) + + def seen(tenant: str) -> str | None: + return _probe(callers[tenant], name, identity)[0] + + assert (seen("t0"), seen("t1"), seen(tenants[-1])) == ( + "", + "Lookup for tenant t1", + f"Lookup for tenant {tenants[-1]}", + ), "the oldest of 257 callers is evicted, the newest 256 keep their own catalog" + assert callers["t0"].list_tools().ok + assert (seen("t0"), seen("t1")) == ("Lookup for tenant t0", ""), "relisting makes t0 newest and evicts t1" + assert _worker(worker) == pid + observed: Final = peer.drain() + assert _forwarded_tenants(observed) == frozenset(tenant.encode() for tenant in tenants), len(observed) + assert tool_calls(observed) == () diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 01bf006a03e..af3bff1fd8b 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -1,16 +1,35 @@ +import asyncio import json import uuid -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Generator, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass -from typing import Final, Literal +from hashlib import sha256 +from pathlib import Path +from typing import Final, Literal, TypeVar import httpx import pytest -from integration._support.client import Gateway, Scenario -from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls +import yaml +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.mcp import ( + McpPeer, + ScriptedTool, + mcp_peer, + register_mcp, + scripted_peer, + text_result, + tool_calls, +) from integration._support.mcp_grants import create_toolset +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionMessageParam +from openai.types.responses.tool_param import Mcp +from pydantic import JsonValue, TypeAdapter Surface = Literal["chat", "responses", "messages", "messages_bridge"] SURFACES: Final[tuple[Surface, ...]] = ("chat", "responses", "messages", "messages_bridge") @@ -18,6 +37,9 @@ ADD: Final = {"a": 2, "b": 3} ANSWER: Final = "the sum is 5" GATEWAY_REF: Final = {"type": "mcp", "server_url": "litellm_proxy", "server_label": "litellm"} AUTO: Final = {**GATEWAY_REF, "require_approval": "never"} +OUTAGE: Final = "bridge-outage" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) def _json(body: Mapping[str, object]) -> Reply: @@ -44,25 +66,94 @@ def _has_tool_result(body: Mapping[str, object]) -> bool: return False -def _model_double(tool: str) -> Callable[[Request], Reply]: - arguments: Final = json.dumps(ADD) +@dataclass(frozen=True, slots=True) +class Turn: + tool: str + arguments: str + answer: str + +def _fixed_turn(tool: str) -> Callable[[Mapping[str, JsonValue]], Turn]: + return lambda _: Turn(tool, json.dumps(ADD), ANSWER) + + +def _echoing_turn(body: Mapping[str, JsonValue]) -> Turn: + names: Final = _tool_names(body) + return Turn(names[0] if names else "", json.dumps({"query": _prompt(body)}), _tool_result_text(body) or "") + + +def _prompt(body: Mapping[str, JsonValue]) -> str: + inputs: Final = body.get("input") + if isinstance(inputs, str): + return inputs + items: Final = inputs if isinstance(inputs, list) else body.get("messages") + first: Final = items[0] if isinstance(items, list) and items else None + return str(first["content"]) if isinstance(first, dict) and isinstance(first.get("content"), str) else "" + + +def _tool_result_text(body: Mapping[str, JsonValue]) -> str | None: + inputs: Final = body.get("input") + messages: Final = body.get("messages") + items: Final = inputs if isinstance(inputs, list) else messages if isinstance(messages, list) else () + for item in items: + if not isinstance(item, dict): + continue + if item.get("type") == "function_call_output": + return str(item["output"]) + if item.get("role") == "tool": + return str(item["content"]) + if (block := _tool_result_block(item)) is not None: + return block + return None + + +def _tool_result_block(item: Mapping[str, JsonValue]) -> str | None: + content: Final = item.get("content") + for block in content if isinstance(content, list) else (): + if isinstance(block, dict) and block.get("type") == "tool_result": + return str(block["content"]) + return None + + +def _responses_stream(response: Mapping[str, JsonValue], item: Mapping[str, JsonValue]) -> Reply: + events: Final = ( + {"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}}, + {"type": "response.in_progress", "sequence_number": 1, "response": {**response, "status": "in_progress"}}, + {"type": "response.output_item.added", "sequence_number": 2, "output_index": 0, "item": item}, + {"type": "response.output_item.done", "sequence_number": 3, "output_index": 0, "item": item}, + {"type": "response.completed", "sequence_number": 4, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _model_double(plan: Callable[[Mapping[str, JsonValue]], Turn]) -> Callable[[Request], Reply]: def respond(request: Request) -> Reply: if request.method == "GET" and request.target.endswith("/models"): return _json({"object": "list", "data": []}) - body: Final = json.loads(request.body) - assert isinstance(body, dict), request.body + body: Final = JSON_OBJECT.validate_json(request.body) + if OUTAGE in _prompt(body): + outage: Final = {"error": {"message": "scripted provider outage", "type": "server_error", "code": None}} + return Reply(status=500, body=json.dumps(outage).encode()) + turn: Final = plan(body) done: Final = _has_tool_result(body) + identity: Final = uuid.uuid4().hex[:12] usage: Final = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} if request.target.endswith("/chat/completions"): message: Final = ( - {"role": "assistant", "content": ANSWER} + {"role": "assistant", "content": turn.answer} if done else { "role": "assistant", "content": None, "tool_calls": [ - {"id": "call_1", "type": "function", "function": {"name": tool, "arguments": arguments}} + { + "id": "call_1", + "type": "function", + "function": {"name": turn.tool, "arguments": turn.arguments}, + } ], } ) @@ -74,7 +165,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: else message ) chunk: Final = { - "id": "chatcmpl-1", + "id": f"chatcmpl-{identity}", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini", @@ -89,7 +180,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: ) return _json( { - "id": "chatcmpl-1", + "id": f"chatcmpl-{identity}", "object": "chat.completion", "created": 1, "model": "gpt-4o-mini", @@ -99,13 +190,13 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: ) if request.target.endswith("/messages"): content: Final = ( - [{"type": "text", "text": ANSWER}] + [{"type": "text", "text": turn.answer}] if done - else [{"type": "tool_use", "id": "toolu_1", "name": tool, "input": ADD}] + else [{"type": "tool_use", "id": "toolu_1", "name": turn.tool, "input": json.loads(turn.arguments)}] ) return _json( { - "id": "msg_1", + "id": f"msg_{identity}", "type": "message", "role": "assistant", "model": "claude", @@ -116,39 +207,34 @@ def _model_double(tool: str) -> Callable[[Request], Reply]: } ) assert request.target.endswith("/responses"), request.target - output: Final = ( - [ - { - "type": "message", - "id": "msg_1", - "role": "assistant", - "status": "completed", - "content": [{"type": "output_text", "text": ANSWER, "annotations": []}], - } - ] - if done - else [ - { - "type": "function_call", - "id": "fc_1", - "call_id": "call_1", - "name": tool, - "arguments": arguments, - "status": "completed", - } - ] - ) - return _json( + item: Final[dict[str, JsonValue]] = ( { - "id": "resp_1", - "object": "response", - "created_at": 1, + "type": "message", + "id": f"msg_{identity}", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": turn.answer, "annotations": []}], + } + if done + else { + "type": "function_call", + "id": f"fc_{identity}", + "call_id": f"call_{identity}", + "name": turn.tool, + "arguments": turn.arguments, "status": "completed", - "model": "gpt-4o-mini", - "output": output, - "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, } ) + response: Final[dict[str, JsonValue]] = { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [item], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + } + return _responses_stream(response, item) if body.get("stream") is True else _json(response) return respond @@ -240,7 +326,7 @@ def _rig(gateway: Gateway, surface: Surface) -> Iterator[Rig]: alias: Final = "llm" + uuid.uuid4().hex[:8] with ( mcp_peer() as peer, - wire_server(_model_double(f"{alias}-add")) as wire, + wire_server(_model_double(_fixed_turn(f"{alias}-add"))) as wire, gateway.scenario() as scenario, ): server_id: Final = register_mcp(scenario, peer, alias) @@ -403,3 +489,460 @@ def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer" assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools() assert response.status_code in (200, 400, 401, 403), response.text + + +Bridge = Literal["chat", "responses", "messages"] +BRIDGES: Final[tuple[Bridge, ...]] = ("chat", "responses") +Client = Literal["sync", "async"] +HOOK_ECHO: Final = "bridge-echo:" +HOOK_PROBE: Final = "bridge-probe" +HOOK_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + " for text in texts:\n" + f' if "{HOOK_PROBE}" in text:\n' + f' return block("{HOOK_ECHO}" + json_stringify(' + '{"description": function.get("description"), "parameters": function.get("parameters")}))\n' + " return allow()\n" +) +LOOKUP: Final[tuple[str, dict[str, JsonValue]]] = ( + "Look up one record", + {"type": "object", "properties": {"query": {"type": "string"}}}, +) +REPORT: Final[tuple[str, dict[str, JsonValue]]] = ( + "Write one report", + {"type": "object", "properties": {"query": {"type": "string"}, "format": {}}}, +) +COLD: Final[tuple[str, dict[str, JsonValue]]] = ( + "", + {"type": "object", "properties": {}, "additionalProperties": False}, +) +RELOAD_FAST: Final = {"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "5"} +CALL_ID: Final = "x-litellm-call-id" +T = TypeVar("T") +Definition = tuple[str, str, JsonValue] +Echo = tuple[JsonValue, JsonValue] + + +def _served(listing: tuple[str, Mapping[str, JsonValue]]) -> Echo: + return listing[0], {**listing[1], "additionalProperties": False} + + +@dataclass(frozen=True, slots=True) +class Hooked: + proxy: Gateway + sink: Wire + + +@pytest.fixture(scope="module") +def hooked(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Hooked]: + directory: Final = tmp_path_factory.mktemp("bridge-hooks") + base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + with wire_server(lambda _: _json({"flagged": False, "session_id": "scripted"})) as sink: + echo: Final = {"guardrail": "custom_code", "mode": "pre_mcp_call", "default_on": True, "custom_code": HOOK_CODE} + pillar: Final = { + "guardrail": "pillar", + "mode": ["pre_call", "pre_mcp_call"], + "default_on": True, + "api_key": "sk-pillar-" + uuid.uuid4().hex, + "api_base": sink.url, + "on_flagged_action": "monitor", + } + guardrails: Final = [ + {"guardrail_name": "bridge-echo-" + uuid.uuid4().hex[:8], "litellm_params": echo}, + {"guardrail_name": "bridge-sink-" + uuid.uuid4().hex[:8], "litellm_params": pillar}, + ] + path: Final = directory / "config.yaml" + path.write_text(yaml.safe_dump({**base, "guardrails": guardrails})) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, directory, RELOAD_FAST, config=path, workers=2) as proxy, + ): + yield Hooked(proxy, sink) + + +@dataclass(frozen=True, slots=True) +class BridgeRig: + hooked: Hooked + scenario: Scenario + peer: McpPeer + wire: Wire + alias: str + server_id: str + model: str + bridge: Bridge + + def tool(self, name: str) -> str: + return f"{self.alias}-{name}" + + def names(self) -> frozenset[str]: + return frozenset(("lookup", "report")) + + def mcp(self, name: str) -> Mcp: + return {**AUTO_MCP, "allowed_tools": [self.tool(name)]} + + def url(self, path: str) -> str: + return str(self.hooked.proxy.client.base_url).rstrip("/") + path + + def post(self, key: str, prompt: str, tools: Sequence[Mcp], **extra: object) -> httpx.Response: + headers: Final = {"Authorization": f"Bearer {key}"} + if self.bridge == "chat": + body: Final = {"model": self.model, "messages": [{"role": "user", "content": prompt}], "tools": list(tools)} + return httpx.post(self.url("/v1/chat/completions"), headers=headers, json={**body, **extra}, timeout=90) + if self.bridge == "responses": + body_r: Final = {"model": self.model, "input": prompt, "tools": list(tools), **extra} + return httpx.post(self.url("/v1/responses"), headers=headers, json=body_r, timeout=90) + body_m: Final = { + "model": self.model, + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt}], + "tools": list(tools), + **extra, + } + return httpx.post(self.url("/v1/messages"), headers=headers, json=body_m, timeout=90) + + def upstream_by_prompt(self) -> Mapping[str, tuple[tuple[Definition, ...], ...]]: + bodies: Final = tuple( + JSON_OBJECT.validate_json(request.body) for request in self.wire.drain() if request.method == "POST" + ) + prompts: Final = frozenset(_prompt(body) for body in bodies) + return {prompt: tuple(_definitions(body) for body in bodies if _prompt(body) == prompt) for prompt in prompts} + + def peer_calls(self) -> tuple[tuple[str, JsonValue], ...]: + return tuple(_peer_call(call) for call in tool_calls(self.peer.drain())) + + def hook_messages(self, marker: str) -> tuple[str, ...]: + posted: Final = tuple( + JSON_OBJECT.validate_json(request.body) for request in self.hooked.sink.drain() if request.method == "POST" + ) + contents: Final = tuple(content for payload in posted for content in _contents(payload)) + return tuple(content for content in contents if _synthetic(content, marker)) + + +AUTO_MCP: Final[Mcp] = { + "type": "mcp", + "server_label": "litellm", + "server_url": "litellm_proxy", + "require_approval": "never", +} + + +def _peer_call(call: Mapping[str, object]) -> tuple[str, JsonValue]: + params: Final = object_value(object_value(JSON_VALUE.validate_python(call["body"]))["params"]) + return str(params["name"]), params["arguments"] + + +def _contents(payload: Mapping[str, JsonValue]) -> Iterator[str]: + messages: Final = payload.get("messages") + for message in messages if isinstance(messages, list) else (): + if isinstance(message, dict): + yield str(message.get("content")) + + +def _function(tool: JsonValue) -> dict[str, JsonValue] | None: + if not isinstance(tool, dict): + return None + function: Final = tool.get("function", tool) + return function if isinstance(function, dict) else None + + +def _definitions(body: Mapping[str, JsonValue]) -> tuple[Definition, ...]: + tools: Final = body.get("tools") + functions: Final = tuple(_function(tool) for tool in tools) if isinstance(tools, list) else () + return tuple( + (str(function["name"]), str(function.get("description", "")), function.get("parameters")) + for function in functions + if function is not None + ) + + +def _uniform(upstream: Mapping[str, tuple[tuple[Definition, ...], ...]]) -> Mapping[str, tuple[Definition, ...] | None]: + return { + prompt: rounds[0] if all(definitions == rounds[0] for definitions in rounds) else None + for prompt, rounds in upstream.items() + } + + +def _strings(value: JsonValue) -> Iterator[str]: + if isinstance(value, str): + yield value + return + children: Final = value.values() if isinstance(value, dict) else value if isinstance(value, list) else () + for child in children: + yield from _strings(child) + + +def _echoed(value: JsonValue) -> Echo: + carrier: Final = next((text for text in _strings(value) if HOOK_ECHO in text), None) + assert carrier is not None, value + payload: Final = carrier.split(HOOK_ECHO, 1)[1] + end: Final = json.JSONDecoder().raw_decode(payload)[1] + echoed: Final = JSON_OBJECT.validate_json(payload[:end]) + return echoed.get("description"), echoed.get("parameters") + + +@contextmanager +def _bridge_rig(hooked: Hooked, bridge: Bridge) -> Generator[BridgeRig, None, None]: + alias: Final = "brg" + uuid.uuid4().hex[:8] + lookup: Final = ScriptedTool( + "lookup", + lambda params: text_result("found:" + json.dumps(JSON_OBJECT.validate_python(params)["arguments"])), + description=LOOKUP[0], + input_schema=LOOKUP[1], + ) + report: Final = ScriptedTool( + "report", lambda _: text_result("reported"), description=REPORT[0], input_schema=REPORT[1] + ) + with ( + scripted_peer(lookup, report) as peer, + wire_server(_model_double(_echoing_turn)) as wire, + hooked.proxy.scenario() as scenario, + ): + server_id: Final = register_mcp(scenario, peer, alias) + model: Final = scenario.model(model=_upstream_model(bridge), api_base=wire.url + "/v1") + rig: Final = BridgeRig(hooked, scenario, peer, wire, alias, server_id, model, bridge) + eventually( + lambda: tuple(_on_worker(hooked.proxy, lambda client: _master_listing(rig, client)) for _ in range(6)), + lambda seen: len({pid for pid, _ in seen}) >= 2 and all(names == rig.names() for _, names in seen), + seconds=45, + ) + peer.drain() + hooked.sink.drain() + yield rig + + +def _bridge_key(rig: BridgeRig) -> str: + return rig.scenario.key(object_permission={"mcp_servers": [rig.server_id]}) + + +@dataclass(frozen=True, slots=True) +class Seen: + response_id: str + call_id: str + text: str + + +def _ask_sync(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen: + sdk: Final = OpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90) + if rig.bridge == "responses": + if stream: + raw_events: Final = sdk.responses.with_raw_response.create( + model=rig.model, input=prompt, tools=tools, stream=True + ) + completed: Final = next( + event.response for event in raw_events.parse() if event.type == "response.completed" + ) + return Seen(completed.id, raw_events.headers[CALL_ID], completed.output_text) + raw_response: Final = sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools) + response: Final = raw_response.parse() + return Seen(response.id, raw_response.headers[CALL_ID], response.output_text) + messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}] + extra: Final = {"tools": list(tools)} + if stream: + raw_chunks: Final = sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, stream=True, extra_body=extra + ) + parts: Final = tuple( + (chunk.id, chunk.choices[0].delta.content or "") for chunk in raw_chunks.parse() if chunk.choices + ) + return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts)) + raw_completion: Final = sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, extra_body=extra + ) + completion: Final = raw_completion.parse() + return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "") + + +async def _ask_async(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen: + sdk: Final = AsyncOpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90) + if rig.bridge == "responses": + if stream: + raw_events: Final = await sdk.responses.with_raw_response.create( + model=rig.model, input=prompt, tools=tools, stream=True + ) + completed: Final = [ + event.response async for event in raw_events.parse() if event.type == "response.completed" + ] + return Seen(completed[0].id, raw_events.headers[CALL_ID], completed[0].output_text) + raw_response: Final = await sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools) + response: Final = raw_response.parse() + return Seen(response.id, raw_response.headers[CALL_ID], response.output_text) + messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}] + extra: Final = {"tools": list(tools)} + if stream: + raw_chunks: Final = await sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, stream=True, extra_body=extra + ) + parts: Final = [ + (chunk.id, chunk.choices[0].delta.content or "") async for chunk in raw_chunks.parse() if chunk.choices + ] + return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts)) + raw_completion: Final = await sdk.chat.completions.with_raw_response.create( + model=rig.model, messages=messages, extra_body=extra + ) + completion: Final = raw_completion.parse() + return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "") + + +def _ask(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool, client: Client) -> Seen: + if client == "async": + return asyncio.run(_ask_async(rig, key, prompt, tools, stream)) + return _ask_sync(rig, key, prompt, tools, stream) + + +def _synthetic(content: str, marker: str) -> bool: + return content.startswith("Tool: lookup\n") and marker in content + + +def _on_worker(gateway: Gateway, act: Callable[[httpx.Client], T]) -> tuple[int, T]: + with httpx.Client(base_url=str(gateway.client.base_url), timeout=30) as client: + summary: Final = client.get("/debug/memory/summary", headers={"Authorization": f"Bearer {gateway.key}"}) + assert summary.status_code == 200, summary.text + pid: Final = JSON_OBJECT.validate_json(summary.content)["worker_pid"] + assert isinstance(pid, int), summary.text + return pid, act(client) + + +def _both_workers(gateway: Gateway) -> frozenset[int]: + return eventually( + lambda: frozenset(_on_worker(gateway, lambda _: None)[0] for _ in range(6)), lambda pids: len(pids) >= 2 + ) + + +def _master_listing(rig: BridgeRig, client: httpx.Client) -> frozenset[str]: + headers: Final = {"x-litellm-api-key": rig.hooked.proxy.key} + response: Final = client.get("/mcp-rest/tools/list", headers=headers, params={"server_id": rig.server_id}) + tools: Final = JSON_OBJECT.validate_json(response.content).get("tools") if response.status_code == 200 else None + return frozenset(str(object_value(tool)["name"]) for tool in tools) if isinstance(tools, list) else frozenset() + + +def _direct_probe(rig: BridgeRig, key: str, name: str, client: httpx.Client) -> Echo: + body: Final = {"server_id": rig.server_id, "name": rig.tool(name), "arguments": {"query": HOOK_PROBE}} + response: Final = client.post("/mcp-rest/tools/call", headers={"x-litellm-api-key": key}, json=body) + return _echoed(JSON_VALUE.validate_json(response.content)) + + +def _spend_row(key: str, call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT api_key, call_type, status, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,) + ), + lambda found: len(found) >= 1, + seconds=70, + ) + assert rows[0]["api_key"] == sha256(key.encode()).hexdigest(), rows + return rows[0] + + +@pytest.mark.parametrize("client", ("sync", "async")) +@pytest.mark.parametrize("stream", (False, True), ids=("plain", "stream")) +@pytest.mark.parametrize("bridge", BRIDGES) +def test_bridge_hook_sees_the_definition_of_the_tool_filtered_for_that_request( + hooked: Hooked, bridge: Bridge, stream: bool, client: Client +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + probe: Final = f"{marker} {HOOK_PROBE}" + found: Final = _ask(rig, key, marker, [rig.mcp("lookup")], stream, client) + assert found.text == "found:" + json.dumps({"query": marker}), found + blocked: Final = _ask(rig, key, probe, [rig.mcp("lookup")], stream, client) + assert _echoed(blocked.text) == _served(LOOKUP), blocked + expected: Final = ((rig.tool("lookup"), *_served(LOOKUP)),) + upstream: Final = rig.upstream_by_prompt() + assert set(upstream) == {marker, probe} and all( + definitions == expected for definitions in upstream[marker] + upstream[probe] + ), upstream + assert rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) + assert _spend_row(key, found.call_id)["status"] == "success" + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_direct_call_after_bridge_only_discovery_stays_cold(hooked: Hooked, bridge: Bridge) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + workers: Final = _both_workers(rig.hooked.proxy) + bridged: Final = _ask(rig, key, HOOK_PROBE, [rig.mcp("lookup")], False, "sync") + assert _echoed(bridged.text) == _served(LOOKUP), bridged + direct: Final = eventually( + lambda: tuple( + _on_worker(rig.hooked.proxy, lambda client: _direct_probe(rig, key, "lookup", client)) for _ in range(6) + ), + lambda seen: frozenset(pid for pid, _ in seen) == workers, + ) + assert all(echo == COLD for _, echo in direct) and frozenset(pid for pid, _ in direct) == workers, ( + direct, + workers, + ) + assert rig.peer_calls() == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_concurrent_requests_of_one_key_with_different_allowed_tools_each_see_their_own_definition( + hooked: Hooked, bridge: Bridge +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + prompts: Final = {name: f"{name} {uuid.uuid4().hex} {HOOK_PROBE}" for name in ("lookup", "report")} + + def ask(name: str) -> Seen: + return _ask(rig, key, prompts[name], [rig.mcp(name)], False, "sync") + + with ThreadPoolExecutor(2) as pool: + lookup, report = pool.map(ask, ("lookup", "report")) + assert (_echoed(lookup.text), _echoed(report.text)) == (_served(LOOKUP), _served(REPORT)), (lookup, report) + upstream: Final = rig.upstream_by_prompt() + assert _uniform(upstream) == { + prompts["lookup"]: ((rig.tool("lookup"), *_served(LOOKUP)),), + prompts["report"]: ((rig.tool("report"), *_served(REPORT)),), + }, upstream + assert rig.peer_calls() == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_provider_outage_reaches_the_caller_and_never_the_peer_or_the_hooks(hooked: Hooked, bridge: Bridge) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + prompt: Final = f"{OUTAGE} {uuid.uuid4().hex}" + response: Final = rig.post(key, prompt, [rig.mcp("lookup")]) + assert response.status_code == 500, response.text + upstream: Final = rig.upstream_by_prompt() + assert set(upstream) == {prompt} and all( + definitions == ((rig.tool("lookup"), *_served(LOOKUP)),) for definitions in upstream[prompt] + ), upstream + assert rig.peer_calls() == () + assert rig.hook_messages(prompt) == () + + +@pytest.mark.parametrize("bridge", BRIDGES) +def test_identical_nonstream_repeat_is_a_cache_hit_without_new_model_peer_or_hook_traffic( + hooked: Hooked, bridge: Bridge +) -> None: + with _bridge_rig(hooked, bridge) as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + first: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync") + assert first.text == "found:" + json.dumps({"query": marker}), first + assert set(rig.upstream_by_prompt()) == {marker} and rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) + assert _spend_row(key, first.call_id)["cache_hit"] != "True" + repeat: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync") + assert repeat.text == first.text, (first, repeat) + assert (rig.upstream_by_prompt(), rig.peer_calls(), rig.hook_messages(marker)) == ({}, (), ()), repeat + assert _spend_row(key, repeat.call_id)["cache_hit"] == "True" + + +def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadata(hooked: Hooked) -> None: + with _bridge_rig(hooked, "messages") as rig: + key: Final = _bridge_key(rig) + marker: Final = "m" + uuid.uuid4().hex + probe: Final = f"{marker} {HOOK_PROBE}" + found: Final = rig.post(key, marker, [rig.mcp("lookup")]) + assert found.status_code == 200, found.text + blocked: Final = rig.post(key, probe, [rig.mcp("lookup")]) + assert blocked.status_code == 200, blocked.text + assert _echoed(JSON_VALUE.validate_json(blocked.content)) == COLD, blocked.text + assert rig.peer_calls() == (("lookup", {"query": marker}),) + assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 60563e7aacd..e95cdb902e8 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,16 +1,20 @@ import base64 import hashlib +import re import secrets +import textwrap import time import uuid +from collections.abc import Iterator from dataclasses import dataclass +from pathlib import Path from typing import Final from urllib.parse import parse_qs, urlsplit import httpx import jwt import pytest -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, eventually, gateway_from_environment from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, @@ -27,6 +31,7 @@ from integration._support.mcp import ( ) from integration._support.mcp_grants import create_toolset from integration._support.oauth_server import AuthorizationServer, oauth_server +from integration._support.process import owned_proxy ADD: Final = {"a": 2, "b": 3} CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb" @@ -503,3 +508,70 @@ def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_a refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {}) assert refused.status == 403, refused.raw assert tool_calls(peer.drain()) == () + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +def test_token_exchange_callers_with_different_subject_tokens_own_separate_listings(echo_rig: Gateway) -> None: + with mcp_peer() as peer, oauth_server() as auth, echo_rig.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=auth.issuer + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + first_subject: Final = "subject-" + uuid.uuid4().hex + first: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": f"Bearer {first_subject}"}) + second: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": "Bearer subject-" + uuid.uuid4().hex}) + auth.drain() + assert first.list_tools().ok + assert [request["subject_token"] for request in auth.token_requests()] == [first_subject] + probe: Final = {"probe": _PROBE} + own: Final = _echoed_description(first.call(f"{alias}-add", probe)) + other: Final = _echoed_description(second.call(f"{alias}-add", probe)) + assert (own, other) == ("Add two integers", _UNLISTED), ( + "the caller bearer is part of the identity on a token-exchange server: one subject, one slot" + ) + assert second.list_tools().ok + assert _echoed_description(second.call(f"{alias}-add", probe)) == "Add two integers" + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" diff --git a/tests/integration/mcp/test_mcp_resilience.py b/tests/integration/mcp/test_mcp_resilience.py index 8efb54a18fd..d821181962f 100644 --- a/tests/integration/mcp/test_mcp_resilience.py +++ b/tests/integration/mcp/test_mcp_resilience.py @@ -1,13 +1,23 @@ +import itertools import uuid +from collections.abc import Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path from typing import Final import pytest +import yaml from integration._support.client import Gateway, eventually +from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, EntryPoint, + JsonRpc, McpCaller, Outcome, + ScriptedTool, disconnecting_tool, echo_tool, listed_tools, @@ -15,8 +25,67 @@ from integration._support.mcp import ( register_mcp, scripted_peer, slow_tool, + text_result, tool_calls, ) +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply +from pydantic import BaseModel, TypeAdapter + +_ECHO: Final = "catalog-echo:" +_PROBE: Final = "catalog-probe" +_DESCRIPTION: Final = "Look up one record" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' texts = list(inputs.get("texts") or [])\n' + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' if "{_PROBE}" in texts:\n' + f' return block("{_ECHO}" + json_stringify({{"description": function.get("description")}}))\n' + " return allow()\n" +) +_BURST: Final = 20 +_OUTAGE: Final = 6 +_SPEND_NONCES: Final = ( + "SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce" + ' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s' +) +_OBJECTS: Final = TypeAdapter(Mapping[str, object]) +_STRINGS: Final = TypeAdapter(Mapping[str, str]) +_ECHOED: Final = TypeAdapter(Mapping[str, str | None]) + + +class _Content(BaseModel): + text: str + + +class _Result(BaseModel): + content: tuple[_Content, ...] + + +class _RpcReply(BaseModel): + id: int + result: _Result + + +class _SessionsReport(BaseModel): + worker_pid: int + + +@pytest.fixture(scope="module") +def echo_config(tmp_path_factory: pytest.TempPathFactory) -> Path: + base: Final = _OBJECTS.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + guardrail: Final = { + "guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_mcp_call", + "default_on": True, + "custom_code": _GUARDRAIL_CODE, + }, + } + path: Final = tmp_path_factory.mktemp("failure-recovery") / "config.yaml" + path.write_text(yaml.safe_dump({**base, "guardrails": [guardrail]})) + return path def _call(caller: McpCaller, name: str, arguments: dict[str, object], entry: EntryPoint, identity: str) -> Outcome: @@ -134,3 +203,151 @@ def test_peer_restart_on_the_same_url_is_picked_up_without_gateway_restart(gatew ) assert back.text == '{"a": 1, "b": 1}', back.raw assert len(tool_calls(replacement.drain())) >= 1 + + +def _rpc_reply(raw: str) -> _RpcReply: + data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:")) + return _RpcReply.model_validate_json(data[-1] if data else raw) + + +def _call_params(call: Mapping[str, object]) -> Mapping[str, object]: + return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"]) + + +def _call_nonce(call: Mapping[str, object]) -> str: + return _STRINGS.validate_python(_call_params(call)["arguments"])["nonce"] + + +def _listed(caller: McpCaller, name: str) -> None: + listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45) + assert listing.error is None, (caller.gateway.client.base_url, listing.raw) + + +@dataclass(frozen=True, slots=True) +class _Worker: + caller: McpCaller + pid: int + + +def _worker(proxy: Gateway, key: str, alias: str) -> _Worker: + sessions: Final = proxy.client.get("/v1/mcp/sessions", headers={"x-litellm-api-key": proxy.key}) + assert sessions.status_code == 200, sessions.text + return _Worker(McpCaller(proxy, key, "mcp", alias), _SessionsReport.model_validate_json(sessions.text).worker_pid) + + +def _served(worker: _Worker, name: str, nonce: str) -> None: + served: Final = worker.caller.call(name, {"nonce": nonce}) + assert served.text == "found", (worker.pid, served.raw) + + +def _probed_description(worker: _Worker, name: str) -> str | None: + """The description the pre_mcp_call guardrail on that worker was handed, recovered from its block reason.""" + blocked: Final = worker.caller.call(name, {"nonce": _PROBE}) + assert blocked.error is not None, (worker.pid, blocked.raw) + carrier: Final = next((item.text for item in _rpc_reply(blocked.raw).result.content if _ECHO in item.text), None) + assert carrier is not None, (worker.pid, blocked.raw) + return _ECHOED.validate_json(carrier.split(_ECHO, 1)[1])["description"] + + +def _catalog_is_cold(worker: _Worker, name: str) -> bool: + return not _probed_description(worker, name) + + +@pytest.mark.timeout(600) +def test_worker_restart_cools_its_listed_catalog_while_the_sibling_worker_keeps_serving( + gateway: Gateway, echo_config: Path, tmp_path: Path +) -> None: + tool: Final = ScriptedTool("lookup", lambda _: text_result("found"), description=_DESCRIPTION) + with scripted_peer(tool) as peer, gateway.scenario() as scenario: + alias: Final = "cold" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-lookup" + with owned_proxy_process(gateway, tmp_path / "sibling", {}, config=echo_config) as sibling_proxy: + sibling: Final = _worker(sibling_proxy.gateway, key, alias) + with owned_proxy_process(gateway, tmp_path / "first", {}, config=echo_config) as first_proxy: + first: Final = _worker(first_proxy.gateway, key, alias) + assert first.pid != sibling.pid + _served(first, name, "first-unlisted") + _served(sibling, name, "sibling-unlisted") + assert _catalog_is_cold(first, name) and _catalog_is_cold(sibling, name) + _listed(first.caller, name) + assert _probed_description(first, name) == _DESCRIPTION + assert _catalog_is_cold(sibling, name), "a listing on one worker warmed its sibling" + _listed(sibling.caller, name) + assert _probed_description(sibling, name) == _DESCRIPTION + _served(sibling, name, "sibling-alone") + with owned_proxy_process(gateway, tmp_path / "restarted", {}, config=echo_config) as restarted_proxy: + restarted: Final = _worker(restarted_proxy.gateway, key, alias) + assert restarted.pid not in (first.pid, sibling.pid) + assert _catalog_is_cold(restarted, name), "a restarted worker kept the old process's catalog" + assert _probed_description(sibling, name) == _DESCRIPTION + _served(restarted, name, "restarted-unlisted") + _listed(restarted.caller, name) + assert _probed_description(restarted, name) == _DESCRIPTION + calls: Final = tool_calls(peer.drain()) + assert [_call_nonce(call) for call in calls] == [ + "first-unlisted", + "sibling-unlisted", + "sibling-alone", + "restarted-unlisted", + ], calls + assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in calls), calls + + +def _outage_echo(name: str, failures: int) -> ScriptedTool: + attempts: Final = itertools.count(1) + + def respond(params: JsonRpc) -> Reply | JsonRpc: + if next(attempts) <= failures: + return Reply(status=503, body=b'{"error": "scripted outage"}') + return text_result(_STRINGS.validate_python(params["arguments"])["nonce"]) + + return ScriptedTool(name, respond) + + +def _echo_call(caller: McpCaller, name: str, nonce: str) -> Outcome: + return caller.call(name, {"nonce": nonce}) + + +def _logged_nonces(key: str, count: int) -> tuple[tuple[str, str], ...]: + rows: Final = eventually( + lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")), + lambda found: len(found) >= count, + seconds=70, + ) + return tuple(sorted((str(row["status"]), str(row["nonce"])) for row in rows)) + + +def test_peer_outage_during_a_bounded_burst_fails_exactly_the_outage_calls_and_lands_each_call_once( + gateway: Gateway, peer: Gateway +) -> None: + with scripted_peer(_outage_echo("echo", _OUTAGE)) as upstream, gateway.scenario() as scenario: + alias: Final = "burst" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = f"{alias}-echo" + callers: Final = (McpCaller(gateway, key, "mcp", alias), McpCaller(peer, key, "mcp", alias)) + for caller in callers: + _listed(caller, name) + upstream.drain() + nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST)) + with ThreadPoolExecutor(max_workers=_BURST) as pool: + outcomes: Final = tuple(pool.map(_echo_call, itertools.cycle(callers), itertools.repeat(name), nonces)) + raws: Final = [outcome.raw for outcome in outcomes] + failed: Final = tuple(nonce for nonce, outcome in zip(nonces, outcomes) if outcome.error is not None) + assert len(failed) == _OUTAGE, raws + assert all(outcome.error is not None or outcome.text == nonce for nonce, outcome in zip(nonces, outcomes)), raws + assert all(_rpc_reply(outcome.raw).id == 1 for outcome in outcomes), raws + burst_calls: Final = tool_calls(upstream.drain()) + assert sorted(_call_nonce(call) for call in burst_calls) == sorted(nonces), burst_calls + assert all(_call_params(call)["arguments"] == {"nonce": _call_nonce(call)} for call in burst_calls), burst_calls + assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in burst_calls), burst_calls + recovered: Final = tuple( + _echo_call(caller, name, nonce) for caller, nonce in zip(itertools.cycle(callers), failed) + ) + assert [outcome.text for outcome in recovered] == list(failed), [outcome.raw for outcome in recovered] + assert sorted(_call_nonce(call) for call in tool_calls(upstream.drain())) == sorted(failed) + assert _logged_nonces(key, _BURST + len(failed)) == tuple( + sorted([("success", nonce) for nonce in nonces] + [("failure", nonce) for nonce in failed]) + ) diff --git a/tests/integration/mcp/test_mcp_toolsets.py b/tests/integration/mcp/test_mcp_toolsets.py index 3dd309db665..b97fa6e72cb 100644 --- a/tests/integration/mcp/test_mcp_toolsets.py +++ b/tests/integration/mcp/test_mcp_toolsets.py @@ -1,5 +1,8 @@ +import re import secrets +import textwrap import uuid +from collections.abc import Iterator from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Final @@ -7,14 +10,17 @@ from typing import Final import httpx import pytest import yaml -from integration._support.client import Gateway, Scenario, object_value +from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value from integration._support.mcp import ( INITIALIZE, Outcome, + ScriptedTool, _outcome_from_rest, _outcome_from_rpc, mcp_peer, register_mcp, + scripted_peer, + text_result, tool_calls, ) from integration._support.mcp_grants import create_toolset @@ -507,3 +513,63 @@ def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_i crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply") assert not crossed.ok, crossed.raw assert tool_calls(peer.drain()) == () + + +_PROBE: Final = "catalog-probe" +_ECHO: Final = "catalog-echo" +_UNLISTED: Final = "" +_GUARDRAIL_CODE: Final = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n' + " return allow()\n" + ' function = inputs.get("tools", [{}])[0].get("function", {})\n' + f' return block("{_ECHO}[" + function.get("description") + "]")\n' +) + + +_ECHO_GUARDRAIL_YAML: Final = ( + "guardrails:\n" + " - guardrail_name: catalog-echo\n" + " litellm_params:\n" + " guardrail: custom_code\n" + " mode: pre_mcp_call\n" + " default_on: true\n" + " custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ") +) + + +@pytest.fixture(scope="module") +def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("catalog-echo") + path: Final = directory / "catalog_echo.yaml" + path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML) + with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig: + yield rig + + +def _echoed_description(outcome: Outcome) -> str: + found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw) + assert found is not None, outcome.raw + return found.group(1) + + +def test_a_team_keys_toolset_route_listing_feeds_its_own_calls_but_not_a_team_mates(echo_rig: Gateway) -> None: + described: Final = "Adds for the team " + uuid.uuid4().hex[:8] + tool: Final = ScriptedTool("add", lambda _: text_result("9"), description=described) + with scripted_peer(tool) as peer, echo_rig.scenario() as scenario: + alias: Final = "lit6029echo" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + team_mate: Final = scenario.key(team_id=team_id) + _assert_team_grants_only(echo_rig, team_id, key, granted_id) + listed: Final = _toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/list", {}) + assert listed.ok and listed.tools == (f"{alias}-add",), listed.raw + probe: Final[dict[str, object]] = {"name": f"{alias}-add", "arguments": {"probe": _PROBE}} + own: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/call", probe)) + mate: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(team_mate), granted_name, "tools/call", probe)) + assert (own, mate) == (described, _UNLISTED), ( + "the slot is keyed by the hashed key, so a team-mate that never listed is handed nothing" + ) + assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"