mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/main' into litellm_shared_rust_pagination
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.test.tsx # ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.tsx # ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx # ui/litellm-dashboard/src/components/lens/traces/list/readFailure.test.ts # ui/litellm-dashboard/src/components/lens/traces/list/readFailure.ts # ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.ts
This commit is contained in:
commit
cc3a1d584d
237 changed files with 6130 additions and 817 deletions
|
|
@ -59,7 +59,7 @@ def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tupl
|
|||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
upstream_model: Final = model.removeprefix(f"{provider}/") if model.startswith(f"{provider}/") else model
|
||||
upstream_model: Final = model.removeprefix(f"{provider}/")
|
||||
if not upstream_model:
|
||||
raise litellm.BadRequestError(
|
||||
message="A model name is required for the Decisions API",
|
||||
|
|
|
|||
|
|
@ -267,11 +267,8 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list
|
||||
list(self._tool_specs)
|
||||
)
|
||||
request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}
|
||||
}
|
||||
kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
**request_kwargs,
|
||||
**{key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}},
|
||||
"messages": messages,
|
||||
**({"tools": tool_specs} if tool_specs else {}),
|
||||
}
|
||||
|
|
@ -287,8 +284,7 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
yield Text(content)
|
||||
tool_calls = message.tool_calls or ()
|
||||
if not tool_calls:
|
||||
final_text = content or ""
|
||||
ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.final_text = content or "" # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink
|
||||
final_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ class DeepAgentsOptions:
|
|||
|
||||
@dataclass(frozen=True)
|
||||
class ToolLoopOptions:
|
||||
completion_kwargs: Mapping[str, Any] = field(default_factory=dict)
|
||||
completion_kwargs: Mapping[str, object] = field(default_factory=dict)
|
||||
|
||||
|
||||
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions | ToolLoopOptions
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""``CustomLogger`` adapter on the OpenTelemetry span engine."""
|
||||
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager, nullcontext
|
||||
|
|
@ -966,10 +967,12 @@ def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2")
|
|||
|
||||
|
||||
def _registered_v2_logger() -> "OpenTelemetryV2 | None":
|
||||
try:
|
||||
from litellm.proxy import proxy_server
|
||||
except Exception:
|
||||
return None
|
||||
"""The proxy's registered V2 logger, read without importing the proxy.
|
||||
|
||||
Request paths call this (the router's ``route`` phase among them), so importing
|
||||
``proxy_server`` here would load the whole proxy on an SDK caller's event loop.
|
||||
"""
|
||||
proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server")
|
||||
logger: Final = getattr(proxy_server, "open_telemetry_logger", None)
|
||||
return logger if isinstance(logger, OpenTelemetryV2) else None
|
||||
|
||||
|
|
|
|||
|
|
@ -125,13 +125,12 @@ class AnthropicFilesConfig(BaseFilesConfig):
|
|||
return self._finalize_headers(headers, auth_header)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_params(
|
||||
litellm_params: dict, api_base: str | None
|
||||
) -> tuple[dict | None, str | None]: # mutable-ok: mirrors the sync validate_environment contract this overrides
|
||||
def _resolve_params(litellm_params: dict, api_base: str | None) -> tuple[Mapping[str, object] | None, str | None]:
|
||||
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
|
||||
if api_base is None and params_mapping is not None:
|
||||
api_base = params_mapping.get("api_base")
|
||||
return params_mapping, api_base
|
||||
resolved_api_base: Final = (
|
||||
api_base if api_base is not None or params_mapping is None else params_mapping.get("api_base")
|
||||
)
|
||||
return params_mapping, resolved_api_base
|
||||
|
||||
@staticmethod
|
||||
def _finalize_headers(headers: dict, auth_header: Mapping[str, str] | None) -> dict: # mutable-ok: out-param
|
||||
|
|
|
|||
|
|
@ -147,6 +147,7 @@ async def anthropic_messages_with_mcp(
|
|||
|
||||
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=list(tool_use_blocks),
|
||||
user_api_key_auth=context.user_api_key_auth,
|
||||
mcp_auth_header=context.mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from contextlib import asynccontextmanager
|
|||
from dataclasses import dataclass, replace
|
||||
from functools import lru_cache
|
||||
from itertools import chain, groupby
|
||||
from types import MappingProxyType
|
||||
from types import EllipsisType, MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
|
|
@ -259,6 +259,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]]
|
|||
_USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds
|
||||
_USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth
|
||||
|
||||
_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]]
|
||||
_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ListedToolsCaller:
|
||||
"""Request inputs that select which upstream catalog a caller was shown by tools/list."""
|
||||
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None
|
||||
mcp_auth_header: str | dict[str, str] | None = None
|
||||
raw_headers: Mapping[str, str] | None = None
|
||||
oauth2_headers: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
# Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the
|
||||
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
|
||||
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
|
||||
|
|
@ -1172,6 +1186,57 @@ def _authorization_is_litellm_admission_credential(
|
|||
return bool(user_api_key_auth and user_api_key_auth.api_key and not admission_header)
|
||||
|
||||
|
||||
def _server_auth_header_for(
|
||||
server: MCPServer,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""Server-specific ``x-mcp-<alias>-authorization`` header, else the deprecated global one."""
|
||||
server_specific: Final = (
|
||||
lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
if mcp_server_auth_headers
|
||||
else None
|
||||
)
|
||||
return mcp_auth_header if server_specific is None else server_specific
|
||||
|
||||
|
||||
def listed_tools_caller_for(
|
||||
server: MCPServer,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
raw_headers: Mapping[str, str] | None,
|
||||
oauth2_headers: Mapping[str, str] | None,
|
||||
) -> ListedToolsCaller:
|
||||
"""The caller a tools/call must look its listed entry up under: the same inputs tools/list keyed by."""
|
||||
return ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
|
||||
|
||||
def _admission_identity(
|
||||
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
|
||||
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
|
||||
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
|
||||
caller admitted with neither a key nor a user."""
|
||||
keyless: Final = auth.api_key is None and auth.user_id is None
|
||||
credential: Final = (
|
||||
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
|
||||
if keyless
|
||||
else None
|
||||
)
|
||||
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
|
|
@ -1297,6 +1362,16 @@ async def _resolve_byok_mcp_auth_header(
|
|||
return mcp_auth_header
|
||||
|
||||
|
||||
def _catalog_auth_header(
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""The header the client supplied, which keys the caller's catalog slot on both tools/list and
|
||||
tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes
|
||||
the client's value explicitly, since the stored credential must never be read to find the slot."""
|
||||
return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
|
|
@ -1925,6 +2000,8 @@ class MCPServerManager:
|
|||
"gmail_send_email": "zapier_mcp_server",
|
||||
}
|
||||
"""
|
||||
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
|
||||
self._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
|
||||
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
|
||||
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
|
||||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
|
|
@ -3242,6 +3319,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Added MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3279,6 +3358,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3754,25 +3835,14 @@ class MCPServerManager:
|
|||
verbose_logger.warning("MCP Server %s not found", server_id)
|
||||
return []
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header: str | dict[str, str] | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
try:
|
||||
tools: Final = await self._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
record_listing=True,
|
||||
)
|
||||
return tools
|
||||
except Exception as e:
|
||||
|
|
@ -3852,7 +3922,7 @@ class MCPServerManager:
|
|||
def _build_stdio_env(
|
||||
self,
|
||||
server: MCPServer,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""Resolve stdio env values, supporting header-driven placeholders."""
|
||||
|
||||
|
|
@ -4147,13 +4217,20 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> list[MCPTool]:
|
||||
*,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
|
||||
record_listing: bool = False,
|
||||
) -> Sequence[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
||||
Args:
|
||||
server (MCPServer): The server to query tools from
|
||||
mcp_auth_header: Optional auth header for MCP server
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header``
|
||||
record_listing: Record the served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
|
|
@ -4169,6 +4246,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("_get_tools_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
listed_caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0)
|
||||
|
||||
try:
|
||||
# Tool *listing* must not be blocked by missing per-user env vars —
|
||||
|
|
@ -4266,8 +4350,12 @@ class MCPServerManager:
|
|||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
unprefixed_tools: Final = guarded_openapi
|
||||
self._record_listed_tools(
|
||||
server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
if not add_prefix:
|
||||
return list(guarded_openapi)
|
||||
return unprefixed_tools
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
|
|
@ -4281,7 +4369,10 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
list(guarded_tools), server, add_prefix=add_prefix
|
||||
guarded_tools, server, add_prefix=add_prefix
|
||||
)
|
||||
self._record_listed_tools(
|
||||
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
|
||||
return prefixed_or_original_tools
|
||||
|
|
@ -4332,8 +4423,99 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._drop_listed_tools(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _drop_listed_tools(self, server_id: str) -> None:
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1
|
||||
|
||||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
"""Key the listed-tool cache by every request input that can change the served catalog.
|
||||
|
||||
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
|
||||
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
|
||||
``_admission_identity``: the hashed key, user, team and organization, plus the admission
|
||||
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
|
||||
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
|
||||
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
|
||||
server-specific auth header also reach upstream and split the slot further. Only unkeyed
|
||||
listings with none of those share the ``None`` slot.
|
||||
"""
|
||||
if caller is None:
|
||||
return None
|
||||
auth: Final = caller.user_api_key_auth
|
||||
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
|
||||
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
|
||||
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
|
||||
caller_bearer: Final = (
|
||||
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
|
||||
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else None
|
||||
)
|
||||
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
|
||||
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
|
||||
return None
|
||||
material: Final = json.dumps(
|
||||
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _forwarded_header_values(
|
||||
server: MCPServer, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
if not raw_headers or not server.extra_headers:
|
||||
return ()
|
||||
forwarded_names: Final = frozenset(name.lower() for name in server.extra_headers)
|
||||
return tuple(
|
||||
sorted((name.lower(), value) for name, value in raw_headers.items() if name.lower() in forwarded_names)
|
||||
)
|
||||
|
||||
def listed_tools_generation(self, server_id: str) -> int:
|
||||
return self._listed_tools_generations.get(server_id, 0)
|
||||
|
||||
def record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing)
|
||||
|
||||
def _record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int | None = None,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
|
||||
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
|
||||
if not record_listing:
|
||||
return
|
||||
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
|
||||
return
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
|
||||
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
|
||||
shared: Final = existing.get(None)
|
||||
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
|
||||
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
|
||||
entries: Final = (
|
||||
*(() if shared is None else ((None, shared),)),
|
||||
*callers[evicted:],
|
||||
(identity, listing),
|
||||
)
|
||||
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries))
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4343,9 +4525,11 @@ class MCPServerManager:
|
|||
stdio_env: dict[str, str] | None,
|
||||
subject_token: str | None,
|
||||
credential_fingerprint: str | None = None,
|
||||
per_caller: bool = False,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
server.requires_per_user_auth
|
||||
per_caller
|
||||
or server.requires_per_user_auth
|
||||
or self._references_per_user_env_var(server)
|
||||
or server.delegate_auth_to_upstream
|
||||
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
|
||||
|
|
@ -5253,7 +5437,12 @@ class MCPServerManager:
|
|||
{seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key}
|
||||
)
|
||||
|
||||
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
|
||||
def _create_prefixed_tools(
|
||||
self,
|
||||
tools: Sequence[MCPTool],
|
||||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Create prefixed tools and update tool mapping.
|
||||
|
||||
|
|
@ -5281,6 +5470,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
return prefixed_tools
|
||||
|
||||
def get_listed_tool(self, server: MCPServer, name: str, caller: ListedToolsCaller | None = None) -> MCPTool | None:
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity)
|
||||
if not listed:
|
||||
return None
|
||||
return listed.get(name)
|
||||
|
||||
def _create_prefixed_prompts(
|
||||
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
|
||||
) -> list[Prompt]:
|
||||
|
|
@ -5514,6 +5710,7 @@ class MCPServerManager:
|
|||
raw_headers: dict[str, str] | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
tool: MCPTool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Run pre-call checks and guardrail hooks for an MCP tool call.
|
||||
|
|
@ -5527,6 +5724,9 @@ class MCPServerManager:
|
|||
``pre_mcp_call`` evaluation (or a block) on the spend-log row the Guardrails
|
||||
Monitor counts. It stays optional so callers that do no logging are unchanged.
|
||||
|
||||
``tool`` is the upstream tool definition when one was listed, so guardrails
|
||||
can see its description and input schema, not just the name and arguments.
|
||||
|
||||
Returns a dict that may contain:
|
||||
- "arguments": hook-modified tool arguments (only if changed)
|
||||
- "extra_headers": headers injected by pre_mcp_call guardrail hooks
|
||||
|
|
@ -5590,6 +5790,8 @@ class MCPServerManager:
|
|||
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"headers": logging_safe_mcp_headers(raw_headers),
|
||||
"tool_description": tool.description if tool is not None else None,
|
||||
"tool_input_schema": tool.input_schema if tool is not None else None,
|
||||
}
|
||||
|
||||
# Create MCP request object for processing
|
||||
|
|
@ -5805,21 +6007,7 @@ class MCPServerManager:
|
|||
GuardrailRaisedException: If guardrails block the call
|
||||
HTTPException: If an HTTP error occurs
|
||||
"""
|
||||
# Get server-specific auth header if available (case-insensitive)
|
||||
# FIX: Added case-insensitive matching to handle auth header keys that may not match
|
||||
# the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway')
|
||||
server_auth_header: dict[str, str] | str | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=mcp_server.alias,
|
||||
server_name=mcp_server.server_name,
|
||||
access_groups=mcp_server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
# Extract subject token for OAuth2 Token Exchange (OBO) and ID-JAG flows
|
||||
subject_token: str | None = None
|
||||
|
|
@ -6242,6 +6430,9 @@ class MCPServerManager:
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None | EllipsisType = ...,
|
||||
listed_tool: MCPTool | None | EllipsisType = ...,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -6253,6 +6444,8 @@ class MCPServerManager:
|
|||
user_api_key_auth: User authentication
|
||||
mcp_auth_header: MCP auth header (deprecated)
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header`` as received, before BYOK resolution
|
||||
proxy_logging_obj: Optional ProxyLogging object for hook integration
|
||||
litellm_logging_obj: Optional request logger the guardrail hooks record
|
||||
their evaluations onto, so MCP guardrail activity reaches the
|
||||
|
|
@ -6264,6 +6457,7 @@ class MCPServerManager:
|
|||
"""
|
||||
start_time: Final = datetime.datetime.now()
|
||||
mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name)
|
||||
client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header)
|
||||
|
||||
# Resolved before any hook runs so a missing BYOK credential (401) never
|
||||
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
|
||||
|
|
@ -6273,6 +6467,9 @@ class MCPServerManager:
|
|||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
)
|
||||
listed_caller: Final = listed_tools_caller_for(
|
||||
mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Pre MCP Tool Call Hook
|
||||
|
|
@ -6289,6 +6486,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
|
|
|||
|
|
@ -81,12 +81,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
_resolve_openapi_tool_auth,
|
||||
_should_strip_caller_authorization,
|
||||
global_mcp_server_manager,
|
||||
listed_tools_caller_for,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
|
|
@ -954,6 +956,8 @@ async def _get_tools_from_mcp_servers(
|
|||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -964,6 +968,8 @@ async def _get_tools_from_mcp_servers(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers
|
||||
oauth2_headers: Optional dict of oauth2 headers
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from filtered servers plus each server's
|
||||
|
|
@ -1111,12 +1117,14 @@ async def _get_tools_from_mcp_servers(
|
|||
prefetched_creds=_prefetched_oauth_creds,
|
||||
)
|
||||
|
||||
catalog_auth_header: Final = server_auth_header
|
||||
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
|
||||
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1127,6 +1135,8 @@ async def _get_tools_from_mcp_servers(
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=False,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1135,6 +1145,21 @@ async def _get_tools_from_mcp_servers(
|
|||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
global_mcp_server_manager.record_listed_tools(
|
||||
server,
|
||||
[
|
||||
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
|
||||
for tool in filtered_tools
|
||||
],
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=catalog_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
),
|
||||
listed_generation,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
|
@ -1458,6 +1483,8 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -1468,6 +1495,8 @@ async def _list_mcp_tools(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
client_ip: Client IP for IP-based server access control
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from all accessible servers plus each server's
|
||||
|
|
@ -1486,6 +1515,7 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source=list_tools_log_source,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
|
||||
return listing
|
||||
|
|
@ -1798,6 +1828,7 @@ async def _list_tools_before_first_call(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
record_listing=False,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
|
||||
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
|
||||
|
|
@ -2001,6 +2032,7 @@ async def _execute_mcp_tool(
|
|||
if mcp_server is None:
|
||||
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
|
||||
client_auth_header: Final = mcp_auth_header
|
||||
if mcp_server:
|
||||
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info")
|
||||
if litellm_logging_obj:
|
||||
|
|
@ -2071,6 +2103,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
# `pre_call_tool_check` may return guardrail-modified
|
||||
# arguments; honor them on the local path too.
|
||||
|
|
@ -2120,6 +2164,7 @@ async def _execute_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=client_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2139,7 +2184,8 @@ async def _execute_mcp_tool(
|
|||
# not in the registry either, `_handle_local_mcp_tool` below reports
|
||||
# 404 and nothing runs, so demanding a server here would turn every
|
||||
# unknown tool name into a misleading 503.
|
||||
if global_mcp_tool_registry.get_tool(original_tool_name) is not None:
|
||||
registered_local_tool: Final = global_mcp_tool_registry.get_tool(original_tool_name)
|
||||
if registered_local_tool is not None:
|
||||
# `mcp_server` is None here because the tool name is not in the
|
||||
# tool -> server mapping, but the name still carries a prefix
|
||||
# that the server-level check above compared against the
|
||||
|
|
@ -2181,6 +2227,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
prefix_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
prefix_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
|
||||
|
|
@ -2583,8 +2641,12 @@ async def _handle_managed_mcp_tool(
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
"""Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client
|
||||
supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved
|
||||
BYOK credential."""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
|
|
@ -2594,6 +2656,7 @@ async def _handle_managed_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2712,6 +2775,7 @@ async def _execute_handle_list_tools(
|
|||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
record_listing=True,
|
||||
)
|
||||
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
|
||||
if not listing.outcomes:
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
ListedToolsCaller,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -701,6 +702,8 @@ if MCP_AVAILABLE:
|
|||
extra_headers: dict[str, str] | None,
|
||||
client_ip: str | None,
|
||||
proxy_logging_obj: "ProxyLogging | None",
|
||||
*,
|
||||
record_listing: bool,
|
||||
) -> list[MCPTool]:
|
||||
return await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
|
|
@ -711,11 +714,12 @@ if MCP_AVAILABLE:
|
|||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
server: MCPServer,
|
||||
server_auth_header: dict[str, str] | str | None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -731,31 +735,44 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tools = await _list_server_tools(
|
||||
server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj
|
||||
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
|
||||
tools: Final = await _list_server_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers,
|
||||
user_api_key_auth,
|
||||
extra_headers,
|
||||
client_ip,
|
||||
proxy_logging_obj,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
return _create_tool_response_objects(tools, server)
|
||||
|
||||
# Always apply allowed_tools/disallowed_tools so the blacklist is
|
||||
# enforced even when no allowlist is set (matches the SSE/HTTP path).
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Filter by the key's effective tool permissions through the same
|
||||
# function the MCP protocol path uses (direct grants, toolset grants,
|
||||
# and team/agent/org ceilings), so REST listing cannot drift from it.
|
||||
# Entries here are tool names on one server, written bare by every
|
||||
# writer, and dispatch compares them bare; matching a wider set of
|
||||
# spellings would advertise a tool that tools/call then refuses
|
||||
if user_api_key_auth:
|
||||
tools = await filter_tools_by_key_team_permissions(
|
||||
tools=tools,
|
||||
server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools
|
||||
served_tools: Final = (
|
||||
await filter_tools_by_key_team_permissions(
|
||||
tools=server_filtered,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if apply_tool_filters and user_api_key_auth
|
||||
else server_filtered
|
||||
)
|
||||
if apply_tool_filters:
|
||||
# Only a listing shaped for the caller's runtime view may set their
|
||||
# listed-tools slot; the admin-only unfiltered configuration view
|
||||
# must not warm it.
|
||||
global_mcp_server_manager.record_listed_tools(
|
||||
server,
|
||||
served_tools,
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
),
|
||||
listed_generation,
|
||||
)
|
||||
|
||||
return _create_tool_response_objects(tools, server)
|
||||
return _create_tool_response_objects(served_tools, server)
|
||||
|
||||
async def fetch_pinnable_tool_catalog(
|
||||
server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth
|
||||
|
|
@ -774,6 +791,7 @@ if MCP_AVAILABLE:
|
|||
await _get_user_oauth_extra_headers(server, user_api_key_dict),
|
||||
IPAddressUtils.get_mcp_client_ip(request),
|
||||
None,
|
||||
record_listing=False,
|
||||
)
|
||||
scan: Final = await scan_tool_descriptions(
|
||||
apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
|
|||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
|
|||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
|
||||
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
|
||||
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_OBO_CACHE_MAX_ENTRIES: Final = 1000
|
||||
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
|
||||
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
|
||||
|
|
@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
|
|||
return ()
|
||||
|
||||
|
||||
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def entra_assertion(value: object) -> str | None:
|
||||
"""``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion.
|
||||
A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``."""
|
||||
|
|
@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False):
|
|||
correlationId: ReadOnly[str]
|
||||
|
||||
|
||||
class _ToolReference(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
name: str
|
||||
description: str | None = None
|
||||
input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema")
|
||||
|
||||
|
||||
class _UnavailableDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
|
@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
arguments: Final = data.get("mcp_arguments")
|
||||
server_name: Final = str(data.get("mcp_server_name") or "litellm")
|
||||
agent_id: Final = user_api_key_dict.key_alias
|
||||
description: Final = data.get("mcp_tool_description")
|
||||
tool_reference: Final = _ToolReference(
|
||||
name=tool_name,
|
||||
description=description if isinstance(description, str) and description else None,
|
||||
input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")),
|
||||
)
|
||||
payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below
|
||||
"tool": {"name": tool_name},
|
||||
"tool": tool_reference.model_dump(by_alias=True, exclude_none=True),
|
||||
"serverName": server_name,
|
||||
"conversationId": self._resolve_conversation_id(data),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1018,6 +1018,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
record_listing=True,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools]
|
||||
|
|
|
|||
|
|
@ -1151,6 +1151,8 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool:
|
|||
|
||||
|
||||
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
|
||||
_MCP_TOOL_DESCRIPTION: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
|
||||
_MCP_TOOL_INPUT_SCHEMA: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1475,9 +1477,14 @@ class ProxyLogging:
|
|||
TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({}))
|
||||
)
|
||||
|
||||
mcp_tool_description: Final = kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = kwargs.get("mcp_input_schema")
|
||||
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
|
||||
mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = (
|
||||
request_obj.tool_input_schema
|
||||
if request_obj.tool_input_schema is not None
|
||||
else kwargs.get("mcp_input_schema")
|
||||
)
|
||||
listing_description: Final = kwargs.get("mcp_tool_description")
|
||||
description_line: Final = f"\nDescription: {listing_description}" if listing_description else ""
|
||||
tool_call_content: Final = (
|
||||
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
|
|
@ -1735,6 +1742,8 @@ class ProxyLogging:
|
|||
tool_name=kwargs.get("name", ""),
|
||||
arguments=kwargs.get("arguments", {}),
|
||||
server_name=kwargs.get("server_name"),
|
||||
tool_description=_MCP_TOOL_DESCRIPTION.validate_python(kwargs.get("tool_description")),
|
||||
tool_input_schema=_MCP_TOOL_INPUT_SCHEMA.validate_python(kwargs.get("tool_input_schema")),
|
||||
user_api_key_auth=user_api_key_auth_dict,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -286,6 +286,7 @@ async def aresponses_api_with_mcp(
|
|||
call_params=call_params,
|
||||
previous_response_id=previous_response_id,
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
**kwargs,
|
||||
)
|
||||
await mcp_streaming_response._create_initial_response_iterator()
|
||||
|
|
@ -339,6 +340,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -395,6 +397,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
final_response = MCPEnhancedStreamingIterator(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
base_iterator=final_response,
|
||||
mcp_events=tool_execution_events,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -435,6 +435,7 @@ async def acompletion_with_mcp(
|
|||
# Execute tool calls
|
||||
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=self.tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
@ -609,6 +610,7 @@ async def acompletion_with_mcp(
|
|||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
tool_calls=tool_calls,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
|
|||
|
|
@ -696,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -860,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
listed_tool=(
|
||||
next((tool for tool in served_tools if tool.name == tool_name), None)
|
||||
if served_tools is not None
|
||||
else ...
|
||||
),
|
||||
)
|
||||
|
||||
if proxy_logging_obj:
|
||||
|
|
@ -1152,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
"""
|
||||
|
|
@ -1181,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
base_iterator=None, # Will be created internally
|
||||
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=served_tools,
|
||||
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
|
||||
user_api_key_auth=kwargs.get("user_api_key_auth")
|
||||
or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"),
|
||||
|
|
|
|||
|
|
@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
original_request_params: dict[str, Any] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
):
|
||||
# MCP setup
|
||||
self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or []
|
||||
|
|
@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.mcp_discovery_generated = True # Events are already generated
|
||||
self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility
|
||||
self.tool_server_map = tool_server_map
|
||||
self.served_tools = tuple(served_tools) if served_tools is not None else None
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
|
|
@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
# Execute the tools
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=self.served_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -429,6 +429,8 @@ class MCPPreCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
user_api_key_auth: dict[str, Any] | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
@ -452,6 +454,8 @@ class MCPDuringCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
start_time: float | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
|
|||
|
|
@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object]
|
|||
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:
|
||||
|
|
@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]:
|
|||
},
|
||||
)
|
||||
if method == "tools/list":
|
||||
return jsonrpc_reply(
|
||||
identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]}
|
||||
)
|
||||
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"])
|
||||
|
|
|
|||
|
|
@ -133,6 +133,7 @@ def wire_server(
|
|||
|
||||
class OwnedHTTPServer(ThreadingHTTPServer):
|
||||
daemon_threads = False
|
||||
request_queue_size = 128
|
||||
|
||||
def server_bind(self) -> None:
|
||||
super().server_bind()
|
||||
|
|
|
|||
|
|
@ -1,20 +1,29 @@
|
|||
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.database import read_rows
|
||||
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"}
|
||||
|
|
@ -135,3 +144,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, {"KEEPALIVE_TIMEOUT": "120"}, 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, keepalive_expiry=120)
|
||||
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,21 +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 = (
|
||||
|
|
@ -192,3 +203,288 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w
|
|||
assert removed.status_code in (200, 204), removed.text
|
||||
eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def _listings(peer: McpPeer) -> tuple[dict[str, object], ...]:
|
||||
return tuple(
|
||||
item
|
||||
for item in peer.drain()
|
||||
if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list"
|
||||
)
|
||||
|
||||
|
||||
def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, oauth_server() as auth, gateway.scenario() as scenario:
|
||||
alias: Final = "cc" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2",
|
||||
oauth2_flow="client_credentials",
|
||||
is_byok=True,
|
||||
token_url=auth.issuer + "/token",
|
||||
credentials={"client_id": "cc-client", "client_secret": "cc-secret-" + uuid.uuid4().hex},
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
auth.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
assert [request["grant_type"] for request in auth.token_requests()] == ["client_credentials"]
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
sent: Final = _header(listings[0], b"authorization")
|
||||
assert sent is not None and auth.is_live(sent.decode().removeprefix("Bearer ")), sent
|
||||
assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2])
|
||||
def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret(
|
||||
gateway: Gateway, auth_type: str, header: bytes, shape: str
|
||||
) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
static: Final = "static-" + uuid.uuid4().hex
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static}
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"]
|
||||
peer.drain()
|
||||
called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD)
|
||||
assert called.status_code == 200, called.text
|
||||
assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode()
|
||||
|
||||
|
||||
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list",
|
||||
params={"server_id": identity},
|
||||
headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
names: Final = {tool["name"] for tool in response.json()["tools"]}
|
||||
assert "add" in names, names
|
||||
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
|
||||
|
|
|
|||
531
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
531
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
|
|
@ -0,0 +1,531 @@
|
|||
"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller.
|
||||
|
||||
One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks
|
||||
``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it
|
||||
blocks and echoes the description and parameters it was handed, which is the only way to observe from outside
|
||||
what metadata the gateway attached to the hook
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import uuid
|
||||
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, eventually, gateway_from_environment
|
||||
from integration._support.mcp import (
|
||||
EntryPoint,
|
||||
JsonRpc,
|
||||
McpCaller,
|
||||
ScriptedTool,
|
||||
listed_tools,
|
||||
openapi_peer,
|
||||
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'
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' if "{_PROBE}" in texts:\n'
|
||||
f' return block("{_ECHO}" + json_stringify('
|
||||
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
|
||||
' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n'
|
||||
" if masked != texts:\n"
|
||||
" return modify(texts=masked)\n"
|
||||
" return allow()\n"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("listed-tool-metadata")
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
|
||||
"litellm_params": {
|
||||
"guardrail": "custom_code",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"custom_code": _GUARDRAIL_CODE,
|
||||
},
|
||||
}
|
||||
]
|
||||
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, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate,
|
||||
ExitStack() as stack,
|
||||
):
|
||||
_two_workers(stack, candidate)
|
||||
yield candidate
|
||||
|
||||
|
||||
def _strings(value: object) -> Iterator[str]:
|
||||
if isinstance(value, str):
|
||||
yield value
|
||||
return
|
||||
children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else ()
|
||||
for child in children:
|
||||
yield from _strings(child)
|
||||
|
||||
|
||||
def _decoded(raw: str) -> object:
|
||||
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
|
||||
return json.loads(data[-1] if data else raw)
|
||||
|
||||
|
||||
def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
"""The (description, parameters) the guardrail was handed, recovered from its block reason."""
|
||||
carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None)
|
||||
assert carrier is not None, raw
|
||||
echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1])
|
||||
assert isinstance(echoed, dict), carrier
|
||||
return echoed.get("description"), echoed.get("parameters")
|
||||
|
||||
|
||||
def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id)
|
||||
assert outcome.error is not None, outcome.raw
|
||||
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
|
||||
) -> None:
|
||||
schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}}
|
||||
tool: Final = ScriptedTool(
|
||||
"lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "meta" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias})
|
||||
assert caller.initialize().ok
|
||||
listed: Final = caller.list_tools(server_id=identity)
|
||||
assert listed.ok, listed.raw
|
||||
name: Final = next(full for full in listed.tools if full.endswith("lookup"))
|
||||
description, parameters = _probe(caller, name, identity)
|
||||
assert description == "Look up one record", (description, parameters)
|
||||
assert parameters is not None and parameters.get("properties") == schema["properties"], parameters
|
||||
|
||||
|
||||
def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool(
|
||||
"report",
|
||||
lambda _: text_result("ok"),
|
||||
description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}",
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "tenant" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"})
|
||||
globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"})
|
||||
acme_listing: Final = acme.list_tools()
|
||||
globex_listing: Final = globex.list_tools()
|
||||
assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw)
|
||||
name: Final = next(full for full in acme_listing.tools if full.endswith("report"))
|
||||
acme_seen, _ = _probe(acme, name, identity)
|
||||
globex_seen, _ = _probe(globex, name, identity)
|
||||
assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), (
|
||||
"each caller's tools/call must be evaluated against the catalog its own headers listed"
|
||||
)
|
||||
|
||||
|
||||
def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note")
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "note" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("read_note"))
|
||||
assert served[name]["description"] == "Read a [MASKED] note", served[name]
|
||||
seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("getpet"))
|
||||
assert served[name]["description"] == "Fetch one [MASKED] pet", served[name]
|
||||
seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served"
|
||||
assert parameters is not None and "petId" in parameters.get("properties", {}), parameters
|
||||
assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the_last_listing(
|
||||
rig: Gateway,
|
||||
) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
guarded: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
opted_out: Final = scenario.key(
|
||||
object_permission={"mcp_servers": [identity]}, metadata={"disable_global_guardrails": True}
|
||||
)
|
||||
guarded_served: Final = listed_tools(rig, guarded, identity)
|
||||
opted_out_served: Final = listed_tools(rig, opted_out, identity)
|
||||
name: Final = next(full for full in guarded_served if full.endswith("getpet"))
|
||||
assert (guarded_served[name]["description"], opted_out_served[name]["description"]) == (
|
||||
"Fetch one [MASKED] pet",
|
||||
"Fetch one SECRET pet",
|
||||
), (guarded_served[name], opted_out_served[name])
|
||||
seen, _ = _probe(McpCaller(rig, guarded, "rest"), name, identity)
|
||||
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) == ()
|
||||
|
||||
|
||||
def test_an_admin_include_disabled_tools_listing_does_not_warm_the_runtime_catalog(rig: Gateway) -> None:
|
||||
"""``include_disabled_tools=true`` is the admin-only configuration view, not a listing the caller
|
||||
runs against: recording it would warm tools/call metadata no runtime listing ever served."""
|
||||
tool: Final = _lookup_tool()
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "adminview" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
admin: Final = rig.key
|
||||
caller: Final = McpCaller(rig, admin, "rest")
|
||||
name: Final = eventually(
|
||||
lambda: caller.call(f"{alias}-lookup", {"probe": _PROBE}, server_id=identity),
|
||||
lambda value: value.error is not None and _ECHO in value.raw,
|
||||
)
|
||||
assert _echoed(name.raw) == _COLD, "before any listing the call is cold"
|
||||
|
||||
view: Final = eventually(
|
||||
lambda: rig.client.get(
|
||||
"/mcp-rest/tools/list",
|
||||
params={"server_id": identity, "include_disabled_tools": "true"},
|
||||
headers={"x-litellm-api-key": admin},
|
||||
),
|
||||
lambda response: response.status_code == 200
|
||||
and any(entry["name"].endswith("lookup") for entry in response.json()["tools"]),
|
||||
)
|
||||
assert view.status_code == 200, view.text
|
||||
|
||||
after_view: Final = _probe(caller, f"{alias}-lookup", identity)
|
||||
assert after_view == _COLD, (
|
||||
"the admin-only include_disabled_tools view must not record the caller's listed-tools slot"
|
||||
)
|
||||
|
||||
runtime: Final = caller.list_tools(server_id=identity)
|
||||
assert runtime.ok, runtime.raw
|
||||
assert _probe(caller, f"{alias}-lookup", identity)[0] == "Look up one record", (
|
||||
"a genuine runtime listing still warms the slot"
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -1343,12 +1343,16 @@ def test_is_prompt_caching_enabled_error_handling():
|
|||
|
||||
def test_is_prompt_caching_enabled_return_default_image_dimensions():
|
||||
"""
|
||||
Assert that `is_prompt_caching_valid_prompt` calls token_counter with use_default_image_token_count=True
|
||||
Assert that `is_prompt_caching_valid_prompt` counts tokens with use_default_image_token_count=True
|
||||
when processing messages containing images
|
||||
|
||||
IMPORTANT: Ensures Get token counter does not make a GET request to the image url
|
||||
"""
|
||||
with patch("litellm.utils.token_counter") as mock_token_counter:
|
||||
mock_token_counter = MagicMock(return_value=False)
|
||||
with patch(
|
||||
"litellm.utils._get_messages_reach_token_count",
|
||||
return_value=mock_token_counter,
|
||||
):
|
||||
litellm.utils.is_prompt_caching_valid_prompt(
|
||||
messages=[
|
||||
{
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@
|
|||
"user": "",
|
||||
"team_id": "",
|
||||
"organization_id": "",
|
||||
"metadata": "{\"actor_agent_id\": null, \"target_agent_id\": null, \"billing_agent_id\": null, \"agent_execution_mode\": null, \"verified_human_user_id\": null, \"used_client_oauth_token\": null, \"applied_guardrails\": [], \"attempted_fallbacks\": null, \"original_model_group\": null, \"batch_models\": null, \"batch_successful_requests\": null, \"batch_failed_requests\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"routing_decision\": null, \"internal_call_origin\": null, \"router_metadata\": null, \"autorouter_savings_estimate\": null, \"autorouter_baseline_observation\": null, \"azure_spillover\": null, \"guardrail_information\": null, \"compression_savings\": null, \"litellm_gateway_injected_cache\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"user_agent\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}",
|
||||
"metadata": "{\"actor_agent_id\": null, \"target_agent_id\": null, \"billing_agent_id\": null, \"agent_execution_mode\": null, \"verified_human_user_id\": null, \"used_client_oauth_token\": null, \"applied_guardrails\": [], \"attempted_fallbacks\": null, \"original_model_group\": null, \"batch_models\": null, \"batch_successful_requests\": null, \"batch_failed_requests\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"routing_decision\": null, \"internal_call_origin\": null, \"litellm_roi_estimator\": false, \"router_metadata\": null, \"autorouter_savings_estimate\": null, \"autorouter_baseline_observation\": null, \"azure_spillover\": null, \"guardrail_information\": null, \"compression_savings\": null, \"litellm_gateway_injected_cache\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"user_agent\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}",
|
||||
"cache_key": "Cache OFF",
|
||||
"spend": 0.00022500000000000002,
|
||||
"total_tokens": 30,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ import lock. These tests pin the import to a single resolution.
|
|||
"""
|
||||
|
||||
import builtins
|
||||
import importlib.abc
|
||||
import sys
|
||||
|
||||
import litellm.integrations.otel.runtime as runtime
|
||||
|
||||
|
|
@ -69,3 +71,25 @@ def test_phase_event_no_ops_when_runtime_absent(monkeypatch):
|
|||
|
||||
assert runtime.phase_event("litellm.request.body_parsed") is None
|
||||
assert runtime.phase_event("litellm.request.body_received", {"litellm.request.body_bytes": 3}) is None
|
||||
|
||||
|
||||
def test_phase_span_does_not_import_the_proxy_in_an_sdk_process(monkeypatch):
|
||||
import litellm.proxy
|
||||
|
||||
monkeypatch.delitem(sys.modules, "litellm.proxy.proxy_server", raising=False)
|
||||
monkeypatch.delattr(litellm.proxy, "proxy_server", raising=False)
|
||||
proxy_imports: list[str] = []
|
||||
|
||||
class _RefuseProxyImport(importlib.abc.MetaPathFinder):
|
||||
def find_spec(self, fullname, path, target=None):
|
||||
if fullname == "litellm.proxy.proxy_server":
|
||||
proxy_imports.append(fullname)
|
||||
raise ImportError(fullname)
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(sys, "meta_path", [_RefuseProxyImport(), *sys.meta_path])
|
||||
|
||||
with runtime.phase_span("route gpt-5-mini") as span:
|
||||
assert span is None
|
||||
|
||||
assert proxy_imports == []
|
||||
|
|
|
|||
|
|
@ -226,6 +226,76 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials(
|
|||
assert execution["guardrail_context"] == {"metadata": {"guardrails": ("block-all",)}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_with_mcp_hands_execution_the_requests_served_tools():
|
||||
"""
|
||||
Regression test: /v1/messages auto-execution must carry this request's
|
||||
resolved tool definitions into execution, matching the Responses and chat
|
||||
completions bridges.
|
||||
|
||||
Given: A request whose MCP reference resolves to a definition carrying a
|
||||
description and input schema
|
||||
When: The model asks for that tool and the gateway executes it
|
||||
Then: _execute_tool_calls receives the definition under served_tools, so
|
||||
pre_mcp_call hooks can judge the call on what the model was shown
|
||||
|
||||
Dropping it does not fail loudly; the call still runs, but the hook sees
|
||||
only the name and arguments, leaving /v1/messages permanently colder than
|
||||
the other two bridges even though all three resolve the same definitions.
|
||||
"""
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.llms.anthropic.pass_through.messages import mcp_handler
|
||||
from litellm.responses.mcp.request_context import MCPRequestContext
|
||||
|
||||
served = [
|
||||
Tool(
|
||||
name="read_wiki_structure",
|
||||
description="Read the structure of a wiki",
|
||||
inputSchema={"type": "object", "properties": {"repoName": {"type": "string"}}},
|
||||
)
|
||||
]
|
||||
|
||||
process = AsyncMock(return_value=(served, {"read_wiki_structure": "deepwiki"}))
|
||||
execute = AsyncMock(
|
||||
return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "read_wiki_structure"}]
|
||||
)
|
||||
responses = [
|
||||
{
|
||||
"stop_reason": "tool_use",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {}}],
|
||||
},
|
||||
{"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]},
|
||||
]
|
||||
|
||||
with (
|
||||
patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")),
|
||||
patch.object(
|
||||
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
|
||||
"_process_mcp_tools_without_openai_transform",
|
||||
new=process,
|
||||
),
|
||||
patch.object(
|
||||
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
|
||||
"_execute_tool_calls",
|
||||
new=execute,
|
||||
),
|
||||
patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)),
|
||||
):
|
||||
await mcp_handler.anthropic_messages_with_mcp(
|
||||
max_tokens=100,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
model="claude-sonnet-4-5",
|
||||
tools=[MCP_REFERENCE],
|
||||
)
|
||||
|
||||
execution = execute.call_args.kwargs
|
||||
assert execution.get("served_tools") == served, (
|
||||
"The request's resolved tool definitions must reach execution so pre_mcp_call "
|
||||
"hooks see the listed description and input schema"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -210,18 +210,18 @@ async def test_persist_and_fetch_round_trip_encrypted_at_rest():
|
|||
stored = {}
|
||||
prisma = _make_prisma(stored)
|
||||
token = _make_id_token()
|
||||
assertion = assertion_from_sso_login(token, "rt_1")
|
||||
assertion = assertion_from_sso_login(token, "refresh.token")
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("user-a", assertion)
|
||||
fetched = await fetch_sso_identity_assertion("user-a")
|
||||
assert fetched is not None
|
||||
assert fetched.id_token.get_secret_value() == token
|
||||
assert fetched.refresh_token is not None
|
||||
assert fetched.refresh_token.get_secret_value() == "rt_1"
|
||||
assert fetched.refresh_token.get_secret_value() == "refresh.token"
|
||||
assert fetched.issuer == assertion.issuer
|
||||
assert fetched.expires_at == assertion.expires_at
|
||||
assert token not in stored["user-a"]
|
||||
assert "rt_1" not in stored["user-a"]
|
||||
assert "refresh.token" not in stored["user-a"]
|
||||
decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug")
|
||||
assert json.loads(decrypted)["id_token"] == token
|
||||
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ def _bare_manager() -> MOD.MCPServerManager:
|
|||
reaches the guardrail hooks; they have their own coverage elsewhere.
|
||||
"""
|
||||
mgr = MOD.MCPServerManager.__new__(MOD.MCPServerManager)
|
||||
mgr._listed_tools_by_server_id = {}
|
||||
mgr.check_allowed_or_banned_tools = lambda name, server: True
|
||||
mgr.validate_allowed_params = lambda tool_name, arguments, server: None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,17 +1,29 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
handle_mcp_proxy_tool,
|
||||
mcp_proxy_tool_id,
|
||||
with_mcp_proxy_identity,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
AUTH = UserAPIKeyAuth(api_key="key")
|
||||
|
||||
|
|
@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke
|
|||
assert hook_payload["arguments"] == arguments
|
||||
assert "raw_headers" not in hook_payload
|
||||
assert "raw-scope-secret" not in recorder.events[1][1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None:
|
||||
"""/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its
|
||||
tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook
|
||||
must see no listed tool for the call."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta")
|
||||
auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult:
|
||||
await asyncio.gather(*tasks)
|
||||
return CallToolResult(content=[TextContent(type="text", text="echoed")])
|
||||
|
||||
with (
|
||||
patch.dict(manager.registry, {server.server_id: server}),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "pre_call_tool_check", pre_call_tool_check),
|
||||
patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_proxy_tool(
|
||||
name="call_tool",
|
||||
arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}},
|
||||
user_api_key_dict=auth,
|
||||
)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "echoed"
|
||||
pre_call_tool_check.assert_awaited_once()
|
||||
assert pre_call_tool_check.await_args.kwargs["name"] == "echo"
|
||||
assert pre_call_tool_check.await_args.kwargs["tool"] is None
|
||||
assert listed is None
|
||||
|
|
|
|||
|
|
@ -963,6 +963,8 @@ async def test_get_tools_from_mcp_servers():
|
|||
user_api_key_auth=None,
|
||||
oauth2_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
catalog_auth_header=None,
|
||||
record_listing=True,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
|
|
@ -1998,6 +2000,7 @@ async def test_get_tools_for_single_server():
|
|||
client_ip=None,
|
||||
user_api_key_auth=None,
|
||||
proxy_logging_obj=ANY,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -9,6 +9,7 @@ from types import SimpleNamespace
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp import ReadResourceResult, Resource
|
||||
|
|
@ -23,18 +24,20 @@ from mcp.types import (
|
|||
TextContent,
|
||||
TextResourceContents,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
MCPTransport,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
|
||||
|
||||
|
||||
def test_mcp_available_on_sdk2():
|
||||
|
|
@ -85,9 +88,6 @@ def cleanup_mcp_global_state():
|
|||
yield
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def _call_tool_params(name, arguments=None):
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
|
|
@ -99,6 +99,7 @@ def _paged_params():
|
|||
|
||||
return PaginatedRequestParams()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx):
|
||||
"""Test that proxy_server_request body contains name and arguments"""
|
||||
|
|
@ -295,7 +296,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r
|
|||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
|
||||
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
|
||||
result = await mcp_server_tool_call(
|
||||
_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})
|
||||
)
|
||||
|
||||
assert result.is_error is True
|
||||
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
|
||||
|
|
@ -1167,20 +1170,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind,
|
|||
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
|
||||
)
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)),
|
||||
),
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
|
||||
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
|
||||
patch.object(
|
||||
operations.global_mcp_server_manager,
|
||||
"read_resource_from_server",
|
||||
AsyncMock(return_value=ReadResourceResult(contents=[content])),
|
||||
),
|
||||
):
|
||||
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
|
||||
|
||||
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
|
||||
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
|
||||
"contents": [{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}],
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
"ttlMs": 0,
|
||||
"contents": [
|
||||
{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1674,7 +1689,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
|
|||
with (
|
||||
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
|
||||
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
|
||||
new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None),
|
||||
new=AsyncMock(
|
||||
return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None
|
||||
),
|
||||
),
|
||||
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
|
||||
"litellm.proxy._experimental.mcp_server.operations._list_mcp_tools",
|
||||
|
|
@ -1913,8 +1930,8 @@ async def test_streamable_http_session_manager_is_stateless():
|
|||
("DELETE", b"", False),
|
||||
),
|
||||
)
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx,
|
||||
debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
|
||||
_mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
) -> None:
|
||||
from starlette.requests import Request
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
|
@ -4056,7 +4073,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock(
|
|||
# parsed, with a nested "method" key in the first bytes to trip a flat
|
||||
# substring heuristic.
|
||||
response_prefix: Final = (
|
||||
'{"jsonrpc":"2.0","id":99,"' + response_field
|
||||
'{"jsonrpc":"2.0","id":99,"'
|
||||
+ response_field
|
||||
+ '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"'
|
||||
).encode()
|
||||
response_body: Final = (
|
||||
|
|
@ -6564,8 +6582,12 @@ class TestGatewayCreateInitializationOptions:
|
|||
yield (None, None)
|
||||
|
||||
async def record_request(
|
||||
serving_server: object, read_stream: object, write_stream: object,
|
||||
*, lifespan_state: object, init_options: InitializationOptions,
|
||||
serving_server: object,
|
||||
read_stream: object,
|
||||
write_stream: object,
|
||||
*,
|
||||
lifespan_state: object,
|
||||
init_options: InitializationOptions,
|
||||
) -> None:
|
||||
captured["server_name"] = init_options.server_name
|
||||
|
||||
|
|
@ -6877,7 +6899,6 @@ async def test_probe_upstream_auth_surfaces_httpx_status_error():
|
|||
returning the response. The probe must catch that specifically (before the
|
||||
fail-open `except Exception`) so the auth check is not silently defeated.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
|
||||
|
||||
|
|
@ -7412,7 +7433,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
|
|||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7658,7 +7680,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
|
|||
return_value=alias_less_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7933,7 +7956,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7989,9 +8013,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
start_time = datetime.now(timezone.utc)
|
||||
litellm_logging_obj, _ = function_setup(
|
||||
|
|
@ -8042,6 +8069,348 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
assert litellm_logging_obj.model == "MCP: list_pets"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing_before_a_listing():
|
||||
"""A local-registry tools/call with no prior tools/list hands the pre-call hooks name and arguments
|
||||
only, as before this metadata existed, so a pre_mcp_call policy never scans a description the caller was
|
||||
not served. Once the caller has listed, the same call hands the entry that listing served."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"list_pets": "ADMIN DESC"},
|
||||
)
|
||||
schema = {"type": "object", "properties": {"limit": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-list_pets", description="List the pets", input_schema=schema, handler=lambda limit: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
pre_call_tool_check = AsyncMock(wraps=manager.pre_call_tool_check)
|
||||
|
||||
async def call() -> tuple[MCPTool | None, dict]:
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-list_pets",
|
||||
arguments={"limit": 10},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
return pre_call_tool_check.call_args.kwargs["tool"], proxy_logging.pre_call_hook.call_args.kwargs["data"]
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
never_listed_tool, never_listed_data = await call()
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
listed_tool, listed_data = await call()
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert never_listed_tool is None
|
||||
assert (never_listed_data.get("mcp_tool_description"), never_listed_data.get("mcp_input_schema")) == (None, None)
|
||||
assert listed_tool is not None and (listed_tool.description, listed_tool.input_schema) == ("ADMIN DESC", schema)
|
||||
assert (listed_data["mcp_tool_description"], listed_data["mcp_input_schema"]) == ("ADMIN DESC", schema)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw():
|
||||
"""When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry
|
||||
call path must hand the pre-call hooks that served entry, not the raw registry one."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"getpetbyid": "Find a SECRET pet"},
|
||||
)
|
||||
registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}}
|
||||
pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid",
|
||||
description="Find pet by ID",
|
||||
input_schema=registry_schema,
|
||||
handler=lambda petId: "ok",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
|
||||
assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), (
|
||||
"the pre-call policy must evaluate the entry tools/list served, not the raw registry entry"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entry():
|
||||
"""Two keys can be shown differently guarded OpenAPI catalogs. The call path must evaluate each key
|
||||
against the entry its own tools/list served, not the entry the most recent listing left behind."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid", description="Find a SECRET pet", input_schema=schema, handler=lambda petId: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice")
|
||||
opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=guarded),
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=opted_out),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
for caller in (guarded, opted_out):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=caller,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed = [call.kwargs["tool"].description for call in pre_call_tool_check.call_args_list]
|
||||
assert handed == ["Find a [MASKED] pet", "Find a SECRET pet"], (
|
||||
"each key's tools/call must be evaluated against the OpenAPI entry its own listing served"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_hooks_no_registry_metadata():
|
||||
"""An OpenAPI operation whose name starts with its own server prefix runs instead of the shorter one, and
|
||||
with no prior listing the pre-call hooks get name and arguments only, never either registry entry."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(name="petstore-get_pet", description="short", input_schema={}, handler=lambda: "short")
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
description="long",
|
||||
input_schema={"type": "object", "properties": {"petId": {"type": "integer"}}},
|
||||
handler=lambda: "long",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation_named_after_a_listed_one():
|
||||
"""After the caller listed ``get_pet``, a call to the never-listed ``petstore-get_pet`` operation hands the
|
||||
pre-call hooks name and arguments only, not the listed sibling's description and schema."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet", description="long", input_schema={}, handler=lambda: "long"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description():
|
||||
"""The listing tools/call runs on its own when this worker does not yet expose the tool is never served
|
||||
to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and
|
||||
arguments only, as on main."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = _never_listed_passthrough_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
upstream = AsyncMock()
|
||||
upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
proxy_logging = _mock_mcp_proxy_logging()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging.during_call_hook = AsyncMock(return_value=None)
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
result = await mcp_operations.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
raw_headers={"authorization": "Bearer caller-token"},
|
||||
)
|
||||
|
||||
assert fetch_tools.await_count == 1
|
||||
assert upstream.call_tool.await_count == 1
|
||||
assert result.content[0].text == "ok"
|
||||
hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin():
|
||||
"""The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description
|
||||
overrides, so it must not become what the admin's own later tools/call is evaluated against."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(
|
||||
server_id="pin-srv",
|
||||
name="pin_srv",
|
||||
transport=MCPTransport.http,
|
||||
url="https://up.example.com/mcp",
|
||||
tool_name_to_description={"add": "Admin wording"},
|
||||
)
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin")
|
||||
request = MagicMock()
|
||||
request.client.host = "10.1.2.3"
|
||||
request.headers = {"x-litellm-api-key": "sk-admin"}
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())),
|
||||
):
|
||||
snapshot = await fetch_pinnable_tool_catalog(server, request, admin)
|
||||
|
||||
assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})}
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server():
|
||||
"""A prefixed REST name that resolves to no tool must still dispatch to the server_id.
|
||||
|
|
@ -8098,7 +8467,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -8597,7 +8967,9 @@ class TestMCPMetaTraceCarrier:
|
|||
|
||||
assert _mcp_meta_trace_carrier(None) is None
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
|
||||
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
|
||||
only_progress = CallToolRequestParams.model_validate(
|
||||
{"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False
|
||||
).meta
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
|
||||
|
||||
|
||||
|
|
@ -10394,7 +10766,9 @@ async def test_mcp_origin_admission_precedes_authentication(
|
|||
patch("litellm.proxy.proxy_server.origins", allowed_origins),
|
||||
patch.object(server, "extract_mcp_auth_context", authenticate),
|
||||
):
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
|
||||
) as client:
|
||||
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
|
||||
|
||||
assert response.status_code == expected_status
|
||||
|
|
@ -10477,12 +10851,15 @@ async def test_streamable_http_rejects_modern_protocol_version(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("handler_name,field", [
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
])
|
||||
@pytest.mark.parametrize(
|
||||
"handler_name,field",
|
||||
[
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
],
|
||||
)
|
||||
async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field):
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
|
|
@ -10500,7 +10877,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
auth = UserAPIKeyAuth(user_id="denied-caller")
|
||||
denial = HTTPException(status_code=403, detail="scope denied")
|
||||
logger = MagicMock()
|
||||
logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None)
|
||||
logger.post_call_failure_hook = AsyncMock(
|
||||
side_effect=RuntimeError("log unavailable") if failure_hook_raises else None
|
||||
)
|
||||
upstream = AsyncMock()
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)),
|
||||
|
|
@ -10509,7 +10888,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream),
|
||||
):
|
||||
with pytest.raises(HTTPException) as rejected:
|
||||
await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True)
|
||||
await operations._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True
|
||||
)
|
||||
assert rejected.value is denial
|
||||
upstream.assert_not_awaited()
|
||||
logger.post_call_failure_hook.assert_awaited_once()
|
||||
|
|
@ -10521,7 +10902,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
|
||||
@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS))
|
||||
async def test_legacy_sse_mount_emits_message_endpoint(
|
||||
prefix: str, suffix: str, opening_protocol: str | None,
|
||||
prefix: str,
|
||||
suffix: str,
|
||||
opening_protocol: str | None,
|
||||
) -> None:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount
|
||||
|
|
@ -10586,16 +10969,20 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
return (await messages.get())["status"]
|
||||
|
||||
if opening_protocol is not None:
|
||||
discover: Final = json.dumps({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}},
|
||||
}).encode()
|
||||
discover: Final = json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {
|
||||
"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}
|
||||
},
|
||||
}
|
||||
).encode()
|
||||
assert await post(discover) == 202
|
||||
discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
|
||||
discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0])
|
||||
|
|
@ -10628,7 +11015,16 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
patch.object(
|
||||
mcp_server,
|
||||
"extract_mcp_auth_context",
|
||||
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
|
||||
AsyncMock(
|
||||
return_value=(
|
||||
post_auth,
|
||||
None,
|
||||
[marker],
|
||||
{marker: {"Authorization": marker}},
|
||||
{"Authorization": marker},
|
||||
{"x-request-marker": marker},
|
||||
)
|
||||
),
|
||||
),
|
||||
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
|
||||
):
|
||||
|
|
@ -10677,7 +11073,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
|
|||
dispatched = AsyncMock(return_value=expected)
|
||||
auth = UserAPIKeyAuth(user_id="discover-caller")
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)),
|
||||
),
|
||||
patch.object(server.operations.GatewayOperations, "execute", dispatched),
|
||||
):
|
||||
result = await server.discover(_mcp_request_ctx(), RequestParams())
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from mcp.types import Tool
|
|||
import litellm
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
AGENT_SEARCH_TOOL_NAME,
|
||||
MCP_TOOL_CALL_TOOL_NAME,
|
||||
|
|
@ -32,12 +33,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
|
|||
ToolSearchResult,
|
||||
coerce_top_k,
|
||||
get_virtual_tool_definitions,
|
||||
handle_mcp_tool_search,
|
||||
search_mcp_tools,
|
||||
search_tools,
|
||||
)
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
|
||||
from litellm.types.mcp import MCPToolSearchSettings
|
||||
from litellm.types.mcp import MCPToolSearchSettings, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
|
||||
|
|
@ -1353,3 +1356,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N
|
|||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The search lists the whole catalog but serves only its hits, so the listing must not fill the
|
||||
caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool."""
|
||||
monkeypatch.setattr(litellm, "mcp_tool_search", None)
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher")
|
||||
upstream = [
|
||||
Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}),
|
||||
]
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user)
|
||||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream]
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"]
|
||||
assert listed == [None, None]
|
||||
|
|
|
|||
|
|
@ -42,9 +42,12 @@ async def test_openapi_local_tool_runs_pre_call_tool_check():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -125,9 +128,12 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "delete_pet"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(
|
||||
side_effect=HTTPException(status_code=403, detail="not allowed")
|
||||
|
|
@ -190,6 +196,8 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -274,6 +282,8 @@ async def test_openapi_local_tool_injects_resolved_oauth_token():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "get_values"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
captured: dict = {}
|
||||
|
||||
async def handle_local(_name, _arguments, _wire_compat):
|
||||
|
|
@ -620,6 +630,8 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc
|
|||
if dispatch_arm == "local_registry":
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server),
|
||||
patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool),
|
||||
|
|
@ -691,6 +703,8 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
fake_tool.handler = raising_handler
|
||||
server = MCPServer(
|
||||
server_id="srv-openapi",
|
||||
|
|
|
|||
|
|
@ -1,15 +1,101 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server import rest_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _CatalogHookCapture(CustomLogger):
|
||||
data: dict[str, object] | None = None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
if call_type == "call_mcp_tool":
|
||||
self.data = data.copy()
|
||||
|
||||
|
||||
async def _served_catalog_tool() -> str:
|
||||
return "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("surface", ["mcp", "rest"])
|
||||
@pytest.mark.parametrize("restriction", ["key", "server"])
|
||||
async def test_listing_records_only_tools_the_caller_received(
|
||||
monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str
|
||||
) -> None:
|
||||
manager: Final = operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="served-catalog", name="served-catalog", transport=MCPTransport.http,
|
||||
spec_path="/catalog.yaml", allow_all_keys=True,
|
||||
allowed_tools=["echo"] if restriction == "server" else None,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-served-catalog", user_id="lister",
|
||||
object_permission={
|
||||
"object_permission_id": "served-permission",
|
||||
"mcp_servers": [server.server_id],
|
||||
"mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None,
|
||||
},
|
||||
)
|
||||
monkeypatch.setitem(manager.registry, server.server_id, server)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id)
|
||||
capture: Final = _CatalogHookCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
for name in ("echo", "status"):
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=f"served-catalog-{name}", description=f"{name} description",
|
||||
input_schema={"type": "object"}, handler=_served_catalog_tool,
|
||||
)
|
||||
try:
|
||||
if surface == "mcp":
|
||||
listing: Final = await operations._list_mcp_tools(
|
||||
user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True,
|
||||
)
|
||||
assert [tool.name for tool in listing.tools] == ["served-catalog-echo"]
|
||||
else:
|
||||
rest_listing: Final = await rest_endpoints._get_tools_for_single_server(
|
||||
server, None, user_api_key_auth=auth,
|
||||
)
|
||||
assert [tool.name for tool in rest_listing] == ["echo"]
|
||||
granted: Final = auth.model_copy(update={"object_permission": None})
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=granted)
|
||||
assert manager.get_listed_tool(server, "status", caller) is None
|
||||
served: Final = manager.get_listed_tool(server, "echo", caller)
|
||||
assert served is not None
|
||||
assert (served.description, served.input_schema) == ("echo description", {"type": "object"})
|
||||
server.allowed_tools = None
|
||||
result: Final = await manager.call_tool(
|
||||
server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
|
||||
)
|
||||
assert result.is_error is False
|
||||
assert capture.data is not None
|
||||
assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}]
|
||||
assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog):
|
||||
from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user
|
||||
|
|
@ -665,3 +751,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
|||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
assert listing.await_args.kwargs["record_listing"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)])
|
||||
async def test_list_mcp_tools_records_the_catalog_only_when_asked(
|
||||
listing_kwargs: dict[str, bool], recorded: bool
|
||||
) -> None:
|
||||
"""The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal
|
||||
caller never serves must not hand a later tools/call a description the caller never saw."""
|
||||
manager = operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
):
|
||||
try:
|
||||
listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
|
||||
assert (listed is not None) is recorded
|
||||
|
|
|
|||
|
|
@ -2218,9 +2218,9 @@ def test_proxy_admin_viewer_can_access_audit_logs(route):
|
|||
# layer, even though the underlying handlers already gate on PROXY_ADMIN_VIEW_ONLY.
|
||||
#
|
||||
# Each route below corresponds to a network call made by the Logs page
|
||||
# (ui/litellm-dashboard/src/components/view_logs/) — see the comment on each.
|
||||
# (ui/litellm-dashboard/src/components/logs/) — see the comment on each.
|
||||
ADMIN_VIEWER_LOGS_PAGE_ROUTES = [
|
||||
# Main paginated log list — uiSpendLogsCall in log_filter_logic.tsx & index.tsx
|
||||
# Main paginated log list — uiSpendLogsCall in request/useLogFilterLogic.ts & index.tsx
|
||||
"/spend/logs/ui",
|
||||
# Single-log detail drawer — fetched on row click in LogDetailsDrawer
|
||||
"/spend/logs/ui/abc-request-id",
|
||||
|
|
|
|||
|
|
@ -346,6 +346,29 @@ class TestAllowFlow:
|
|||
assert evaluate_call.json["conversationId"] == "sess-123"
|
||||
assert evaluate_call.json["agentId"] == "my-agent-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_payload_includes_listed_tool_metadata(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {
|
||||
"name": "send_email",
|
||||
"description": "Send an email",
|
||||
"inputSchema": schema,
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("description", "schema"),
|
||||
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
|
||||
)
|
||||
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {"name": "send_email"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_mcp_call_type_skipped(self):
|
||||
handler: Final = FakeHandler([])
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional
|
|||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
|
@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
|
|||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
|
||||
ParallelSlotAcquisition,
|
||||
|
|
@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPPreCallRequestObject
|
||||
from litellm.types.utils import (
|
||||
EmbeddingResponse,
|
||||
ModelResponse,
|
||||
|
|
@ -108,6 +111,159 @@ def test_api_key_descriptor_applies_budget_throttle(
|
|||
assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(
|
||||
description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 25
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 25
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)])
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_preserves_project_input_and_output_reservations(
|
||||
description: str | None, itpm_limit: int, otpm_limit: int,
|
||||
arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
base_data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}]
|
||||
}
|
||||
expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool")
|
||||
expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit)
|
||||
expected_combined: Final = handler._estimate_tokens_for_request(
|
||||
base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool"
|
||||
)
|
||||
caller: Final = UserAPIKeyAuth(
|
||||
api_key=hash_token("sk-mcp-project-reservation"),
|
||||
tpm_limit=4096,
|
||||
project_id="mcp-project-reservation",
|
||||
project_metadata={
|
||||
"model_itpm_limit": {"mcp-tool-call": itpm_limit},
|
||||
"model_otpm_limit": {"mcp-tool-call": otpm_limit},
|
||||
},
|
||||
)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == (
|
||||
expected_combined,
|
||||
expected_input,
|
||||
expected_output,
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_input
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_output
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None:
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "x" * 400}],
|
||||
"max_tokens": 1,
|
||||
"mcp_tool_name": "echo",
|
||||
"mcp_arguments": {},
|
||||
}
|
||||
assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unconverted_mcp_request_keeps_its_reservation() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64)
|
||||
data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"}
|
||||
|
||||
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 16
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 16
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky(reruns=3)
|
||||
@pytest.mark.asyncio
|
||||
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):
|
||||
|
|
|
|||
|
|
@ -403,6 +403,26 @@ def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api
|
|||
assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"}
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_flows_from_kwargs_to_synthetic_data(proxy_logging):
|
||||
schema = {"type": "object", "properties": {"x": {"type": "integer"}}}
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(
|
||||
kwargs={
|
||||
"name": "calc",
|
||||
"arguments": {"x": 1},
|
||||
"tool_description": "Adds numbers",
|
||||
"tool_input_schema": schema,
|
||||
}
|
||||
)
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert (out["mcp_tool_description"], out["mcp_input_schema"]) == ("Adds numbers", schema)
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}})
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert "mcp_tool_description" not in out and "mcp_input_schema" not in out
|
||||
|
||||
|
||||
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={})
|
||||
snapshot = {
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from typing import Any, Final, Literal, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -12,15 +13,26 @@ from mcp.types import CallToolResult, TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.responses import main as responses_main
|
||||
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
|
||||
|
||||
|
||||
class _DummyMCPResult:
|
||||
|
|
@ -1310,6 +1322,167 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("real_listing", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_tools", "expected_names"),
|
||||
[
|
||||
([], ["responses_slot-echo", "responses_slot-status"]),
|
||||
(["echo"], ["responses_slot-echo"]),
|
||||
(["responses_slot-echo"], ["responses_slot-echo"]),
|
||||
(["absent"], []),
|
||||
],
|
||||
)
|
||||
async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool
|
||||
) -> None:
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http
|
||||
)
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder")
|
||||
upstream: Final = [
|
||||
MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
MCPTool(name="status", description="Report status", inputSchema={"type": "object"}),
|
||||
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
|
||||
]
|
||||
fake_manager: Final = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
fake_manager,
|
||||
)
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
if real_listing:
|
||||
await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True)
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=user)
|
||||
before: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert bool(before) is real_listing
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy/mcp/responses-slot",
|
||||
"allowed_tools": allowed_tools,
|
||||
}
|
||||
],
|
||||
)
|
||||
recorded: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert recorded == before
|
||||
assert (
|
||||
manager.get_listed_tool(
|
||||
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
|
||||
)
|
||||
is None
|
||||
)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert [tool.name for tool in tools] == expected_names
|
||||
|
||||
|
||||
class _BridgeMetadataGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True)
|
||||
self.calls: tuple[tuple[object, object], ...] = ()
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Logging | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if request_data.get("mcp_arguments") == {"probe": "bridge"}:
|
||||
self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),)
|
||||
return inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream")
|
||||
manager.registry = {server.server_id: server}
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user")
|
||||
upstream: Final = [
|
||||
MCPTool(
|
||||
name="echo",
|
||||
description="Echo text",
|
||||
inputSchema={"type": "object", "properties": {"text": {"type": "string"}}},
|
||||
),
|
||||
MCPTool(name="status", description="Read status", inputSchema={"type": "object"}),
|
||||
]
|
||||
client: Final = AsyncMock()
|
||||
client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")])
|
||||
manager._create_mcp_client = AsyncMock(return_value=client)
|
||||
manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream)
|
||||
guardrail: Final = _BridgeMetadataGuardrail()
|
||||
logger: Final = ProxyLogging(user_api_key_cache=DualCache())
|
||||
monkeypatch.setattr(litellm, "callbacks", [guardrail])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger)
|
||||
monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
|
||||
monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager)
|
||||
first_listed: Final = asyncio.Event()
|
||||
second_listed: Final = asyncio.Event()
|
||||
|
||||
async def bridge(name: str, first: bool) -> None:
|
||||
if not first:
|
||||
await first_listed.wait()
|
||||
tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]}
|
||||
],
|
||||
)
|
||||
(first_listed if first else second_listed).set()
|
||||
await second_listed.wait()
|
||||
result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=server_map,
|
||||
tool_calls=[
|
||||
{"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name}
|
||||
],
|
||||
user_api_key_auth=user,
|
||||
served_tools=tools,
|
||||
)
|
||||
assert [entry["result"] for entry in result] == ["ok"]
|
||||
|
||||
try:
|
||||
await asyncio.gather(bridge("echo", True), bridge("status", False))
|
||||
assert sorted(guardrail.calls, key=str) == sorted(
|
||||
((tool.description, tool.input_schema) for tool in upstream), key=str
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (None, None)
|
||||
await manager._get_tools_from_server(
|
||||
server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
return types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
|
|
|
|||
|
|
@ -2266,12 +2266,12 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/EvalViewer/EvalViewer.tsx": {
|
||||
"src/components/logs/detail/sections/EvalViewer/EvalViewer.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/GuardrailViewer/CompliancePanel.tsx": {
|
||||
"src/components/logs/detail/sections/GuardrailViewer/CompliancePanel.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 2
|
||||
},
|
||||
|
|
@ -2279,17 +2279,17 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/GuardrailViewer/ContentFilterDetails.tsx": {
|
||||
"src/components/logs/detail/sections/GuardrailViewer/ContentFilterDetails.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx": {
|
||||
"src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 3
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": {
|
||||
"src/components/logs/detail/LogDetailsDrawer.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 2
|
||||
},
|
||||
|
|
@ -2297,22 +2297,22 @@
|
|||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts": {
|
||||
"src/components/logs/detail/useKeyboardNavigation.ts": {
|
||||
"react-hooks/immutability": {
|
||||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/columns.tsx": {
|
||||
"src/components/logs/types.ts": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/log_filter_logic.tsx": {
|
||||
"src/components/logs/request/useLogFilterLogic.ts": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/logs_utils.tsx": {
|
||||
"src/components/logs/request/timeRange.ts": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
}
|
||||
|
|
|
|||
|
|
@ -100,7 +100,7 @@ const eslintConfig = [
|
|||
rules: { "local/no-ad-hoc-z-index": ["error", { allowPopupLayer: true }] },
|
||||
},
|
||||
{
|
||||
files: ["src/components/view_logs/TraceView/**/*.tsx", "src/components/lens/**/*.tsx"],
|
||||
files: ["src/components/lens/**/*.tsx"],
|
||||
ignores: ["src/**/*.test.tsx"],
|
||||
rules: { "local/no-arbitrary-design-value": "error" },
|
||||
},
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
|||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { LOG_ID_QUERY_PARAM } from "@/components/view_logs/logDetailRouting";
|
||||
import { LOG_ID_QUERY_PARAM } from "@/components/logs/request/logDetailRouting";
|
||||
import type { paths } from "@/lib/http/schema";
|
||||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
import { uiHref } from "@/utils/uiHref";
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({
|
|||
useOrganizations: useOrganizationsMock,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/view_logs/RequestLogsPanel", () => ({
|
||||
vi.mock("@/components/logs/request/RequestLogsPanel", () => ({
|
||||
default: function RequestLogsPanelMock() {
|
||||
return <div data-testid="request-logs-panel" />;
|
||||
},
|
||||
|
|
|
|||
|
|
@ -17,13 +17,13 @@ vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({
|
|||
useOrganizations: useOrganizationsMock,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/view_logs/RequestLogsPanel", () => ({
|
||||
vi.mock("@/components/logs/request/RequestLogsPanel", () => ({
|
||||
default: function RequestLogsPanelMock({ isActive }: { isActive: boolean }) {
|
||||
return <div data-testid="request-logs-panel">{isActive ? "active" : "inactive"}</div>;
|
||||
},
|
||||
}));
|
||||
|
||||
vi.mock("@/components/view_logs/AuditLogsPanel", () => ({
|
||||
vi.mock("@/components/logs/audit/AuditLogsPanel", () => ({
|
||||
default: function AuditLogsPanelMock({ isActive }: { isActive: boolean }) {
|
||||
return <div data-testid="audit-logs-panel">{isActive ? "active" : "inactive"}</div>;
|
||||
},
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
|||
import useCan from "@/app/(dashboard)/hooks/useCan";
|
||||
import DeletedKeysPage from "@/components/DeletedKeysPage/DeletedKeysPage";
|
||||
import DeletedTeamsPage from "@/components/DeletedTeamsPage/DeletedTeamsPage";
|
||||
import AuditLogsPanel from "@/components/view_logs/AuditLogsPanel";
|
||||
import RequestLogsPanel from "@/components/view_logs/RequestLogsPanel";
|
||||
import AuditLogsPanel from "@/components/logs/audit/AuditLogsPanel";
|
||||
import RequestLogsPanel from "@/components/logs/request/RequestLogsPanel";
|
||||
import { Page, PageTabs, PageTabsList, PageTabsTrigger, PageTabsContent } from "@/components/shared/Page";
|
||||
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import React from "react";
|
|||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import { renderWithProviders, screen, testQueryClient, waitFor, within } from "../../../tests/test-utils";
|
||||
import type { LogEntry as SpendLogEntry } from "@/components/view_logs/columns";
|
||||
import type { LogEntry as SpendLogEntry } from "@/components/logs/types";
|
||||
import { LogViewer } from "./LogViewer";
|
||||
|
||||
vi.mock("@/components/networking", async (importOriginal) => {
|
||||
|
|
@ -11,7 +11,7 @@ vi.mock("@/components/networking", async (importOriginal) => {
|
|||
return { ...actual, uiSpendLogsCall: vi.fn() };
|
||||
});
|
||||
|
||||
vi.mock("@/components/view_logs/LogDetailsDrawer", () => ({
|
||||
vi.mock("@/components/logs/detail", () => ({
|
||||
LogDetailsDrawer: function LogDetailsDrawerMock({
|
||||
open,
|
||||
logEntry,
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ import React, { useState } from "react";
|
|||
import { Button } from "@/components/ui/button";
|
||||
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
|
||||
import { uiSpendLogsCall } from "@/components/networking";
|
||||
import { LogDetailsDrawer } from "@/components/view_logs/LogDetailsDrawer";
|
||||
import type { LogEntry as ViewLogsLogEntry } from "@/components/view_logs/columns";
|
||||
import { LogDetailsDrawer } from "@/components/logs/detail";
|
||||
import type { LogEntry as ViewLogsLogEntry } from "@/components/logs/types";
|
||||
import type { LogEntry } from "./mockData";
|
||||
|
||||
const actionConfig: Record<
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import { TriangleAlert } from "lucide-react";
|
|||
import { Badge } from "@/components/ui/badge";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import RoutingDecisionCard from "@/components/view_logs/LogDetailsDrawer/RoutingDecisionCard";
|
||||
import RoutingDecisionCard from "@/components/logs/detail/RoutingDecisionCard";
|
||||
import { AutoRouterRoutingTestResult, testAutoRouterRouting } from "../networking";
|
||||
import { ComplexityRouterConfigPayload, getHeuristicV2SuccessThresholdError } from "./build_complexity_router_config";
|
||||
import { buildAutoRouterRoutingTestRequest } from "./build_auto_router_routing_test_request";
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ interface SimpleTableProps<T> {
|
|||
|
||||
/**
|
||||
* Simple table component for forms and settings pages
|
||||
* For complex tables with sorting/filtering, use DataTable from view_logs
|
||||
* For complex tables with sorting/filtering, use DataTable from shared/DataTable
|
||||
*/
|
||||
export function SimpleTable<T>({
|
||||
data,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import LensPage from "@/app/(dashboard)/lens/page";
|
|||
|
||||
const { auth } = vi.hoisted(() => ({ auth: vi.fn() }));
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: auth }));
|
||||
vi.mock("@/components/view_logs/TraceView/AgentTracesPage", () => ({
|
||||
vi.mock("@/components/lens/traces/list/AgentTracesPage", () => ({
|
||||
default: ({ isActive }: { isActive: boolean }) => <div>Trace polling {isActive ? "active" : "paused"}</div>,
|
||||
}));
|
||||
vi.mock("./investigations/InvestigationsView", () => ({
|
||||
|
|
|
|||
|
|
@ -3,13 +3,13 @@
|
|||
import { useId, useState } from "react";
|
||||
import { useQuery } from "@tanstack/react-query";
|
||||
import { Aperture, ArrowUpRight } from "lucide-react";
|
||||
import AgentTracesPage from "@/components/view_logs/TraceView/AgentTracesPage";
|
||||
import AgentTracesPage from "@/components/lens/traces/list/AgentTracesPage";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import type { TraceSummary } from "@/components/view_logs/TraceView/traceTypes";
|
||||
import type { TraceSummary } from "@/components/lens/traces/types";
|
||||
import { Switch } from "@/components/ui/switch";
|
||||
import { Tabs, TabsContent } from "@/components/ui/tabs";
|
||||
import { LensServicesProvider, useLensAccessToken, useLensApi, useLiveLensServices } from "./data/LensServices";
|
||||
import { LensPreviewContext } from "@/components/view_logs/TraceView/LensPreviewButton";
|
||||
import { LensPreviewContext } from "@/components/lens/ui/LensPreviewButton";
|
||||
import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles";
|
||||
import { InvestigationsView } from "./investigations/InvestigationsView";
|
||||
import { LensSettings } from "./settings/LensSettings";
|
||||
|
|
@ -22,7 +22,7 @@ import { cn } from "@/lib/cva.config";
|
|||
import { useDialogRoute, useLensRoute, type LensDialog, type LensTab } from "./route";
|
||||
import { LensIntroDialog, useLensIntro } from "./onboarding/LensIntroDialog";
|
||||
import { OnboardingProvider, type Onboarding } from "./onboarding/OnboardingContext";
|
||||
import { traceRefOf, useOpenTraceRouting } from "@/components/view_logs/TraceView/traceRouting";
|
||||
import { traceRefOf, useOpenTraceRouting } from "@/components/lens/traces/routing";
|
||||
|
||||
type WorkspaceProps = { accessToken: string; userRole: string; readOnly: boolean };
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
import { createContext, useContext, useMemo, type ReactNode } from "react";
|
||||
import { apiClient } from "@/components/networking";
|
||||
import { liveTracesApi, TracesApiContext, type TracesApi } from "@/components/view_logs/TraceView/tracesApi";
|
||||
import { liveTracesApi, TracesApiContext, type TracesApi } from "@/components/lens/traces/api";
|
||||
import { liveLensApi, type LensApi } from "./service";
|
||||
|
||||
export interface LensServices {
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import { ApiError } from "@/lib/http/client";
|
||||
import type { TracesApi } from "@/components/view_logs/TraceView/tracesApi";
|
||||
import type { TracesApi } from "@/components/lens/traces/api";
|
||||
import type { LensServices } from "../LensServices";
|
||||
import type { LensApi } from "../service";
|
||||
import { createLensDemoData, type LensDemoData } from "./fixtures";
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import type { Trace, Span, SpanDetail } from "@/components/view_logs/TraceView/traceTypes";
|
||||
import type { Trace, Span, SpanDetail } from "@/components/lens/traces/types";
|
||||
import type { Lens, Finding, Job, Settings } from "../../model/types";
|
||||
import { withReleaseCases } from "./lensDemoLongTrace";
|
||||
import { scenarios, type Scenario } from "./scenarios";
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import type { Span, SpanDetail, Trace } from "@/components/view_logs/TraceView/traceTypes";
|
||||
import type { Span, SpanDetail, Trace } from "@/components/lens/traces/types";
|
||||
|
||||
export function withReleaseCases(run: { trace: Trace; details: SpanDetail[] }) {
|
||||
const { trace } = run;
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"use client";
|
||||
|
||||
import { useQuery } from "@tanstack/react-query";
|
||||
import { isTracingNotEnabled, useTraceAvailability } from "@/components/view_logs/TraceView/useAgentTraces";
|
||||
import { isTracingNotEnabled, useTraceAvailability } from "@/components/lens/traces/list/useAgentTraces";
|
||||
import { useLensAccessToken, useLensApi } from "../data/LensServices";
|
||||
import { lensQueries } from "../data/queries";
|
||||
import { readiness, type Readiness, type ReadinessInput } from "../model/readiness";
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ import { useQuery } from "@tanstack/react-query";
|
|||
|
||||
import { Inspector } from "@/components/shared/Inspector";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { RunView } from "@/components/view_logs/TraceView/TraceDrawer";
|
||||
import { useLocalRunSelection } from "@/components/view_logs/TraceView/traceRouting";
|
||||
import { RunView } from "@/components/lens/traces/detail/TraceDrawer";
|
||||
import { useLocalRunSelection } from "@/components/lens/traces/routing";
|
||||
|
||||
import { lensQueries } from "../data/queries";
|
||||
import { useLensAccessToken, useLensApi } from "../data/LensServices";
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import { fireEvent, screen, within } from "@testing-library/react";
|
|||
import userEvent from "@testing-library/user-event";
|
||||
import { afterEach, expect, it, vi } from "vitest";
|
||||
|
||||
import type { RunSelection } from "@/components/view_logs/TraceView/traceRouting";
|
||||
import type { RunSelection } from "@/components/lens/traces/routing";
|
||||
|
||||
import { renderWithLens } from "@/../tests/lens-test-utils";
|
||||
import { Inspector } from "@/components/shared/Inspector";
|
||||
|
|
@ -11,7 +11,7 @@ import type { OwnedFinding } from "../model/inbox";
|
|||
import type { Finding, Lens } from "../model/types";
|
||||
import { FindingPanel, ownedFindingKey } from "./FindingDetails";
|
||||
|
||||
vi.mock("@/components/view_logs/TraceView/TraceDrawer", () => ({
|
||||
vi.mock("@/components/lens/traces/detail/TraceDrawer", () => ({
|
||||
RunView: ({ traceId, selection }: { traceId: string; selection: RunSelection }) => (
|
||||
<div data-testid="run-view">
|
||||
{traceId} at {selection.spanId}
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import { ApiError } from "@/lib/http/client";
|
|||
import { apiClient } from "@/components/networking";
|
||||
import { lensKeys } from "../data/queries";
|
||||
import { InvestigationsView } from "./InvestigationsView";
|
||||
import { LensPreviewContext } from "@/components/view_logs/TraceView/LensPreviewButton";
|
||||
import { LensPreviewContext } from "@/components/lens/ui/LensPreviewButton";
|
||||
import { briefMarkdown } from "../model/findings";
|
||||
import { findingKey } from "../model/inbox";
|
||||
import { runTime } from "../model/format";
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import { useQuery } from "@tanstack/react-query";
|
|||
import { Plus } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { LensPreviewButton } from "@/components/view_logs/TraceView/LensPreviewButton";
|
||||
import { LensPreviewButton } from "@/components/lens/ui/LensPreviewButton";
|
||||
|
||||
import { useInvalidateLenses } from "../data/mutations";
|
||||
import { lensQueries } from "../data/queries";
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"use client";
|
||||
|
||||
import { createContext, useContext } from "react";
|
||||
import type { TraceSummary } from "@/components/view_logs/TraceView/traceTypes";
|
||||
import type { TraceSummary } from "@/components/lens/traces/types";
|
||||
|
||||
export interface Onboarding {
|
||||
readonly readOnly: boolean;
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
import { useId, useRef, useState, type ReactNode } from "react";
|
||||
import { ArrowRight, ChevronDown } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { TracingSetupFields } from "@/components/view_logs/TraceView/TracingSetupCard";
|
||||
import { TracingSetupFields } from "@/components/lens/onboarding/tracing/TracingSetupCard";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { useLensAccessToken } from "../data/LensServices";
|
||||
import type { LensReadiness } from "../hooks/useLensReadiness";
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ import userEvent from "@testing-library/user-event";
|
|||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { chooseSelectOption, renderWithProviders } from "@/../tests/test-utils";
|
||||
import { copyToClipboard } from "@/utils/dataUtils";
|
||||
import { LensPreviewContext } from "./LensPreviewButton";
|
||||
import { agentTraceCall, apiClient, sendOtlpTraceCall } from "../../networking";
|
||||
import { LensPreviewContext } from "../../ui/LensPreviewButton";
|
||||
import { agentTraceCall, apiClient, sendOtlpTraceCall } from "../../../networking";
|
||||
import {
|
||||
codingAgentCommand,
|
||||
codingAgentPrompt,
|
||||
|
|
@ -14,9 +14,9 @@ import {
|
|||
TracingSetupCard,
|
||||
} from "./TracingSetupCard";
|
||||
import { FRAMEWORKS } from "./tracingSetupGuides";
|
||||
import type { Trace } from "./traceTypes";
|
||||
import type { Trace } from "../../traces/types";
|
||||
|
||||
vi.mock("../../networking", () => ({
|
||||
vi.mock("../../../networking", () => ({
|
||||
getProxyBaseUrl: () => "http://proxy.test/",
|
||||
sendOtlpTraceCall: vi.fn(),
|
||||
agentTraceCall: vi.fn(),
|
||||
|
|
@ -5,19 +5,19 @@ import { useState } from "react";
|
|||
import { useTimeout } from "usehooks-ts";
|
||||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { LensPreviewButton } from "./LensPreviewButton";
|
||||
import { LensPreviewButton } from "../../ui/LensPreviewButton";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { copyToClipboard } from "@/utils/dataUtils";
|
||||
|
||||
import anthropicLogo from "../../../../public/assets/logos/anthropic.svg";
|
||||
import openaiLogo from "../../../../public/assets/logos/openai_small.svg";
|
||||
import otelLogo from "../../../../public/assets/logos/opentelemetry.svg";
|
||||
import { agentTraceCall, apiClient, getProxyBaseUrl, sendOtlpTraceCall } from "../../networking";
|
||||
import { ActiveDot } from "./ActiveDot";
|
||||
import anthropicLogo from "../../../../../public/assets/logos/anthropic.svg";
|
||||
import openaiLogo from "../../../../../public/assets/logos/openai_small.svg";
|
||||
import otelLogo from "../../../../../public/assets/logos/opentelemetry.svg";
|
||||
import { agentTraceCall, apiClient, getProxyBaseUrl, sendOtlpTraceCall } from "../../../networking";
|
||||
import { ActiveDot } from "../../traces/ui/ActiveDot";
|
||||
import { sampleTraceExport } from "./sampleTrace";
|
||||
import { FRAMEWORKS, frameworkSnippet, type FrameworkGuide } from "./tracingSetupGuides";
|
||||
import type { TraceSummary } from "./traceTypes";
|
||||
import type { TraceSummary } from "../../traces/types";
|
||||
|
||||
const COPIED_RESET_MS = 1500;
|
||||
const DOCS_URL = "https://docs.litellm.ai/docs/proxy/lens";
|
||||
|
|
@ -1,17 +1,17 @@
|
|||
import langgraphLogo from "../../../../public/assets/logos/langgraph-color.svg";
|
||||
import langchainLogo from "../../../../public/assets/logos/langchain.svg";
|
||||
import openaiAgentsLogo from "../../../../public/assets/logos/openai-agents.svg";
|
||||
import anthropicLogo from "../../../../public/assets/logos/anthropic.svg";
|
||||
import crewaiLogo from "../../../../public/assets/logos/crewai-color.svg";
|
||||
import pydanticAiLogo from "../../../../public/assets/logos/pydantic-ai-color.svg";
|
||||
import llamaindexLogo from "../../../../public/assets/logos/llamaindex-color.svg";
|
||||
import vercelLogo from "../../../../public/assets/logos/vercel.svg";
|
||||
import otelLogo from "../../../../public/assets/logos/opentelemetry.svg";
|
||||
import langgraphLogo from "../../../../../public/assets/logos/langgraph-color.svg";
|
||||
import langchainLogo from "../../../../../public/assets/logos/langchain.svg";
|
||||
import openaiAgentsLogo from "../../../../../public/assets/logos/openai-agents.svg";
|
||||
import anthropicLogo from "../../../../../public/assets/logos/anthropic.svg";
|
||||
import crewaiLogo from "../../../../../public/assets/logos/crewai-color.svg";
|
||||
import pydanticAiLogo from "../../../../../public/assets/logos/pydantic-ai-color.svg";
|
||||
import llamaindexLogo from "../../../../../public/assets/logos/llamaindex-color.svg";
|
||||
import vercelLogo from "../../../../../public/assets/logos/vercel.svg";
|
||||
import otelLogo from "../../../../../public/assets/logos/opentelemetry.svg";
|
||||
|
||||
import adkLogo from "../../../../public/assets/logos/google-adk.png";
|
||||
import strandsLogo from "../../../../public/assets/logos/strands.svg";
|
||||
import hermesLogo from "../../../../public/assets/logos/hermes.png";
|
||||
import openclawLogo from "../../../../public/assets/logos/openclaw.png";
|
||||
import adkLogo from "../../../../../public/assets/logos/google-adk.png";
|
||||
import strandsLogo from "../../../../../public/assets/logos/strands.svg";
|
||||
import hermesLogo from "../../../../../public/assets/logos/hermes.png";
|
||||
import openclawLogo from "../../../../../public/assets/logos/openclaw.png";
|
||||
|
||||
export interface FrameworkGuide {
|
||||
id: string;
|
||||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
import { parseAsBoolean, parseAsString, parseAsStringLiteral, useQueryStates } from "nuqs";
|
||||
import { useCallback } from "react";
|
||||
import { OPEN_TRACE_PARSERS, RUN_FILTER_PARSERS } from "@/components/view_logs/TraceView/traceRouting";
|
||||
import { OPEN_TRACE_PARSERS, RUN_FILTER_PARSERS } from "@/components/lens/traces/routing";
|
||||
|
||||
export const LENS_TABS = { traces: "Traces", investigations: "Investigations", settings: "Settings" } as const;
|
||||
export type LensTab = keyof typeof LENS_TABS;
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import {
|
|||
apiClient,
|
||||
getProxyBaseUrl,
|
||||
} from "../../networking";
|
||||
import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./traceTypes";
|
||||
import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./types";
|
||||
|
||||
export interface TraceWindow {
|
||||
readonly startMs: number;
|
||||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
import { type KeyValue, KeyValueRows } from "./KeyValueRows";
|
||||
import { Card } from "./MessageCard";
|
||||
import type { Span } from "./traceTypes";
|
||||
import type { Span } from "../types";
|
||||
|
||||
interface AttributesDetailProps {
|
||||
traceId: string;
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
"use client";
|
||||
import { useTracesApi } from "./tracesApi";
|
||||
import { useTracesApi } from "../api";
|
||||
|
||||
import { useQuery, type UseQueryOptions } from "@tanstack/react-query";
|
||||
import { useState } from "react";
|
||||
|
|
@ -10,9 +10,9 @@ import { cn } from "@/lib/cva.config";
|
|||
|
||||
import { type KeyValue, KeyValueRows, objectEntries } from "./KeyValueRows";
|
||||
import { Card, MessageCard, Section, ToolResultCard } from "./MessageCard";
|
||||
import type { ErrorSource } from "./traceTree";
|
||||
import type { Span, SpanDetail, SpanErrorPage, TraceMessage, UIContent, UIMessage } from "./traceTypes";
|
||||
import { errorSource, parseJson, parseMessages, prettyPayload } from "./traceUtils";
|
||||
import type { ErrorSource } from "../tree";
|
||||
import type { Span, SpanDetail, SpanErrorPage, TraceMessage, UIContent, UIMessage } from "../types";
|
||||
import { errorSource, parseJson, parseMessages, prettyPayload } from "../utils";
|
||||
|
||||
const ERROR_SOURCE_LABEL: Record<ErrorSource, string> = { tool: "Tool", model: "Model", litellm: "LiteLLM" };
|
||||
const TRACEBACK_MARKER = "Traceback (most recent call last):";
|
||||
|
|
@ -3,20 +3,20 @@ import userEvent from "@testing-library/user-event";
|
|||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { type ComponentProps, useState } from "react";
|
||||
|
||||
import { renderWithProviders, testQueryClient } from "../../../../tests/test-utils";
|
||||
import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils";
|
||||
import { DetailPane } from "./DetailPane";
|
||||
import type { SpanTab } from "./traceRouting";
|
||||
import type { SpanTab } from "../routing";
|
||||
import { absoluteTime, SpanHoverCard, spanFacts } from "./SpanHoverCard";
|
||||
import type { GroupRowData, SpanRowData } from "./traceTree";
|
||||
import type { Span, SpanDetail, SpanErrorPage, Trace } from "./traceTypes";
|
||||
import type { GroupRowData, SpanRowData } from "../tree";
|
||||
import type { Span, SpanDetail, SpanErrorPage, Trace } from "../types";
|
||||
|
||||
vi.mock("../../networking", () => ({
|
||||
vi.mock("../../../networking", () => ({
|
||||
agentTraceSpanCall: vi.fn(),
|
||||
agentTraceSpanErrorCall: vi.fn(),
|
||||
getProxyBaseUrl: () => "http://proxy.test/",
|
||||
}));
|
||||
|
||||
import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../networking";
|
||||
import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../../networking";
|
||||
|
||||
type SpanFields = Partial<Span> & Pick<Span, "span_id">;
|
||||
|
||||
|
|
@ -6,17 +6,17 @@ import { Button } from "@/components/ui/button";
|
|||
import { Tabs, TabsList, TabsTrigger, TabsContent } from "@/components/ui/tabs";
|
||||
|
||||
import { AttributesDetail } from "./AttributesDetail";
|
||||
import { CopyButton } from "./CopyButton";
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
import { DetailContent, errorHeadline, useSpanDetail } from "./DetailContent";
|
||||
import { IdChip } from "./IdChip";
|
||||
import { PaneBar } from "./PaneBar";
|
||||
import { IdChip } from "../ui/IdChip";
|
||||
import { PaneBar } from "../ui/PaneBar";
|
||||
import { RequestDetail } from "./RequestDetail";
|
||||
import { SpanIcon } from "./SpanIcon";
|
||||
import { useTracesApi } from "./tracesApi";
|
||||
import { SPAN_TABS, type SpanTab } from "./traceRouting";
|
||||
import type { GroupRowData, TreeRow } from "./traceTree";
|
||||
import type { Span, SpanType, Trace } from "./traceTypes";
|
||||
import { fmtMs, fmtTok } from "./traceUtils";
|
||||
import { SpanIcon } from "../ui/SpanIcon";
|
||||
import { useTracesApi } from "../api";
|
||||
import { SPAN_TABS, type SpanTab } from "../routing";
|
||||
import type { GroupRowData, TreeRow } from "../tree";
|
||||
import type { Span, SpanType, Trace } from "../types";
|
||||
import { fmtMs, fmtTok } from "../utils";
|
||||
|
||||
interface SpanTabProps {
|
||||
spanTab: SpanTab;
|
||||
|
|
@ -144,7 +144,7 @@ function SpanPane({
|
|||
</TabsContent>
|
||||
</Tabs>
|
||||
<PaneFooter>
|
||||
<CopyButton value={handoff.text} label="Copy step" copiedLabel={handoff.copied} />
|
||||
<CopyButton variant="action" value={handoff.text} label="Copy step" copiedLabel={handoff.copied} />
|
||||
<div className="ml-auto flex items-center gap-3 text-xs text-muted-foreground tabular-nums">
|
||||
<Meta label="time" value={fmtMs(span.duration_ms)} />
|
||||
{tokens > 0 && <Meta label="tokens" value={fmtTok(tokens)} />}
|
||||
|
|
@ -209,7 +209,7 @@ function GroupPane({
|
|||
)}
|
||||
</div>
|
||||
<PaneFooter>
|
||||
<CopyButton value={handoff.text} label="Copy group sample" copiedLabel={handoff.copied} />
|
||||
<CopyButton variant="action" value={handoff.text} label="Copy group sample" copiedLabel={handoff.copied} />
|
||||
</PaneFooter>
|
||||
</aside>
|
||||
);
|
||||
|
|
@ -4,7 +4,7 @@ import { useState } from "react";
|
|||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
import { FoldChevron } from "./Collapse";
|
||||
import { FoldChevron } from "../ui/Collapse";
|
||||
|
||||
const LONG_VALUE_CHARS = 90;
|
||||
const ID_KEY = /(^|_)id$/i;
|
||||
|
|
@ -4,7 +4,7 @@ import { describe, expect, it, vi } from "vitest";
|
|||
|
||||
import { MessageCard, Section, ToolResultCard } from "./MessageCard";
|
||||
|
||||
vi.mock("./spanProvider", () => ({ useSpanProvider: () => null }));
|
||||
vi.mock("../ui/spanProvider", () => ({ useSpanProvider: () => null }));
|
||||
|
||||
const LONG_QUERY = "Find every invoice for the customer that was billed twice. ".repeat(3).trim();
|
||||
|
||||
|
|
@ -7,10 +7,10 @@ import remarkGfm from "remark-gfm";
|
|||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
import { FoldChevron } from "./Collapse";
|
||||
import { CopyButton } from "./CopyButton";
|
||||
import { FoldChevron } from "../ui/Collapse";
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
import { displayValue, type KeyValue, KeyValueRows, objectEntries } from "./KeyValueRows";
|
||||
import type { TraceMessage, TraceToolCall } from "./traceTypes";
|
||||
import type { TraceMessage, TraceToolCall } from "../types";
|
||||
|
||||
const ROLE_LABEL: Record<string, string> = { user: "User", system: "System", assistant: "Assistant", tool: "Tool" };
|
||||
|
||||
|
|
@ -94,7 +94,13 @@ export function ToolCallBlock({ call }: { call: TraceToolCall }) {
|
|||
<div className="flex items-center gap-2 py-1">
|
||||
<RoleTile role="tool" />
|
||||
<span className="min-w-0 truncate text-sm font-medium">{call.name}</span>
|
||||
<CopyButton value={call.name} label={`Copy ${call.name} name`} iconOnly className={INLINE_COPY} />
|
||||
<CopyButton
|
||||
variant="action"
|
||||
value={call.name}
|
||||
label={`Copy ${call.name} name`}
|
||||
iconOnly
|
||||
className={INLINE_COPY}
|
||||
/>
|
||||
</div>
|
||||
<KeyValueRows entries={entries} />
|
||||
</div>
|
||||
|
|
@ -114,7 +120,7 @@ export function MessageCard({ message }: { message: TraceMessage; model: string
|
|||
<div className={cn(HEADER, "sticky top-10 z-sticky", expanded ? "rounded-t-sm border-b-0" : "rounded-sm")}>
|
||||
<FoldTile label={label} open={open} onToggle={() => setOpen((v) => !v)} />
|
||||
<span className={LABEL}>{label}</span>
|
||||
<CopyButton value={copyValue} label={`Copy ${label}`} iconOnly className={CARD_COPY} />
|
||||
<CopyButton variant="action" value={copyValue} label={`Copy ${label}`} iconOnly className={CARD_COPY} />
|
||||
</div>
|
||||
{expanded && (
|
||||
<div
|
||||
|
|
@ -146,7 +152,7 @@ export function ToolResultCard({ name, result, failed = false }: { name: string;
|
|||
<RoleTile role="tool" failed={failed} />
|
||||
<span className="flex min-w-0 max-w-[50%] shrink-0 items-center gap-1.5">
|
||||
<span className={cn(LABEL, failed && "text-destructive")}>{name}</span>
|
||||
<CopyButton value={name} label={`Copy ${name} name`} iconOnly className={INLINE_COPY} />
|
||||
<CopyButton variant="action" value={name} label={`Copy ${name} name`} iconOnly className={INLINE_COPY} />
|
||||
</span>
|
||||
{expandable ? (
|
||||
<button
|
||||
|
|
@ -162,7 +168,7 @@ export function ToolResultCard({ name, result, failed = false }: { name: string;
|
|||
) : (
|
||||
<span className={cn("min-w-0 break-words text-sm leading-5", tone)}>{result || "No output"}</span>
|
||||
)}
|
||||
<CopyButton value={result} label={`Copy ${name} result`} iconOnly className={CARD_COPY} />
|
||||
<CopyButton variant="action" value={result} label={`Copy ${name} result`} iconOnly className={CARD_COPY} />
|
||||
</div>
|
||||
{open && (
|
||||
<pre
|
||||
|
|
@ -5,13 +5,13 @@ import { useState } from "react";
|
|||
|
||||
import { Button } from "@/components/ui/button";
|
||||
|
||||
import { LogDetailsDrawer } from "../LogDetailsDrawer";
|
||||
import { formatCost } from "./AgentTracesTable";
|
||||
import { LogDetailsDrawer } from "../../../logs/detail";
|
||||
import { formatCost } from "../list/AgentTracesTable";
|
||||
import { DetailGroup } from "./AttributesDetail";
|
||||
import { CopyButton } from "./CopyButton";
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
import { type KeyValue, KeyValueRows } from "./KeyValueRows";
|
||||
import type { Span } from "./traceTypes";
|
||||
import { fmtMs, fmtTok } from "./traceUtils";
|
||||
import type { Span } from "../types";
|
||||
import { fmtMs, fmtTok } from "../utils";
|
||||
import { useSpanRequestLog } from "./useSpanRequestLog";
|
||||
|
||||
interface RequestDetailProps {
|
||||
|
|
@ -52,7 +52,7 @@ export function RequestDetail({ span, accessToken, traceStartMs }: RequestDetail
|
|||
<div className="flex flex-col gap-2.5">
|
||||
<div className="flex min-w-0 items-center gap-1">
|
||||
<KeyValueRows entries={[["request_id", span.litellm_request_id]]} mono className="min-w-0 flex-1" />
|
||||
<CopyButton value={span.litellm_request_id} label="Copy request ID" iconOnly />
|
||||
<CopyButton variant="action" value={span.litellm_request_id} label="Copy request ID" iconOnly />
|
||||
</div>
|
||||
<Button
|
||||
variant="outline"
|
||||
|
|
@ -4,12 +4,12 @@ import { Check } from "lucide-react";
|
|||
|
||||
import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card";
|
||||
|
||||
import { formatCost } from "./AgentTracesTable";
|
||||
import { formatCost } from "../list/AgentTracesTable";
|
||||
import { errorHeadline } from "./DetailContent";
|
||||
import { SpanIcon } from "./SpanIcon";
|
||||
import type { GroupRowData } from "./traceTree";
|
||||
import type { Span, SpanType } from "./traceTypes";
|
||||
import { fmtMs, fmtTok } from "./traceUtils";
|
||||
import { SpanIcon } from "../ui/SpanIcon";
|
||||
import type { GroupRowData } from "../tree";
|
||||
import type { Span, SpanType } from "../types";
|
||||
import { fmtMs, fmtTok } from "../utils";
|
||||
|
||||
export const HOVER_OPEN_DELAY_MS = 300;
|
||||
const HOVER_CLOSE_DELAY_MS = 100;
|
||||
|
|
@ -15,13 +15,13 @@ import {
|
|||
} from "@/components/ui/dropdown-menu";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
import { FoldChevron } from "./Collapse";
|
||||
import { PaneBar } from "./PaneBar";
|
||||
import { FoldChevron } from "../ui/Collapse";
|
||||
import { PaneBar } from "../ui/PaneBar";
|
||||
import { groupFacts, SpanHoverCard, spanFacts } from "./SpanHoverCard";
|
||||
import { SpanIcon } from "./SpanIcon";
|
||||
import type { GroupRowData, SpanRowData, TreeRow } from "./traceTree";
|
||||
import type { TraceSummary } from "./traceTypes";
|
||||
import { fmtMs, previewText, type TreeGuide, treeGuides } from "./traceUtils";
|
||||
import { SpanIcon } from "../ui/SpanIcon";
|
||||
import type { GroupRowData, SpanRowData, TreeRow } from "../tree";
|
||||
import type { TraceSummary } from "../types";
|
||||
import { fmtMs, previewText, type TreeGuide, treeGuides } from "../utils";
|
||||
|
||||
interface SpanTreeProps {
|
||||
rows: TreeRow[];
|
||||
|
|
@ -2,19 +2,19 @@ import { act, screen, waitFor, within } from "@testing-library/react";
|
|||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import type { ComponentProps } from "react";
|
||||
import { renderWithProviders, testQueryClient } from "../../../../tests/test-utils";
|
||||
import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils";
|
||||
import { RunView } from "./TraceDrawer";
|
||||
import { useOpenTraceRouting } from "./traceRouting";
|
||||
import { useOpenTraceRouting } from "../routing";
|
||||
import { TraceConversation } from "./TraceConversation";
|
||||
import type { SpanDetail, Trace } from "./traceTypes";
|
||||
import research from "./__fixtures__/research_trace.json";
|
||||
import type { SpanDetail, Trace } from "../types";
|
||||
import research from "../__fixtures__/research_trace.json";
|
||||
|
||||
vi.mock("../../networking", () => ({
|
||||
vi.mock("../../../networking", () => ({
|
||||
agentTraceCall: vi.fn(),
|
||||
agentTraceSpanCall: vi.fn(),
|
||||
getProxyBaseUrl: () => "http://proxy.test",
|
||||
}));
|
||||
import { agentTraceCall, agentTraceSpanCall } from "../../networking";
|
||||
import { agentTraceCall, agentTraceSpanCall } from "../../../networking";
|
||||
|
||||
function RoutedRunView(props: Omit<ComponentProps<typeof RunView>, "selection">) {
|
||||
const { selection } = useOpenTraceRouting();
|
||||
|
|
@ -4,14 +4,14 @@ import { useQueries } from "@tanstack/react-query";
|
|||
import { useState } from "react";
|
||||
import { ChevronRight, Wrench } from "lucide-react";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { CopyButton } from "./CopyButton";
|
||||
import { useTracesApi } from "./tracesApi";
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
import { useTracesApi } from "../api";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { buildConversation, conversationSteps, CONVERSATION_PAGE_SIZE, type ConversationItem } from "./conversation";
|
||||
import { ErrorBlock } from "./DetailContent";
|
||||
import { Markdown, ToolCallBlock } from "./MessageCard";
|
||||
import type { SpanDetail, Trace, TraceMessage } from "./traceTypes";
|
||||
import { fmtMs } from "./traceUtils";
|
||||
import type { SpanDetail, Trace, TraceMessage } from "../types";
|
||||
import { fmtMs } from "../utils";
|
||||
|
||||
export function TraceConversation({
|
||||
trace,
|
||||
|
|
@ -144,7 +144,12 @@ function ConversationTool({ item }: { item: ConversationItem }) {
|
|||
)}
|
||||
<div className="flex items-center justify-between text-xs text-muted-foreground">
|
||||
<span>Result</span>
|
||||
<CopyButton value={item.toolResult ?? ""} label={`Copy ${item.span.name} result`} iconOnly />
|
||||
<CopyButton
|
||||
variant="action"
|
||||
value={item.toolResult ?? ""}
|
||||
label={`Copy ${item.span.name} result`}
|
||||
iconOnly
|
||||
/>
|
||||
</div>
|
||||
<pre className="max-h-80 overflow-auto whitespace-pre-wrap break-words text-xs leading-5">
|
||||
{item.toolResult || "No output recorded"}
|
||||
|
|
@ -5,19 +5,19 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
|
|||
|
||||
import { ApiError } from "@/lib/http/client";
|
||||
|
||||
import { renderWithProviders, testQueryClient } from "../../../../tests/test-utils";
|
||||
import researchTrace from "./__fixtures__/research_trace.json";
|
||||
import swarmTrace from "./__fixtures__/swarm_trace.json";
|
||||
import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils";
|
||||
import researchTrace from "../__fixtures__/research_trace.json";
|
||||
import swarmTrace from "../__fixtures__/swarm_trace.json";
|
||||
import type { ComponentProps } from "react";
|
||||
import { ShortcutHints } from "@/components/shared/ShortcutHints";
|
||||
import { initialRunSelection, RunView } from "./TraceDrawer";
|
||||
import { useOpenTraceRouting } from "./traceRouting";
|
||||
import { agentHandoffText } from "./tracesApi";
|
||||
import type { Span } from "./traceTypes";
|
||||
import type { Trace } from "./traceTypes";
|
||||
import { traceDisplayName } from "./traceUtils";
|
||||
import { useOpenTraceRouting } from "../routing";
|
||||
import { agentHandoffText } from "../api";
|
||||
import type { Span } from "../types";
|
||||
import type { Trace } from "../types";
|
||||
import { traceDisplayName } from "../utils";
|
||||
|
||||
vi.mock("../../networking", () => ({
|
||||
vi.mock("../../../networking", () => ({
|
||||
agentTraceCall: vi.fn(),
|
||||
agentTraceSpanCall: vi.fn(),
|
||||
getProxyBaseUrl: () => "http://proxy.test/",
|
||||
|
|
@ -38,7 +38,7 @@ vi.mock("@/utils/dataUtils", () => ({ copyToClipboard: vi.fn().mockResolvedValue
|
|||
|
||||
import { copyToClipboard } from "@/utils/dataUtils";
|
||||
|
||||
import { agentTraceCall } from "../../networking";
|
||||
import { agentTraceCall } from "../../../networking";
|
||||
|
||||
const swarm = swarmTrace as Trace;
|
||||
const research = researchTrace as Trace;
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
"use client";
|
||||
import { type TraceHandoff, useTracesApi } from "./tracesApi";
|
||||
import { type TraceHandoff, useTracesApi } from "../api";
|
||||
|
||||
import { QueryErrorResetBoundary, useQueryClient, useSuspenseInfiniteQuery } from "@tanstack/react-query";
|
||||
import { ArrowLeft, Check, Copy } from "lucide-react";
|
||||
|
|
@ -15,16 +15,16 @@ import { cn } from "@/lib/cva.config";
|
|||
import { copyToClipboard } from "@/utils/dataUtils";
|
||||
|
||||
import { DetailPane } from "./DetailPane";
|
||||
import { IdChip } from "./IdChip";
|
||||
import { formatCost } from "./AgentTracesTable";
|
||||
import { SpanIcon } from "./SpanIcon";
|
||||
import { restartsTraversal } from "./readFailure";
|
||||
import { IdChip } from "../ui/IdChip";
|
||||
import { formatCost } from "../list/AgentTracesTable";
|
||||
import { SpanIcon } from "../ui/SpanIcon";
|
||||
import { restartsTraversal } from "../readFailure";
|
||||
import { SpanTree } from "./SpanTree";
|
||||
import { TraceConversation } from "./TraceConversation";
|
||||
import { FrameworkLogo, traceFramework } from "./TraceFramework";
|
||||
import { type RunSelection, traceKey } from "./traceRouting";
|
||||
import type { SpanTreeState, TreeRow } from "./traceTree";
|
||||
import type { Trace } from "./traceTypes";
|
||||
import { FrameworkLogo, traceFramework } from "../ui/TraceFramework";
|
||||
import { type RunSelection, traceKey } from "../routing";
|
||||
import type { SpanTreeState, TreeRow } from "../tree";
|
||||
import type { Trace } from "../types";
|
||||
import {
|
||||
buildTreeRows,
|
||||
firstErrorSpan,
|
||||
|
|
@ -37,7 +37,7 @@ import {
|
|||
revealSpanInState,
|
||||
traceAgentNames,
|
||||
traceDisplayName,
|
||||
} from "./traceUtils";
|
||||
} from "../utils";
|
||||
|
||||
const INITIAL_STATE: SpanTreeState = {
|
||||
hideFramework: true,
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import { buildConversation, conversationSteps, newConversationMessages } from "./conversation";
|
||||
import type { Span, SpanDetail, TraceMessage } from "./traceTypes";
|
||||
import research from "./__fixtures__/research_trace.json";
|
||||
import type { Span, SpanDetail, TraceMessage } from "../types";
|
||||
import research from "../__fixtures__/research_trace.json";
|
||||
|
||||
const root = { ...research.spans[0], span_id: "root", parent_span_id: null, type: "agent" } as Span;
|
||||
const user: TraceMessage = { role: "user", content: "Find my order" };
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
import type { Span, SpanDetail, TraceMessage, TraceToolCall, UIContent } from "./traceTypes";
|
||||
import { isFrameworkSpan, parseJson, parseMessages, prettyPayload } from "./traceUtils";
|
||||
import type { Span, SpanDetail, TraceMessage, TraceToolCall, UIContent } from "../types";
|
||||
import { isFrameworkSpan, parseJson, parseMessages, prettyPayload } from "../utils";
|
||||
|
||||
export const CONVERSATION_PAGE_SIZE = 20;
|
||||
|
||||
|
|
@ -3,8 +3,8 @@
|
|||
import { useQuery, type UseQueryOptions } from "@tanstack/react-query";
|
||||
import moment from "moment";
|
||||
|
||||
import { uiSpendLogsCall } from "../../networking";
|
||||
import type { LogEntry } from "../columns";
|
||||
import { uiSpendLogsCall } from "../../../networking";
|
||||
import type { LogEntry } from "../../../logs/types";
|
||||
|
||||
/** Spend-log timestamps are written when the call finishes, so pad the span start on both sides. */
|
||||
const LOOKUP_PAD_MINUTES = 30;
|
||||
|
|
@ -4,8 +4,8 @@ import moment from "moment";
|
|||
import { useMemo, useState } from "react";
|
||||
|
||||
import { AgentTracesSection } from "./AgentTracesSection";
|
||||
import { useTracesLive } from "./tracesApi";
|
||||
import { useRangeHoursRouting } from "./traceRouting";
|
||||
import { useTracesLive } from "../api";
|
||||
import { useRangeHoursRouting } from "../routing";
|
||||
|
||||
const TIME_FORMAT = "YYYY-MM-DDTHH:mm:ss";
|
||||
|
||||
|
|
@ -5,16 +5,16 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
|||
|
||||
import { ApiError } from "@/lib/http/client";
|
||||
|
||||
import { renderWithProviders, testQueryClient } from "../../../../tests/test-utils";
|
||||
import traceList from "./__fixtures__/trace_list.json";
|
||||
import { LensPreviewContext } from "./LensPreviewButton";
|
||||
import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils";
|
||||
import traceList from "../__fixtures__/trace_list.json";
|
||||
import { LensPreviewContext } from "../../ui/LensPreviewButton";
|
||||
import AgentTracesPage from "./AgentTracesPage";
|
||||
import { filterRuns } from "./runSearch/runQuery";
|
||||
import { AgentTracesSection, type TimeControls } from "./AgentTracesSection";
|
||||
import type { TracePage, TraceSummary } from "./traceTypes";
|
||||
import type { TracePage, TraceSummary } from "../types";
|
||||
import { RESULTS_CHANGED_MESSAGE } from "./useAgentTraces";
|
||||
|
||||
vi.mock("../../networking", () => ({
|
||||
vi.mock("../../../networking", () => ({
|
||||
apiClient: { get: vi.fn(), post: vi.fn() },
|
||||
agentTraceListCall: vi.fn(),
|
||||
sendOtlpTraceCall: vi.fn(),
|
||||
|
|
@ -23,7 +23,7 @@ vi.mock("../../networking", () => ({
|
|||
getProxyBaseUrl: () => "http://localhost:4000",
|
||||
}));
|
||||
|
||||
vi.mock("./TraceDrawer", () => ({
|
||||
vi.mock("../detail/TraceDrawer", () => ({
|
||||
RunView: ({ traceId, onBack }: { traceId: string; onBack: () => void }) => (
|
||||
<div data-testid="run-view">
|
||||
run {traceId}
|
||||
|
|
@ -32,7 +32,7 @@ vi.mock("./TraceDrawer", () => ({
|
|||
),
|
||||
}));
|
||||
|
||||
import { agentTraceListCall, apiClient } from "../../networking";
|
||||
import { agentTraceListCall, apiClient } from "../../../networking";
|
||||
|
||||
const runs = (traceList as TracePage).items as TraceSummary[];
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue