From 0b74ae9c5cfc9d3e0fab8c4a87a4dcce5d28c1f3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 00:53:58 -0700 Subject: [PATCH 1/8] feat(mcp): hand listed-tool description and input schema to pre-call hooks per caller (#41162) * feat(mcp): hand listed-tool metadata to pre-call hooks with per-caller catalog identity Track the tools each MCP server listed per caller identity so pre_mcp_call and during_mcp_call hooks receive the tool description and input schema the client saw. Servers with no caller-dependent inputs share one slot; user identity, forwarded headers, stdio env, relayed bearers, and server-specific auth get their own. Local registry and OpenAPI paths pass the registered metadata and admin description overrides. The Agent 365 guardrail reads the new fields into its evaluate payload. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): drop the listed-tools empty sentinel and routine test docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): mark the listed-tools cache digest as a non-security hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key the listed-tools cache by the OBO subject token token_exchange servers list upstream with the caller's own Entra bearer, so two callers on one LiteLLM key with different subjects were sharing a catalog slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): resolve the BYOK credential before keying the listed-tools slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): drop the OAuth discovery cache when a server definition changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): drop a diff-narrating comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): never validate a supplied header on the tools/list BYOK path The pre-listing resolver ran the tool-call byok_auth_required check even when the caller already supplied x-mcp-auth, and it ran outside the per-server error boundary, so a single deprecated-header caller dropped the server from the aggregate list. Listing now returns a supplied header unchanged and falls back to the stored credential without raising Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): assert the BYOK listing lands in the caller's listed-tool slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): cover the deprecated string x-mcp-auth header on a BYOK tools/list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key the per-caller listed-tool slot by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key discovery cache by the hashed token Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): derive the listed-tool caller identity from the discovery cache key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep OAuth metadata generations only while a fetch is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): return one masked text per scanned string in the selected-guardrail REST test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): fold the signed caller into the discovery digest instead of a second key hash Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): satisfy type discipline gate on listed-tool identity Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): hand tools/call hooks the exact catalog entry tools/list served get_listed_tool re-applied the admin description override on top of the cached listing, so a guardrail-masked description was restored to its original wording at call time, and the OpenAPI / local-registry call path built its metadata from the registry instead of the guarded caller catalog. Both paths now return the cached entry as served, falling back to the registry only when no listing was recorded Adds tests/integration/mcp/test_mcp_listed_tool_metadata.py (red on the prior head for the two regressions, red on the merge base for the feature, green on this head) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): key OpenAPI listed-tool entries per caller so tools/call reads its own guarded listing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): align listed-tool slot tests with per-caller keying Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep oauth2 listing on the minted or signed credential, not the stored BYOK secret The listing helper that keys the per-caller catalog by the stored BYOK credential also handed that credential to the upstream client, which on an oauth2 server short-circuited the client_credentials mint and the MCPJWTSigner gate. Split the two: the catalog identity keeps the stored credential so tools/call finds the caller's slot, while an oauth2 server's tools/list sends only the per-request header, letting the M2M mint or signed JWT proceed as on main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): oauth2 BYOK listing sends the minted token, not the stored secret, through the real proxy Integration cell for the listing fix: a client_credentials BYOK server with a stored user credential, one tools/list as that user, the peer must see a live minted bearer and one /token mint. Red at the pre-fix tip (zero mints, stored secret upstream), green at the fixed head and at the merge base Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list Listing used the resolved stored credential both to key the caller's catalog slot and as the upstream transport header, so REST api_key and bearer_token listings sent the user's secret instead of the server's static token and the MCPJWTSigner gate went quiet. The upstream client and the signer gate now read the caller-supplied mcp_auth_header for every auth type, exactly as before the catalog existed, and the stored credential only names the slot tools/call reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): type the listed-tool metadata read from pre-call kwargs for the basedpyright gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): hand never-listed tools/call hooks name and arguments only The local-registry call path fell back to the registry entry with the admin description override when no tools/list had been recorded for the caller, so a pre_mcp_call guardrail scanned a description the caller was never served and blocked OpenAPI calls that passed before, and base's own selected-guardrail REST test failed on the two-text redaction. _registered_tool_metadata now returns the listed entry or None, so a tools/call with no prior listing sends name and arguments only as promised, and that REST test double goes back to its base shape * fix(mcp): keep during_mcp_call hooks on name and arguments only call_tool handed the caller's listed entry to the during-hook task as well, so during_mcp_call guardrails scanned the description line and schema leaves of any listed tool after the upstream call had already run, blocking calls that passed before whenever the policy matched the description, returned a fixed-length texts list, or hit the depth guard on a deep schema. The listed entry is only disclosed for pre_mcp_call, so the during task no longer receives it and its request object carries no description or schema, as before * fix(mcp): key the BYOK catalog slot by the client's header, not the stored credential tools/list resolved the stored BYOK credential to pick the caller's catalog slot, which read the credential store before the classified try block. With Postgres down and a cold per-worker cache that made every REST tools/list on an is_byok server fail with tools=[] and no upstream call, and the read seeded the per-worker cache (including a negative entry), so a tools/call on another worker after a store, rotate or revoke on this one kept using the stale value. The slot is now keyed by what the client supplied plus the caller's hashed key, on both sides. _get_tools_from_server and call_tool take a keyword-only catalog_auth_header that defaults to mcp_auth_header as received (the default is the builtin Ellipsis so it survives a module reload). The /mcp fan-out and execute_mcp_tool, which swap the resolved credential into mcp_auth_header, pass the client's value explicitly. What goes upstream is unchanged. _byok_catalog_auth_header is gone. * fix(mcp): drop a listed catalog recorded across a server save _record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save description next to the post-save definition until the next listing, instead of name and arguments only. The manager now keeps a per-server listed-tools generation, bumped by _invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and _record_listed_tools skips the write when it moved; the next listing records normally. * fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt add_server and update_server publish the saved definition before the OpenAPI registry entries are rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the new generation. The generation is bumped a second time after the registry refresh. The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is inlined at its two call sites, and the per-server generation map is a plain dict. * fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read add_server and update_server ran the full server-definition invalidation a second time after the awaited OpenAPI spec fetch, which also dropped the prompts/resources/templates discovery entries and the OAuth protected-resource metadata filled under the already-published definition, so the next request went upstream again. Only the listed-tool catalog recorded during the fetch holds pre-save entries, so the post-fetch pass now drops just that catalog and bumps its generation via the new _drop_listed_tools helper, which the full invalidation also calls. * fix(mcp): look a called tool up in the listed catalog by its bare name only get_listed_tool stripped the server prefix a second time when the exact name was absent from the caller's listing, so a never-listed upstream tool whose bare name starts with the server prefix resolved to the listed sibling and that sibling's description and input schema reached the pre-call hooks for a call to a different tool. Every caller already passes the once-stripped bare name, so the lookup is now exact. Tests that looked the catalog up by a prefixed name now use the bare name the callers pass; two new tests pin the never-listed sibling case at the manager and at the tools/call path. * fix(mcp): record a listed-tool catalog only for a listing the caller is served _get_tools_from_server now records the catalog into the caller's listed-tools slot only when asked (record_listing=True), which the served listings pass: the /mcp and Responses API tools/list handlers via _get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST listing via _list_server_tools. Four internal listings stop recording, so a later tools/call hands pre_mcp_call hooks name and arguments only, as on main: - _list_tools_before_first_call, the implicit listing inside tools/call when this worker does not yet expose the tool - fetch_pinnable_tool_catalog, the admin pin snapshot listed without the catalog guard and without description overrides - _initialize_tool_name_to_mcp_server_name_mapping, the startup fill - get_tools_for_server, used by the semantic tool filter _create_prefixed_tools returns to its tool-name mapping job only; the record follows it in _get_tools_from_server. * fix(mcp): opt every listing out of catalog recording unless it is served The aggregate listing and _list_mcp_tools now default to record_listing=False, so a catalog fetched inside a tools/call no longer fills the caller's listed-tools slot. The /mcp/proxy meta-tools (call_tool, search_tools, get_tool_schema) and the tool-search virtual tool stop recording: /mcp/proxy serves only the meta-tools and the search serves only its hits, so a later pre_mcp_call hook was reading a description the caller never listed. The tools/list handler, the Responses MCP handler and the /v1/mcp/tools management listing opt in with record_listing=True, since each serves the catalog to the caller. * fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer The slot a tools/list records for a later tools/call was keyed by (user_id, api_key) only, so every team-only JWT caller shared one slot and one JWT user acting in two teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another caller was served. The slot is now keyed by the hashed key, user, team and organization, plus the admission credential of a caller admitted with neither a key nor a user. The caller bearer split the slot only on client-forwarded-token and token-exchange servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client credentials) also forwards it upstream and served a different catalog per bearer into one slot. The bearer now splits the slot on every server whose egress forwards it (_consumes_caller_authorization) or exchanges it. * fix(mcp): record only tools served by the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep bridge tool metadata request-local * refactor(mcp): centralize listed catalog recording guard * fix(mcp): preserve base TPM reservations for listed tool calls * fix(mcp): preserve project token reservations for listed calls * fix(mcp): record served catalogs and preserve call message bytes * test(mcp): align listing expectation with deferred recording * test(mcp): audit listed metadata across callers and bridge lifecycles * test(mcp): preserve guardrail fixture worker affinity * refactor(mcp): expose listed catalog recording API * feat(mcp): pass served_tools through the anthropic messages bridge /v1/messages auto-execution now hands this request's resolved tool definitions to _execute_tool_calls, matching the Responses and chat completions bridges: the pre_mcp_call hook receives the description and input schema the model was shown for that call. Request-local only; the shared listed-tools catalog is untouched. * test(mcp): pin served_tools handoff on the anthropic messages bridge Mirrors the credentials-forwarding test: the request's resolved tool definitions must reach _execute_tool_calls under served_tools so pre_mcp_call hooks judge the call on the description and input schema the model was shown. Fails without the previous commit's one-liner. * style(mcp): sort the local import block ruff flagged * fix(mcp): keep the admin include_disabled_tools view off the listed-tools catalog GET /mcp-rest/tools/list?include_disabled_tools=true is the admin-only configuration view: apply_tool_filters is False, so it serves the full server catalog. Recording that response into the caller's listed-tools slot warmed tools/call metadata no runtime listing ever served, breaking the only-a-served-listing-records invariant (Bugbot). The record is now gated on apply_tool_filters; disabled tools stay unreachable (the call-time allowlist 403 fires before hooks), so the observable fix is the slot no longer warming from a settings view. Verified live: the new test fails on the unfixed head and passes here, and the rest of the listed-tool-metadata suite is unchanged. --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../pass_through/messages/mcp_handler.py | 1 + .../mcp_server/mcp_server_manager.py | 268 ++- .../_experimental/mcp_server/operations.py | 68 +- .../mcp_server/rest_endpoints.py | 60 +- .../guardrail_hooks/agent_365/agent_365.py | 26 +- .../mcp_management_endpoints.py | 1 + litellm/proxy/utils.py | 15 +- litellm/responses/main.py | 3 + .../responses/mcp/chat_completions_handler.py | 2 + .../mcp/litellm_proxy_mcp_handler.py | 8 + .../responses/mcp/mcp_streaming_iterator.py | 3 + litellm/types/mcp.py | 4 + tests/integration/_support/mcp.py | 14 +- tests/integration/_support/wire.py | 1 + .../integration/mcp/test_mcp_access_matrix.py | 95 +- .../mcp/test_mcp_accounting_guardrails.py | 491 ++++- tests/integration/mcp/test_mcp_credentials.py | 298 ++- tests/integration/mcp/test_mcp_lifecycle.py | 71 +- .../mcp/test_mcp_listed_tool_metadata.py | 531 +++++ .../integration/mcp/test_mcp_llm_endpoints.py | 633 +++++- tests/integration/mcp/test_mcp_oauth_flows.py | 74 +- tests/integration/mcp/test_mcp_resilience.py | 217 ++ tests/integration/mcp/test_mcp_toolsets.py | 68 +- .../pass_through/messages/test_mcp_handler.py | 70 + .../test_mcp_guardrail_usage_monitor.py | 1 + .../mcp_server/test_mcp_proxy_mode.py | 55 + .../mcp_server/test_mcp_server.py | 3 + .../mcp_server/test_mcp_server_manager.py | 1874 +++++++++++++++-- .../test_mcp_server_tool_calls_and_headers.py | 496 ++++- .../mcp_server/test_mcp_tool_search.py | 35 +- .../mcp_server/test_openapi_tool_auth.py | 14 + .../mcp_server/test_operations.py | 113 + .../guardrail_hooks/test_agent_365.py | 23 + .../hooks/test_parallel_request_limiter_v3.py | 156 ++ .../utils/proxy_logging/test_mcp_bridging.py | 20 + .../mcp/test_litellm_proxy_mcp_handler.py | 179 +- 36 files changed, 5585 insertions(+), 406 deletions(-) create mode 100644 tests/integration/mcp/test_mcp_listed_tool_metadata.py diff --git a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py index ab722a60ca5..7961a38dedf 100644 --- a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py +++ b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a776470e2ab..49cd8cf6fb9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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--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"] diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index f98a8c0ea5a..d83a4a72e2f 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7f936ec9269..845dcaa1f16 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 6621d94ed96..c683e1f2c5d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -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), } diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index be4e46f247e..e1504172d9e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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] diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8af90e19ecb..1c738aefb97 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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(), ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 14c12fc571d..5e2793ed40d 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index b17e0befba6..8aa4181cc8a 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -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, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index fb7882adcb4..1d6c0695621 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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"), diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index c0bb92cfb2a..b1f12233f33 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -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, diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index da7401e2a2e..ce86e79a6f0 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -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() diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index a3693433de4..d39d7e4029e 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -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"]) diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index c807852ef6b..10bb7787cc9 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -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() diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index e458911cc68..bf8adbc7cd7 100644 --- a/tests/integration/mcp/test_mcp_access_matrix.py +++ b/tests/integration/mcp/test_mcp_access_matrix.py @@ -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" diff --git a/tests/integration/mcp/test_mcp_accounting_guardrails.py b/tests/integration/mcp/test_mcp_accounting_guardrails.py index 323afad40db..a276718c280 100644 --- a/tests/integration/mcp/test_mcp_accounting_guardrails.py +++ b/tests/integration/mcp/test_mcp_accounting_guardrails.py @@ -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] diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index 95a5e46646b..16dcaab274a 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -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" diff --git a/tests/integration/mcp/test_mcp_lifecycle.py b/tests/integration/mcp/test_mcp_lifecycle.py index fa253f03520..eed3cbf658e 100644 --- a/tests/integration/mcp/test_mcp_lifecycle.py +++ b/tests/integration/mcp/test_mcp_lifecycle.py @@ -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 diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py new file mode 100644 index 00000000000..d09864da4b1 --- /dev/null +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -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" + ) diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 01bf006a03e..af3bff1fd8b 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -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)}",) diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 60563e7aacd..e95cdb902e8 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -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" diff --git a/tests/integration/mcp/test_mcp_resilience.py b/tests/integration/mcp/test_mcp_resilience.py index 8efb54a18fd..d821181962f 100644 --- a/tests/integration/mcp/test_mcp_resilience.py +++ b/tests/integration/mcp/test_mcp_resilience.py @@ -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]) + ) diff --git a/tests/integration/mcp/test_mcp_toolsets.py b/tests/integration/mcp/test_mcp_toolsets.py index 3dd309db665..b97fa6e72cb 100644 --- a/tests/integration/mcp/test_mcp_toolsets.py +++ b/tests/integration/mcp/test_mcp_toolsets.py @@ -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" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py index 37db61031c9..fc9d77a10a2 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py @@ -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(): """ diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index e1e4cd3d161..8e87837611a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py index ed5d67164bd..e52a86d76af 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index d3679506a2f..316988ef175 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cf6cf93f35b..cb017afbea5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6,6 +6,7 @@ import json import logging import os import sys +import time from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path @@ -41,8 +42,10 @@ from mcp.types import Tool as MCPTool from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT +from litellm.proxy._experimental.mcp_server import discoverable_endpoints from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + ListedToolsCaller, MCPServerManager, _deserialize_json_dict, _flow_endpoints_missing, @@ -53,6 +56,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _obo_retry_applies, _resolve_openapi_tool_auth, _should_strip_caller_authorization, + listed_tools_caller_for, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -116,7 +120,6 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} - @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -218,8 +221,6 @@ def _reload_mcp_manager_module(): return reloaded - - @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1576,7 +1577,9 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( + self, config_only_mcp_manager_factory, monkeypatch, caplog + ): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4862,7 +4865,9 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): + async def test_openapi_health_loads_spec_without_mcp_handshake( + self, respx_mock, monkeypatch, auth_type, is_byok, scheme + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4912,14 +4917,28 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), + ( + httpx.Response(401, text="secret response content"), + "unhealthy", + "OpenAPI specification request failed (HTTP 401)", + ), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), - (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), + ( + httpx.ConnectError("secret network details"), + "unhealthy", + "OpenAPI specification could not be loaded (ConnectError)", + ), + ( + httpx.Response(200, text="secret invalid JSON body"), + "unhealthy", + "OpenAPI specification could not be loaded (JSONDecodeError)", + ), ], ) - async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): + async def test_openapi_health_reports_safe_failures( + self, respx_mock, monkeypatch, failure, expected_status, expected_error + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -5077,7 +5096,10 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"]) async def test_health_check_server_oauth2_reports_reachability( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + oauth2_flow: Literal["authorization_code", "client_credentials"] | None, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() @@ -5106,14 +5128,28 @@ class TestMCPServerManager: assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type", [ - MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token, - MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, - ]) + @pytest.mark.parametrize( + "auth_type", + [ + MCPAuth.bearer_token, + MCPAuth.api_key, + MCPAuth.basic, + MCPAuth.authorization, + MCPAuth.token, + MCPAuth.oauth2_token_exchange, + MCPAuth.oauth2_id_jag, + MCPAuth.true_passthrough, + MCPAuth.oauth_delegate, + ], + ) @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse]) @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503]) async def test_health_check_without_credentials_accepts_any_http_response( - self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse], + self, + monkeypatch: pytest.MonkeyPatch, + respx_mock: MockRouter, + auth_type: MCPAuthType, + transport: Literal[MCPTransport.http, MCPTransport.sse], response_code: int, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") @@ -5144,6 +5180,7 @@ class TestMCPServerManager: self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + class UnreadBody(httpx.AsyncByteStream): def __init__(self) -> None: self.read = False @@ -5158,17 +5195,28 @@ class TestMCPServerManager: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events", + server_id="streaming-health", + name="streaming-health", + transport=MCPTransport.sse, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/events", ) manager.registry[server.server_id] = server bodies: Final = (UnreadBody(), UnreadBody()) - route: Final = respx_mock.get(server.url).mock(side_effect=[ - httpx.Response(response_code, stream=body, headers={ - "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/", - "Location": "http://127.0.0.1/private", - }) for body in bodies - ]) + route: Final = respx_mock.get(server.url).mock( + side_effect=[ + httpx.Response( + response_code, + stream=body, + headers={ + "Content-Type": "text/event-stream", + "Set-Cookie": "health=secret; Path=/", + "Location": "http://127.0.0.1/private", + }, + ) + for body in bodies + ] + ) first: Final = await manager.health_check_server(server.server_id) second: Final = await manager.health_check_server(server.server_id) @@ -5179,19 +5227,28 @@ class TestMCPServerManager: assert all("cookie" not in call.request.headers for call in route.calls) @pytest.mark.asyncio - @pytest.mark.parametrize(("transport", "url"), [ - (MCPTransport.stdio, "https://mcp.example.test"), - (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"), - (MCPTransport.http, "ftp://mcp.example.test"), - (MCPTransport.http, "https://user:secret@mcp.example.test"), - (MCPTransport.http, "https://mcp.example.test:bad/mcp"), - ]) + @pytest.mark.parametrize( + ("transport", "url"), + [ + (MCPTransport.stdio, "https://mcp.example.test"), + (MCPTransport.http, None), + (MCPTransport.http, ""), + (MCPTransport.http, "not-a-url"), + (MCPTransport.http, "ftp://mcp.example.test"), + (MCPTransport.http, "https://user:secret@mcp.example.test"), + (MCPTransport.http, "https://mcp.example.test:bad/mcp"), + ], + ) async def test_health_reachability_rejects_unprobeable_urls_without_requests( self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None ) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url, + server_id="unprobeable", + name="unprobeable", + transport=transport, + auth_type=MCPAuth.oauth2, + url=url, ) manager.registry[server.server_id] = server @@ -5202,19 +5259,26 @@ class TestMCPServerManager: assert not respx_mock.calls @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [ - httpx.ConnectError("TLS/connection failure with secret details"), - httpx.ReadTimeout("secret timeout details"), - httpx.RemoteProtocolError("secret malformed response"), - ]) + @pytest.mark.parametrize( + "failure", + [ + httpx.ConnectError("TLS/connection failure with secret details"), + httpx.ReadTimeout("secret timeout details"), + httpx.RemoteProtocolError("secret malformed response"), + ], + ) async def test_health_reachability_reports_no_response_without_secret_details( self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="failed-health", name="failed-health", transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret", + server_id="failed-health", + name="failed-health", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + is_byok=True, + url="https://mcp.example.test/secret?token=secret", ) manager.registry[server.server_id] = server route: Final = respx_mock.get(server.url).mock(side_effect=failure) @@ -5230,8 +5294,11 @@ class TestMCPServerManager: monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher") manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="bad-tls", name="bad-tls", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test", + server_id="bad-tls", + name="bad-tls", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test", ) manager.registry[server.server_id] = server @@ -5249,8 +5316,11 @@ class TestMCPServerManager: monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1) manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="slow-health", name="slow-health", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow", + server_id="slow-health", + name="slow-health", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/slow", ) manager.registry[server.server_id] = server started: Final = asyncio.Event() @@ -5303,8 +5373,11 @@ class TestMCPServerManager: server_ids: Final = [f"health-{index}" for index in range(server_count)] manager.registry = { server_id: MCPServer( - server_id=server_id, name=server_id, transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}", + server_id=server_id, + name=server_id, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=f"https://health.example.test/{server_id}", ) for server_id in server_ids } @@ -5631,8 +5704,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5717,8 +5797,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers @@ -7067,8 +7154,7 @@ class TestMCPServerManager: # Mock _create_mcp_client to return our mock client manager._create_mcp_client = AsyncMock(return_value=mock_client) - # Mock user auth with no restrictions - user_api_key_auth: Final = UserAPIKeyAuth() + user_api_key_auth: Final = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = MagicMock() @@ -7094,6 +7180,1051 @@ class TestMCPServerManager: # Verify the MCP client call was awaited exactly once assert mock_client.call_tool.await_count == 1 + @staticmethod + def _manager_ready_for_call_tool( + listed_tools: list[MCPTool], caller: ListedToolsCaller | None = None + ) -> tuple[MCPServerManager, MagicMock]: + from mcp.types import CallToolResult + + manager = MCPServerManager() + server = MCPServer( + server_id="test-server", + name="test-server", + transport=MCPTransport.http, + url="http://test-server.com", + ) + manager.registry = {"test-server": server} + manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" + manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" + manager._create_prefixed_tools(listed_tools, server) + manager._record_listed_tools(server, listed_tools, caller) + + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + return manager, proxy_logging_obj + + @staticmethod + def _unrestricted_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test") + + @pytest.mark.asyncio + async def test_call_tool_hands_listed_tool_description_and_schema_to_pre_call_hooks(self): + schema = {"type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"]} + listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] + auth = self._unrestricted_auth() + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + listed, caller=ListedToolsCaller(user_api_key_auth=auth) + ) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema) + + @pytest.mark.asyncio + async def test_call_tool_hands_during_call_hooks_name_and_arguments_only_even_for_a_listed_tool(self): + """A during_mcp_call guardrail evaluates the call in flight, so it keeps seeing only the name and + arguments it always did; the listed description and schema go to the pre-call hooks alone.""" + schema = {"type": "object", "properties": {"param": {"type": "string"}}} + listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)] + auth = UserAPIKeyAuth(api_key="sk-test") + manager, _ = self._manager_ready_for_call_tool(listed, caller=ListedToolsCaller(user_api_key_auth=auth)) + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + during_data = proxy_logging_obj.during_call_hook.call_args.kwargs["data"] + assert during_data["mcp_arguments"] == {"param": "value"} + assert (during_data.get("mcp_tool_description"), during_data.get("mcp_input_schema")) == (None, None) + assert "Description:" not in during_data["messages"][0]["content"] + + @pytest.mark.asyncio + async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self): + auth = self._unrestricted_auth() + manager, proxy_logging_obj = self._manager_ready_for_call_tool( + [MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})], + caller=ListedToolsCaller(user_api_key_auth=auth), + ) + + await manager.call_tool( + server_name="test-server", + name="test_tool", + arguments={"param": "value"}, + user_api_key_auth=auth, + proxy_logging_obj=proxy_logging_obj, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) + + def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) + + latest = manager.get_listed_tool(server, "echo") + assert latest is not None and latest.description == "v2" + assert manager.get_listed_tool(server, "missing") is None + + def test_get_listed_tool_never_strips_the_bare_name_it_is_given(self): + """The lookup is exact: a never-listed tool whose bare name starts with the server prefix is not the + listed sibling that stripping the prefix again would name.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) + manager._record_listed_tools( + server, + [ + MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), + MCPTool(name="bar", description="Fetches bar records", inputSchema={"type": "object"}), + ], + caller, + ) + + assert manager.get_listed_tool(server, "srv-foo", caller) is None + listed = manager.get_listed_tool(server, "foo", caller) + assert listed is not None and listed.description == "Fetches foo records" + + @pytest.mark.asyncio + async def test_get_listed_tool_uses_admin_description_override_clients_saw(self): + schema = {"type": "object", "properties": {"text": {"type": "string"}}} + manager = _catalog_manager( + MCPTool(name="echo", description="Upstream wording", inputSchema=schema), + MCPTool(name="ping", description="Untouched", inputSchema={}), + ) + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + tool_name_to_description={"echo": "Admin wording"}, + ) + await manager._get_tools_from_server(server, add_prefix=True, record_listing=True) + + overridden = manager.get_listed_tool(server, "echo") + assert overridden is not None + assert (overridden.name, overridden.description, overridden.input_schema) == ("echo", "Admin wording", schema) + untouched = manager.get_listed_tool(server, "ping") + assert untouched is not None and untouched.description == "Untouched" + + @pytest.mark.asyncio + async def test_get_listed_tool_keeps_the_masked_description_over_the_admin_override(self, catalog_guardrail): + """A discovery guardrail masked the admin override in tools/list, so the tool-call hooks must see + the masked wording, not the original override the caller never saw.""" + _, proxy_logging_obj = catalog_guardrail + manager = _catalog_manager(MCPTool(name="read_note", description="Read a note", inputSchema={"type": "object"})) + server = MCPServer( + server_id="notes", + name="notes", + transport=MCPTransport.http, + tool_name_to_description={"read_note": "Read a SECRET note"}, + ) + served = await manager._get_tools_from_server( + server, add_prefix=True, proxy_logging_obj=proxy_logging_obj, record_listing=True + ) + assert [tool.description for tool in served] == ["Read a [MASKED] note"] + + listed = manager.get_listed_tool(server, "read_note") + assert listed is not None and listed.description == "Read a [MASKED] note", ( + "tools/call must be evaluated against the description tools/list served" + ) + + def test_server_definition_change_drops_listed_tools(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") + manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) + + manager._invalidate_server_definition_caches(server.server_id) + + assert manager.get_listed_tool(server, "echo") is None + kept = manager.get_listed_tool(other, "ping") + assert kept is not None and kept.description == "kept" + + @pytest.mark.asyncio + async def test_server_save_during_an_in_flight_listing_is_not_undone_by_the_stale_record(self): + """A PUT /v1/mcp/server that lands while a listing awaits its upstream fetch drops the server's + catalog; the fetch completing afterwards must not write the pre-save catalog back, or hooks see + the old description next to the new definition until the next listing.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="saver") + fetch_started = asyncio.Event() + release_fetch = asyncio.Event() + + async def fetch(client, name): + fetch_started.set() + await release_fetch.wait() + return [MCPTool(name="turn", description="before save", inputSchema={})] + + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = fetch + caller = ListedToolsCaller(user_api_key_auth=user) + + async def list_tools() -> None: + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + + listing = asyncio.create_task(list_tools()) + await fetch_started.wait() + manager._invalidate_server_definition_caches(server.server_id) + release_fetch.set() + await listing + + assert manager.get_listed_tool(server, "turn", caller) is None + + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="after save", inputSchema={})] + ) + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "after save" + + @pytest.mark.asyncio + async def test_update_server_refreshing_openapi_tools_drops_a_listing_recorded_during_the_spec_fetch(self): + """An OpenAPI server's registry entries are rebuilt after the save is published, so a listing that + records while the spec is fetched holds the pre-save entries; the catalog is dropped again once the + registry is current.""" + manager = MCPServerManager() + old = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json" + ) + manager.registry[old.server_id] = old + new = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json" + ) + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) + + async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None: + manager._record_listed_tools( + server, + [MCPTool(name="search", description="pre-save", inputSchema={})], + caller, + manager._listed_tools_generations.get(server.server_id, 0), + ) + + manager.build_mcp_server_from_table = AsyncMock(return_value=new) + manager._maybe_register_openapi_tools = register_while_a_listing_records + manager.prime_oauth_metadata_discovery = MagicMock() + record = LiteLLM_MCPServerTable( + server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http + ) + + await manager.update_server(record) + + assert manager.registry["srv"] is new + assert manager.get_listed_tool(new, "search", caller) is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("already_registered", [False, True], ids=["add_server", "update_server"]) + async def test_openapi_spec_re_read_keeps_discovery_and_oauth_metadata_filled_during_the_fetch( + self, already_registered: bool + ): + """The listed-tool catalog recorded during the spec fetch holds pre-save entries, but a prompts + discovery or OAuth protected-resource fetch answered in that window already saw the published + definition; dropping those too sends the next request upstream again.""" + manager = MCPServerManager() + if already_registered: + manager.registry["srv"] = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json" + ) + new = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json" + ) + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) + metadata_key: Final = (new.server_id, new.url) + prompt_fetches = 0 + + async def fetch_prompts() -> list[Prompt]: + nonlocal prompt_fetches + prompt_fetches += 1 + return [Prompt(name="greet")] + + async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: + manager._record_listed_tools( + server, + [MCPTool(name="search", description="pre-save", inputSchema={})], + caller, + manager._listed_tools_generations.get(server.server_id, 0), + ) + await manager._prompt_discovery_cache.get((server.server_id, None), fetch_prompts) + discoverable_endpoints._OAUTH_METADATA_CACHE[metadata_key] = (time.time() + 300, {"resource": new.url}) + + manager.build_mcp_server_from_table = AsyncMock(return_value=new) + manager._maybe_register_openapi_tools = register_while_discovery_fills + manager.prime_oauth_metadata_discovery = MagicMock() + record = LiteLLM_MCPServerTable( + server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http + ) + save = manager.update_server if already_registered else manager.add_server + + try: + await save(record) + + assert manager.registry["srv"] is new + assert manager.get_listed_tool(new, "search", caller) is None + prompts = await manager._prompt_discovery_cache.get((new.server_id, None), fetch_prompts) + assert [prompt.name for prompt in prompts] == ["greet"] + assert prompt_fetches == 1, "the prompts list filled after the save was published went upstream again" + cached_metadata = discoverable_endpoints._OAUTH_METADATA_CACHE.get(metadata_key) + assert cached_metadata is not None and cached_metadata[1] == {"resource": new.url} + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(metadata_key, None) + + @pytest.mark.asyncio + async def test_user_oauth_refresh_keeps_listed_tools(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) + + await manager.invalidate_user_oauth_token_cache("alice", server.server_id) + + listed = manager.get_listed_tool(server, "echo") + assert listed is not None and listed.description == "shared" + + def test_per_caller_server_keeps_listed_tools_per_identity(self): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + auth_type=MCPAuth.oauth2_token_exchange, + ) + alice = UserAPIKeyAuth(user_id="alice", token="hashed-alice") + bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") + alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} + bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} + manager._record_listed_tools( + server, + [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], + ListedToolsCaller(user_api_key_auth=alice), + ) + manager._record_listed_tools( + server, + [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], + ListedToolsCaller(user_api_key_auth=bob), + ) + + alice_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=alice)) + bob_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=bob)) + assert alice_tool is not None and (alice_tool.description, alice_tool.input_schema) == ( + "alice view", + alice_schema, + ) + assert bob_tool is not None and (bob_tool.description, bob_tool.input_schema) == ("bob view", bob_schema) + carol = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="carol", token="k")) + assert manager.get_listed_tool(server, "read", carol) is None + + shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") + manager._record_listed_tools( + shared, + [MCPTool(name="echo", description="everyone", inputSchema={})], + ListedToolsCaller(user_api_key_auth=alice), + ) + for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob)) + assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth" + anonymous = manager.get_listed_tool(shared, "echo") + assert anonymous is None + + @pytest.mark.parametrize( + ("server_kwargs", "caller_a", "caller_b"), + [ + pytest.param( + {"extra_headers": ["X-Workspace"]}, + ListedToolsCaller(raw_headers={"x-workspace": "A"}), + ListedToolsCaller(raw_headers={"X-Workspace": "B"}), + id="forwarded-header", + ), + pytest.param( + {"auth_type": MCPAuth.true_passthrough}, + ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-a"}), + ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-b"}), + id="anonymous-passthrough-bearer", + ), + pytest.param( + {"auth_type": MCPAuth.bearer_token}, + ListedToolsCaller(mcp_auth_header="byok-a"), + ListedToolsCaller(mcp_auth_header="byok-b"), + id="per-server-auth-header", + ), + pytest.param( + {"auth_type": MCPAuth.oauth2_token_exchange}, + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-alice"}, + ), + ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(user_id="team-bot", token="hashed-shared"), + raw_headers={"x-litellm-api-key": "sk-shared", "authorization": "Bearer entra-bob"}, + ), + id="shared-key-different-obo-subjects", + ), + pytest.param( + {"transport": MCPTransport.stdio, "command": "srv", "env": {"WS": "${X-WS}"}}, + ListedToolsCaller(raw_headers={"X-WS": "A"}), + ListedToolsCaller(raw_headers={"X-WS": "B"}), + id="header-driven-stdio-env", + ), + ], + ) + def test_upstream_identity_inputs_keep_listed_tools_apart(self, server_kwargs, caller_a, caller_b): + manager = MCPServerManager() + server = MCPServer( + **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} + ) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) + + for_a = manager.get_listed_tool(server, "turn", caller_a) + for_b = manager.get_listed_tool(server, "turn", caller_b) + assert for_a is not None and for_a.description == "Catalog A" + assert for_b is not None and for_b.description == "Catalog B" + assert manager.get_listed_tool(server, "turn", ListedToolsCaller()) is None + + def test_shared_server_ignores_headers_it_never_forwards(self): + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager._record_listed_tools( + server, + [MCPTool(name="turn", description="everyone", inputSchema={})], + ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), + ) + + other = ListedToolsCaller(raw_headers={"authorization": "Bearer sk-other", "x-workspace": "B"}) + listed = manager.get_listed_tool(server, "turn", other) + assert listed is not None and listed.description == "everyone" + + @pytest.mark.asyncio + async def test_byok_listing_never_reads_the_credential_store(self): + """tools/list keys the caller's catalog slot by what the client supplied plus the caller's key. + Resolving the stored BYOK credential for that would fail every REST listing while the DB is + down and would seed a per-worker cache the next tools/call trusts over the store.""" + manager = MCPServerManager() + server = MCPServer( + server_id="byok-cold", + name="byok_cold", + transport=MCPTransport.http, + url="http://byok-cold", + is_byok=True, + auth_type=MCPAuth.api_key, + ) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-cold-user") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="listed while db down", inputSchema={})] + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.proxy._experimental.mcp_server.db.get_user_credential", + AsyncMock(side_effect=RuntimeError("DB DOWN")), + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + + listed = manager.get_listed_tool(server, "turn", listed_tools_caller_for(server, user, None, None, None, None)) + assert listed is not None and listed.description == "listed while db down" + + @pytest.mark.parametrize( + ("list_header", "call_kwargs"), + [ + pytest.param( + None, {"mcp_auth_header": "stored-secret", "catalog_auth_header": None}, id="execute-mcp-tool" + ), + pytest.param(None, {"mcp_auth_header": None}, id="responses-api"), + pytest.param("Bearer hdr", {"mcp_auth_header": "Bearer hdr"}, id="client-supplied-header"), + ], + ) + @pytest.mark.asyncio + async def test_byok_tools_call_reads_the_slot_the_clients_own_header_listed( + self, list_header: str | None, call_kwargs: dict[str, str | None] + ): + """A REST listing records under the header the client sent (none here). tools/call then swaps the + stored credential in, either before reaching ``call_tool`` (``execute_mcp_tool``) or inside it (the + Responses API), and must still read that slot rather than one keyed by the credential.""" + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache_key, + cache_byok_credential, + ) + from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache + + manager = MCPServerManager() + server = MCPServer( + server_id="byok-catalog", + name="byok_catalog", + transport=MCPTransport.http, + url="http://byok-catalog", + is_byok=True, + ) + manager.registry = {"byok-catalog": server} + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] + ) + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + cache_byok_credential("byok-user", "byok-catalog", "stored-secret") + try: + await manager._get_tools_from_server( + server=server, mcp_auth_header=list_header, user_api_key_auth=user, record_listing=True + ) + listed = manager.get_listed_tool( + server, "turn", listed_tools_caller_for(server, user, list_header, None, None, None) + ) + assert listed is not None and listed.description == "stored cred catalog" + + await manager.call_tool( + server_name="byok_catalog", + name="turn", + arguments={}, + user_api_key_auth=user, + proxy_logging_obj=proxy_logging_obj, + **call_kwargs, + ) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert hook_kwargs["tool_description"] == "stored cred catalog" + + @pytest.mark.asyncio + async def test_byok_supplied_header_lists_without_credential_validation(self): + manager = MCPServerManager() + server = MCPServer( + server_id="byok-catalog", + name="byok_catalog", + transport=MCPTransport.http, + url="http://byok-catalog", + is_byok=True, + ) + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="t", inputSchema={})] + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + await manager._get_tools_from_server( + server=server, + mcp_auth_header="Bearer hdr", + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), + record_listing=True, + ) + + caller: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), mcp_auth_header="Bearer hdr" + ) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "t" + assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + + @pytest.mark.parametrize( + "server_auth", + [ + pytest.param( + { + "auth_type": MCPAuth.oauth2, + "client_id": "cid", + "client_secret": "csec", + "token_url": "http://cc1/token", + }, + id="oauth2", + ), + pytest.param({"auth_type": MCPAuth.api_key, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="api_key"), + pytest.param( + {"auth_type": MCPAuth.bearer_token, "authentication_token": "STATIC-ADMIN-TOKEN"}, id="bearer_token" + ), + pytest.param({"auth_type": MCPAuth.none}, id="none"), + ], + ) + @pytest.mark.asyncio + async def test_byok_listing_keys_the_catalog_by_the_caller_and_never_touches_the_stored_secret( + self, server_auth: dict[str, object] + ): + """The caller's key plus what the caller supplied (nothing here) keys the catalog slot tools/call + reads, even with the stored BYOK secret at hand in the cache, and tools/list sends upstream exactly + what the caller supplied, so the static token, the M2M mint and MCPJWTSigner all behave as they + did before the catalog existed, whatever the auth_type.""" + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache_key, + cache_byok_credential, + ) + from litellm.proxy._experimental.mcp_server.operations import byok_credential_cache + + manager = MCPServerManager() + server = MCPServer( + server_id="cc1", + name="cc1", + transport=MCPTransport.http, + url="http://cc1", + is_byok=True, + **server_auth, + ) + alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="echo", description="listed catalog", inputSchema={})] + ) + signer_headers = AsyncMock(return_value={"Authorization": "Bearer signed-jwt"}) + cache_byok_credential("alice", "cc1", "BYOK-ALICE-SECRET") + try: + with ( + patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ), + patch( # test-quality-ok: same singleton's header injection, asserted on by call + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.inject_mcp_jwt_headers_for_upstream", + signer_headers, + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=alice, record_listing=True) + finally: + byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) + + client_kwargs = manager._create_mcp_client.await_args.kwargs + assert client_kwargs["mcp_auth_header"] is None, client_kwargs + assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} + signer_headers.assert_awaited_once() + listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=alice)) + assert listed is not None and listed.description == "listed catalog" + + @pytest.mark.parametrize( + ("signer", "static_headers"), + [ + pytest.param(MagicMock(), None, id="signer"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, id="static-authorization"), + pytest.param(None, None, id="no-signer"), + ], + ) + def test_keyed_callers_always_list_into_their_own_slot(self, signer, static_headers): + """The catalog is guardrail-shaped per key, so a keyed caller never reads another caller's + listing regardless of the signer or static authorization configuration.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers + ) + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="alice", token="hashed-alice")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", token="hashed-bob")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=signer, + ): + manager._record_listed_tools( + server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice + ) + assert manager.get_listed_tool(server, "turn", bob) is None + + for_alice = manager.get_listed_tool(server, "turn", alice) + assert for_alice is not None and for_alice.description == "alice view" + + def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self): + """Two keys sharing a user_id get different signed JWTs, so they split; the same key + presented again lands on its own slot.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-beta")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=MagicMock(), + ): + manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) + assert manager.get_listed_tool(server, "turn", bob) is None + + same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) + listed = manager.get_listed_tool(server, "turn", same_key) + + assert listed is not None and listed.description == "slot a" + + def test_listed_tools_slot_is_split_per_team_for_keyless_callers(self): + """A team-only JWT admits a caller with neither a key nor a user, so the team keys the slot.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + team_one: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one") + ) + team_two: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two") + ) + manager._record_listed_tools( + server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one + ) + + assert manager.get_listed_tool(server, "foo", team_two) is None + listed: Final = manager.get_listed_tool(server, "foo", team_one) + assert listed is not None and listed.description == "Fetch rows FLAGWORD" + + def test_listed_tools_slot_is_split_per_team_for_the_same_keyless_user(self): + """One JWT user acting in two teams is served two team-shaped catalogs, so each team is a slot.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice_in_one: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-one") + ) + alice_in_two: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two") + ) + manager._record_listed_tools( + server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one + ) + + assert manager.get_listed_tool(server, "foo", alice_in_two) is None + listed: Final = manager.get_listed_tool(server, "foo", alice_in_one) + assert listed is not None and listed.description == "Fetch rows FLAGWORD" + + def test_listed_tools_slot_is_split_by_the_admission_bearer_of_keyless_callers_without_a_user(self): + """Two team-only JWT callers of one team differ only in the JWT they were admitted with, so that + credential keys the slot, on a server that never forwards it.""" + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + alice: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), + raw_headers={"authorization": "Bearer jwt-alice"}, + ) + bob: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), + raw_headers={"authorization": "Bearer jwt-bob"}, + ) + manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) + + assert manager.get_listed_tool(server, "foo", bob) is None + listed: Final = manager.get_listed_tool(server, "foo", alice) + assert listed is not None and listed.description == "alice view" + + @pytest.mark.parametrize( + ("server_kwargs", "forwards_bearer"), + [ + pytest.param( + {"auth_type": MCPAuth.oauth2, "delegate_auth_to_upstream": True, "oauth2_flow": "authorization_code"}, + True, + id="oauth2-delegated-to-upstream", + ), + pytest.param({"auth_type": MCPAuth.oauth_delegate}, True, id="oauth-delegate"), + pytest.param({"auth_type": MCPAuth.true_passthrough}, True, id="true-passthrough"), + pytest.param({"auth_type": MCPAuth.oauth2_token_exchange}, True, id="token-exchange"), + pytest.param( + {"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True}, + True, + id="oauth-passthrough", + ), + pytest.param({}, False, id="plain"), + pytest.param( + {"auth_type": MCPAuth.oauth2, "oauth2_flow": "authorization_code"}, + False, + id="oauth2-gateway-managed", + ), + pytest.param( + { + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "client_credentials", + "delegate_auth_to_upstream": True, + "client_id": "gateway", + "client_secret": "secret", + "token_url": "http://idp/token", + }, + False, + id="oauth2-client-credentials", + ), + ], + ) + def test_listed_tools_slot_is_split_by_the_forwarded_bearer_on_servers_that_forward_it( + self, server_kwargs: dict[str, object], forwards_bearer: bool + ): + """Two callers sharing one key but carrying different upstream bearers are served two upstream + catalogs exactly on the servers whose egress forwards or exchanges that bearer.""" + manager: Final = MCPServerManager() + server: Final = MCPServer( + **{"server_id": "dg", "name": "dg", "transport": MCPTransport.http, "url": "http://dg", **server_kwargs} + ) + caller_a: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), + raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-A"}, + ) + caller_b: Final = ListedToolsCaller( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), + raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"}, + ) + manager._record_listed_tools( + server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a + ) + + for_b: Final = manager.get_listed_tool(server, "lookup", caller_b) + assert (for_b is None) is forwards_bearer + for_a: Final = manager.get_listed_tool(server, "lookup", caller_a) + assert for_a is not None and for_a.description == "Workspace A lookup FLAGWORD" + + @pytest.mark.asyncio + async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): + manager = MCPServerManager() + server = MCPServer( + server_id="catalog", + name="catalog", + transport=MCPTransport.http, + url="http://catalog", + extra_headers=["X-Workspace"], + ) + manager.registry = {"catalog": server} + catalogs = { + "A": [ + MCPTool( + name="turn", description="Catalog A", inputSchema={"properties": {"turn": {"description": "A"}}} + ) + ], + "B": [ + MCPTool( + name="turn", description="Catalog B", inputSchema={"properties": {"turn": {"description": "B"}}} + ) + ], + } + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager._fetch_tools_with_timeout = AsyncMock(side_effect=lambda client, name: catalogs[client.workspace]) + for workspace in ("A", "B"): + manager._create_mcp_client.return_value.workspace = workspace + await manager._get_tools_from_server( + server=server, + extra_headers={"X-Workspace": workspace}, + raw_headers={"x-workspace": workspace, "authorization": "Bearer sk-litellm"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"), + record_listing=True, + ) + + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + await manager.call_tool( + server_name="catalog", + name="turn", + arguments={"turn": "A-1"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"), + proxy_logging_obj=proxy_logging_obj, + raw_headers={"x-workspace": "A", "authorization": "Bearer sk-litellm"}, + ) + + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ( + "Catalog A", + {"properties": {"turn": {"description": "A"}}}, + ) + + def test_per_caller_listed_tools_evict_oldest_caller_and_keep_shared(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _LISTED_TOOLS_CALLERS_PER_SERVER + + manager = MCPServerManager() + server = MCPServer( + server_id="srv", + name="srv", + transport=MCPTransport.http, + url="http://srv", + auth_type=MCPAuth.oauth2_token_exchange, + ) + manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) + callers = [ + ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) + for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) + ] + for caller in callers: + manager._record_listed_tools( + server, + [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], + caller, + ) + manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) + + assert manager.get_listed_tool(server, "read", callers[0]) is None + second = manager.get_listed_tool(server, "read", callers[1]) + assert second is not None and second.description == "u1 again" + newest = manager.get_listed_tool(server, "read", callers[-1]) + assert newest is not None and newest.description == callers[-1].user_api_key_auth.user_id + assert len(manager._listed_tools_by_server_id[server.server_id]) == _LISTED_TOOLS_CALLERS_PER_SERVER + 1 + shared = manager.get_listed_tool(server, "read") + assert shared is not None and shared.description == "shared" + + @pytest.mark.asyncio + @pytest.mark.parametrize("add_prefix", [True, False]) + async def test_openapi_listing_records_listed_tools(self, add_prefix): + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + server = MCPServer( + server_id="petstore-id", + name="petstore", + alias="petstore", + transport=MCPTransport.http, + url=None, + spec_path="/spec.yaml", + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + global_mcp_tool_registry.register_tool( + name="petstore-list_pets", + description="List pets", + input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}}, + handler=_handler, + ) + try: + listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix, record_listing=True) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("petstore-") + + assert [t.name for t in listed] == ["petstore-list_pets" if add_prefix else "list_pets"] + tool = manager.get_listed_tool(server, "list_pets") + assert tool is not None and tool.description == "List pets" + assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + + @pytest.mark.asyncio + async def test_openapi_listing_ignores_overlapping_server_prefix(self): + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + server = MCPServer( + server_id="pet-id", + name="pet", + alias="pet", + transport=MCPTransport.http, + url=None, + spec_path="/spec.yaml", + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + + async def _handler(**kwargs): + return None + + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + global_mcp_tool_registry.register_tool( + name="pet-petstore-list", + description="Local pet tool", + input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}}, + handler=_handler, + ) + global_mcp_tool_registry.register_tool( + name="petstore-list", + description="Foreign petstore tool", + input_schema={"type": "object", "properties": {"status": {"type": "string"}}}, + handler=_handler, + ) + try: + listed = await manager._get_tools_from_server(server=server, add_prefix=True, record_listing=True) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in listed] == ["pet-petstore-list"] + tool = manager.get_listed_tool(server, "petstore-list") + assert tool is not None and tool.description == "Local pet tool" + assert tool.input_schema["properties"] == {"limit": {"type": "integer"}} + + @pytest.mark.asyncio + @pytest.mark.parametrize("openapi", [False, True], ids=["remote", "openapi"]) + async def test_get_tools_from_server_records_the_catalog_only_when_asked_to(self, openapi): + """The startup fill, the implicit pre-call listing and the pin snapshot reuse this fetch without + serving its result, so only a listing that asks to be recorded sets what tools/call hooks see.""" + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + if openapi: + server = MCPServer( + server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + global_mcp_tool_registry.register_tool( + name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None + ) + else: + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + try: + listed = await manager._get_tools_from_server(server=server, user_api_key_auth=user) + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + + await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("srv-") + + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_list_tools_records_the_served_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + manager.get_allowed_mcp_servers = AsyncMock(return_value=["srv"]) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister") + + listed = await manager.list_tools(user_api_key_auth=user) + + assert [t.name for t in listed] == ["srv-echo"] + recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) + assert recorded is not None and recorded.description == "Echoes" + + @pytest.mark.asyncio + async def test_startup_tool_name_mapping_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + await manager._initialize_tool_name_to_mcp_server_name_mapping() + + assert manager.server_exposes_tool(server, "echo") is True + assert server.server_id not in manager._listed_tools_by_server_id + assert manager.get_listed_tool(server, "echo") is None + + @pytest.mark.asyncio + async def test_get_tools_for_server_records_no_listed_catalog(self): + manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"})) + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + manager.registry = {"srv": server} + + listed = await manager.get_tools_for_server("srv") + + assert [t.name for t in listed] == ["srv-echo"] + assert server.server_id not in manager._listed_tools_by_server_id + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_with_user_api_key_auth(self): """ @@ -9739,15 +10870,11 @@ class TestGetPublicMCPServers: if registered_in == "both" else server ) - manager.config_mcp_servers = ( - {server.server_id: config_server} if registered_in in ("config", "both") else {} - ) + manager.config_mcp_servers = {server.server_id: config_server} if registered_in in ("config", "both") else {} manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} original_server: Final = server.model_dump() original_config_server: Final = config_server.model_dump() - expected_public: Final = registered_in != "neither" and ( - public_ids == [server.server_id] or implicitly_public - ) + expected_public: Final = registered_in != "neither" and (public_ids == [server.server_id] or implicitly_public) with ( patch("litellm.public_mcp_servers", public_ids), @@ -9755,9 +10882,7 @@ class TestGetPublicMCPServers: ): public_servers: Final = manager.get_public_mcp_servers() assert manager.is_mcp_server_public(server.server_id) is expected_public - assert [item.server_id for item in public_servers] == ( - [server.server_id] if expected_public else [] - ) + assert [item.server_id for item in public_servers] == ([server.server_id] if expected_public else []) assert manager.is_mcp_server_public("server-alias") is False assert manager.is_mcp_server_public("missing-server") is False assert manager.is_mcp_server_public(server.server_id, public_ids=frozenset()) is ( @@ -10657,7 +11782,9 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): + async def call_tool( + self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False + ): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -12992,7 +14119,9 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( + self, config_only_mcp_manager_factory, aliasing_entry_first: bool + ): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -13008,7 +14137,9 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( + self, config_only_mcp_manager_factory + ): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -13096,7 +14227,9 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( + self, config_only_mcp_manager_factory, caplog + ): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -13267,7 +14400,9 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another( + self, config_only_mcp_manager_factory, caplog + ): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13639,7 +14774,8 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner( + _mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -13758,12 +14894,16 @@ async def test_debug_reports_legacy_signing_and_non_http_transport( async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -13783,13 +14923,18 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + server_id="repeated-stale", + name="stale", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -13809,13 +14954,20 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + server_id="resolved-replacement", + name="replacement", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy( + update={ + "url": "https://new.example.com/mcp", + "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + } ) - replacement: Final = original.model_copy(update={ - "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -13823,8 +14975,11 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", name="publication", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="stale-publication", + name="publication", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -13840,9 +14995,13 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + server_id="expiring-session", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13943,7 +15102,9 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + assert ( + result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + ) assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13955,8 +15116,11 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, + server_id="cancelled-cache", + name="cancelled-cache", + transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", + auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -14085,7 +15249,9 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + return MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http + ) @pytest.mark.asyncio @@ -14253,7 +15419,9 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) +@pytest.mark.parametrize( + "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) +) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -14570,26 +15738,45 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,credential", [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ]) + @pytest.mark.parametrize( + "auth_type,credential", + [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ], + ) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, credential: str | None, dispatch: str, + self, + tmp_path: Path, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, + credential: str | None, + dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}})) + spec_path.write_text( + json.dumps( + { + "openapi": "3.0.0", + "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}, + } + ) + ) server: Final = MCPServer( - server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, + server_id="dispatch-auth", + name="dispatch-auth", + url="https://upstream.example", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14612,14 +15799,21 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", - transport=transport, auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", + server_id="incomplete-obo", + name="incomplete-obo", + url="https://upstream.example/mcp", + transport=transport, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", + client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", + authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header="Bearer override", subject_token=subject, + server, + mcp_auth_header="Bearer override", + subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -14632,8 +15826,11 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-static", + name="empty-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -14641,16 +15838,22 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,headers", [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ]) + @pytest.mark.parametrize( + "auth_type,headers", + [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ], + ) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", name="header-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="header-static", + name="header-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -14659,29 +15862,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="openapi-empty", + name="openapi-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, - user_api_key_auth=None, forwarded_headers=None, + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=None, + forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,slot,value", [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ]) + @pytest.mark.parametrize( + "auth_type,slot,value", + [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ], + ) async def test_raw_static_credentials_are_forwarded_unchanged( - self, auth_type: MCPAuthType, slot: str, value: str, + self, + auth_type: MCPAuthType, + slot: str, + value: str, ) -> None: - server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) + server = MCPServer( + server_id="raw-key", + name="raw-key", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, + ) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -14695,17 +15917,24 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, respx_mock: MockRouter, value: str, source: str, + self, + respx_mock: MockRouter, + value: str, + source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.authorization, + server_id="raw-empty", + name="raw-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, + server, + mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -14713,9 +15942,15 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, - token_exchange_endpoint="https://idp.example/token") + server = MCPServer( + server_id="obo-byok", + name="obo-byok", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + is_byok=True, + token_exchange_endpoint="https://idp.example/token", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -14723,41 +15958,66 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) + server = MCPServer( + server_id="override", + name="override", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=configured, + ) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) + server = MCPServer( + server_id="empty-header", + name="empty-header", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=token, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", authentication_token="key") + server = MCPServer( + server_id="custom", + name="custom", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", + authentication_token="key", + ) client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize("static_headers,accepted", [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ]) + @pytest.mark.parametrize( + "static_headers,accepted", + [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ], + ) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, + server_id="static-slot", + name="static-slot", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -14769,21 +16029,36 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize("static,forwarded,caller", [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ]) + @pytest.mark.parametrize( + "static,forwarded,caller", + [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ], + ) async def test_openapi_static_credentials_remain_supported( - self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None + self, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], + forwarded: dict[str, str] | None, + caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, _request_extra_headers, create_tool_function, + _request_auth_header, + _request_extra_headers, + create_tool_function, ) + tool: Final = create_tool_function( - "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, + "/echo", + "get", + {}, + "https://upstream.example", + headers=static, + auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -14817,8 +16092,13 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key) + server = MCPServer( + server_id="cancel", + name="cancel", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + ) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -14827,8 +16107,14 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") + server = MCPServer( + server_id="blank-static", + name="blank-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=" ", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -14836,8 +16122,13 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic) + server = MCPServer( + server_id="bad-basic", + name="bad-basic", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -14846,34 +16137,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None) + server = MCPServer( + server_id="basic-scheme", + name="basic-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,default_slot", [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ]) + @pytest.mark.parametrize( + "auth_type,value,default_slot", + [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", name="alternate", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", + server_id="alternate", + name="alternate", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, + server, + mcp_auth_header=value if source == "caller" else None, + extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -14882,8 +16187,12 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", + server_id="both-empty", + name="both-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -14896,12 +16205,17 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, + server_id="caller-auth", + name="caller-auth", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=headers if source == "caller" else None, + server, + mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -14910,14 +16224,29 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize("value", [ - "", " ", "Bearer", "Basic", "token", "ApiKey", - "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", - ]) + @pytest.mark.parametrize( + "value", + [ + "", + " ", + "Bearer", + "Basic", + "token", + "ApiKey", + "Bearer Bearer", + "ApiKey ApiKey", + "token token", + "bEaReR BEARER", + "aPiKeY\tAPIKEY", + ], + ) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, + server_id="caller-empty", + name="caller-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -14928,8 +16257,11 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, + server_id="basic-pair", + name="basic-pair", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14942,8 +16274,12 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, + server_id="basic-valid", + name="basic-valid", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14952,17 +16288,27 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value", [ - (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), - ]) + @pytest.mark.parametrize( + "auth_type,value", + [ + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.bearer_token, "Bearer "), + (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), + (MCPAuth.token, "token "), + (MCPAuth.token, "TOKEN"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-scheme", + name="empty-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14970,17 +16316,24 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,expected", [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ]) + @pytest.mark.parametrize( + "auth_type,value,expected", + [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ], + ) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", name="real-token", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, + server_id="real-token", + name="real-token", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -15019,16 +16372,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = {"observer": MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", - )} + manager.registry = { + "observer": MCPServer( + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", + ) + } manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for(manager.call_tool( - server_name="observer", name="execute", arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), - ), timeout=5) + result = await asyncio.wait_for( + manager.call_tool( + server_name="observer", + name="execute", + arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context( + {"metadata": {"guardrails": ["observe"] if selected else []}} + ), + ), + timeout=5, + ) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -15057,11 +16425,21 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + upstream = MCPServer( + server_id="explicit-empty", + name="explicit_empty", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + legacy_server.set_auth_context( + UserAPIKeyAuth(user_id="unrelated"), + raw_headers={"authorization": "unrelated-credential"}, + client_ip="192.0.2.99", + ) with ( patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -15069,7 +16447,9 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await MCPServerManager()._create_mcp_client( + upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None + ) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -15092,16 +16472,28 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-a", + server_name="alpha", + alias="shared", + url="https://a.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-b", + server_name="beta", + alias="Shared", + url="https://b.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-c", + server_name="gamma", + alias="lonely", + url="https://c.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] @@ -15193,7 +16585,9 @@ async def test_reload_warns_once_about_a_blocked_stdio_row_that_is_rebuilt_every @pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): manager = config_only_mcp_manager_factory() - await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}}) + await manager.load_servers_from_config( + {"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}} + ) server = next(iter(manager.config_mcp_servers.values())) client = await manager._create_mcp_client(server) assert server.protocol_version == revision @@ -15205,11 +16599,15 @@ async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_m def test_runtime_protocol_metadata_preserves_explicit_precedence( revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None ) -> None: - server: Final = MCPServer.model_validate({ - "server_id": "preview", "name": "preview", "transport": "http", - "mcp_info": {"protocol_version": revision}, - **({"protocol_version": explicit} if explicit is not None else {}), - }) + server: Final = MCPServer.model_validate( + { + "server_id": "preview", + "name": "preview", + "transport": "http", + "mcp_info": {"protocol_version": revision}, + **({"protocol_version": explicit} if explicit is not None else {}), + } + ) assert server.protocol_version == (explicit if explicit is not None else revision) @@ -15387,9 +16785,7 @@ class TestToolCatalogGuard: proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=asyncio.CancelledError) with pytest.raises(asyncio.CancelledError): - await manager._get_tools_from_server( - _notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj - ) + await manager._get_tools_from_server(_notes_server(), add_prefix=False, proxy_logging_obj=proxy_logging_obj) proxy_logging_obj.pre_call_hook.assert_awaited_once() proxy_logging_obj.slack_alerting_instance.send_alert.assert_not_awaited() @@ -15446,7 +16842,10 @@ class TestToolCatalogGuard: guardrail, proxy_logging_obj = catalog_guardrail monkeypatch.setattr(signer_module, "_mcp_jwt_signer_instance", None) signer = signer_module.MCPJWTSigner( - guardrail_name="jwt-signer", event_hook="pre_mcp_call", default_on=True, issuer="https://litellm.example.com" + guardrail_name="jwt-signer", + event_hook="pre_mcp_call", + default_on=True, + issuer="https://litellm.example.com", ) monkeypatch.setattr(litellm, "callbacks", [signer, guardrail]) manager = _catalog_manager(LIST_NOTES, POISONED_DELETE) @@ -15666,7 +17065,10 @@ class TestToolCatalogGuard: widened = MCPTool( name="read_note", description="Read a note", - inputSchema={"type": "object", "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}}, + inputSchema={ + "type": "object", + "properties": {"id": {"type": "string"}, "callback_url": {"type": "string"}}, + }, ) manager = _catalog_manager(widened) @@ -15728,8 +17130,12 @@ class TestToolCatalogGuard: return "ok" with patch.dict(global_mcp_tool_registry.tools, {}, clear=True): - global_mcp_tool_registry.register_tool("petstore-list_pets", "List pets, newest first", {"type": "object"}, handler) - global_mcp_tool_registry.register_tool("petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler) + global_mcp_tool_registry.register_tool( + "petstore-list_pets", "List pets, newest first", {"type": "object"}, handler + ) + global_mcp_tool_registry.register_tool( + "petstore-delete_pets", POISONED_DELETE.description, {"type": "object"}, handler + ) global_mcp_tool_registry.register_tool("petstore-find_pet", "Find a pet", {"type": "object"}, handler) served = await manager._get_tools_from_server( server, add_prefix=add_prefix, proxy_logging_obj=proxy_logging_obj diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 4cc7794d4ad..545b2757ffd 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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()) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1cfe6198b6f..1c987778da6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -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] diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index a8ec7be55f0..60157193a16 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -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", diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index b16b27ac919..f710bc7f1d7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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 diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py index 23131321938..e72b716665c 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -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([]) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index 8b9439e25d2..d3a723fe1d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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): diff --git a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py index c5bd89645b4..25be3b5de6b 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py @@ -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 = { diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 2699f9445c9..73bee304fc2 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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={}), From bb4f7211d71e0c757b6a3f16ba3885fc34376028 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 01:31:23 -0700 Subject: [PATCH 2/8] refactor: clean up fresh tech debt from 2026-10-03 (#44484) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/decisions/main.py | 2 +- litellm/harness/handlers/tool_loop_handler.py | 8 ++------ litellm/harness/options.py | 2 +- litellm/llms/anthropic/files/transformation.py | 11 +++++------ 4 files changed, 9 insertions(+), 14 deletions(-) diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py index da037f1d8cb..4864fcaa0e4 100644 --- a/litellm/decisions/main.py +++ b/litellm/decisions/main.py @@ -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", diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py index 33efcba4ff7..d3acb50eadb 100644 --- a/litellm/harness/handlers/tool_loop_handler.py +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -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", diff --git a/litellm/harness/options.py b/litellm/harness/options.py index b0014fea393..18865359ba8 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -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 diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index 04d057ed3e9..da21f130f1a 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -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 From b4fcc5c1bc3e92a21bd291ddb74185e78d1c637c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 03:17:07 -0700 Subject: [PATCH 3/8] fix(otel): read the registered v2 logger without importing the proxy (#44485) * fix(otel): read the registered v2 logger without importing the proxy phase_span, which the router enters on every deployment pick since #44150, looked up the proxy's OTel logger by importing litellm.proxy.proxy_server. In an SDK process that import loads the whole proxy synchronously on the caller's event loop during its first request, which stalled litellm_router_unit_testing past its 5s wait. Read the module from sys.modules instead: when the proxy was never imported it has no registered logger. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): pin the no-proxy-import guarantee in-process Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/otel/logger.py | 11 +++++---- tests/unit/integrations/otel/test_runtime.py | 24 ++++++++++++++++++++ 2 files changed, 31 insertions(+), 4 deletions(-) diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 96df9728a73..ca8e507bff5 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -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 diff --git a/tests/unit/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py index 285759c7553..5747be27656 100644 --- a/tests/unit/integrations/otel/test_runtime.py +++ b/tests/unit/integrations/otel/test_runtime.py @@ -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 == [] From a2bf67a03707e474be41e07806ee4fd9791ca7cd Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 03:18:54 -0700 Subject: [PATCH 4/8] test(mcp): keep the SSO assertion round trip from matching its refresh token inside random ciphertext (#44494) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../outbound_credentials/test_sso_assertion_store.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py index 5d6d47b8c38..088b232f3b6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -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 From 6532dcb73ba57f09bb465280964b0a51c83f244a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 08:40:23 -0700 Subject: [PATCH 5/8] test(ci): pin the ROI estimator flag and the prompt-cache counter in two drifted tests (#44499) Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/litellm_utils_tests/test_utils.py | 8 ++++++-- .../gcs_pub_sub_body/spend_logs_payload.json | 2 +- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index 92947fbf6fe..64402b5c016 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -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=[ { diff --git a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json index 63baadaaf31..f7ebf357d83 100644 --- a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json +++ b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json @@ -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, From 05f1c73a3c5bc98b5073ae730dc372f556a0c243 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 10:00:48 -0700 Subject: [PATCH 6/8] refactor(ui): move TraceView into components/lens/traces (#44501) Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/eslint.config.mjs | 2 +- .../lens/LensPage.integration.test.tsx | 2 +- .../src/components/lens/LensWorkspace.tsx | 8 +++--- .../src/components/lens/data/LensServices.tsx | 2 +- .../lens/data/demo/createLensDemo.ts | 2 +- .../src/components/lens/data/demo/fixtures.ts | 2 +- .../lens/data/demo/lensDemoLongTrace.ts | 2 +- .../components/lens/hooks/useLensReadiness.ts | 2 +- .../lens/investigations/Evidence.tsx | 4 +-- .../FindingDetails.integration.test.tsx | 4 +-- .../InvestigationsView.integration.test.tsx | 2 +- .../investigations/InvestigationsView.tsx | 2 +- .../lens/onboarding/OnboardingContext.tsx | 2 +- .../lens/onboarding/OnboardingSteps.tsx | 2 +- .../TracingSetupCard.integration.test.tsx | 8 +++--- .../onboarding/tracing}/TracingSetupCard.tsx | 14 +++++----- .../onboarding/tracing}/previewTrace.json | 0 .../onboarding/tracing}/sampleTrace.test.ts | 0 .../onboarding/tracing}/sampleTrace.ts | 0 .../onboarding/tracing}/tracingSetupGuides.ts | 26 ++++++++--------- .../src/components/lens/route.ts | 2 +- .../__fixtures__/deep_agent_trace.json | 0 .../traces}/__fixtures__/research_trace.json | 0 .../traces}/__fixtures__/swarm_trace.json | 0 .../traces}/__fixtures__/trace_list.json | 0 .../tracesApi.ts => lens/traces/api.ts} | 2 +- .../traces/detail}/AttributesDetail.tsx | 2 +- .../traces/detail}/DetailContent.tsx | 8 +++--- .../detail}/DetailPane.integration.test.tsx | 12 ++++---- .../traces/detail}/DetailPane.tsx | 18 ++++++------ .../traces/detail}/KeyValueRows.tsx | 2 +- .../traces/detail}/MessageCard.test.tsx | 2 +- .../traces/detail}/MessageCard.tsx | 6 ++-- .../traces/detail}/RequestDetail.tsx | 10 +++---- .../traces/detail}/SpanHoverCard.tsx | 10 +++---- .../traces/detail}/SpanTree.tsx | 12 ++++---- .../TraceConversation.integration.test.tsx | 12 ++++---- .../traces/detail}/TraceConversation.tsx | 8 +++--- .../traces/detail}/TraceDrawer.test.tsx | 20 ++++++------- .../traces/detail}/TraceDrawer.tsx | 18 ++++++------ .../traces/detail}/conversation.test.ts | 4 +-- .../traces/detail}/conversation.ts | 4 +-- .../traces/detail}/useSpanRequestLog.ts | 4 +-- .../traces/list}/AgentTracesPage.tsx | 4 +-- .../AgentTracesSection.integration.test.tsx | 14 +++++----- .../traces/list}/AgentTracesSection.tsx | 10 +++---- .../traces/list}/AgentTracesTable.test.tsx | 8 +++--- .../traces/list}/AgentTracesTable.tsx | 12 ++++---- .../traces/list}/TimeRangeControls.tsx | 0 .../traces/list}/TracesTimeline.test.ts | 2 +- .../traces/list}/TracesTimeline.tsx | 4 +-- .../traces/list}/runSearch/RunSearch.test.tsx | 0 .../traces/list}/runSearch/RunSearch.tsx | 2 +- .../traces/list}/runSearch/RunsToolbar.tsx | 2 +- .../list}/runSearch/__fixtures__/runs.ts | 2 +- .../traces/list}/runSearch/runQuery.test.ts | 0 .../traces/list}/runSearch/runQuery.ts | 4 +-- .../traces/list}/runSearch/runSql.test.ts | 0 .../traces/list}/runSearch/runSql.ts | 0 .../traces/list}/useAgentTraces.test.ts | 2 +- .../traces/list}/useAgentTraces.ts | 9 +++--- .../traces/routing.ts} | 6 ++-- .../traceTree.ts => lens/traces/tree.ts} | 4 +-- .../traceTypes.ts => lens/traces/types.ts} | 0 .../traces/ui}/ActiveDot.tsx | 0 .../TraceView => lens/traces/ui}/Collapse.tsx | 0 .../traces/ui}/CopyButton.tsx | 0 .../TraceView => lens/traces/ui}/IdChip.tsx | 0 .../TraceView => lens/traces/ui}/PaneBar.tsx | 0 .../TraceView => lens/traces/ui}/SpanIcon.tsx | 2 +- .../traces/ui}/StatusMark.tsx | 0 .../traces/ui}/TraceFramework.test.ts | 0 .../traces/ui}/TraceFramework.tsx | 28 +++++++++---------- .../traces/ui}/spanProvider.test.ts | 0 .../traces/ui}/spanProvider.ts | 2 +- .../traces/utils.test.ts} | 6 ++-- .../traceUtils.ts => lens/traces/utils.ts} | 4 +-- .../ui}/LensPreviewButton.tsx | 0 .../src/components/networking.tsx | 2 +- .../src/lib/http/api.test-d.ts | 2 +- 80 files changed, 187 insertions(+), 186 deletions(-) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/TracingSetupCard.integration.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/TracingSetupCard.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/previewTrace.json (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/sampleTrace.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/sampleTrace.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/onboarding/tracing}/tracingSetupGuides.ts (93%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces}/__fixtures__/deep_agent_trace.json (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces}/__fixtures__/research_trace.json (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces}/__fixtures__/swarm_trace.json (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces}/__fixtures__/trace_list.json (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/tracesApi.ts => lens/traces/api.ts} (99%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/AttributesDetail.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/DetailContent.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/DetailPane.integration.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/DetailPane.tsx (94%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/KeyValueRows.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/MessageCard.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/MessageCard.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/RequestDetail.tsx (92%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/SpanHoverCard.tsx (96%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/SpanTree.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/TraceConversation.integration.test.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/TraceConversation.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/TraceDrawer.test.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/TraceDrawer.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/conversation.test.ts (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/conversation.ts (99%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/detail}/useSpanRequestLog.ts (93%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/AgentTracesPage.tsx (93%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/AgentTracesSection.integration.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/AgentTracesSection.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/AgentTracesTable.test.tsx (91%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/AgentTracesTable.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/TimeRangeControls.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/TracesTimeline.test.ts (98%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/TracesTimeline.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/RunSearch.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/RunSearch.tsx (96%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/RunsToolbar.tsx (94%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/__fixtures__/runs.ts (94%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/runQuery.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/runQuery.ts (93%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/runSql.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/runSearch/runSql.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/useAgentTraces.test.ts (95%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/list}/useAgentTraces.ts (95%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/traceRouting.ts => lens/traces/routing.ts} (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/traceTree.ts => lens/traces/tree.ts} (86%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/traceTypes.ts => lens/traces/types.ts} (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/ActiveDot.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/Collapse.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/CopyButton.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/IdChip.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/PaneBar.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/SpanIcon.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/StatusMark.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/TraceFramework.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/TraceFramework.tsx (68%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/spanProvider.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/traces/ui}/spanProvider.ts (97%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/traceUtils.test.ts => lens/traces/utils.test.ts} (99%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView/traceUtils.ts => lens/traces/utils.ts} (99%) rename ui/litellm-dashboard/src/components/{view_logs/TraceView => lens/ui}/LensPreviewButton.tsx (100%) diff --git a/ui/litellm-dashboard/eslint.config.mjs b/ui/litellm-dashboard/eslint.config.mjs index ad2923bf2a0..ce2d933292d 100644 --- a/ui/litellm-dashboard/eslint.config.mjs +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -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" }, }, diff --git a/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx index 320dbef4142..32e4e069cc1 100644 --- a/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx @@ -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 }) =>
Trace polling {isActive ? "active" : "paused"}
, })); vi.mock("./investigations/InvestigationsView", () => ({ diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index caa8576bdd7..ccd2f665564 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -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 }; diff --git a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx index 8abeab16a78..9ef2c316225 100644 --- a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx +++ b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx @@ -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 { diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts index 97ab6c84e91..c0b3aa50f41 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts b/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts index 51dc89e6fa6..becdff35d40 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts b/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts index e30cff6ee1d..866f1042b73 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts b/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts index a41c9b909c3..36dc613d5ab 100644 --- a/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts +++ b/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx b/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx index 1023e636537..05017e76b68 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx index a8abcf9bc25..93942cb8fa9 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx @@ -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 }) => (
{traceId} at {selection.spanId} diff --git a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx index 5d90281a080..fded81f230a 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx index c668bba3799..a99133d6ca2 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingContext.tsx b/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingContext.tsx index c7a1a05c6f4..bcbf8ebaa6e 100644 --- a/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingContext.tsx +++ b/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingContext.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingSteps.tsx b/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingSteps.tsx index 2b321a46c2e..165908bc77d 100644 --- a/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingSteps.tsx +++ b/ui/litellm-dashboard/src/components/lens/onboarding/OnboardingSteps.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.integration.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.integration.test.tsx rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.integration.test.tsx index 26e0be89104..abb535a5cb2 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.integration.test.tsx @@ -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(), diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.tsx index dcb458fc367..5d608e2908b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TracingSetupCard.tsx +++ b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/TracingSetupCard.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/previewTrace.json b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/previewTrace.json similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/previewTrace.json rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/previewTrace.json diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/sampleTrace.test.ts b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/sampleTrace.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/sampleTrace.test.ts rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/sampleTrace.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/sampleTrace.ts b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/sampleTrace.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/sampleTrace.ts rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/sampleTrace.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/tracingSetupGuides.ts b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/tracingSetupGuides.ts similarity index 93% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/tracingSetupGuides.ts rename to ui/litellm-dashboard/src/components/lens/onboarding/tracing/tracingSetupGuides.ts index af85c38cb9b..474d80d0d48 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/tracingSetupGuides.ts +++ b/ui/litellm-dashboard/src/components/lens/onboarding/tracing/tracingSetupGuides.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/lens/route.ts b/ui/litellm-dashboard/src/components/lens/route.ts index cd16165d5a5..e4b05c9ca71 100644 --- a/ui/litellm-dashboard/src/components/lens/route.ts +++ b/ui/litellm-dashboard/src/components/lens/route.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/deep_agent_trace.json b/ui/litellm-dashboard/src/components/lens/traces/__fixtures__/deep_agent_trace.json similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/deep_agent_trace.json rename to ui/litellm-dashboard/src/components/lens/traces/__fixtures__/deep_agent_trace.json diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/research_trace.json b/ui/litellm-dashboard/src/components/lens/traces/__fixtures__/research_trace.json similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/research_trace.json rename to ui/litellm-dashboard/src/components/lens/traces/__fixtures__/research_trace.json diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/swarm_trace.json b/ui/litellm-dashboard/src/components/lens/traces/__fixtures__/swarm_trace.json similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/swarm_trace.json rename to ui/litellm-dashboard/src/components/lens/traces/__fixtures__/swarm_trace.json diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/trace_list.json b/ui/litellm-dashboard/src/components/lens/traces/__fixtures__/trace_list.json similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/__fixtures__/trace_list.json rename to ui/litellm-dashboard/src/components/lens/traces/__fixtures__/trace_list.json diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/tracesApi.ts b/ui/litellm-dashboard/src/components/lens/traces/api.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/tracesApi.ts rename to ui/litellm-dashboard/src/components/lens/traces/api.ts index 33f6e7166a5..03ea4d068bb 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/tracesApi.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/api.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AttributesDetail.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/AttributesDetail.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/AttributesDetail.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/AttributesDetail.tsx index 00371bacfc4..4461a3ca429 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AttributesDetail.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/AttributesDetail.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailContent.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/DetailContent.tsx index 187fa5ca83f..b9c0d9bd0b3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailContent.tsx @@ -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 = { tool: "Tool", model: "Model", litellm: "LiteLLM" }; const TRACEBACK_MARKER = "Traceback (most recent call last):"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.integration.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.integration.test.tsx index 1bc6d371841..907a5a67129 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.integration.test.tsx @@ -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 & Pick; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx similarity index 94% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx index f8232d28192..80f9b3da47c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx @@ -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 "../ui/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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/KeyValueRows.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/KeyValueRows.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/KeyValueRows.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/KeyValueRows.tsx index aaf6a935ffe..6eb8217df3b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/KeyValueRows.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/KeyValueRows.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.test.tsx index 9cb94ee969c..5583e937ed7 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.test.tsx @@ -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(); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx index 310c6278ee4..cd3ea968d1a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/MessageCard.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx @@ -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 "../ui/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 = { user: "User", system: "System", assistant: "Assistant", tool: "Tool" }; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx similarity index 92% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx index 6664516cff4..58177ede37a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx @@ -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 "../../../view_logs/LogDetailsDrawer"; +import { formatCost } from "../list/AgentTracesTable"; import { DetailGroup } from "./AttributesDetail"; -import { CopyButton } from "./CopyButton"; +import { CopyButton } from "../ui/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 { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanHoverCard.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/SpanHoverCard.tsx similarity index 96% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/SpanHoverCard.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/SpanHoverCard.tsx index 1eaa1124f47..6d16a8fa713 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanHoverCard.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/SpanHoverCard.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanTree.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/SpanTree.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/SpanTree.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/SpanTree.tsx index 69143bf1b26..f225d92c6a1 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanTree.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/SpanTree.tsx @@ -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[]; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.integration.test.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.integration.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.integration.test.tsx index 8a1249472b8..1e8dc0b7c9a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.integration.test.tsx @@ -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, "selection">) { const { selection } = useOpenTraceRouting(); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.tsx index 6d41be9a262..d8dc950767e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceConversation.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceConversation.tsx @@ -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 "../ui/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, diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.test.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.test.tsx index 0eb9889f2d8..c2cc01353e2 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.test.tsx @@ -3,19 +3,19 @@ import { focusManager, onlineManager } from "@tanstack/react-query"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -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/", @@ -36,7 +36,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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx rename to ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.tsx index e93fe5dc9c0..4047218177d 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/TraceDrawer.tsx @@ -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,15 +15,15 @@ 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 { IdChip } from "../ui/IdChip"; +import { formatCost } from "../list/AgentTracesTable"; +import { SpanIcon } from "../ui/SpanIcon"; 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, @@ -36,7 +36,7 @@ import { revealSpanInState, traceAgentNames, traceDisplayName, -} from "./traceUtils"; +} from "../utils"; const INITIAL_STATE: SpanTreeState = { hideFramework: true, diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.test.ts b/ui/litellm-dashboard/src/components/lens/traces/detail/conversation.test.ts similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/detail/conversation.test.ts index 055cdf28288..729fdaffaf4 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.test.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/conversation.test.ts @@ -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" }; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.ts b/ui/litellm-dashboard/src/components/lens/traces/detail/conversation.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.ts rename to ui/litellm-dashboard/src/components/lens/traces/detail/conversation.ts index 9631bc4a12b..af80191f698 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/conversation.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/conversation.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/useSpanRequestLog.ts b/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts similarity index 93% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/useSpanRequestLog.ts rename to ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts index 083dceac790..c654ca55793 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/useSpanRequestLog.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts @@ -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 "../../../view_logs/columns"; /** Spend-log timestamps are written when the call finishes, so pad the span start on both sides. */ const LOOKUP_PAD_MINUTES = 30; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesPage.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesPage.tsx similarity index 93% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesPage.tsx rename to ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesPage.tsx index 8d8e78866d1..dfc4b770994 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesPage.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesPage.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx index bf68f31f0cb..be295533651 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx @@ -5,15 +5,15 @@ 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"; -vi.mock("../../networking", () => ({ +vi.mock("../../../networking", () => ({ apiClient: { get: vi.fn(), post: vi.fn() }, agentTraceListCall: vi.fn(), sendOtlpTraceCall: vi.fn(), @@ -22,7 +22,7 @@ vi.mock("../../networking", () => ({ getProxyBaseUrl: () => "http://localhost:4000", })); -vi.mock("./TraceDrawer", () => ({ +vi.mock("../detail/TraceDrawer", () => ({ RunView: ({ traceId, onBack }: { traceId: string; onBack: () => void }) => (
run {traceId} @@ -31,7 +31,7 @@ vi.mock("./TraceDrawer", () => ({ ), })); -import { agentTraceListCall, apiClient } from "../../networking"; +import { agentTraceListCall, apiClient } from "../../../networking"; const runs = (traceList as TracePage).data as TraceSummary[]; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx rename to ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx index 7d04a709825..04fce64b551 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx @@ -16,13 +16,13 @@ import { useOpenTraceRouting, useRunFilterRouting, useZoomRouting, -} from "./traceRouting"; -import type { TraceSummary } from "./traceTypes"; -import { RunView } from "./TraceDrawer"; +} from "../routing"; +import type { TraceSummary } from "../types"; +import { RunView } from "../detail/TraceDrawer"; import { TimeRangeControls } from "./TimeRangeControls"; import { TracesTimeline, type TimeWindow } from "./TracesTimeline"; -import { TracingSetupCard } from "./TracingSetupCard"; -import { useTracesLive } from "./tracesApi"; +import { TracingSetupCard } from "../../onboarding/tracing/TracingSetupCard"; +import { useTracesLive } from "../api"; import { type AgentTracesResult, traceWindowStartMs, useAgentTraces, useTraceAvailability } from "./useAgentTraces"; const DRAWER_WIDTH_KEY = "litellm.agentTraces.drawerWidth"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx similarity index 91% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.test.tsx rename to ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx index b03cd8a07dd..dff968334f7 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx @@ -1,12 +1,12 @@ import { fireEvent, render, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; -import { renderWithProviders } from "../../../../tests/test-utils"; +import { renderWithProviders } from "../../../../../tests/test-utils"; import { Inspector } from "@/components/shared/Inspector"; -import traceList from "./__fixtures__/trace_list.json"; +import traceList from "../__fixtures__/trace_list.json"; import { AgentTracesTable } from "./AgentTracesTable"; -import { traceKey } from "./traceRouting"; -import type { TracePage, TraceSummary } from "./traceTypes"; +import { traceKey } from "../routing"; +import type { TracePage, TraceSummary } from "../types"; const inList = (table: React.ReactElement) => ( ): TraceSummary => ({ agent_count: 1, diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runQuery.test.ts b/ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runQuery.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runQuery.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runQuery.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runQuery.ts b/ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runQuery.ts similarity index 93% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runQuery.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runQuery.ts index 08adfa98656..3a69345797c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runQuery.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runQuery.ts @@ -1,7 +1,7 @@ import { Bot, Box, Braces, CircleDashed, Hash, SquareChevronRight } from "lucide-react"; -import type { TraceSummary } from "../traceTypes"; -import { previewText, traceAgentNames } from "../traceUtils"; +import type { TraceSummary } from "../../types"; +import { previewText, traceAgentNames } from "../../utils"; import { type ClientIndex, filterItems } from "@/components/shared/search/evaluate"; import type { FieldSpec, QueryLanguage } from "@/components/shared/search/language"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runSql.test.ts b/ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runSql.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runSql.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runSql.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runSql.ts b/ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runSql.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/runSearch/runSql.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/runSearch/runSql.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.test.ts b/ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.test.ts similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.test.ts index cbced08785f..039ef950768 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.test.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.test.ts @@ -1,7 +1,7 @@ import { describe, expect, it } from "vitest"; import { traceWindowStartMs } from "./useAgentTraces"; -import { spanLogWindow } from "./useSpanRequestLog"; +import { spanLogWindow } from "../detail/useSpanRequestLog"; const HOUR = 3600 * 1000; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts b/ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.ts similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts rename to ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.ts index 82d190a21ff..d82f224ef58 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/list/useAgentTraces.ts @@ -1,13 +1,14 @@ -import { useTracesApi } from "./tracesApi"; +import { useTracesApi } from "../api"; import { useInfiniteQuery, useQuery, type UseQueryOptions } from "@tanstack/react-query"; import moment from "moment"; import { useMemo } from "react"; import { ApiError } from "@/lib/http/client"; -import { LIVE_TAIL_INTERVAL_MS } from "../log_filter_logic"; -import type { TracePage, TraceSummary } from "./traceTypes"; -import type { TraceWindow } from "./tracesApi"; +import type { TracePage, TraceSummary } from "../types"; +import type { TraceWindow } from "../api"; + +const LIVE_TAIL_INTERVAL_MS = 15000; interface LoadedTracePage extends TracePage { window: TraceWindow; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceRouting.ts b/ui/litellm-dashboard/src/components/lens/traces/routing.ts similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/traceRouting.ts rename to ui/litellm-dashboard/src/components/lens/traces/routing.ts index 17551ab4110..744be6bba2b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceRouting.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/routing.ts @@ -9,9 +9,9 @@ import { } from "nuqs"; import { useCallback, useState } from "react"; -import { RANGE_PRESETS } from "./TimeRangeControls"; -import type { TimeWindow } from "./TracesTimeline"; -import type { TraceSummary } from "./traceTypes"; +import { RANGE_PRESETS } from "./list/TimeRangeControls"; +import type { TimeWindow } from "./list/TracesTimeline"; +import type { TraceSummary } from "./types"; export const TRACE_VIEWS = ["steps", "conversation"] as const; export type TraceView = (typeof TRACE_VIEWS)[number]; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTree.ts b/ui/litellm-dashboard/src/components/lens/traces/tree.ts similarity index 86% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/traceTree.ts rename to ui/litellm-dashboard/src/components/lens/traces/tree.ts index 2f0e66833a4..d18324d75f4 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTree.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/tree.ts @@ -1,8 +1,8 @@ /** * Shared contract for the agent run view (span tree + detail pane). - * Rows are produced by `buildTreeRows` in traceUtils.ts and rendered by SpanTree / DetailPane. + * Rows are produced by `buildTreeRows` in utils.ts and rendered by SpanTree / DetailPane. */ -import type { Span, SpanType } from "./traceTypes"; +import type { Span, SpanType } from "./types"; export type ErrorSource = "model" | "tool" | "litellm"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts b/ui/litellm-dashboard/src/components/lens/traces/types.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts rename to ui/litellm-dashboard/src/components/lens/traces/types.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/ActiveDot.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/ActiveDot.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/ActiveDot.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/ActiveDot.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/Collapse.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/Collapse.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/Collapse.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/Collapse.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/CopyButton.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/CopyButton.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/IdChip.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/IdChip.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/IdChip.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/IdChip.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/PaneBar.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/PaneBar.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/PaneBar.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/PaneBar.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanIcon.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/SpanIcon.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/SpanIcon.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/SpanIcon.tsx index 569842f8fdc..4374e535802 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/SpanIcon.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/ui/SpanIcon.tsx @@ -18,7 +18,7 @@ import { Logo } from "@/components/molecules/logo/Logo"; import { cn } from "@/lib/cva.config"; import { useSpanProvider } from "./spanProvider"; -import type { SpanType } from "./traceTypes"; +import type { SpanType } from "../types"; const SIZE = { sm: "size-4 rounded-sm", diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/StatusMark.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/StatusMark.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/StatusMark.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/StatusMark.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceFramework.test.ts b/ui/litellm-dashboard/src/components/lens/traces/ui/TraceFramework.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceFramework.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/ui/TraceFramework.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceFramework.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/TraceFramework.tsx similarity index 68% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/TraceFramework.tsx rename to ui/litellm-dashboard/src/components/lens/traces/ui/TraceFramework.tsx index f0bb8437abf..9c0ff2ee274 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceFramework.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/ui/TraceFramework.tsx @@ -1,20 +1,20 @@ -import anthropicLogo from "../../../../public/assets/logos/anthropic.svg"; -import crewaiLogo from "../../../../public/assets/logos/crewai-color.svg"; -import cursorLogo from "../../../../public/assets/logos/cursor.svg"; -import copilotLogo from "../../../../public/assets/logos/github_copilot.svg"; -import googleAdkLogo from "../../../../public/assets/logos/google-adk.png"; -import langchainLogo from "../../../../public/assets/logos/langchain.svg"; -import langgraphLogo from "../../../../public/assets/logos/langgraph-color.svg"; -import llamaIndexLogo from "../../../../public/assets/logos/llamaindex-color.svg"; -import openaiLogo from "../../../../public/assets/logos/openai_small.svg"; -import openaiAgentsLogo from "../../../../public/assets/logos/openai-agents.svg"; -import pydanticAiLogo from "../../../../public/assets/logos/pydantic-ai-color.svg"; -import strandsLogo from "../../../../public/assets/logos/strands.svg"; -import vercelLogo from "../../../../public/assets/logos/vercel.svg"; +import anthropicLogo from "../../../../../public/assets/logos/anthropic.svg"; +import crewaiLogo from "../../../../../public/assets/logos/crewai-color.svg"; +import cursorLogo from "../../../../../public/assets/logos/cursor.svg"; +import copilotLogo from "../../../../../public/assets/logos/github_copilot.svg"; +import googleAdkLogo from "../../../../../public/assets/logos/google-adk.png"; +import langchainLogo from "../../../../../public/assets/logos/langchain.svg"; +import langgraphLogo from "../../../../../public/assets/logos/langgraph-color.svg"; +import llamaIndexLogo from "../../../../../public/assets/logos/llamaindex-color.svg"; +import openaiLogo from "../../../../../public/assets/logos/openai_small.svg"; +import openaiAgentsLogo from "../../../../../public/assets/logos/openai-agents.svg"; +import pydanticAiLogo from "../../../../../public/assets/logos/pydantic-ai-color.svg"; +import strandsLogo from "../../../../../public/assets/logos/strands.svg"; +import vercelLogo from "../../../../../public/assets/logos/vercel.svg"; import { Logo } from "@/components/molecules/logo/Logo"; import { cn } from "@/lib/cva.config"; -import type { TraceSummary } from "./traceTypes"; +import type { TraceSummary } from "../types"; export interface TraceFramework { readonly id: string; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/spanProvider.test.ts b/ui/litellm-dashboard/src/components/lens/traces/ui/spanProvider.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/spanProvider.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/ui/spanProvider.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/spanProvider.ts b/ui/litellm-dashboard/src/components/lens/traces/ui/spanProvider.ts similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/spanProvider.ts rename to ui/litellm-dashboard/src/components/lens/traces/ui/spanProvider.ts index 2b74d0c44f4..8e5e595f9e9 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/spanProvider.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/ui/spanProvider.ts @@ -1,4 +1,4 @@ -import { useTracesLive } from "./tracesApi"; +import { useTracesLive } from "../api"; import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; import { getProviderLogoAndName } from "@/components/provider_info_helpers"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts b/ui/litellm-dashboard/src/components/lens/traces/utils.test.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts rename to ui/litellm-dashboard/src/components/lens/traces/utils.test.ts index c6bf860300c..1bb5f9ec3c1 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/utils.test.ts @@ -3,8 +3,8 @@ import { describe, expect, it } from "vitest"; import deepAgentTrace from "./__fixtures__/deep_agent_trace.json"; import researchTrace from "./__fixtures__/research_trace.json"; import swarmTrace from "./__fixtures__/swarm_trace.json"; -import { type SpanTreeState, type TreeRow } from "./traceTree"; -import type { Span, Trace } from "./traceTypes"; +import { type SpanTreeState, type TreeRow } from "./tree"; +import type { Span, Trace } from "./types"; import { buildTreeRows, buildVisibleTree, @@ -22,7 +22,7 @@ import { revealSpanInState, ROOT_KEY, treeGuides, -} from "./traceUtils"; +} from "./utils"; const swarm = swarmTrace as Trace; const research = researchTrace as Trace; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts b/ui/litellm-dashboard/src/components/lens/traces/utils.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts rename to ui/litellm-dashboard/src/components/lens/traces/utils.ts index 066d30bceed..a1706d11b9e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/utils.ts @@ -2,8 +2,8 @@ * Pure helpers for the agent trace views. No React in here: everything the span tree * computes lives here so it can be unit-tested directly. */ -import { type ErrorSource, type SpanTreeState, type TreeRow } from "./traceTree"; -import type { Span, TraceMessage, TraceSummary } from "./traceTypes"; +import { type ErrorSource, type SpanTreeState, type TreeRow } from "./tree"; +import type { Span, TraceMessage, TraceSummary } from "./types"; /* ------------------------------------------------------------------ */ /* Formatting */ diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/LensPreviewButton.tsx b/ui/litellm-dashboard/src/components/lens/ui/LensPreviewButton.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/LensPreviewButton.tsx rename to ui/litellm-dashboard/src/components/lens/ui/LensPreviewButton.tsx diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 4e606222eae..c297ca72141 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -117,7 +117,7 @@ import type { ComplexityRouterConfigPayload } from "./add_model/build_complexity import type { AutoRouterPresetsResponse } from "@/lib/autorouter_presets"; import type { VectorStoreIndex } from "@/app/(dashboard)/vector-stores/_components/IndexesTab"; import type { RoutingDecision } from "./view_logs/LogDetailsDrawer/RoutingDecisionCard"; -import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./view_logs/TraceView/traceTypes"; +import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./lens/traces/types"; import { createApiClient, deriveErrorMessage, diff --git a/ui/litellm-dashboard/src/lib/http/api.test-d.ts b/ui/litellm-dashboard/src/lib/http/api.test-d.ts index 7d88865358b..47ef63554dd 100644 --- a/ui/litellm-dashboard/src/lib/http/api.test-d.ts +++ b/ui/litellm-dashboard/src/lib/http/api.test-d.ts @@ -1,7 +1,7 @@ import { expectTypeOf, test } from "vitest"; import type { components } from "./schema"; import type { getClaudeCodePluginsList, userListCall } from "@/components/networking"; -import type { TraceMessage, UIMessage } from "@/components/view_logs/TraceView/traceTypes"; +import type { TraceMessage, UIMessage } from "@/components/lens/traces/types"; import type { TagNewRequest, TagUpdateRequest } from "@/components/tag_management/types"; test("API read functions expose the generated response contracts", () => { From 1459e004300c6f7d356add1870df329745673cc4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 10:00:49 -0700 Subject: [PATCH 7/8] refactor(ui): rename view_logs to logs and split request, audit and detail (#44505) Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/unit/proxy/auth/test_route_checks.py | 4 ++-- ui/litellm-dashboard/eslint-suppressions.json | 18 +++++++++--------- .../_components/PromptCachingRequestsTable.tsx | 2 +- .../(dashboard)/logs/page.integration.test.tsx | 2 +- .../src/app/(dashboard)/logs/page.test.tsx | 4 ++-- .../src/app/(dashboard)/logs/page.tsx | 4 ++-- .../GuardrailsMonitor/LogViewer.test.tsx | 4 ++-- .../components/GuardrailsMonitor/LogViewer.tsx | 4 ++-- .../add_model/AutoRouterRoutingTest.tsx | 2 +- .../common_components/simple_table.tsx | 2 +- .../lens/traces/detail/RequestDetail.tsx | 2 +- .../lens/traces/detail/useSpanRequestLog.ts | 2 +- .../AuditLogDrawer/AuditLogDrawer.test.tsx | 2 +- .../audit}/AuditLogDrawer/AuditLogDrawer.tsx | 2 +- .../audit}/AuditLogsPanel.test.tsx | 8 ++++---- .../audit}/AuditLogsPanel.tsx | 2 +- .../audit}/AuditLogsTable.test.tsx | 0 .../audit}/AuditLogsTable.tsx | 0 .../audit}/AuditLogsTableColumns.tsx | 2 +- .../{view_logs => logs}/batchLogUtils.test.ts | 0 .../{view_logs => logs}/batchLogUtils.ts | 0 .../{view_logs => logs}/constants.ts | 0 .../detail}/ClassifierAuditView.test.tsx | 0 .../detail}/ClassifierAuditView.tsx | 0 .../detail}/ClassifyTag.test.tsx | 0 .../detail}/ClassifyTag.tsx | 0 .../detail}/CollapsibleMessage.test.tsx | 0 .../detail}/CollapsibleMessage.tsx | 0 .../detail}/DrawerHeader.test.tsx | 2 +- .../detail}/DrawerHeader.tsx | 2 +- .../detail}/HistoryTree.test.tsx | 0 .../detail}/HistoryTree.tsx | 0 .../detail}/InputCard.test.tsx | 0 .../detail}/InputCard.tsx | 0 .../detail}/JsonViewer.test.tsx | 0 .../detail}/JsonViewer.tsx | 0 .../LogDetailContent.integration.test.tsx | 4 ++-- .../detail}/LogDetailContent.tsx | 14 +++++++------- .../detail}/LogDetailsDrawer.test.tsx | 2 +- .../detail}/LogDetailsDrawer.tsx | 4 ++-- .../detail}/OutputCard.test.tsx | 0 .../detail}/OutputCard.tsx | 0 .../detail}/PrettyMessagesView.test.tsx | 0 .../detail}/PrettyMessagesView.tsx | 0 .../detail}/RealtimePrettyView.test.tsx | 0 .../detail}/RealtimePrettyView.tsx | 0 .../detail}/RoutingDecisionCard.test.tsx | 0 .../detail}/RoutingDecisionCard.tsx | 0 .../detail}/SectionHeader.test.tsx | 0 .../detail}/SectionHeader.tsx | 0 .../detail}/SidebarToggle.tsx | 0 .../detail}/SimpleMessageBlock.test.tsx | 0 .../detail}/SimpleMessageBlock.tsx | 0 .../detail}/SimpleToolCallBlock.test.tsx | 0 .../detail}/SimpleToolCallBlock.tsx | 0 .../detail}/TokenFlow.test.tsx | 0 .../detail}/TokenFlow.tsx | 0 .../detail}/TruncatedValue.test.tsx | 0 .../detail}/TruncatedValue.tsx | 0 .../detail}/constants.ts | 0 .../detail/eventDisplayName.ts} | 2 +- .../LogDetailsDrawer => logs/detail}/index.ts | 0 .../detail}/prettyMessagesTypes.ts | 0 .../detail}/prettyMessagesUtils.ts | 0 .../sections}/ConfigInfoMessage.test.tsx | 0 .../detail/sections}/ConfigInfoMessage.tsx | 0 .../sections}/CostBreakdownViewer.test.tsx | 2 +- .../detail/sections}/CostBreakdownViewer.tsx | 0 .../detail/sections}/EvalViewer/EvalViewer.tsx | 0 .../BedrockGuardrailDetails.test.tsx | 6 +++--- .../BedrockGuardrailDetails.tsx | 0 .../GuardrailViewer/CompliancePanel.tsx | 0 .../GuardrailViewer/ContentFilterDetails.tsx | 0 .../GuardrailViewer/GuardrailViewer.test.tsx | 16 ++++++++-------- .../GuardrailViewer/GuardrailViewer.tsx | 2 +- .../PresidioDetectedEntities.test.tsx | 6 +++--- .../PresidioDetectedEntities.tsx | 0 .../GuardrailViewer/__tests__/fixtures.ts | 2 +- .../ToolsSection/FormattedToolView.tsx | 0 .../sections}/ToolsSection/JsonToolView.tsx | 0 .../ToolsSection/ToolExpandedContent.tsx | 0 .../detail/sections}/ToolsSection/ToolItem.tsx | 0 .../ToolsSection/ToolsSection.test.tsx | 2 +- .../sections}/ToolsSection/ToolsSection.tsx | 2 +- .../detail/sections}/ToolsSection/index.ts | 0 .../detail/sections}/ToolsSection/types.ts | 0 .../sections}/ToolsSection/utils.test.ts | 2 +- .../detail/sections}/ToolsSection/utils.ts | 2 +- .../detail/sections}/VectorStoreViewer.tsx | 2 +- .../detail}/useKeyboardNavigation.test.tsx | 2 +- .../detail}/useKeyboardNavigation.ts | 2 +- .../detail}/utils.test.ts | 0 .../LogDetailsDrawer => logs/detail}/utils.ts | 0 .../request}/LogsTableToolbar.tsx | 4 ++-- .../request}/RequestLogsFilters.test.tsx | 6 +++--- .../request}/RequestLogsFilters.tsx | 6 +++--- .../request}/RequestLogsPanel.test.tsx | 12 ++++++------ .../request}/RequestLogsPanel.tsx | 12 ++++++------ .../request}/RequestLogsTable.tsx | 8 ++++---- .../request}/RequestLogsTableColumns.test.tsx | 2 +- .../request}/RequestLogsTableColumns.tsx | 8 ++++---- .../request}/TypeBadges.test.tsx | 0 .../{view_logs => logs/request}/TypeBadges.tsx | 0 .../request}/logDetailRouting.ts | 0 .../request/timeRange.test.ts} | 2 +- .../request/timeRange.ts} | 0 .../request/useLogFilterLogic.test.tsx} | 8 ++++---- .../request/useLogFilterLogic.ts} | 12 ++++++------ .../{view_logs/columns.tsx => logs/types.ts} | 0 .../src/components/networking.tsx | 2 +- .../tests/CreateKeyPage.expiredToken.test.tsx | 2 +- 111 files changed, 114 insertions(+), 114 deletions(-) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogDrawer/AuditLogDrawer.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogDrawer/AuditLogDrawer.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogsPanel.test.tsx (95%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogsPanel.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogsTable.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogsTable.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/audit}/AuditLogsTableColumns.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs => logs}/batchLogUtils.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs}/batchLogUtils.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs}/constants.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/ClassifierAuditView.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/ClassifierAuditView.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/ClassifyTag.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/ClassifyTag.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/CollapsibleMessage.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/CollapsibleMessage.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/DrawerHeader.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/DrawerHeader.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/HistoryTree.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/HistoryTree.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/InputCard.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/InputCard.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/JsonViewer.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/JsonViewer.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/LogDetailContent.integration.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/LogDetailContent.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/LogDetailsDrawer.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/LogDetailsDrawer.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/OutputCard.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/OutputCard.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/PrettyMessagesView.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/PrettyMessagesView.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/RealtimePrettyView.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/RealtimePrettyView.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/RoutingDecisionCard.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/RoutingDecisionCard.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SectionHeader.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SectionHeader.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SidebarToggle.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SimpleMessageBlock.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SimpleMessageBlock.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SimpleToolCallBlock.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/SimpleToolCallBlock.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/TokenFlow.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/TokenFlow.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/TruncatedValue.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/TruncatedValue.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/constants.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/utils.ts => logs/detail/eventDisplayName.ts} (93%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/index.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/prettyMessagesTypes.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/prettyMessagesUtils.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ConfigInfoMessage.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ConfigInfoMessage.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/CostBreakdownViewer.test.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/CostBreakdownViewer.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/EvalViewer/EvalViewer.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/BedrockGuardrailDetails.test.tsx (95%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/BedrockGuardrailDetails.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/CompliancePanel.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/ContentFilterDetails.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/GuardrailViewer.test.tsx (95%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/GuardrailViewer.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/PresidioDetectedEntities.test.tsx (87%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/PresidioDetectedEntities.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/GuardrailViewer/__tests__/fixtures.ts (97%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/FormattedToolView.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/JsonToolView.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/ToolExpandedContent.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/ToolItem.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/ToolsSection.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/ToolsSection.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/index.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/types.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/utils.test.ts (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/ToolsSection/utils.ts (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/detail/sections}/VectorStoreViewer.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/useKeyboardNavigation.test.tsx (95%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/useKeyboardNavigation.ts (96%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/utils.test.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/LogDetailsDrawer => logs/detail}/utils.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/LogsTableToolbar.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsFilters.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsFilters.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsPanel.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsPanel.tsx (97%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsTable.tsx (96%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsTableColumns.test.tsx (99%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/RequestLogsTableColumns.tsx (98%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/TypeBadges.test.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/TypeBadges.tsx (100%) rename ui/litellm-dashboard/src/components/{view_logs => logs/request}/logDetailRouting.ts (100%) rename ui/litellm-dashboard/src/components/{view_logs/logs_utils.test.tsx => logs/request/timeRange.test.ts} (97%) rename ui/litellm-dashboard/src/components/{view_logs/logs_utils.tsx => logs/request/timeRange.ts} (100%) rename ui/litellm-dashboard/src/components/{view_logs/log_filter_logic.test.tsx => logs/request/useLogFilterLogic.test.tsx} (98%) rename ui/litellm-dashboard/src/components/{view_logs/log_filter_logic.tsx => logs/request/useLogFilterLogic.ts} (96%) rename ui/litellm-dashboard/src/components/{view_logs/columns.tsx => logs/types.ts} (100%) diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 6b1d8bc3a2c..d0d52dd6566 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -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", diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 61ab85d613a..74672cf950e 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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 } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx index ba9cfd8ca22..c6900a3b2ab 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx index dc8a57e50ba..c463a04bf05 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx @@ -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
; }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx index c7e0a82a363..0e86e8bb9d2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx @@ -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
{isActive ? "active" : "inactive"}
; }, })); -vi.mock("@/components/view_logs/AuditLogsPanel", () => ({ +vi.mock("@/components/logs/audit/AuditLogsPanel", () => ({ default: function AuditLogsPanelMock({ isActive }: { isActive: boolean }) { return
{isActive ? "active" : "inactive"}
; }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx index 51603dec29e..a9672b088ba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx index 083b1e5f3e2..c1fb0431e8e 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx @@ -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, diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx index 2abd699ba86..1e2538100ae 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx @@ -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< diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx index 2b6c06e9a96..44ef0e61340 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx index a4b84d28801..8d174e49cf2 100644 --- a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx +++ b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx @@ -28,7 +28,7 @@ interface SimpleTableProps { /** * 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({ data, diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx index 58177ede37a..242d89b21a0 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/RequestDetail.tsx @@ -5,7 +5,7 @@ import { useState } from "react"; import { Button } from "@/components/ui/button"; -import { LogDetailsDrawer } from "../../../view_logs/LogDetailsDrawer"; +import { LogDetailsDrawer } from "../../../logs/detail"; import { formatCost } from "../list/AgentTracesTable"; import { DetailGroup } from "./AttributesDetail"; import { CopyButton } from "../ui/CopyButton"; diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts b/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts index c654ca55793..0bbc5aea1e2 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/useSpanRequestLog.ts @@ -4,7 +4,7 @@ import { useQuery, type UseQueryOptions } from "@tanstack/react-query"; import moment from "moment"; import { uiSpendLogsCall } from "../../../networking"; -import type { LogEntry } from "../../../view_logs/columns"; +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; diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.test.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.test.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.test.tsx index b0672cee296..da158fb8450 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.test.tsx @@ -5,7 +5,7 @@ import moment from "moment"; import { AuditLogDrawer } from "./AuditLogDrawer"; import { AuditLogEntry } from "../AuditLogsTableColumns"; -vi.mock("../../common_components/DefaultProxyAdminTag", () => ({ +vi.mock("../../../common_components/DefaultProxyAdminTag", () => ({ default: ({ userId }: { userId: string }) => {userId}, })); diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.tsx index bb50d162181..fae3c35b80c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogDrawer/AuditLogDrawer.tsx +++ b/ui/litellm-dashboard/src/components/logs/audit/AuditLogDrawer/AuditLogDrawer.tsx @@ -2,7 +2,7 @@ import { Check, Copy } from "lucide-react"; import { useState, useCallback } from "react"; import moment from "moment"; import { AuditLogEntry, AUDIT_TABLE_NAME_DISPLAY } from "../AuditLogsTableColumns"; -import DefaultProxyAdminTag from "../../common_components/DefaultProxyAdminTag"; +import DefaultProxyAdminTag from "../../../common_components/DefaultProxyAdminTag"; import CopyButton from "@/components/shared/CopyButton"; import { StatusBadge, type StatusTone } from "@/components/shared/table_cells/status_badge"; import { Button } from "@/components/ui/button"; diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.test.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.test.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.test.tsx index 3b27f663b8e..0c83d5a2d89 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.test.tsx @@ -3,11 +3,11 @@ import { fireEvent, render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { chooseSelectOption } from "../../../tests/test-utils"; +import { chooseSelectOption } from "../../../../tests/test-utils"; import AuditLogsPanel from "./AuditLogsPanel"; -vi.mock("../networking", async (importOriginal) => { - const actual = await importOriginal(); +vi.mock("../../networking", async (importOriginal) => { + const actual = await importOriginal(); return { ...actual, uiAuditLogsCall: vi.fn() }; }); @@ -16,7 +16,7 @@ vi.mock("@tanstack/react-pacer/debouncer", () => ({ useDebouncedValue: (value: unknown) => [value, { cancel: vi.fn(), flush: vi.fn() }], })); -import { uiAuditLogsCall } from "../networking"; +import { uiAuditLogsCall } from "../../networking"; type AuditLogsParams = NonNullable[0]["params"]>; diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.tsx index 6480cb59682..f174f45a791 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsPanel.tsx @@ -4,7 +4,7 @@ import { useQuery, keepPreviousData } from "@tanstack/react-query"; import { ColumnFiltersState, OnChangeFn, PaginationState } from "@tanstack/react-table"; import { resolveLogoSrc } from "@/lib/assetPaths"; import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import { uiAuditLogsCall } from "../networking"; +import { uiAuditLogsCall } from "../../networking"; import { AuditLogEntry } from "./AuditLogsTableColumns"; import { AuditLogsTable } from "./AuditLogsTable"; import { AuditLogDrawer } from "./AuditLogDrawer/AuditLogDrawer"; diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsTable.test.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsTable.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogsTable.test.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogsTable.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsTable.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsTable.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogsTable.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogsTable.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/AuditLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsTableColumns.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/AuditLogsTableColumns.tsx rename to ui/litellm-dashboard/src/components/logs/audit/AuditLogsTableColumns.tsx index 10e19e25a77..17e057a1343 100644 --- a/ui/litellm-dashboard/src/components/view_logs/AuditLogsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/logs/audit/AuditLogsTableColumns.tsx @@ -4,7 +4,7 @@ import { ColumnDef } from "@tanstack/react-table"; import { DateCell, IdCell, IdentityCell, StatusBadge, type StatusTone } from "@/components/shared/table_cells"; -import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; +import DefaultProxyAdminTag from "../../common_components/DefaultProxyAdminTag"; export type AuditLogEntry = { id: string; diff --git a/ui/litellm-dashboard/src/components/view_logs/batchLogUtils.test.ts b/ui/litellm-dashboard/src/components/logs/batchLogUtils.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/batchLogUtils.test.ts rename to ui/litellm-dashboard/src/components/logs/batchLogUtils.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/batchLogUtils.ts b/ui/litellm-dashboard/src/components/logs/batchLogUtils.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/batchLogUtils.ts rename to ui/litellm-dashboard/src/components/logs/batchLogUtils.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/constants.ts b/ui/litellm-dashboard/src/components/logs/constants.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/constants.ts rename to ui/litellm-dashboard/src/components/logs/constants.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifierAuditView.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/ClassifierAuditView.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifierAuditView.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/ClassifierAuditView.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifierAuditView.tsx b/ui/litellm-dashboard/src/components/logs/detail/ClassifierAuditView.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifierAuditView.tsx rename to ui/litellm-dashboard/src/components/logs/detail/ClassifierAuditView.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifyTag.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/ClassifyTag.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifyTag.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/ClassifyTag.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifyTag.tsx b/ui/litellm-dashboard/src/components/logs/detail/ClassifyTag.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/ClassifyTag.tsx rename to ui/litellm-dashboard/src/components/logs/detail/ClassifyTag.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/CollapsibleMessage.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/CollapsibleMessage.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.tsx b/ui/litellm-dashboard/src/components/logs/detail/CollapsibleMessage.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.tsx rename to ui/litellm-dashboard/src/components/logs/detail/CollapsibleMessage.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.test.tsx index ef320f04f58..983d6650478 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.test.tsx @@ -2,7 +2,7 @@ import { screen, within } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { render } from "../../../../tests/test-utils"; -import type { LogEntry } from "../columns"; +import type { LogEntry } from "../types"; import { DrawerHeader } from "./DrawerHeader"; const logEntry = (overrides: Partial): LogEntry => diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.tsx b/ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.tsx rename to ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.tsx index a4afdb68fae..bc24614a744 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/DrawerHeader.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/DrawerHeader.tsx @@ -4,7 +4,7 @@ import moment from "moment"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../types"; import { AutoRouterTag } from "@/components/shared/table_cells"; import { ClassifyTag } from "./ClassifyTag"; import { SidebarToggle } from "./SidebarToggle"; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/HistoryTree.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/HistoryTree.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/HistoryTree.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/HistoryTree.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/HistoryTree.tsx b/ui/litellm-dashboard/src/components/logs/detail/HistoryTree.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/HistoryTree.tsx rename to ui/litellm-dashboard/src/components/logs/detail/HistoryTree.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/InputCard.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/InputCard.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/InputCard.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/InputCard.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/InputCard.tsx b/ui/litellm-dashboard/src/components/logs/detail/InputCard.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/InputCard.tsx rename to ui/litellm-dashboard/src/components/logs/detail/InputCard.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/JsonViewer.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/JsonViewer.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/JsonViewer.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/JsonViewer.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/JsonViewer.tsx b/ui/litellm-dashboard/src/components/logs/detail/JsonViewer.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/JsonViewer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/JsonViewer.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.integration.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.integration.test.tsx index e0551c61062..6c2569710e6 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.integration.test.tsx @@ -2,9 +2,9 @@ import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import { GuardrailJumpLink, LogDetailContent } from "./LogDetailContent"; -import type { LogEntry } from "../columns"; +import type { LogEntry } from "../types"; -vi.mock("../GuardrailViewer/GuardrailViewer", () => ({ +vi.mock("./sections/GuardrailViewer/GuardrailViewer", () => ({ default: ({ data }: { data: unknown }) =>
{JSON.stringify(data)}
, })); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx b/ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx rename to ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.tsx index f3762223e96..0da6a5e790b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/LogDetailContent.tsx @@ -8,11 +8,11 @@ import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/component import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../types"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { PROMPT_CACHE_CREATION_TOOLTIP, PROMPT_CACHE_READ_TOOLTIP } from "@/utils/promptCacheUsage"; -import GuardrailViewer from "../GuardrailViewer/GuardrailViewer"; -import EvalViewer from "../EvalViewer/EvalViewer"; +import GuardrailViewer from "./sections/GuardrailViewer/GuardrailViewer"; +import EvalViewer from "./sections/EvalViewer/EvalViewer"; import { getBatchIdFromRequestId, getBatchModels, @@ -20,9 +20,9 @@ import { getReasoningTokens, isBatchCallType, } from "../batchLogUtils"; -import { CostBreakdownViewer } from "../CostBreakdownViewer"; -import { ConfigInfoMessage } from "../ConfigInfoMessage"; -import { VectorStoreViewer } from "../VectorStoreViewer"; +import { CostBreakdownViewer } from "./sections/CostBreakdownViewer"; +import { ConfigInfoMessage } from "./sections/ConfigInfoMessage"; +import { VectorStoreViewer } from "./sections/VectorStoreViewer"; import { CREDENTIAL_LABELS } from "../constants"; import { TruncatedValue } from "./TruncatedValue"; import { TokenFlow } from "./TokenFlow"; @@ -47,7 +47,7 @@ import { FONT_FAMILY_MONO, SPACING_XLARGE, } from "./constants"; -import { ToolsSection } from "../ToolsSection"; +import { ToolsSection } from "./sections/ToolsSection"; import { PrettyMessagesView } from "./PrettyMessagesView"; import { ClassifierAuditView } from "./ClassifierAuditView"; import { AUTOROUTER_CLASSIFIER_ORIGIN } from "./ClassifyTag"; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.test.tsx index f5bf7cda951..1321a4a152a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.test.tsx @@ -3,7 +3,7 @@ import { fireEvent, render, screen, waitFor } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { LogDetailsDrawer } from "./LogDetailsDrawer"; import { sessionSpendLogsCall } from "../../networking"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../types"; import { AutoRouterModelGroupsProvider } from "@/components/shared/table_cells"; vi.mock("../../networking", () => ({ diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.tsx index dc9207d59eb..3b4debd7e0b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/LogDetailsDrawer.tsx @@ -2,10 +2,10 @@ import { useEffect, useMemo, useState } from "react"; import { Bot, Check, Copy, Sparkles, Wrench } from "lucide-react"; import { Sheet, SheetContent, SheetTitle } from "@/components/ui/sheet"; import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../types"; import { AutoRouterIcon, useIsAutoRoutedModelGroup } from "@/components/shared/table_cells"; import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "../constants"; -import { getEventDisplayName } from "../utils"; +import { getEventDisplayName } from "./eventDisplayName"; import { ClassifyTag } from "./ClassifyTag"; import { DrawerHeader } from "./DrawerHeader"; import { SidebarToggle } from "./SidebarToggle"; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/OutputCard.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/OutputCard.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/OutputCard.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/OutputCard.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/OutputCard.tsx b/ui/litellm-dashboard/src/components/logs/detail/OutputCard.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/OutputCard.tsx rename to ui/litellm-dashboard/src/components/logs/detail/OutputCard.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/PrettyMessagesView.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/PrettyMessagesView.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/PrettyMessagesView.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/PrettyMessagesView.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/PrettyMessagesView.tsx b/ui/litellm-dashboard/src/components/logs/detail/PrettyMessagesView.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/PrettyMessagesView.tsx rename to ui/litellm-dashboard/src/components/logs/detail/PrettyMessagesView.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/RealtimePrettyView.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/RealtimePrettyView.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.tsx b/ui/litellm-dashboard/src/components/logs/detail/RealtimePrettyView.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.tsx rename to ui/litellm-dashboard/src/components/logs/detail/RealtimePrettyView.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/RoutingDecisionCard.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/RoutingDecisionCard.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx b/ui/litellm-dashboard/src/components/logs/detail/RoutingDecisionCard.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx rename to ui/litellm-dashboard/src/components/logs/detail/RoutingDecisionCard.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SectionHeader.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/SectionHeader.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SectionHeader.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SectionHeader.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SectionHeader.tsx b/ui/litellm-dashboard/src/components/logs/detail/SectionHeader.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SectionHeader.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SectionHeader.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SidebarToggle.tsx b/ui/litellm-dashboard/src/components/logs/detail/SidebarToggle.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SidebarToggle.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SidebarToggle.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/SimpleMessageBlock.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SimpleMessageBlock.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.tsx b/ui/litellm-dashboard/src/components/logs/detail/SimpleMessageBlock.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SimpleMessageBlock.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/SimpleToolCallBlock.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SimpleToolCallBlock.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.tsx b/ui/litellm-dashboard/src/components/logs/detail/SimpleToolCallBlock.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.tsx rename to ui/litellm-dashboard/src/components/logs/detail/SimpleToolCallBlock.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TokenFlow.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/TokenFlow.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TokenFlow.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/TokenFlow.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TokenFlow.tsx b/ui/litellm-dashboard/src/components/logs/detail/TokenFlow.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TokenFlow.tsx rename to ui/litellm-dashboard/src/components/logs/detail/TokenFlow.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TruncatedValue.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/TruncatedValue.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TruncatedValue.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/TruncatedValue.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TruncatedValue.tsx b/ui/litellm-dashboard/src/components/logs/detail/TruncatedValue.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/TruncatedValue.tsx rename to ui/litellm-dashboard/src/components/logs/detail/TruncatedValue.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/constants.ts b/ui/litellm-dashboard/src/components/logs/detail/constants.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/constants.ts rename to ui/litellm-dashboard/src/components/logs/detail/constants.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/utils.ts b/ui/litellm-dashboard/src/components/logs/detail/eventDisplayName.ts similarity index 93% rename from ui/litellm-dashboard/src/components/view_logs/utils.ts rename to ui/litellm-dashboard/src/components/logs/detail/eventDisplayName.ts index 37a82a7523e..09883b427d5 100644 --- a/ui/litellm-dashboard/src/components/view_logs/utils.ts +++ b/ui/litellm-dashboard/src/components/logs/detail/eventDisplayName.ts @@ -1,4 +1,4 @@ -import { MCP_CALL_TYPES } from "./constants"; +import { MCP_CALL_TYPES } from "../constants"; /** * Derive a short, human-readable display name for a log entry. diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/index.ts b/ui/litellm-dashboard/src/components/logs/detail/index.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/index.ts rename to ui/litellm-dashboard/src/components/logs/detail/index.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/prettyMessagesTypes.ts b/ui/litellm-dashboard/src/components/logs/detail/prettyMessagesTypes.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/prettyMessagesTypes.ts rename to ui/litellm-dashboard/src/components/logs/detail/prettyMessagesTypes.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/prettyMessagesUtils.ts b/ui/litellm-dashboard/src/components/logs/detail/prettyMessagesUtils.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/prettyMessagesUtils.ts rename to ui/litellm-dashboard/src/components/logs/detail/prettyMessagesUtils.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/ConfigInfoMessage.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ConfigInfoMessage.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ConfigInfoMessage.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ConfigInfoMessage.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/ConfigInfoMessage.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ConfigInfoMessage.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ConfigInfoMessage.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ConfigInfoMessage.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/CostBreakdownViewer.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/CostBreakdownViewer.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/CostBreakdownViewer.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/CostBreakdownViewer.test.tsx index d484185f873..ee7535ca6ff 100644 --- a/ui/litellm-dashboard/src/components/view_logs/CostBreakdownViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/CostBreakdownViewer.test.tsx @@ -1,7 +1,7 @@ import React from "react"; import { describe, it, expect } from "vitest"; import userEvent from "@testing-library/user-event"; -import { renderWithProviders, screen } from "../../../tests/test-utils"; +import { renderWithProviders, screen } from "../../../../../tests/test-utils"; import { CostBreakdownViewer, CostBreakdown } from "./CostBreakdownViewer"; async function expandCostBreakdown() { diff --git a/ui/litellm-dashboard/src/components/view_logs/CostBreakdownViewer.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/CostBreakdownViewer.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/CostBreakdownViewer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/CostBreakdownViewer.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/EvalViewer/EvalViewer.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/EvalViewer/EvalViewer.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/EvalViewer/EvalViewer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/EvalViewer/EvalViewer.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/BedrockGuardrailDetails.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails.test.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/BedrockGuardrailDetails.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails.test.tsx index a6b75499448..a7aae99c452 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/BedrockGuardrailDetails.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails.test.tsx @@ -2,14 +2,14 @@ import React from "react"; import { describe, it, expect } from "vitest"; import BedrockGuardrailDetails, { BedrockGuardrailResponse, -} from "@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails"; -import { renderWithProviders, screen } from "../../../../tests/test-utils"; +} from "@/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails"; +import { renderWithProviders, screen } from "../../../../../../tests/test-utils"; import { makeAssessment, makeBedrockCoverage, makeBedrockResponse, makeBedrockUsage, -} from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"; +} from "@/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures"; describe("BedrockGuardrailDetails", () => { it("returns null when response is falsy", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/BedrockGuardrailDetails.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/BedrockGuardrailDetails.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/CompliancePanel.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/CompliancePanel.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/CompliancePanel.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/CompliancePanel.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/ContentFilterDetails.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/ContentFilterDetails.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/ContentFilterDetails.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/ContentFilterDetails.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.test.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.test.tsx index ff5e736306a..62a5e61dfe3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.test.tsx @@ -1,19 +1,19 @@ import React from "react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import userEvent from "@testing-library/user-event"; -import { renderWithProviders, screen, waitFor, within } from "../../../../tests/test-utils"; +import { renderWithProviders, screen, waitFor, within } from "../../../../../../tests/test-utils"; import { GuardrailInformation, makeBedrockResponse, makeEntity, makeGuardrailInformation, -} from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"; -import GuardrailViewer from "@/components/view_logs/GuardrailViewer/GuardrailViewer"; +} from "@/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures"; +import GuardrailViewer from "@/components/logs/detail/sections/GuardrailViewer/GuardrailViewer"; // We will mock child components selectively for some tests to assert prop passthrough, // but also run an integration-style render without mocks. -const PresidioPath = "@/components/view_logs/GuardrailViewer/PresidioDetectedEntities"; -const BedrockPath = "@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails"; +const PresidioPath = "@/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities"; +const BedrockPath = "@/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails"; const skippedPreCall: Partial = { guardrail_status: "not_run", @@ -245,7 +245,7 @@ describe("GuardrailViewer", () => { __esModule: true, default: ({ entities }: any) =>
presidio {entities?.length}
, })); - const { default: Component } = await import("@/components/view_logs/GuardrailViewer/GuardrailViewer"); + const { default: Component } = await import("@/components/logs/detail/sections/GuardrailViewer/GuardrailViewer"); const data = makeGuardrailInformation({ guardrail_provider: undefined, @@ -264,7 +264,7 @@ describe("GuardrailViewer", () => { __esModule: true, default: ({ entities }: any) =>
count:{entities?.length}
, })); - const { default: Component } = await import("@/components/view_logs/GuardrailViewer/GuardrailViewer"); + const { default: Component } = await import("@/components/logs/detail/sections/GuardrailViewer/GuardrailViewer"); const data = makeGuardrailInformation({ guardrail_provider: "presidio", @@ -283,7 +283,7 @@ describe("GuardrailViewer", () => { __esModule: true, default: ({ response }: any) =>
{response?.action ?? "no-action"}
, })); - const { default: Component } = await import("@/components/view_logs/GuardrailViewer/GuardrailViewer"); + const { default: Component } = await import("@/components/logs/detail/sections/GuardrailViewer/GuardrailViewer"); const data = makeGuardrailInformation({ guardrail_provider: "bedrock", diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.tsx index 58076ef6c00..0c5ab90f454 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.tsx @@ -3,7 +3,7 @@ import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/comp import PresidioDetectedEntities from "./PresidioDetectedEntities"; import BedrockGuardrailDetails, { BedrockGuardrailResponse, -} from "@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails"; +} from "@/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails"; import ContentFilterDetails from "./ContentFilterDetails"; import CompliancePanel from "./CompliancePanel"; import { getSpendString } from "@/utils/dataUtils"; diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/PresidioDetectedEntities.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities.test.tsx similarity index 87% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/PresidioDetectedEntities.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities.test.tsx index 77cbe3584e1..9a31349b5bc 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/PresidioDetectedEntities.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities.test.tsx @@ -1,9 +1,9 @@ import React from "react"; import { describe, it, expect } from "vitest"; import userEvent from "@testing-library/user-event"; -import PresidioDetectedEntities from "@/components/view_logs/GuardrailViewer/PresidioDetectedEntities"; -import { renderWithProviders, screen } from "../../../../tests/test-utils"; -import { makeEntity } from "@/components/view_logs/GuardrailViewer/__tests__/fixtures"; +import PresidioDetectedEntities from "@/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities"; +import { renderWithProviders, screen } from "../../../../../../tests/test-utils"; +import { makeEntity } from "@/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures"; describe("PresidioDetectedEntities", () => { it("renders null when entities empty", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/PresidioDetectedEntities.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/PresidioDetectedEntities.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/PresidioDetectedEntities.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/__tests__/fixtures.ts b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures.ts similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/__tests__/fixtures.ts rename to ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures.ts index ab121adf6b4..6511d4babd6 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/__tests__/fixtures.ts +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/GuardrailViewer/__tests__/fixtures.ts @@ -3,7 +3,7 @@ import type { BedrockAssessment, BedrockGuardrailCoverage, BedrockGuardrailUsage, -} from "@/components/view_logs/GuardrailViewer/BedrockGuardrailDetails"; +} from "@/components/logs/detail/sections/GuardrailViewer/BedrockGuardrailDetails"; export interface RecognitionMetadata { recognizer_name: string; diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/FormattedToolView.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/FormattedToolView.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/FormattedToolView.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/FormattedToolView.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/JsonToolView.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/JsonToolView.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/JsonToolView.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/JsonToolView.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolExpandedContent.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolExpandedContent.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolExpandedContent.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolExpandedContent.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolItem.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolItem.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolItem.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolItem.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.test.tsx index e146f38b42b..5885a5b736d 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.test.tsx @@ -7,7 +7,7 @@ import userEvent from "@testing-library/user-event"; import { describe, it, expect } from "vitest"; import { parseToolsFromLog } from "./utils"; import { ToolsSection } from "./ToolsSection"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../../../types"; const logWithTools = (toolNames: string[], calledName?: string): LogEntry => ({ request_id: "render-1", diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.tsx index e3965f182fd..a5ee65ceaa1 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/ToolsSection.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/ToolsSection.tsx @@ -6,7 +6,7 @@ import { useState } from "react"; import { ChevronDown, ChevronRight } from "lucide-react"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../../../types"; import { parseToolsFromLog } from "./utils"; import { ToolItem } from "./ToolItem"; diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/index.ts b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/index.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/index.ts rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/index.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/types.ts b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/types.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/types.ts rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/types.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.test.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.test.ts index b3c118ef368..78372fe56d9 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.test.ts +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.test.ts @@ -4,7 +4,7 @@ import { describe, it, expect } from "vitest"; import { parseToolsFromLog, hasTools } from "./utils"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../../../types"; describe("ToolsSection utils", () => { describe("parseToolsFromLog", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.ts similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts rename to ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.ts index d1d1698232a..5e76a09c155 100644 --- a/ui/litellm-dashboard/src/components/view_logs/ToolsSection/utils.ts +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/ToolsSection/utils.ts @@ -2,7 +2,7 @@ * Utility functions for parsing and processing tool data from log entries */ -import { LogEntry } from "../columns"; +import { LogEntry } from "../../../types"; import { ParsedTool, ToolDefinition, ToolCall } from "./types"; /** diff --git a/ui/litellm-dashboard/src/components/view_logs/VectorStoreViewer.tsx b/ui/litellm-dashboard/src/components/logs/detail/sections/VectorStoreViewer.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/VectorStoreViewer.tsx rename to ui/litellm-dashboard/src/components/logs/detail/sections/VectorStoreViewer.tsx index 8480a123373..6b560acb90c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/VectorStoreViewer.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/sections/VectorStoreViewer.tsx @@ -1,7 +1,7 @@ import React, { useState } from "react"; import { ChevronDown, ChevronRight } from "lucide-react"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; -import { getProviderLogoAndName } from "../provider_info_helpers"; +import { getProviderLogoAndName } from "../../../provider_info_helpers"; interface VectorStoreContent { text: string; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.test.tsx b/ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.test.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.test.tsx rename to ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.test.tsx index aec56296917..dfc527a6b0d 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.test.tsx @@ -1,6 +1,6 @@ import { fireEvent, renderHook } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; -import type { LogEntry } from "../columns"; +import type { LogEntry } from "../types"; import { useKeyboardNavigation } from "./useKeyboardNavigation"; const logs = [{ request_id: "first" }, { request_id: "second" }] as LogEntry[]; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts b/ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.ts similarity index 96% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts rename to ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.ts index 6d30b75f312..b8f1c037cc0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts +++ b/ui/litellm-dashboard/src/components/logs/detail/useKeyboardNavigation.ts @@ -1,6 +1,6 @@ import { useShortcut } from "@/components/shared/useShortcut"; -import { LogEntry } from "../columns"; +import { LogEntry } from "../types"; interface UseKeyboardNavigationProps { isOpen: boolean; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts b/ui/litellm-dashboard/src/components/logs/detail/utils.test.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts rename to ui/litellm-dashboard/src/components/logs/detail/utils.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts b/ui/litellm-dashboard/src/components/logs/detail/utils.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts rename to ui/litellm-dashboard/src/components/logs/detail/utils.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx b/ui/litellm-dashboard/src/components/logs/request/LogsTableToolbar.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx rename to ui/litellm-dashboard/src/components/logs/request/LogsTableToolbar.tsx index 3496b709d9b..4705a9057f5 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/LogsTableToolbar.tsx @@ -10,8 +10,8 @@ import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover import { Switch } from "@/components/ui/switch"; import { cn } from "@/lib/cva.config"; -import { QUICK_SELECT_OPTIONS } from "./constants"; -import { getTimeRangeDisplay } from "./logs_utils"; +import { QUICK_SELECT_OPTIONS } from "../constants"; +import { getTimeRangeDisplay } from "./timeRange"; const DATETIME_LOCAL_FORMAT = "YYYY-MM-DDTHH:mm"; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.test.tsx index 38542f033dc..f13e24b268a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.test.tsx @@ -3,9 +3,9 @@ import userEvent from "@testing-library/user-event"; import { useState } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { chooseSelectOption, renderWithProviders, testQueryClient } from "../../../tests/test-utils"; -import { ERROR_CODE_OPTIONS } from "./constants"; -import { LOG_FILTER_IDS } from "./log_filter_logic"; +import { chooseSelectOption, renderWithProviders, testQueryClient } from "../../../../tests/test-utils"; +import { ERROR_CODE_OPTIONS } from "../constants"; +import { LOG_FILTER_IDS } from "./useLogFilterLogic"; import { RequestLogsFilters } from "./RequestLogsFilters"; vi.mock("@/app/(dashboard)/hooks/keys/useKeyAliases", () => ({ diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.tsx index e0c3205c80f..2bf893820ce 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsFilters.tsx @@ -20,9 +20,9 @@ import { import { Input } from "@/components/ui/input"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; -import type { Team } from "../key_team_helpers/key_list"; -import { CREDENTIAL_LABELS, ERROR_CODE_OPTIONS } from "./constants"; -import { LOG_FILTER_IDS, type LogsWindow } from "./log_filter_logic"; +import type { Team } from "../../key_team_helpers/key_list"; +import { CREDENTIAL_LABELS, ERROR_CODE_OPTIONS } from "../constants"; +import { LOG_FILTER_IDS, type LogsWindow } from "./useLogFilterLogic"; const ALL_VALUE = "all"; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.test.tsx index aefbf68d662..d950859f22e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.test.tsx @@ -5,12 +5,12 @@ import moment from "moment"; import { NuqsTestingAdapter, type UrlUpdateEvent } from "nuqs/adapters/testing"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { render, renderWithProviders, testQueryClient } from "../../../tests/test-utils"; -import type { LogEntry } from "./columns"; +import { render, renderWithProviders, testQueryClient } from "../../../../tests/test-utils"; +import type { LogEntry } from "../types"; import RequestLogsPanel from "./RequestLogsPanel"; -vi.mock("../networking", async (importOriginal) => { - const actual = await importOriginal(); +vi.mock("../../networking", async (importOriginal) => { + const actual = await importOriginal(); return { ...actual, uiSpendLogsCall: vi.fn(), @@ -22,7 +22,7 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ fetchAllTeams: vi.fn().mockResolvedValue([]), })); -vi.mock("./LogDetailsDrawer", () => ({ +vi.mock("../detail", () => ({ LogDetailsDrawer: function LogDetailsDrawerMock({ open, logEntry, @@ -61,7 +61,7 @@ vi.mock("@tanstack/react-pacer/debouncer", () => ({ import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import { uiSpendLogsCall } from "../networking"; +import { uiSpendLogsCall } from "../../networking"; const logEntry = (overrides: Partial): LogEntry => ({ request_id: "req-1", diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.tsx index ff26e30baca..605ff373c7c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsPanel.tsx @@ -10,10 +10,10 @@ import { DEFAULT_PAGE_SIZE_OPTIONS } from "@/components/shared/DataTable"; import { Button } from "@/components/ui/button"; import { AutoRouterModelGroupsProvider } from "@/components/shared/table_cells"; import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import type { KeyResponse } from "../key_team_helpers/key_list"; -import { keyInfoV1Call, uiSpendLogsCall } from "../networking"; -import KeyInfoView from "../templates/key_info_view"; -import type { LogEntry } from "./columns"; +import type { KeyResponse } from "../../key_team_helpers/key_list"; +import { keyInfoV1Call, uiSpendLogsCall } from "../../networking"; +import KeyInfoView from "../../templates/key_info_view"; +import type { LogEntry } from "../types"; import { DEFAULT_LOGS_SORTING, formatLogsWindow, @@ -21,9 +21,9 @@ import { LOG_FILTER_IDS, type PaginatedResponse, useLogFilterLogic, -} from "./log_filter_logic"; +} from "./useLogFilterLogic"; import { useLogDetailRouting } from "./logDetailRouting"; -import { LogDetailsDrawer } from "./LogDetailsDrawer"; +import { LogDetailsDrawer } from "../detail"; import { defaultLogsTimeRange, type LogsTimeRange, diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTable.tsx similarity index 96% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsTable.tsx index 0b834f43e40..822bf73302b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTable.tsx @@ -7,10 +7,10 @@ import { useMemo, useState, type ReactNode } from "react"; import { useUserEmailLookup } from "@/app/(dashboard)/hooks/users/useUsers"; import { DataTable, DataTableFilterDrawer, DataTableToolbar } from "@/components/shared/DataTable"; -import type { Team } from "../key_team_helpers/key_list"; -import type { LogEntry } from "./columns"; -import { CREDENTIAL_LABELS, SPAN_TYPE_LABELS } from "./constants"; -import { LOG_FILTER_IDS, LOG_FILTER_LABELS, type LogsWindow } from "./log_filter_logic"; +import type { Team } from "../../key_team_helpers/key_list"; +import type { LogEntry } from "../types"; +import { CREDENTIAL_LABELS, SPAN_TYPE_LABELS } from "../constants"; +import { LOG_FILTER_IDS, LOG_FILTER_LABELS, type LogsWindow } from "./useLogFilterLogic"; import { RequestLogsFilters } from "./RequestLogsFilters"; import { getRequestLogsTableColumns } from "./RequestLogsTableColumns"; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.test.tsx index c6c2714bc49..6d4a6e789b4 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.test.tsx @@ -4,7 +4,7 @@ import { describe, expect, it, vi } from "vitest"; import { DataTable } from "@/components/shared/DataTable"; -import type { LogEntry } from "./columns"; +import type { LogEntry } from "../types"; import { getRequestLogsTableColumns } from "./RequestLogsTableColumns"; const { copyToClipboardMock } = vi.hoisted(() => ({ copyToClipboardMock: vi.fn() })); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx rename to ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.tsx index 80e7b512471..5d5918e4657 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/RequestLogsTableColumns.tsx @@ -7,10 +7,10 @@ import { DataTableSortHeader } from "@/components/shared/DataTable"; import { CellTooltip, DateCell, IdCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; import { copyToClipboard, getSpendString } from "@/utils/dataUtils"; -import { getProviderLogoAndName } from "../provider_info_helpers"; -import { getBatchIdFromRequestId, getBatchRequestCounts, isBatchCallType } from "./batchLogUtils"; -import type { LogEntry } from "./columns"; -import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; +import { getProviderLogoAndName } from "../../provider_info_helpers"; +import { getBatchIdFromRequestId, getBatchRequestCounts, isBatchCallType } from "../batchLogUtils"; +import type { LogEntry } from "../types"; +import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "../constants"; import { AgentBadge, AgentIcon, BatchBadge, LlmBadge, McpBadge, SparkleIcon, WrenchIcon } from "./TypeBadges"; export interface RequestLogsTableColumnsDeps { diff --git a/ui/litellm-dashboard/src/components/view_logs/TypeBadges.test.tsx b/ui/litellm-dashboard/src/components/logs/request/TypeBadges.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TypeBadges.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/TypeBadges.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/TypeBadges.tsx b/ui/litellm-dashboard/src/components/logs/request/TypeBadges.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/TypeBadges.tsx rename to ui/litellm-dashboard/src/components/logs/request/TypeBadges.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts b/ui/litellm-dashboard/src/components/logs/request/logDetailRouting.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts rename to ui/litellm-dashboard/src/components/logs/request/logDetailRouting.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx b/ui/litellm-dashboard/src/components/logs/request/timeRange.test.ts similarity index 97% rename from ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/timeRange.test.ts index b0b74df4c24..69f4f914e7d 100644 --- a/ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/timeRange.test.ts @@ -1,6 +1,6 @@ import moment from "moment"; import { describe, expect, it } from "vitest"; -import { getTimeRangeDisplay } from "./logs_utils"; +import { getTimeRangeDisplay } from "./timeRange"; // startTime built relative to "now"; getTimeRangeDisplay computes now() internally. const ago = (amount: number, unit: moment.unitOfTime.DurationConstructor) => diff --git a/ui/litellm-dashboard/src/components/view_logs/logs_utils.tsx b/ui/litellm-dashboard/src/components/logs/request/timeRange.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/logs_utils.tsx rename to ui/litellm-dashboard/src/components/logs/request/timeRange.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx rename to ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.test.tsx index baf713b1537..9c266759f53 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.test.tsx @@ -15,9 +15,9 @@ import { LOGS_WINDOW_TICK_MS, useLogFilterLogic, type PaginatedResponse, -} from "./log_filter_logic"; +} from "./useLogFilterLogic"; -vi.mock("../networking", () => ({ +vi.mock("../../networking", () => ({ uiSpendLogsCall: vi.fn(), })); @@ -25,9 +25,9 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ fetchAllTeams: vi.fn().mockResolvedValue([]), })); -import { uiSpendLogsCall } from "../networking"; +import { uiSpendLogsCall } from "../../networking"; import { fetchAllTeams } from "@/components/key_team_helpers/filter_helpers"; -import type { Team } from "../key_team_helpers/key_list"; +import type { Team } from "../../key_team_helpers/key_list"; const emptyResponse: PaginatedResponse = { data: [], diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.ts similarity index 96% rename from ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx rename to ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.ts index 8b2d9f22c15..17d7eb4b860 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/logs/request/useLogFilterLogic.ts @@ -1,12 +1,12 @@ import moment from "moment"; import { keepPreviousData, useQuery, type UseQueryOptions } from "@tanstack/react-query"; import type { ColumnFiltersState, PaginationState, SortingState } from "@tanstack/react-table"; -import { uiSpendLogsCall } from "../networking"; -import { Team } from "../key_team_helpers/key_list"; -import { fetchAllTeams } from "../../components/key_team_helpers/filter_helpers"; -import { teamListScopeUserId } from "../../utils/roles"; -import { defaultPageSize } from "../constants"; -import { LOGS_SORT_FIELD_MAP, type LogEntry, type LogsSortField } from "./columns"; +import { uiSpendLogsCall } from "../../networking"; +import { Team } from "../../key_team_helpers/key_list"; +import { fetchAllTeams } from "../../key_team_helpers/filter_helpers"; +import { teamListScopeUserId } from "../../../utils/roles"; +import { defaultPageSize } from "../../constants"; +import { LOGS_SORT_FIELD_MAP, type LogEntry, type LogsSortField } from "../types"; export interface PaginatedResponse { data: LogEntry[]; diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.tsx b/ui/litellm-dashboard/src/components/logs/types.ts similarity index 100% rename from ui/litellm-dashboard/src/components/view_logs/columns.tsx rename to ui/litellm-dashboard/src/components/logs/types.ts diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index c297ca72141..98f3a120a49 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -116,7 +116,7 @@ import { MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE } from "./mcp_tools/constants"; import type { ComplexityRouterConfigPayload } from "./add_model/build_complexity_router_config"; import type { AutoRouterPresetsResponse } from "@/lib/autorouter_presets"; import type { VectorStoreIndex } from "@/app/(dashboard)/vector-stores/_components/IndexesTab"; -import type { RoutingDecision } from "./view_logs/LogDetailsDrawer/RoutingDecisionCard"; +import type { RoutingDecision } from "./logs/detail/RoutingDecisionCard"; import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./lens/traces/types"; import { createApiClient, diff --git a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx index 6a4fa2e32c3..2c57833102d 100644 --- a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx +++ b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx @@ -119,7 +119,7 @@ vi.mock("@/app/(dashboard)/router-settings/_components/general_settings", () => })); vi.mock("@/components/pass_through_settings", () => ({ default: stub("pass-through-settings") })); vi.mock("@/components/budgets/budget_panel", () => ({ default: stub("budget-panel") })); -vi.mock("@/components/view_logs", () => ({ default: stub("spend-logs") })); +vi.mock("@/components/logs", () => ({ default: stub("spend-logs") })); vi.mock("@/components/model_hub_table", () => ({ default: stub("model-hub-table") })); vi.mock("@/components/new_usage", () => ({ default: stub("new-usage") })); vi.mock("@/components/api_ref", () => ({ default: stub("api-ref") })); From e1d16f51d14849c1b3decf17cd81a3bcb4863dca Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 10:00:49 -0700 Subject: [PATCH 8/8] refactor(ui): share CopyButton between Lens traces and logs (#44513) Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../lens/traces/detail/DetailPane.tsx | 6 +- .../lens/traces/detail/MessageCard.tsx | 16 ++- .../lens/traces/detail/RequestDetail.tsx | 4 +- .../lens/traces/detail/TraceConversation.tsx | 9 +- .../components/lens/traces/ui/CopyButton.tsx | 59 -------- .../src/components/shared/CopyButton.test.tsx | 126 ++++++++++++++++-- .../src/components/shared/CopyButton.tsx | 67 ++++++++-- 7 files changed, 197 insertions(+), 90 deletions(-) delete mode 100644 ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx index 80f9b3da47c..e030d393ab3 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/DetailPane.tsx @@ -6,7 +6,7 @@ import { Button } from "@/components/ui/button"; import { Tabs, TabsList, TabsTrigger, TabsContent } from "@/components/ui/tabs"; import { AttributesDetail } from "./AttributesDetail"; -import { CopyButton } from "../ui/CopyButton"; +import CopyButton from "@/components/shared/CopyButton"; import { DetailContent, errorHeadline, useSpanDetail } from "./DetailContent"; import { IdChip } from "../ui/IdChip"; import { PaneBar } from "../ui/PaneBar"; @@ -144,7 +144,7 @@ function SpanPane({ - +
{tokens > 0 && } @@ -209,7 +209,7 @@ function GroupPane({ )}
- + ); diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx index cd3ea968d1a..f27023dd1e1 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/MessageCard.tsx @@ -8,7 +8,7 @@ import remarkGfm from "remark-gfm"; import { cn } from "@/lib/cva.config"; import { FoldChevron } from "../ui/Collapse"; -import { CopyButton } from "../ui/CopyButton"; +import CopyButton from "@/components/shared/CopyButton"; import { displayValue, type KeyValue, KeyValueRows, objectEntries } from "./KeyValueRows"; import type { TraceMessage, TraceToolCall } from "../types"; @@ -94,7 +94,13 @@ export function ToolCallBlock({ call }: { call: TraceToolCall }) {
{call.name} - +
@@ -114,7 +120,7 @@ export function MessageCard({ message }: { message: TraceMessage; model: string
setOpen((v) => !v)} /> {label} - +
{expanded && (
{name} - + {expandable ? (
{open && (
             
- +
             {item.toolResult || "No output recorded"}
diff --git a/ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx b/ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx
deleted file mode 100644
index fecd817f7b6..00000000000
--- a/ui/litellm-dashboard/src/components/lens/traces/ui/CopyButton.tsx
+++ /dev/null
@@ -1,59 +0,0 @@
-"use client";
-
-import { Check, Copy } from "lucide-react";
-import { useEffect, useState } from "react";
-
-import { Button } from "@/components/ui/button";
-import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
-import { cn } from "@/lib/cva.config";
-import { copyToClipboard } from "@/utils/dataUtils";
-
-const COPIED_RESET_MS = 1600;
-
-interface CopyButtonProps {
-  value: string;
-  label?: string;
-  copiedLabel?: string;
-  iconOnly?: boolean;
-  className?: string;
-}
-
-/** Copy → check for a moment. `iconOnly` renders a bare icon button with a tooltip. */
-export function CopyButton({
-  value,
-  label = "Copy",
-  copiedLabel = "Copied",
-  iconOnly = false,
-  className,
-}: CopyButtonProps) {
-  const [copied, setCopied] = useState(false);
-
-  useEffect(() => {
-    if (!copied) return;
-    const timeout = window.setTimeout(() => setCopied(false), COPIED_RESET_MS);
-    return () => window.clearTimeout(timeout);
-  }, [copied]);
-
-  const button = (
-    
-  );
-
-  if (!iconOnly) return button;
-  return (
-    
-      
-        
-        {copied ? copiedLabel : label}
-      
-    
-  );
-}
diff --git a/ui/litellm-dashboard/src/components/shared/CopyButton.test.tsx b/ui/litellm-dashboard/src/components/shared/CopyButton.test.tsx
index 07e115d7a39..c2c2ee199b4 100644
--- a/ui/litellm-dashboard/src/components/shared/CopyButton.test.tsx
+++ b/ui/litellm-dashboard/src/components/shared/CopyButton.test.tsx
@@ -1,12 +1,31 @@
+import { act, fireEvent } from "@testing-library/react";
 import userEvent from "@testing-library/user-event";
-import { describe, expect, it, vi } from "vitest";
+import { afterEach, describe, expect, it, vi } from "vitest";
+import type { ComponentProps } from "react";
+import { copyToClipboard } from "@/utils/dataUtils";
 import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils";
 import CopyButton from "./CopyButton";
 
+vi.mock("@/utils/dataUtils", () => ({ copyToClipboard: vi.fn() }));
+vi.mock("lucide-react", async (importOriginal) => ({
+  ...(await importOriginal()),
+  Copy: (props: ComponentProps<"svg">) => ,
+  Check: (props: ComponentProps<"svg">) => ,
+}));
+
+const clipboardDescriptor = Object.getOwnPropertyDescriptor(navigator, "clipboard");
+
+afterEach(() => {
+  vi.useRealTimers();
+  vi.resetAllMocks();
+  if (clipboardDescriptor) Object.defineProperty(navigator, "clipboard", clipboardDescriptor);
+  else Reflect.deleteProperty(navigator, "clipboard");
+});
+
 describe("CopyButton", () => {
   it("renders nothing when there is no value", () => {
-    const { container } = renderWithProviders();
-    expect(container.querySelector("button")).toBeNull();
+    renderWithProviders();
+    expect(screen.queryByRole("button")).not.toBeInTheDocument();
   });
 
   it("shows the confirmation checkmark only after a successful write", async () => {
@@ -17,12 +36,12 @@ describe("CopyButton", () => {
     renderWithProviders();
 
     const button = screen.getByRole("button", { name: "Copy value" });
-    expect(button.querySelector(".lucide-copy")).toBeInTheDocument();
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
 
     await user.click(button);
 
     expect(writeText).toHaveBeenCalledWith("secret-value");
-    await waitFor(() => expect(button.querySelector(".lucide-check")).toBeInTheDocument());
+    expect(await screen.findByRole("img", { name: "Copied icon" })).toBeInTheDocument();
   });
 
   it("does not show the checkmark when the clipboard write is rejected", async () => {
@@ -36,8 +55,8 @@ describe("CopyButton", () => {
     await user.click(button);
 
     await waitFor(() => expect(writeText).toHaveBeenCalledWith("secret-value"));
-    expect(button.querySelector(".lucide-check")).not.toBeInTheDocument();
-    expect(button.querySelector(".lucide-copy")).toBeInTheDocument();
+    expect(screen.queryByRole("img", { name: "Copied icon" })).not.toBeInTheDocument();
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
   });
 
   it("does not show the checkmark when the clipboard API is unavailable", async () => {
@@ -49,7 +68,96 @@ describe("CopyButton", () => {
     const button = screen.getByRole("button", { name: "Copy value" });
     await user.click(button);
 
-    expect(button.querySelector(".lucide-check")).not.toBeInTheDocument();
-    expect(button.querySelector(".lucide-copy")).toBeInTheDocument();
+    expect(screen.queryByRole("img", { name: "Copied icon" })).not.toBeInTheDocument();
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
+    expect(copyToClipboard).not.toHaveBeenCalled();
+  });
+
+  it("keeps the default icon styling, custom classes, title and 1200ms confirmation", async () => {
+    vi.useFakeTimers();
+    const writeText = vi.fn().mockResolvedValue(undefined);
+    Object.defineProperty(navigator, "clipboard", { value: { writeText }, configurable: true });
+    renderWithProviders(
+      ,
+    );
+    const button = screen.getByRole("button", { name: "Copy request ID" });
+    expect(button).toHaveAttribute("title", "Copy request ID");
+    expect(button).toHaveClass("hover:text-primary", "shrink-0");
+    expect(button).not.toHaveTextContent(/.+/);
+    expect(screen.getByRole("img", { name: "Copy icon" })).toHaveClass("size-3");
+
+    await act(async () => fireEvent.click(button));
+    expect(writeText).toHaveBeenCalledWith("request-id");
+    expect(copyToClipboard).not.toHaveBeenCalled();
+    act(() => vi.advanceTimersByTime(1199));
+    expect(screen.getByRole("img", { name: "Copied icon" })).toBeInTheDocument();
+    act(() => vi.advanceTimersByTime(1));
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
+  });
+
+  it("shows an outlined action with custom confirmation text for 1600ms", async () => {
+    vi.useFakeTimers();
+    vi.mocked(copyToClipboard).mockResolvedValue(true);
+    renderWithProviders(
+      ,
+    );
+    const button = screen.getByRole("button", { name: "Copy step" });
+    expect(button).toHaveTextContent("Copy step");
+    expect(button).toHaveClass("border", "gap-1.5", "hover:text-foreground");
+    expect(screen.getByRole("img", { name: "Copy icon" })).toHaveClass("size-3");
+
+    await act(async () => fireEvent.click(button));
+    expect(copyToClipboard).toHaveBeenCalledWith("step-context", "Step copied");
+    expect(button).toHaveTextContent("Step copied");
+    expect(screen.getByRole("img", { name: "Copied icon" })).toBeInTheDocument();
+    act(() => vi.advanceTimersByTime(1599));
+    expect(button).toHaveTextContent("Step copied");
+    act(() => vi.advanceTimersByTime(1));
+    expect(button).toHaveTextContent("Copy step");
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
+  });
+
+  it("shows a tooltip and no text for an icon-only action", async () => {
+    const user = userEvent.setup();
+    vi.mocked(copyToClipboard).mockResolvedValue(true);
+    renderWithProviders();
+    const button = screen.getByRole("button", { name: "Copy tool result" });
+    expect(button).not.toHaveTextContent(/.+/);
+    expect(button).not.toHaveAttribute("title");
+    await user.hover(button);
+    expect(await screen.findByText("Copy tool result")).toBeVisible();
+    await user.click(button);
+    expect(copyToClipboard).toHaveBeenCalledWith("tool-result", "Copied");
+    await user.unhover(button);
+    await user.hover(button);
+    expect(await screen.findByText("Copied")).toBeVisible();
+  });
+
+  it("keeps the action label and copy icon when copying fails", async () => {
+    const user = userEvent.setup();
+    vi.mocked(copyToClipboard).mockResolvedValue(false);
+    renderWithProviders();
+    const button = screen.getByRole("button", { name: "Copy step" });
+    await user.click(button);
+    expect(copyToClipboard).toHaveBeenCalledWith("step-context", "Copied");
+    expect(button).toHaveTextContent("Copy step");
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
+    expect(screen.queryByRole("img", { name: "Copied icon" })).not.toBeInTheDocument();
+  });
+
+  it("retains an empty action while default empty values remain hidden", async () => {
+    const user = userEvent.setup();
+    vi.mocked(copyToClipboard).mockResolvedValue(false);
+    renderWithProviders(
+      <>
+        
+        
+      ,
+    );
+    expect(screen.queryByRole("button", { name: "Copy value" })).not.toBeInTheDocument();
+    const button = screen.getByRole("button", { name: "Copy empty result" });
+    await user.click(button);
+    expect(copyToClipboard).toHaveBeenCalledWith("", "Copied");
+    expect(screen.getByRole("img", { name: "Copy icon" })).toBeInTheDocument();
   });
 });
diff --git a/ui/litellm-dashboard/src/components/shared/CopyButton.tsx b/ui/litellm-dashboard/src/components/shared/CopyButton.tsx
index 052d03df319..57ae86a26d0 100644
--- a/ui/litellm-dashboard/src/components/shared/CopyButton.tsx
+++ b/ui/litellm-dashboard/src/components/shared/CopyButton.tsx
@@ -1,5 +1,9 @@
+"use client";
+
 import { Button } from "@/components/ui/button";
+import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
 import { cn } from "@/lib/cva.config";
+import { copyToClipboard } from "@/utils/dataUtils";
 import { Check, Copy } from "lucide-react";
 import React, { useEffect, useState } from "react";
 
@@ -8,20 +12,52 @@ interface CopyButtonProps {
   label: string;
   className?: string;
   iconClassName?: string;
+  variant?: "icon" | "action";
+  iconOnly?: boolean;
+  copiedLabel?: string;
 }
 
-const CopyButton: React.FC = ({ value, label, className, iconClassName = "size-[15px]" }) => {
+const VARIANTS = {
+  icon: {
+    resetMs: 1200,
+    iconClassName: "size-[15px]",
+    className: "text-muted-foreground hover:text-primary",
+  },
+  action: {
+    resetMs: 1600,
+    iconClassName: "size-3",
+    className: "text-xs text-muted-foreground hover:text-foreground",
+  },
+} as const;
+
+const CopyButton: React.FC = ({
+  value,
+  label,
+  className,
+  iconClassName,
+  variant = "icon",
+  iconOnly = false,
+  copiedLabel = "Copied",
+}) => {
   const [copied, setCopied] = useState(false);
+  const isAction = variant === "action";
+  const showLabel = isAction && !iconOnly;
+  const { resetMs, className: variantClassName, iconClassName: defaultIconClassName } = VARIANTS[variant];
+  const iconSize = iconClassName ?? defaultIconClassName;
 
   useEffect(() => {
     if (!copied) return;
-    const timer = setTimeout(() => setCopied(false), 1200);
+    const timer = setTimeout(() => setCopied(false), resetMs);
     return () => clearTimeout(timer);
-  }, [copied]);
+  }, [copied, resetMs]);
 
-  if (!value) return null;
+  if (value == null || (!value && !isAction)) return null;
 
   const handleCopy = async () => {
+    if (isAction) {
+      setCopied(await copyToClipboard(value, copiedLabel));
+      return;
+    }
     if (!navigator.clipboard) return;
     try {
       await navigator.clipboard.writeText(value);
@@ -31,19 +67,30 @@ const CopyButton: React.FC = ({ value, label, className, iconCl
     }
   };
 
-  return (
+  const button = (
     
   );
+
+  if (!isAction || !iconOnly) return button;
+  return (
+    
+      
+        
+        {copied ? copiedLabel : label}
+      
+    
+  );
 };
 
 export default CopyButton;