mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
test(mcp): use advertised fixture names without warming discovery
This commit is contained in:
parent
d77751f95b
commit
9d25f89408
3 changed files with 48 additions and 5 deletions
|
|
@ -1,4 +1,5 @@
|
|||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import signal
|
||||
|
|
@ -115,7 +116,7 @@ def reference_server(root: Path, directory: Path, port: int) -> Iterator["McpPee
|
|||
|
||||
|
||||
@contextmanager
|
||||
def authenticated_endpoint(target: str, key: str) -> Iterator[str]:
|
||||
def authenticated_endpoint(target: str, key: str, alias: str) -> Iterator[str]:
|
||||
from integration._support.asgi import asgi_server
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
|
|
@ -125,10 +126,13 @@ def authenticated_endpoint(target: str, key: str) -> Iterator[str]:
|
|||
|
||||
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)
|
||||
headers: Final = tuple(
|
||||
(name, value) for name, value in request.headers.raw if name.lower() != b"authorization"
|
||||
(name, value)
|
||||
for name, value in request.headers.raw
|
||||
if name.lower() != b"authorization" and (name.lower() != b"content-length" or body == original)
|
||||
) + ((b"authorization", f"Bearer {key}".encode()),)
|
||||
body: Final = await request.body()
|
||||
async with httpx.AsyncClient(timeout=35, trust_env=False) as client:
|
||||
async with client.stream(request.method, target, headers=headers, content=body) as response:
|
||||
streamed: Final = StreamingResponse(response.aiter_raw(), status_code=response.status_code)
|
||||
|
|
@ -141,3 +145,16 @@ def authenticated_endpoint(target: str, key: str) -> Iterator[str]:
|
|||
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ def test_official_scenario_through_gateway(gateway: Gateway, tmp_path: Path, unu
|
|||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
direct: Final = run_scenario(root, reference.url, name, output / "direct")
|
||||
endpoint: Final = str(gateway.client.base_url).rstrip("/") + f"/{alias}/mcp"
|
||||
with authenticated_endpoint(endpoint, key) as authenticated:
|
||||
with authenticated_endpoint(endpoint, key, alias) as authenticated:
|
||||
proxied: Final = run_scenario(root, authenticated, name, output / "gateway")
|
||||
assert tuple(check.status for check in direct if check.id == name) == ("SUCCESS",)
|
||||
assert tuple(check.status for check in proxied if check.id == name) == ("SUCCESS",)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from pathlib import Path
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from tests.integration._support.conformance import read_checks, verify_archive
|
||||
from tests.integration._support.conformance import read_checks, verify_archive, prefixed_request
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
||||
|
|
@ -97,3 +97,29 @@ def test_official_text_success_requires_the_reference_payload(tmp_path: Path) ->
|
|||
)
|
||||
)
|
||||
assert read_checks(tmp_path, "tools-call-simple-text")[0].status == "SUCCESS"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("method", ("tools/call", "prompts/get"))
|
||||
def test_only_fixture_name_is_translated_to_advertised_name(method: str) -> None:
|
||||
request: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 4,
|
||||
"method": method,
|
||||
"params": {"name": "test_tool", "_meta": {"progressToken": "keep"}},
|
||||
}
|
||||
changed: Final = json.loads(prefixed_request(json.dumps(request).encode(), "official"))
|
||||
assert changed == {**request, "params": {"name": "official-test_tool", "_meta": {"progressToken": "keep"}}}
|
||||
assert "arguments" not in changed["params"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
(
|
||||
b"malformed JSON",
|
||||
b"[]",
|
||||
b'{"method":"initialize","params":{"protocolVersion":"2025-03-26"}}',
|
||||
b'{"method":"tools/call","params":{"name":17}}',
|
||||
),
|
||||
)
|
||||
def test_name_integration_preserves_other_requests(body: bytes) -> None:
|
||||
assert prefixed_request(body, "official") == body
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue