test(mcp): audit listed metadata across callers and bridge lifecycles

This commit is contained in:
Devin AI 2026-10-03 19:11:59 +00:00
parent cbc12d1e1c
commit 982a49d498
10 changed files with 2087 additions and 57 deletions

View file

@ -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()

View file

@ -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"

View file

@ -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]

View file

@ -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"

View file

@ -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

View file

@ -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) == ()

View file

@ -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)}",)

View file

@ -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"

View file

@ -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])
)

View file

@ -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"