mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
* feat(mcp): hand listed-tool metadata to pre-call hooks with per-caller catalog identity Track the tools each MCP server listed per caller identity so pre_mcp_call and during_mcp_call hooks receive the tool description and input schema the client saw. Servers with no caller-dependent inputs share one slot; user identity, forwarded headers, stdio env, relayed bearers, and server-specific auth get their own. Local registry and OpenAPI paths pass the registered metadata and admin description overrides. The Agent 365 guardrail reads the new fields into its evaluate payload. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): drop the listed-tools empty sentinel and routine test docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): mark the listed-tools cache digest as a non-security hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key the listed-tools cache by the OBO subject token token_exchange servers list upstream with the caller's own Entra bearer, so two callers on one LiteLLM key with different subjects were sharing a catalog slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): resolve the BYOK credential before keying the listed-tools slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): drop the OAuth discovery cache when a server definition changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): drop a diff-narrating comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): never validate a supplied header on the tools/list BYOK path The pre-listing resolver ran the tool-call byok_auth_required check even when the caller already supplied x-mcp-auth, and it ran outside the per-server error boundary, so a single deprecated-header caller dropped the server from the aggregate list. Listing now returns a supplied header unchanged and falls back to the stored credential without raising Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): assert the BYOK listing lands in the caller's listed-tool slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): cover the deprecated string x-mcp-auth header on a BYOK tools/list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key the per-caller listed-tool slot by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key discovery cache by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): derive the listed-tool caller identity from the discovery cache key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep OAuth metadata generations only while a fetch is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): return one masked text per scanned string in the selected-guardrail REST test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): fold the signed caller into the discovery digest instead of a second key hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): satisfy type discipline gate on listed-tool identity Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): hand tools/call hooks the exact catalog entry tools/list served get_listed_tool re-applied the admin description override on top of the cached listing, so a guardrail-masked description was restored to its original wording at call time, and the OpenAPI / local-registry call path built its metadata from the registry instead of the guarded caller catalog. Both paths now return the cached entry as served, falling back to the registry only when no listing was recorded Adds tests/integration/mcp/test_mcp_listed_tool_metadata.py (red on the prior head for the two regressions, red on the merge base for the feature, green on this head) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key OpenAPI listed-tool entries per caller so tools/call reads its own guarded listing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): align listed-tool slot tests with per-caller keying Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep oauth2 listing on the minted or signed credential, not the stored BYOK secret The listing helper that keys the per-caller catalog by the stored BYOK credential also handed that credential to the upstream client, which on an oauth2 server short-circuited the client_credentials mint and the MCPJWTSigner gate. Split the two: the catalog identity keeps the stored credential so tools/call finds the caller's slot, while an oauth2 server's tools/list sends only the per-request header, letting the M2M mint or signed JWT proceed as on main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): oauth2 BYOK listing sends the minted token, not the stored secret, through the real proxy Integration cell for the listing fix: a client_credentials BYOK server with a stored user credential, one tools/list as that user, the peer must see a live minted bearer and one /token mint. Red at the pre-fix tip (zero mints, stored secret upstream), green at the fixed head and at the merge base Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list Listing used the resolved stored credential both to key the caller's catalog slot and as the upstream transport header, so REST api_key and bearer_token listings sent the user's secret instead of the server's static token and the MCPJWTSigner gate went quiet. The upstream client and the signer gate now read the caller-supplied mcp_auth_header for every auth type, exactly as before the catalog existed, and the stored credential only names the slot tools/call reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): type the listed-tool metadata read from pre-call kwargs for the basedpyright gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): hand never-listed tools/call hooks name and arguments only The local-registry call path fell back to the registry entry with the admin description override when no tools/list had been recorded for the caller, so a pre_mcp_call guardrail scanned a description the caller was never served and blocked OpenAPI calls that passed before, and base's own selected-guardrail REST test failed on the two-text redaction. _registered_tool_metadata now returns the listed entry or None, so a tools/call with no prior listing sends name and arguments only as promised, and that REST test double goes back to its base shape * fix(mcp): keep during_mcp_call hooks on name and arguments only call_tool handed the caller's listed entry to the during-hook task as well, so during_mcp_call guardrails scanned the description line and schema leaves of any listed tool after the upstream call had already run, blocking calls that passed before whenever the policy matched the description, returned a fixed-length texts list, or hit the depth guard on a deep schema. The listed entry is only disclosed for pre_mcp_call, so the during task no longer receives it and its request object carries no description or schema, as before * fix(mcp): key the BYOK catalog slot by the client's header, not the stored credential tools/list resolved the stored BYOK credential to pick the caller's catalog slot, which read the credential store before the classified try block. With Postgres down and a cold per-worker cache that made every REST tools/list on an is_byok server fail with tools=[] and no upstream call, and the read seeded the per-worker cache (including a negative entry), so a tools/call on another worker after a store, rotate or revoke on this one kept using the stale value. The slot is now keyed by what the client supplied plus the caller's hashed key, on both sides. _get_tools_from_server and call_tool take a keyword-only catalog_auth_header that defaults to mcp_auth_header as received (the default is the builtin Ellipsis so it survives a module reload). The /mcp fan-out and execute_mcp_tool, which swap the resolved credential into mcp_auth_header, pass the client's value explicitly. What goes upstream is unchanged. _byok_catalog_auth_header is gone. * fix(mcp): drop a listed catalog recorded across a server save _record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save description next to the post-save definition until the next listing, instead of name and arguments only. The manager now keeps a per-server listed-tools generation, bumped by _invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and _record_listed_tools skips the write when it moved; the next listing records normally. * fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt add_server and update_server publish the saved definition before the OpenAPI registry entries are rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the new generation. The generation is bumped a second time after the registry refresh. The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is inlined at its two call sites, and the per-server generation map is a plain dict. * fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read add_server and update_server ran the full server-definition invalidation a second time after the awaited OpenAPI spec fetch, which also dropped the prompts/resources/templates discovery entries and the OAuth protected-resource metadata filled under the already-published definition, so the next request went upstream again. Only the listed-tool catalog recorded during the fetch holds pre-save entries, so the post-fetch pass now drops just that catalog and bumps its generation via the new _drop_listed_tools helper, which the full invalidation also calls. * fix(mcp): look a called tool up in the listed catalog by its bare name only get_listed_tool stripped the server prefix a second time when the exact name was absent from the caller's listing, so a never-listed upstream tool whose bare name starts with the server prefix resolved to the listed sibling and that sibling's description and input schema reached the pre-call hooks for a call to a different tool. Every caller already passes the once-stripped bare name, so the lookup is now exact. Tests that looked the catalog up by a prefixed name now use the bare name the callers pass; two new tests pin the never-listed sibling case at the manager and at the tools/call path. * fix(mcp): record a listed-tool catalog only for a listing the caller is served _get_tools_from_server now records the catalog into the caller's listed-tools slot only when asked (record_listing=True), which the served listings pass: the /mcp and Responses API tools/list handlers via _get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST listing via _list_server_tools. Four internal listings stop recording, so a later tools/call hands pre_mcp_call hooks name and arguments only, as on main: - _list_tools_before_first_call, the implicit listing inside tools/call when this worker does not yet expose the tool - fetch_pinnable_tool_catalog, the admin pin snapshot listed without the catalog guard and without description overrides - _initialize_tool_name_to_mcp_server_name_mapping, the startup fill - get_tools_for_server, used by the semantic tool filter _create_prefixed_tools returns to its tool-name mapping job only; the record follows it in _get_tools_from_server. * fix(mcp): opt every listing out of catalog recording unless it is served The aggregate listing and _list_mcp_tools now default to record_listing=False, so a catalog fetched inside a tools/call no longer fills the caller's listed-tools slot. The /mcp/proxy meta-tools (call_tool, search_tools, get_tool_schema) and the tool-search virtual tool stop recording: /mcp/proxy serves only the meta-tools and the search serves only its hits, so a later pre_mcp_call hook was reading a description the caller never listed. The tools/list handler, the Responses MCP handler and the /v1/mcp/tools management listing opt in with record_listing=True, since each serves the catalog to the caller. * fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer The slot a tools/list records for a later tools/call was keyed by (user_id, api_key) only, so every team-only JWT caller shared one slot and one JWT user acting in two teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another caller was served. The slot is now keyed by the hashed key, user, team and organization, plus the admission credential of a caller admitted with neither a key nor a user. The caller bearer split the slot only on client-forwarded-token and token-exchange servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client credentials) also forwards it upstream and served a different catalog per bearer into one slot. The bearer now splits the slot on every server whose egress forwards it (_consumes_caller_authorization) or exchanges it. * fix(mcp): record only tools served by the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep bridge tool metadata request-local * refactor(mcp): centralize listed catalog recording guard * fix(mcp): preserve base TPM reservations for listed tool calls * fix(mcp): preserve project token reservations for listed calls * fix(mcp): record served catalogs and preserve call message bytes * test(mcp): align listing expectation with deferred recording * test(mcp): audit listed metadata across callers and bridge lifecycles * test(mcp): preserve guardrail fixture worker affinity * refactor(mcp): expose listed catalog recording API * feat(mcp): pass served_tools through the anthropic messages bridge /v1/messages auto-execution now hands this request's resolved tool definitions to _execute_tool_calls, matching the Responses and chat completions bridges: the pre_mcp_call hook receives the description and input schema the model was shown for that call. Request-local only; the shared listed-tools catalog is untouched. * test(mcp): pin served_tools handoff on the anthropic messages bridge Mirrors the credentials-forwarding test: the request's resolved tool definitions must reach _execute_tool_calls under served_tools so pre_mcp_call hooks judge the call on the description and input schema the model was shown. Fails without the previous commit's one-liner. * style(mcp): sort the local import block ruff flagged * fix(mcp): keep the admin include_disabled_tools view off the listed-tools catalog GET /mcp-rest/tools/list?include_disabled_tools=true is the admin-only configuration view: apply_tool_filters is False, so it serves the full server catalog. Recording that response into the caller's listed-tools slot warmed tools/call metadata no runtime listing ever served, breaking the only-a-served-listing-records invariant (Bugbot). The record is now gated on apply_tool_filters; disabled tools stay unreachable (the call-time allowlist 403 fires before hooks), so the observable fix is the slot no longer warming from a settings view. Verified live: the new test fails on the unfixed head and passes here, and the rest of the listed-tool-metadata suite is unchanged. --------- Co-authored-by: yucheng <yucheng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
624 lines
24 KiB
Python
624 lines
24 KiB
Python
import asyncio
|
|
import json
|
|
import os
|
|
import queue
|
|
import sys
|
|
import time
|
|
from collections.abc import Callable, Iterator, Mapping
|
|
from contextlib import contextmanager
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Final, Literal
|
|
|
|
import httpx
|
|
from integration._support.asgi import asgi_server
|
|
from integration._support.client import Gateway, Scenario
|
|
from integration._support.database import read_rows
|
|
from integration._support.wire import Reply, Request, wire_server
|
|
from mcp import ClientSession
|
|
from mcp.client.sse import sse_client
|
|
from mcp.client.streamable_http import streamable_http_client
|
|
from mcp.server.mcpserver import Context, MCPServer
|
|
from mcp.server.transport_security import TransportSecuritySettings
|
|
from mcp.types import SamplingMessage, TextContent
|
|
from mcp_tests.mcp_e2e_upstream_server import add, multiply
|
|
from pydantic import BaseModel
|
|
from sse_starlette.sse import AppStatus
|
|
from starlette.requests import Request as StarletteRequest
|
|
from starlette.responses import Response
|
|
from starlette.types import Message, Receive, Scope, Send
|
|
|
|
Transport = Literal["http", "sse", "stdio"]
|
|
STDIO_PEER: Final = Path(__file__).with_name("mcp_stdio_peer.py")
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class McpPeer:
|
|
url: str
|
|
calls: queue.Queue[dict[str, object]]
|
|
transport: Transport = "http"
|
|
command: str | None = None
|
|
args: tuple[str, ...] = ()
|
|
record: Path | None = None
|
|
spec_path: Path | None = None
|
|
consumed: list[int] = field(default_factory=lambda: [0])
|
|
|
|
def drain(self) -> tuple[dict[str, object], ...]:
|
|
if self.record is not None:
|
|
lines: Final = self.record.read_text().splitlines() if self.record.exists() else []
|
|
fresh: Final = tuple(json.loads(line) for line in lines[self.consumed[0] :])
|
|
self.consumed[0] = len(lines)
|
|
return fresh
|
|
return tuple(self.calls.get_nowait() for _ in range(self.calls.qsize()))
|
|
|
|
def registration(self) -> dict[str, object]:
|
|
if self.transport == "stdio":
|
|
return {"transport": "stdio", "command": self.command, "args": list(self.args)}
|
|
if self.spec_path is not None:
|
|
return {"transport": "http", "url": self.url, "spec_path": str(self.spec_path)}
|
|
return {"transport": self.transport, "url": self.url}
|
|
|
|
|
|
class Confirmation(BaseModel):
|
|
confirmed: bool
|
|
|
|
|
|
def math_service(name: str = "integration-math", *, rich: bool = False) -> MCPServer:
|
|
service: Final = MCPServer(name)
|
|
service.add_tool(add)
|
|
service.add_tool(multiply)
|
|
|
|
@service.tool()
|
|
def fail() -> str:
|
|
raise ValueError("synthetic tool failure")
|
|
|
|
if not rich:
|
|
return service
|
|
|
|
@service.tool()
|
|
async def slow(seconds: float) -> str:
|
|
await asyncio.sleep(seconds)
|
|
return "slept"
|
|
|
|
@service.tool()
|
|
async def progress(steps: int, ctx: Context) -> str:
|
|
for step in range(steps):
|
|
await ctx.report_progress(step + 1, steps, f"step {step + 1}")
|
|
return f"{steps} steps"
|
|
|
|
@service.tool()
|
|
async def sample(prompt: str, ctx: Context) -> str:
|
|
result: Final = await ctx.session.create_message(
|
|
messages=[SamplingMessage(role="user", content=TextContent(type="text", text=prompt))],
|
|
max_tokens=32,
|
|
)
|
|
return "sampled:" + (result.content.text if isinstance(result.content, TextContent) else "")
|
|
|
|
@service.tool()
|
|
async def elicit(question: str, ctx: Context) -> str:
|
|
result: Final = await ctx.elicit(message=question, schema=Confirmation)
|
|
return f"elicited:{result.action}"
|
|
|
|
@service.prompt()
|
|
def greeting(name: str) -> str:
|
|
return f"Hello, {name}"
|
|
|
|
@service.resource("status://ready")
|
|
def status() -> str:
|
|
return "ready"
|
|
|
|
@service.resource("greeting://{name}")
|
|
def greeting_resource(name: str) -> str:
|
|
return f"Hello, {name}"
|
|
|
|
return service
|
|
|
|
|
|
def _capturing(app: Callable[[Scope, Receive, Send], object], observed: queue.Queue[dict[str, object]]):
|
|
async def capture(scope: Scope, receive: Receive, send: Send) -> None:
|
|
if scope["type"] != "http":
|
|
await app(scope, receive, send)
|
|
return
|
|
if scope["method"] == "GET" and scope["path"].endswith("/mcp"):
|
|
await Response(status_code=405, headers={"Allow": "POST, DELETE"})(scope, receive, send)
|
|
return
|
|
body: Final = await StarletteRequest(scope, receive).body()
|
|
assert len(body) <= 65536
|
|
if body:
|
|
observed.put({"body": json.loads(body), "headers": dict(scope["headers"]), "path": scope["path"]})
|
|
message: Final = {"type": "http.request", "body": body, "more_body": False}
|
|
pending: Final = iter((message,))
|
|
|
|
async def replay() -> Message:
|
|
buffered: Final = next(pending, None)
|
|
if buffered is not None:
|
|
return buffered
|
|
return await receive()
|
|
|
|
await app(scope, replay, send)
|
|
|
|
return capture
|
|
|
|
|
|
def _drain_sse_streams() -> None:
|
|
AppStatus.should_exit = True
|
|
|
|
|
|
def _draining_sse_watcher(app: Callable[[Scope, Receive, Send], object]):
|
|
"""sse_starlette parks a per-loop watcher that only stops once AppStatus.should_exit flips."""
|
|
|
|
async def lifespan(scope: Scope, receive: Receive, send: Send) -> None:
|
|
while True:
|
|
message: Final = await receive()
|
|
if message["type"] == "lifespan.startup":
|
|
AppStatus.should_exit = False
|
|
await send({"type": "lifespan.startup.complete"})
|
|
elif message["type"] == "lifespan.shutdown":
|
|
_drain_sse_streams()
|
|
watchers: Final = tuple(
|
|
task for task in asyncio.all_tasks() if "_shutdown_watcher" in repr(task.get_coro())
|
|
)
|
|
await asyncio.gather(*watchers)
|
|
await send({"type": "lifespan.shutdown.complete"})
|
|
return
|
|
|
|
async def wrapped(scope: Scope, receive: Receive, send: Send) -> None:
|
|
if scope["type"] == "lifespan":
|
|
await lifespan(scope, receive, send)
|
|
return
|
|
starts: Final = [0]
|
|
|
|
async def send_once(message: Message) -> None:
|
|
if message["type"] == "http.response.start":
|
|
starts[0] += 1
|
|
if starts[0] == 2:
|
|
await send({"type": "http.response.body", "body": b"", "more_body": False})
|
|
if starts[0] > 1:
|
|
return
|
|
await send(message)
|
|
|
|
await app(scope, receive, send_once)
|
|
|
|
return wrapped
|
|
|
|
|
|
@contextmanager
|
|
def mcp_peer(transport: Literal["http", "sse"] = "http", *, rich: bool = False) -> Iterator[McpPeer]:
|
|
service: Final = math_service(rich=rich)
|
|
security: Final = TransportSecuritySettings(enable_dns_rebinding_protection=False)
|
|
app: Final = (
|
|
_draining_sse_watcher(service.sse_app(transport_security=security))
|
|
if transport == "sse"
|
|
else service.streamable_http_app(stateless_http=True, json_response=True, transport_security=security)
|
|
)
|
|
observed: Final[queue.Queue[dict[str, object]]] = queue.Queue()
|
|
with asgi_server(_capturing(app, observed), before_stop=_drain_sse_streams if transport == "sse" else None) as url:
|
|
yield McpPeer(url + ("/sse" if transport == "sse" else "/mcp"), observed, transport)
|
|
|
|
|
|
@contextmanager
|
|
def stdio_peer(directory: Path, *, rich: bool = False) -> Iterator[McpPeer]:
|
|
record: Final = directory / f"stdio-{os.getpid()}-{time.monotonic_ns()}.jsonl"
|
|
yield McpPeer(
|
|
"",
|
|
queue.Queue(),
|
|
"stdio",
|
|
sys.executable,
|
|
(str(STDIO_PEER), str(record), "rich" if rich else "plain"),
|
|
record,
|
|
)
|
|
|
|
|
|
JsonRpc = Mapping[str, object]
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class ScriptedTool:
|
|
name: str
|
|
respond: Callable[[JsonRpc], Reply | JsonRpc]
|
|
description: str | Callable[[Mapping[str, str]], str] | None = None
|
|
input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"})
|
|
|
|
def listing(self, headers: Mapping[str, str]) -> JsonRpc:
|
|
described: Final = self.description(headers) if callable(self.description) else self.description
|
|
return {
|
|
"name": self.name,
|
|
"inputSchema": self.input_schema,
|
|
**({} if described is None else {"description": described}),
|
|
}
|
|
|
|
|
|
def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply:
|
|
return Reply(body=json.dumps({"jsonrpc": "2.0", "id": identity, "result": result}).encode())
|
|
|
|
|
|
def jsonrpc_error(identity: object, code: int, message: str) -> Reply:
|
|
return Reply(
|
|
body=json.dumps({"jsonrpc": "2.0", "id": identity, "error": {"code": code, "message": message}}).encode()
|
|
)
|
|
|
|
|
|
@contextmanager
|
|
def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]:
|
|
"""Raw JSON-RPC peer for shapes the SDK server cannot produce: half-written bodies, stalls, wire errors."""
|
|
observed: Final[queue.Queue[dict[str, object]]] = queue.Queue()
|
|
by_name: Final = {tool.name: tool for tool in tools}
|
|
|
|
def provider(request: Request) -> Reply:
|
|
if request.method != "POST":
|
|
return Reply(status=405)
|
|
body: Final = json.loads(request.body)
|
|
observed.put({"body": body, "headers": dict(request.headers), "path": request.target})
|
|
if "id" not in body:
|
|
return Reply(status=202)
|
|
identity: Final = body["id"]
|
|
method: Final = body["method"]
|
|
if method == "initialize":
|
|
return jsonrpc_reply(
|
|
identity,
|
|
{
|
|
"protocolVersion": body["params"]["protocolVersion"],
|
|
"capabilities": {"tools": {}},
|
|
"serverInfo": {"name": "integration-scripted-peer", "version": "1"},
|
|
},
|
|
)
|
|
if method == "tools/list":
|
|
return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]})
|
|
if method != "tools/call":
|
|
return jsonrpc_error(identity, -32601, f"unsupported method {method}")
|
|
tool: Final = by_name.get(body["params"]["name"])
|
|
if tool is None:
|
|
return jsonrpc_error(identity, -32602, "unknown tool")
|
|
produced: Final = tool.respond(body["params"])
|
|
return produced if isinstance(produced, Reply) else jsonrpc_reply(identity, produced)
|
|
|
|
with wire_server(provider) as wire:
|
|
yield McpPeer(wire.url + "/mcp", observed)
|
|
|
|
|
|
def text_result(text: str) -> JsonRpc:
|
|
return {"content": [{"type": "text", "text": text}], "isError": False}
|
|
|
|
|
|
def slow_tool(name: str, seconds: float) -> ScriptedTool:
|
|
def respond(params: JsonRpc) -> JsonRpc:
|
|
time.sleep(seconds)
|
|
return text_result("slept")
|
|
|
|
return ScriptedTool(name, respond)
|
|
|
|
|
|
def disconnecting_tool(name: str) -> ScriptedTool:
|
|
return ScriptedTool(name, lambda params: Reply(chunks=(b'{"jsonrpc":"2.0",', b'"id":1}'), abort_after=1))
|
|
|
|
|
|
def echo_tool(name: str) -> ScriptedTool:
|
|
return ScriptedTool(name, lambda params: text_result(json.dumps(params.get("arguments", {}), sort_keys=True)))
|
|
|
|
|
|
@contextmanager
|
|
def openapi_peer() -> Iterator[McpPeer]:
|
|
"""OpenAPI-described HTTP service plus the spec file the proxy turns into MCP tools."""
|
|
observed: Final[queue.Queue[dict[str, object]]] = queue.Queue()
|
|
|
|
def provider(request: Request) -> Reply:
|
|
observed.put(
|
|
{
|
|
"body": json.loads(request.body) if request.body else None,
|
|
"headers": dict(request.headers),
|
|
"path": request.target,
|
|
"method": request.method,
|
|
}
|
|
)
|
|
if request.target.startswith("/pets/") and request.method == "GET":
|
|
return Reply(body=json.dumps({"id": request.target.rsplit("/", 1)[1], "name": "integration-pet"}).encode())
|
|
if request.target == "/pets" and request.method == "POST":
|
|
return Reply(status=201, body=json.dumps({"created": json.loads(request.body)}).encode())
|
|
return Reply(status=404, body=b'{"error":"synthetic not found"}')
|
|
|
|
with wire_server(provider) as wire:
|
|
spec: Final = {
|
|
"openapi": "3.0.0",
|
|
"info": {"title": "integration pets", "version": "1"},
|
|
"servers": [{"url": wire.url}],
|
|
"paths": {
|
|
"/pets/{petId}": {
|
|
"get": {
|
|
"operationId": "getPet",
|
|
"summary": "Fetch one pet",
|
|
"parameters": [{"name": "petId", "in": "path", "required": True, "schema": {"type": "string"}}],
|
|
"responses": {"200": {"description": "pet"}},
|
|
}
|
|
},
|
|
"/pets": {
|
|
"post": {
|
|
"operationId": "createPet",
|
|
"summary": "Create a pet",
|
|
"requestBody": {
|
|
"required": True,
|
|
"content": {
|
|
"application/json": {
|
|
"schema": {
|
|
"type": "object",
|
|
"properties": {"name": {"type": "string"}},
|
|
"required": ["name"],
|
|
}
|
|
}
|
|
},
|
|
},
|
|
"responses": {"201": {"description": "created"}},
|
|
}
|
|
},
|
|
},
|
|
}
|
|
yield McpPeer(wire.url, observed, spec_path=_spec_file(spec))
|
|
|
|
|
|
def scratch_directory() -> Path:
|
|
path: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", "/tmp")) / "mcp-peers"
|
|
path.mkdir(parents=True, exist_ok=True)
|
|
return path
|
|
|
|
|
|
def _spec_file(spec: JsonRpc) -> Path:
|
|
path: Final = scratch_directory() / f"openapi-{time.monotonic_ns()}.json"
|
|
path.write_text(json.dumps(spec))
|
|
return path
|
|
|
|
|
|
PeerKind = Literal["http", "sse", "stdio", "openapi"]
|
|
PEER_KINDS: Final[tuple[PeerKind, ...]] = ("http", "sse", "stdio", "openapi")
|
|
|
|
|
|
@contextmanager
|
|
def peer_of(kind: PeerKind, *, rich: bool = False) -> Iterator[McpPeer]:
|
|
if kind == "openapi":
|
|
with openapi_peer() as candidate:
|
|
yield candidate
|
|
elif kind == "stdio":
|
|
with stdio_peer(scratch_directory(), rich=rich) as candidate:
|
|
yield candidate
|
|
else:
|
|
with mcp_peer(kind, rich=rich) as candidate:
|
|
yield candidate
|
|
|
|
|
|
def register_mcp(scenario: Scenario, peer: McpPeer, alias: str, **fields: object) -> str:
|
|
response: Final = scenario.gateway.request(
|
|
"POST", "/v1/mcp/server", {"server_name": alias, "alias": alias, **peer.registration(), **fields}
|
|
)
|
|
identity: Final = response.json()["server_id"]
|
|
scenario.cleanups.callback(forget_mcp, scenario.gateway, identity)
|
|
assert response.status_code == 201, response.text
|
|
return identity
|
|
|
|
|
|
def forget_mcp(gateway: Gateway, identity: str) -> None:
|
|
response: Final = gateway.request("DELETE", f"/v1/mcp/server/{identity}")
|
|
assert response.status_code in (202, 404), response.text
|
|
|
|
|
|
def delete_mcp(gateway: Gateway, identity: str) -> None:
|
|
response: Final = gateway.request("DELETE", f"/v1/mcp/server/{identity}")
|
|
assert response.status_code == 202, response.text
|
|
assert read_rows('SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) == []
|
|
|
|
|
|
def listed_tools(gateway: Gateway, key: str, identity: str | None = None) -> dict[str, dict[str, object]]:
|
|
response: Final = gateway.client.get(
|
|
"/mcp-rest/tools/list",
|
|
headers={"x-litellm-api-key": key},
|
|
params={"server_id": identity} if identity else None,
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
return {
|
|
tool["name"]: tool
|
|
for tool in response.json()["tools"]
|
|
if identity is None or tool.get("mcp_info", {}).get("server_id") == identity
|
|
}
|
|
|
|
|
|
def tool_names(gateway: Gateway, key: str, identity: str) -> dict[str, str]:
|
|
return {
|
|
name: full
|
|
for full in listed_tools(gateway, key, identity)
|
|
for name in ("add", "multiply", "fail", "slow", "progress", "sample", "elicit")
|
|
if full.endswith(name)
|
|
}
|
|
|
|
|
|
def call_tool(gateway: Gateway, key: str, identity: str, name: str, arguments: dict[str, object]) -> httpx.Response:
|
|
return gateway.client.post(
|
|
"/mcp-rest/tools/call",
|
|
headers={"x-litellm-api-key": key},
|
|
json={"server_id": identity, "name": name, "arguments": arguments},
|
|
)
|
|
|
|
|
|
EntryPoint = Literal["mcp", "server_mcp", "root", "sse", "rest"]
|
|
ENTRY_POINTS: Final[tuple[EntryPoint, ...]] = ("mcp", "server_mcp", "root", "sse", "rest")
|
|
INITIALIZE: Final = {
|
|
"protocolVersion": "2025-06-18",
|
|
"capabilities": {},
|
|
"clientInfo": {"name": "integration", "version": "1"},
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class Outcome:
|
|
"""What a caller saw from one MCP operation, normalised across entry points."""
|
|
|
|
status: int
|
|
error: str | None
|
|
tools: tuple[str, ...] = ()
|
|
text: str | None = None
|
|
raw: str = ""
|
|
|
|
@property
|
|
def ok(self) -> bool:
|
|
return self.status == 200 and self.error is None
|
|
|
|
|
|
def _parse_rpc_body(response: httpx.Response) -> Mapping[str, object] | None:
|
|
if response.headers.get("content-type", "").startswith("text/event-stream"):
|
|
data: Final = tuple(line[5:].strip() for line in response.text.splitlines() if line.startswith("data:"))
|
|
return json.loads(data[-1]) if data else None
|
|
try:
|
|
return json.loads(response.text)
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def _outcome_from_rpc(response: httpx.Response) -> Outcome:
|
|
body: Final = _parse_rpc_body(response)
|
|
if response.status_code != 200 or body is None:
|
|
return Outcome(response.status_code, response.text or f"HTTP {response.status_code}", raw=response.text)
|
|
if "error" in body:
|
|
return Outcome(response.status_code, json.dumps(body["error"]), raw=response.text)
|
|
result: Final = body.get("result", {})
|
|
assert isinstance(result, dict)
|
|
if "tools" in result:
|
|
return Outcome(200, None, tuple(tool["name"] for tool in result["tools"]), raw=response.text)
|
|
content: Final = result.get("content", [])
|
|
text: Final = content[0].get("text") if content else None
|
|
if result.get("isError"):
|
|
return Outcome(200, text or "isError", text=text, raw=response.text)
|
|
return Outcome(200, None, text=text, raw=response.text)
|
|
|
|
|
|
def _outcome_from_rest(response: httpx.Response) -> Outcome:
|
|
if response.status_code != 200:
|
|
return Outcome(response.status_code, response.text, raw=response.text)
|
|
body: Final = response.json()
|
|
if "tools" in body:
|
|
return Outcome(200, None, tuple(tool["name"] for tool in body["tools"]), raw=response.text)
|
|
content: Final = body.get("content", [])
|
|
text: Final = content[0].get("text") if content else None
|
|
if body.get("isError"):
|
|
return Outcome(200, text or "isError", text=text, raw=response.text)
|
|
return Outcome(200, None, text=text, raw=response.text)
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class McpCaller:
|
|
"""One caller's view of the gateway through a specific entry point."""
|
|
|
|
gateway: Gateway
|
|
key: str | None
|
|
entry: EntryPoint
|
|
alias: str | None = None
|
|
headers: Mapping[str, str] = field(default_factory=dict)
|
|
|
|
def _path(self) -> str:
|
|
if self.entry == "server_mcp":
|
|
assert self.alias is not None
|
|
return f"/{self.alias}/mcp"
|
|
return {"mcp": "/mcp", "root": "/mcp/", "sse": "/mcp/sse", "rest": "/mcp-rest"}[self.entry]
|
|
|
|
def _headers(self) -> dict[str, str]:
|
|
return {
|
|
**({"x-litellm-api-key": self.key} if self.key is not None else {}),
|
|
"Accept": "application/json, text/event-stream",
|
|
**self.headers,
|
|
}
|
|
|
|
def rpc(self, method: str, params: JsonRpc | None = None) -> httpx.Response:
|
|
if self.entry == "sse":
|
|
return _legacy_sse_rpc(self.gateway, self._headers(), method, params)
|
|
return self.gateway.client.post(
|
|
self._path(),
|
|
json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params or {})},
|
|
headers=self._headers(),
|
|
)
|
|
|
|
def initialize(self) -> Outcome:
|
|
if self.entry == "rest":
|
|
return Outcome(200, None)
|
|
return _outcome_from_rpc(self.rpc("initialize", INITIALIZE))
|
|
|
|
def list_tools(self, server_id: str | None = None) -> Outcome:
|
|
if self.entry == "rest":
|
|
return _outcome_from_rest(
|
|
self.gateway.client.get(
|
|
"/mcp-rest/tools/list",
|
|
headers=self._headers(),
|
|
params={"server_id": server_id} if server_id else None,
|
|
)
|
|
)
|
|
return _outcome_from_rpc(self.rpc("tools/list"))
|
|
|
|
def call(self, name: str, arguments: JsonRpc, server_id: str | None = None) -> Outcome:
|
|
if self.entry == "rest":
|
|
return _outcome_from_rest(
|
|
self.gateway.client.post(
|
|
"/mcp-rest/tools/call",
|
|
headers=self._headers(),
|
|
json={
|
|
"name": name,
|
|
"arguments": dict(arguments),
|
|
**({"server_id": server_id} if server_id else {}),
|
|
},
|
|
)
|
|
)
|
|
return _outcome_from_rpc(self.rpc("tools/call", {"name": name, "arguments": dict(arguments)}))
|
|
|
|
|
|
def _legacy_sse_rpc(
|
|
gateway: Gateway, headers: Mapping[str, str], method: str, params: JsonRpc | None
|
|
) -> httpx.Response:
|
|
"""Drive the legacy GET /mcp/sse + POST /mcp/sse/messages pair for one request and synthesise a JSON response."""
|
|
with gateway.client.stream("GET", "/mcp/sse", headers=headers, timeout=15) as stream:
|
|
if stream.status_code != 200:
|
|
stream.read()
|
|
return httpx.Response(stream.status_code, text=stream.text)
|
|
lines: Final = stream.iter_lines()
|
|
endpoint: Final = next(line[5:].strip() for line in lines if line.startswith("data:"))
|
|
init: Final = gateway.client.post(
|
|
endpoint,
|
|
json={"jsonrpc": "2.0", "id": 0, "method": "initialize", "params": INITIALIZE},
|
|
headers=headers,
|
|
)
|
|
assert init.status_code in (200, 202), init.text
|
|
gateway.client.post(endpoint, json={"jsonrpc": "2.0", "method": "notifications/initialized"}, headers=headers)
|
|
posted: Final = gateway.client.post(
|
|
endpoint, json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params or {})}, headers=headers
|
|
)
|
|
if posted.status_code not in (200, 202):
|
|
return httpx.Response(posted.status_code, text=posted.text)
|
|
for line in lines:
|
|
if line.startswith("data:") and '"id": 1' in line.replace('"id":1', '"id": 1'):
|
|
return httpx.Response(200, text=line[5:].strip(), headers={"content-type": "application/json"})
|
|
return httpx.Response(599, text="legacy SSE stream ended without a reply")
|
|
|
|
|
|
def official_client_outcomes(
|
|
gateway: Gateway, key: str, path: str, name: str, arguments: JsonRpc, *, legacy_sse: bool = False
|
|
) -> tuple[Outcome, Outcome]:
|
|
"""List then call through the official MCP client session, returning both outcomes."""
|
|
url: Final = str(gateway.client.base_url).rstrip("/") + path
|
|
headers: Final = {"x-litellm-api-key": key}
|
|
|
|
async def run() -> tuple[Outcome, Outcome]:
|
|
transport: Final = (
|
|
sse_client(url, headers=headers)
|
|
if legacy_sse
|
|
else streamable_http_client(url, http_client=httpx.AsyncClient(headers=headers, timeout=30))
|
|
)
|
|
async with transport as streams, ClientSession(streams[0], streams[1]) as session:
|
|
await session.initialize()
|
|
listed: Final = await session.list_tools()
|
|
result: Final = await session.call_tool(name, dict(arguments))
|
|
content: Final = result.content[0] if result.content else None
|
|
text: Final = content.text if isinstance(content, TextContent) else None
|
|
return (
|
|
Outcome(200, None, tuple(tool.name for tool in listed.tools)),
|
|
Outcome(200, (text or "isError") if result.is_error else None, text=text),
|
|
)
|
|
|
|
return asyncio.run(run())
|
|
|
|
|
|
def tool_calls(observed: tuple[dict[str, object], ...]) -> tuple[dict[str, object], ...]:
|
|
return tuple(
|
|
item for item in observed if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/call"
|
|
)
|