mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test(mcp): audit listed metadata across callers and bridge lifecycles
This commit is contained in:
parent
cbc12d1e1c
commit
982a49d498
10 changed files with 2087 additions and 57 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) == ()
|
||||
|
|
|
|||
|
|
@ -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)}",)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue