mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
338 lines
14 KiB
Python
338 lines
14 KiB
Python
import hashlib
|
|
import json
|
|
import os
|
|
import queue
|
|
import signal
|
|
import subprocess
|
|
from collections.abc import AsyncIterator, Iterator
|
|
from contextlib import contextmanager
|
|
from itertools import product
|
|
from pathlib import Path
|
|
from typing import TYPE_CHECKING, Final, Literal
|
|
|
|
import httpx
|
|
import psutil
|
|
from pydantic import BaseModel, Field, JsonValue, TypeAdapter
|
|
|
|
if TYPE_CHECKING:
|
|
from integration._support.mcp import McpPeer
|
|
|
|
|
|
class ConformanceCheck(BaseModel):
|
|
id: str
|
|
status: Literal["SUCCESS", "FAILURE", "WARNING", "INFO"]
|
|
errorMessage: str | None = None
|
|
details: dict[str, JsonValue] = Field(default_factory=dict)
|
|
|
|
|
|
def verify_archive(archive: Path, expected_sha256: str) -> None:
|
|
actual: Final = hashlib.sha256(archive.read_bytes()).hexdigest()
|
|
assert actual == expected_sha256, f"SHA-256 mismatch for {archive.name}: {actual}"
|
|
|
|
|
|
def read_checks(directory: Path, scenario: str) -> tuple[ConformanceCheck, ...]:
|
|
reports: Final = tuple(directory.rglob("checks.json"))
|
|
assert len(reports) == 1, f"conformance {scenario}: expected one fresh report, found {len(reports)}"
|
|
checks: Final = TypeAdapter(tuple[ConformanceCheck, ...]).validate_json(reports[0].read_bytes())
|
|
assert len({check.id for check in checks}) == len(checks), f"conformance {scenario}: duplicate checks"
|
|
required: Final = {
|
|
"server-session-lifecycle": (
|
|
"server-session-initialized-accepted",
|
|
"server-session-delete-accepted",
|
|
"server-session-terminated-returns-404",
|
|
),
|
|
"server-sse-multiple-streams": (
|
|
"server-accepts-multiple-post-streams",
|
|
"server-sse-streams-functional",
|
|
"wire-schema-valid",
|
|
),
|
|
"server-initialize": ("server-initialize", "server-session-id-visible-ascii", "wire-schema-valid"),
|
|
"tools-list": ("tools-list", "tools-name-format", "wire-schema-valid"),
|
|
"tools-call-image": ("tools-call-image", "wire-schema-valid"),
|
|
}.get(scenario, (scenario, "wire-schema-valid") if scenario in OFFICIAL_SCENARIOS else (scenario,))
|
|
for identity in required:
|
|
assert tuple(check.status for check in checks if check.id == identity) == ("SUCCESS",), (
|
|
f"conformance {scenario}: required check {identity} did not pass: {checks}"
|
|
)
|
|
assert all(check.status != "FAILURE" for check in checks), f"conformance {scenario}: failed checks: {checks}"
|
|
if scenario == "tools-call-simple-text":
|
|
result: Final = next(check for check in checks if check.id == scenario).details.get("result")
|
|
# The official scenario accepts any nonempty text, including tool errors.
|
|
assert isinstance(result, dict) and result.get("isError", False) is False, f"reference payload: {result}"
|
|
assert result.get("content") == [{"type": "text", "text": "This is a simple text response for testing."}], (
|
|
f"reference payload: {result}"
|
|
)
|
|
return checks
|
|
|
|
|
|
def run_scenario(root: Path, url: str, scenario: str, directory: Path) -> tuple[ConformanceCheck, ...]:
|
|
directory.mkdir(parents=True, exist_ok=False)
|
|
with (directory / "runner.log").open("w") as log:
|
|
result: Final = subprocess.run(
|
|
[
|
|
"node",
|
|
"--import",
|
|
"tsx",
|
|
"src/index.ts",
|
|
"server",
|
|
"--url",
|
|
url,
|
|
"--scenario",
|
|
scenario,
|
|
"--spec-version",
|
|
"2025-11-25",
|
|
"--output-dir",
|
|
str(directory.resolve()),
|
|
"--timeout",
|
|
"30000",
|
|
],
|
|
cwd=root,
|
|
stdout=log,
|
|
stderr=subprocess.STDOUT,
|
|
timeout=45,
|
|
check=False,
|
|
)
|
|
assert result.returncode == 0, (directory / "runner.log").read_text()
|
|
return read_checks(directory, scenario)
|
|
|
|
|
|
@contextmanager
|
|
def reference_server(root: Path, directory: Path, port: int) -> Iterator["McpPeer"]:
|
|
from integration._support.client import eventually
|
|
from integration._support.mcp import McpPeer
|
|
from integration._support.process import group_members, signal_group, stop_root_process
|
|
|
|
directory.mkdir(parents=True, exist_ok=True)
|
|
url: Final = f"http://127.0.0.1:{port}"
|
|
with (directory / "reference.log").open("w") as log:
|
|
process: Final = subprocess.Popen(
|
|
["node", "--import", "tsx", "everything-server.ts"],
|
|
cwd=root / "examples/servers/typescript",
|
|
env={"PATH": os.environ["PATH"], "PORT": str(port)},
|
|
stdout=log,
|
|
stderr=subprocess.STDOUT,
|
|
start_new_session=True,
|
|
)
|
|
try:
|
|
with httpx.Client(timeout=1, trust_env=False) as client:
|
|
|
|
def ready() -> bool:
|
|
assert process.poll() is None, (directory / "reference.log").read_text()
|
|
try:
|
|
return client.get(url + "/mcp").status_code == 400
|
|
except httpx.TransportError:
|
|
return False
|
|
|
|
eventually(ready, bool, seconds=15)
|
|
yield McpPeer(url + "/mcp", queue.Queue())
|
|
finally:
|
|
stopped: Final = stop_root_process(process)
|
|
residual: Final = group_members(process.pid)
|
|
if residual:
|
|
signal_group(process.pid, signal.SIGTERM)
|
|
psutil.wait_procs(residual, timeout=5)
|
|
remaining: Final = group_members(process.pid)
|
|
if remaining:
|
|
signal_group(process.pid, signal.SIGKILL)
|
|
psutil.wait_procs(remaining, timeout=3)
|
|
process.wait(timeout=3)
|
|
assert not group_members(process.pid), "Official reference child survived cleanup"
|
|
assert stopped and not remaining, "Official reference required forced cleanup"
|
|
|
|
|
|
@contextmanager
|
|
def authenticated_endpoint(
|
|
target: str,
|
|
key: str | None,
|
|
alias: str | None,
|
|
negotiations: queue.Queue[tuple[str, str]] | None = None,
|
|
) -> Iterator[str]:
|
|
from integration._support.asgi import asgi_server
|
|
from starlette.applications import Starlette
|
|
from starlette.requests import Request
|
|
from starlette.responses import StreamingResponse
|
|
from starlette.routing import Mount
|
|
from starlette.types import Receive, Scope, Send
|
|
|
|
async def forward(scope: Scope, receive: Receive, send: Send) -> None:
|
|
request: Final = Request(scope, receive)
|
|
original: Final = await request.body()
|
|
body: Final = prefixed_request(original, alias) if alias is not None else original
|
|
headers: Final = tuple(
|
|
(name, value)
|
|
for name, value in request.headers.raw
|
|
if (key is None or name.lower() != b"authorization")
|
|
and (name.lower() != b"content-length" or body == original)
|
|
) + (((b"authorization", f"Bearer {key}".encode()),) if key is not None else ())
|
|
async with httpx.AsyncClient(timeout=35, trust_env=False) as client:
|
|
async with client.stream(request.method, target, headers=headers, content=body) as response:
|
|
|
|
async def observed_chunks() -> AsyncIterator[bytes]:
|
|
capture: Final = (
|
|
negotiations is not None and response.status_code == 200 and is_initialization(body)
|
|
)
|
|
chunks: Final[list[bytes]] = []
|
|
async for chunk in response.aiter_raw():
|
|
if capture:
|
|
chunks.append(chunk)
|
|
yield chunk
|
|
if capture and negotiations is not None:
|
|
negotiations.put(read_negotiation(body, b"".join(chunks), response.headers["content-type"]))
|
|
|
|
streamed: Final = StreamingResponse(observed_chunks(), status_code=response.status_code)
|
|
streamed.raw_headers = [
|
|
(name, value)
|
|
for name, value in response.headers.raw
|
|
if name.lower() not in (b"transfer-encoding", b"connection")
|
|
]
|
|
await streamed(scope, receive, send)
|
|
|
|
with asgi_server(Starlette(routes=[Mount("/mcp", app=forward)])) as url:
|
|
yield url + "/mcp/"
|
|
|
|
|
|
def prefixed_request(body: bytes, alias: str) -> bytes:
|
|
try:
|
|
request: Final = json.loads(body)
|
|
except (ValueError, UnicodeDecodeError):
|
|
return body
|
|
if not isinstance(request, dict) or request.get("method") not in ("tools/call", "prompts/get"):
|
|
return body
|
|
params: Final = request.get("params")
|
|
if not isinstance(params, dict) or not isinstance(params.get("name"), str):
|
|
return body
|
|
return json.dumps({**request, "params": {**params, "name": f"{alias}-{params['name']}"}}).encode()
|
|
|
|
|
|
def require_negotiations(expected: str, observed: tuple[tuple[str, str], ...]) -> None:
|
|
assert observed and all(pair == (expected, expected) for pair in observed), (
|
|
f"conformance negotiation: expected {expected} on the wire, observed {observed}"
|
|
)
|
|
|
|
|
|
class ExecutionReport(BaseModel):
|
|
collected: tuple[str, ...]
|
|
passed: tuple[str, ...]
|
|
skipped: tuple[str, ...]
|
|
complete: bool
|
|
|
|
|
|
def require_passes(directory: Path, required: tuple[str, ...]) -> None:
|
|
report: Final = ExecutionReport.model_validate_json((directory / "execution.json").read_bytes())
|
|
assert required and report.complete, "conformance execution was incomplete"
|
|
assert set(required) <= set(report.collected), "conformance cases were not collected"
|
|
assert set(required) <= set(report.passed), "conformance cases did not pass"
|
|
assert not set(required).intersection(report.skipped), "conformance cases were skipped"
|
|
|
|
|
|
def translation_cases() -> tuple[tuple[str, str, str, str], ...]:
|
|
from litellm.proxy._experimental.mcp_server.capabilities import REVISION_SUPPORT, TRANSLATION_PAIRS
|
|
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPTransport
|
|
|
|
completed: Final = {revision for revision, support in REVISION_SUPPORT.items() if support.completed}
|
|
assert completed == set(MCP_LEGACY_VERSIONS), "New revisions need conformance cases before activation"
|
|
assert TRANSLATION_PAIRS == frozenset(product(completed, repeat=2)), "Declared translation coverage changed"
|
|
return tuple(
|
|
(downstream, upstream, peer, ingress)
|
|
for (downstream, upstream), peer, ingress in product(
|
|
sorted(TRANSLATION_PAIRS), ("http", "sse", "stdio"), ("http", "sse")
|
|
)
|
|
if MCPTransport(peer) in REVISION_SUPPORT[upstream].transports
|
|
and MCPTransport(ingress) in REVISION_SUPPORT[downstream].transports
|
|
)
|
|
|
|
|
|
# Applicable legacy server scenarios for the operations advertised by capabilities.py.
|
|
# The simple-text runner/reference mismatch has an explicit SDK gap case instead.
|
|
OFFICIAL_SCENARIOS: Final = (
|
|
"server-initialize",
|
|
"server-sse-multiple-streams",
|
|
"ping",
|
|
"tools-list",
|
|
"tools-call-image",
|
|
"tools-call-audio",
|
|
"tools-call-embedded-resource",
|
|
"tools-call-mixed-content",
|
|
"tools-call-error",
|
|
"tools-call-with-progress",
|
|
"resources-list",
|
|
"resources-read-text",
|
|
"resources-read-binary",
|
|
"resources-templates-read",
|
|
"prompts-list",
|
|
"prompts-get-simple",
|
|
"prompts-get-with-args",
|
|
"prompts-get-embedded-resource",
|
|
"prompts-get-with-image",
|
|
)
|
|
|
|
|
|
def official_cases() -> tuple[tuple[str, str], ...]:
|
|
from litellm.types.mcp import MCP_LEGACY_VERSIONS
|
|
|
|
return tuple(product(OFFICIAL_SCENARIOS, MCP_LEGACY_VERSIONS))
|
|
|
|
|
|
def required_conformance_nodes() -> tuple[str, ...]:
|
|
official: Final = tuple(
|
|
"tests/integration/mcp/test_mcp_official_conformance.py::test_official_scenario_through_gateway["
|
|
+ "-".join(case)
|
|
+ "]"
|
|
for case in official_cases()
|
|
)
|
|
matrix: Final = tuple(
|
|
"tests/integration/mcp/test_mcp_transports.py::test_pinned_revision_pairs_list_and_call_through_gateway["
|
|
+ "-".join(case)
|
|
+ "]"
|
|
for case in translation_cases()
|
|
)
|
|
return (
|
|
official
|
|
+ matrix
|
|
+ ("tests/integration/mcp/test_mcp_protocol_errors.py::test_omitted_tool_arguments_reach_the_upstream",)
|
|
+ tuple(
|
|
"tests/integration/mcp/test_mcp_protocol_errors.py::test_configured_origin_policy_rejects_before_tool_execution["
|
|
+ ingress
|
|
+ "]"
|
|
for ingress in ("server_mcp", "sse")
|
|
)
|
|
+ tuple(
|
|
"tests/integration/mcp/test_mcp_official_conformance.py::" + name
|
|
for name in (
|
|
"test_conformance_bridge_preserves_headers_payload_and_error_status[200]",
|
|
"test_conformance_bridge_preserves_headers_payload_and_error_status[403]",
|
|
"test_official_runner_rejects_unknown_scenario",
|
|
"test_official_gateway_session_lifecycle",
|
|
"test_stalled_reference_is_killed_and_cannot_report_clean_teardown",
|
|
"test_reference_children_are_stopped_after_the_root_exits",
|
|
)
|
|
)
|
|
)
|
|
|
|
|
|
def read_negotiation(request: bytes, response: bytes, content_type: str) -> tuple[str, str]:
|
|
from httpx_sse import EventSource
|
|
from mcp.types import InitializeRequestParams, InitializeResult
|
|
|
|
body: Final = json.loads(request)
|
|
received: Final = httpx.Response(200, content=response, headers={"content-type": content_type})
|
|
messages: Final = (
|
|
tuple(event.json() for event in EventSource(received).iter_sse() if event.data)
|
|
if content_type.startswith("text/event-stream")
|
|
else (received.json(),)
|
|
)
|
|
replies: Final = tuple(
|
|
message["result"] for message in messages if message.get("id") == body["id"] and "result" in message
|
|
)
|
|
assert len(replies) == 1, "conformance negotiation: missing or duplicated initialize response"
|
|
return InitializeRequestParams.model_validate(body["params"]).protocol_version, InitializeResult.model_validate(
|
|
replies[0]
|
|
).protocol_version
|
|
|
|
|
|
def is_initialization(body: bytes) -> bool:
|
|
try:
|
|
parsed: Final = json.loads(body)
|
|
except (ValueError, UnicodeDecodeError):
|
|
return False
|
|
return isinstance(parsed, dict) and parsed.get("method") == "initialize"
|