diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index e9e5dd3d66b..d56e29fb627 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -7,6 +7,7 @@ legacy_flags=( caching-local enterprise-package enterprise-routing + integrations llm-other-providers llm-vertex-ai mcp-integration @@ -52,6 +53,7 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py echo tests/unit/enterprise/proxy/test_managed_files_access_check.py echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; + integrations) echo tests/unit/integrations ;; llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;; llm-vertex-ai) echo tests/unit/llms/vertex_ai ;; mcp-integration) @@ -75,6 +77,7 @@ legacy_paths() { echo tests/unit/messages echo tests/unit/rag echo tests/unit/rerank_api + echo tests/unit/secret_managers echo tests/unit/vector_stores echo tests/unit/videos ;; proxy-db-auth-checks) diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 994d67da64d..41e9f11cefa 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -369,6 +369,13 @@ workflows: reruns: 2 base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-integrations + flag: integrations + shards: 2 + reruns: 3 + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - unit: name: unit-misc flag: misc diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 2fa05879350..4dca8075440 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -80,7 +80,8 @@ jobs: - shard: integrations artifact-name: integrations - test-path: "tests/test_litellm/integrations" + test-path: "" + unit-flag: integrations workers: 2 reruns: 3 timeout-minutes: 20 @@ -107,7 +108,6 @@ jobs: - shard: misc artifact-name: misc test-path: >- - tests/test_litellm/secret_managers tests/test_litellm/interactions tests/test_litellm/ocr tests/test_litellm/passthrough diff --git a/Makefile b/Makefile index e86047b1987..f27525b58ff 100644 --- a/Makefile +++ b/Makefile @@ -326,13 +326,13 @@ test-unit-proxy-misc: install-test-deps $(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20 test-unit-integrations: install-test-deps - $(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20 test-unit-core-utils: install-test-deps $(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20 test-unit-other: install-test-deps - $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 test-unit-root: install-test-deps $(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md index 70562b18aa6..a3869213be2 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md @@ -2,7 +2,7 @@ Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about -Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard +Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs index 30e46f93248..e1797784b84 100644 --- a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs @@ -118,7 +118,7 @@ pub(super) struct ParityCase { pub(super) fn parity_cases() -> Vec { serde_json::from_str(include_str!(concat!( env!("CARGO_MANIFEST_DIR"), - "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" + "/../../../tests/unit/secret_managers/hashicorp_vault_parity.json" ))) .unwrap() } diff --git a/litellm-rust/crates/secrets/PARITY.md b/litellm-rust/crates/secrets/PARITY.md index aeed4ba4b83..9439da93d65 100644 --- a/litellm-rust/crates/secrets/PARITY.md +++ b/litellm-rust/crates/secrets/PARITY.md @@ -32,7 +32,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo `_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed -## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py) +## [tests/unit/secret_managers/test_aws_secret_manager_replication.py](../../../tests/unit/secret_managers/test_aws_secret_manager_replication.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -48,7 +48,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) | | `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) | -## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py) +## [tests/unit/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/unit/secret_managers/test_aws_secret_manager_rotation.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -59,7 +59,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) | | `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) | -## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py) +## [tests/unit/secret_managers/test_aws_secret_manager_v2.py](../../../tests/unit/secret_managers/test_aws_secret_manager_v2.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -70,14 +70,14 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | | `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py) +## [tests/unit/secret_managers/test_base_secret_manager.py](../../../tests/unit/secret_managers/test_base_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | | `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) | | `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) | -## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py) +## [tests/unit/secret_managers/test_custom_secret_manager.py](../../../tests/unit/secret_managers/test_custom_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -89,7 +89,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) | | `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | -## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py) +## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -97,7 +97,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) | | `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py) +## [tests/unit/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/unit/secret_managers/test_get_azure_ad_token_provider.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -115,7 +115,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | | `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | -## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py) +## [tests/unit/secret_managers/test_hashicorp_secret_manager.py](../../../tests/unit/secret_managers/test_hashicorp_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -130,13 +130,13 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) | | `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py) +## [tests/unit/secret_managers/test_secret_manager_handler.py](../../../tests/unit/secret_managers/test_secret_manager_handler.py) | Python test | Rust coverage or boundary | | --- | --- | | `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) | -## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py) +## [tests/unit/secret_managers/test_secret_managers_main.py](../../../tests/unit/secret_managers/test_secret_managers_main.py) | Python test | Rust coverage or boundary | | --- | --- | diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 01670be74c8..4e3b92edc89 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -10,6 +10,7 @@ import os from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence from contextlib import AbstractAsyncContextManager from functools import partial +from importlib.metadata import version from types import MappingProxyType from typing import Final, TypeAlias, TypeVar, cast @@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams] from mcp.types import ( METHOD_NOT_FOUND, REQUEST_TIMEOUT, + ClientCapabilities, + ElicitationCapability, + FormElicitationCapability, GetPromptRequestParams, GetPromptResult, + Implementation, + InitializedNotification, + InitializeRequest, + InitializeRequestParams, + InitializeResult, InputRequiredResult, ListPromptsResult, ListResourcesResult, @@ -44,12 +53,14 @@ from mcp.types import ( PaginatedResult, Prompt, ResourceTemplate, + SamplingCapability, ServerNotification, + UrlElicitationCapability, ) from mcp.types import CallToolRequestParams as MCPCallToolRequestParams from mcp.types import CallToolResult as MCPCallToolResult from mcp.types import Tool as MCPTool -from pydantic import AnyUrl +from pydantic import AnyUrl, TypeAdapter from litellm._logging import verbose_logger from litellm.constants import ( @@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result from litellm.types.llms.custom_http import VerifyTypes from litellm.types.mcp import ( + MCP_LEGACY_VERSIONS, MCPAuth, MCPAuthType, MCPStdioConfig, MCPTransport, MCPTransportType, + MCPUpstreamProtocol, credential_redirect_hook, has_header, without_header, @@ -386,7 +399,9 @@ class MCPClient: sampling_callback: Callable | None = None, elicitation_callback: Callable | None = None, logging_callback: Callable | None = None, + protocol_version: MCPUpstreamProtocol = "auto", ): + self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version) self.server_url: str = server_url self.transport_type: MCPTransport = transport_type self.auth_type: MCPAuthType = auth_type @@ -525,6 +540,35 @@ class MCPClient: return safe_env + async def _initialize_session(self, session: ClientSession) -> InitializeResult: + if self.protocol_version == "auto": + automatic: Final = await session.initialize() + if automatic.protocol_version not in MCP_LEGACY_VERSIONS: + raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version") + return automatic + result: Final = await session.send_request( + InitializeRequest( + params=InitializeRequestParams( + protocol_version=self.protocol_version, + client_info=Implementation(name="litellm", version=version("litellm")), + capabilities=ClientCapabilities( + sampling=SamplingCapability() if self._sampling_callback is not None else None, + elicitation=ElicitationCapability( + form=FormElicitationCapability(), url=UrlElicitationCapability() + ) + if self._elicitation_callback is not None + else None, + ), + ) + ), + InitializeResult, + ) + if result.protocol_version != self.protocol_version: + raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version") + session.adopt(result) + await session.send_notification(InitializedNotification()) + return result + async def _execute_session_operation( self, transport_ctx: _TransportContext, @@ -579,7 +623,7 @@ class MCPClient: ) session: Final = await session_ctx.__aenter__() try: - init_result: Final = await session.initialize() + init_result: Final = await self._initialize_session(session) instructions: Final = getattr(init_result, "instructions", None) self._last_initialize_instructions = ( instructions.strip() or None if isinstance(instructions, str) else None diff --git a/litellm/integrations/levo/README.md b/litellm/integrations/levo/README.md index 5296acb7ff4..1fbd202d9a5 100644 --- a/litellm/integrations/levo/README.md +++ b/litellm/integrations/levo/README.md @@ -92,7 +92,7 @@ litellm/integrations/levo/ ## Testing -See the test files in `tests/test_litellm/integrations/levo/`: +See the test files in `tests/unit/integrations/levo/`: - `test_levo.py`: Unit tests for configuration - `test_levo_integration.py`: Integration tests for callback registration diff --git a/litellm/proxy/_experimental/mcp_server/capabilities.py b/litellm/proxy/_experimental/mcp_server/capabilities.py new file mode 100644 index 00000000000..bfd00327eb4 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/capabilities.py @@ -0,0 +1,145 @@ +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from itertools import product +from types import MappingProxyType +from typing import Final, Literal + +from mcp.server.context import CallNext, HandlerResult, ServerRequestContext +from mcp.shared.exceptions import MCPError +from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities +from mcp_types.methods import CLIENT_REQUESTS +from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION +from pydantic import TypeAdapter + +from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport + +GATEWAY_OPERATIONS: Final = frozenset( + { + "tools/list", + "tools/call", + "prompts/list", + "prompts/get", + "resources/list", + "resources/read", + "resources/templates/list", + } +) + + +@dataclass(frozen=True, slots=True) +class RevisionSupport: + transports: frozenset[MCPTransport] + operations: frozenset[str] + results: frozenset[Literal["complete", "input_required"]] + extensions: frozenset[str] + completed: bool + + +REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType( + { + version.value: RevisionSupport( + transports=frozenset(MCPTransport) + if version.value in HANDSHAKE_PROTOCOL_VERSIONS + else frozenset({MCPTransport.http, MCPTransport.stdio}), + operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS), + results=frozenset({"complete"}) + if version.value in HANDSHAKE_PROTOCOL_VERSIONS + else frozenset({"complete", "input_required"}), + extensions=frozenset(), + completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS, + ) + for version in MCPSpecVersion + } +) +_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed) +TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2)) +_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions) + + +def configured_versions() -> tuple[str, ...]: + from litellm.proxy.proxy_server import general_settings_view + + configured: Final = general_settings_view().get("mcp_advertised_versions") + return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured) + + +def build_discovery( + *, + configured: tuple[str, ...], + revision: str, + transport: MCPTransport, + authorized_operations: frozenset[str], + upstream_versions: frozenset[str], + capabilities: ServerCapabilities, + client_extensions: frozenset[str] = frozenset(), + upstream_extensions: frozenset[str] = frozenset(), + instructions: str | None = None, +) -> DiscoverResult: + supported: Final = tuple( + version + for version, support in REVISION_SUPPORT.items() + if version in configured and support.completed and transport in support.transports + ) + revision_support: Final = REVISION_SUPPORT.get(revision) + operations: Final[frozenset[str]] = ( + authorized_operations & revision_support.operations + if revision in supported + and revision_support is not None + and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions) + else frozenset() + ) + extensions: Final[frozenset[str]] = ( + revision_support.extensions & client_extensions & upstream_extensions + if operations and revision_support is not None + else frozenset() + ) + caller_capabilities: Final = capabilities.model_copy(deep=True) + return DiscoverResult( + supported_versions=list(supported), + capabilities=ServerCapabilities( + tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None, + prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None, + resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None, + extensions={ + key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions + } + or None, + ), + instructions=instructions, + cache_scope="private", + ttl_ms=0, + ) + + +class GatewayVersionPolicy: + def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None: + self._versions = versions + + async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult: + versions: Final = self._versions() + requested: Final = ( + InitializeRequestParams.model_validate(ctx.params or {}).protocol_version + if ctx.method == "initialize" + else ctx.protocol_version + ) + negotiated: Final = ( + (requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION) + if ctx.method == "initialize" + else requested + ) + if negotiated not in versions: + raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)}) + result: Final = await call_next(ctx) + if ctx.method != "initialize": + return result + initialized: Final = InitializeResult.model_validate(result) + discovery: Final = build_discovery( + configured=versions, + revision=initialized.protocol_version, + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=initialized.capabilities, + instructions=initialized.instructions, + ) + return initialized.model_copy(update={"capabilities": discovery.capabilities}) diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index a88d400282c..a0a08dc08ce 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -28,6 +28,7 @@ class OperationContext: client_ip: str | None = None mcp_proxy_mode: bool = False wire_compat: WireCompat = WireCompat.LEGACY + protocol_version: str | None = None def __post_init__(self) -> None: object.__setattr__(self, "_caller", copy_caller(self._caller)) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 24cae976174..6baa695433c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -193,6 +193,7 @@ from litellm.types.mcp import ( MCPAuth, MCPStdioConfig, MCPTokenEndpointAuthMethod, + MCPUpstreamProtocol, has_header, without_header, ) @@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False): whatever the admin wrote, and each read applies its own default.""" server_id: ReadOnly[str] + protocol_version: ReadOnly[MCPUpstreamProtocol] alias: str description: str mcp_info: MCPInfo @@ -2549,6 +2551,9 @@ class MCPServerManager: new_server = MCPServer( server_id=server_id, name=name_for_prefix, + protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python( + server_config.get("protocol_version", mcp_info.get("protocol_version", "auto")) + ), alias=alias, server_name=server_name, spec_path=server_config.get("spec_path", None), @@ -3109,6 +3114,9 @@ class MCPServerManager: new_server: Final = MCPServer( server_id=mcp_server.server_id, name=name_for_prefix, + protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python( + _mcp_info.get("protocol_version", "auto") + ), alias=getattr(mcp_server, "alias", None), server_name=getattr(mcp_server, "server_name", None), url=mcp_server.url, @@ -4145,6 +4153,7 @@ class MCPServerManager: cred_provider: UpstreamCredentialProvider | None = None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, + protocol_version_override: MCPUpstreamProtocol | None = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -4168,6 +4177,9 @@ class MCPServerManager: """ record_auth_resolution(server.server_id, AuthResolution.unresolved) resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) + protocol_version: Final = ( + protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version + ) transport: Final = resolved_server.transport or MCPTransport.sse spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server) provider: Final = cred_provider or self._cred_provider @@ -4249,6 +4261,7 @@ class MCPServerManager: return MCPClient( server_url="", # Not used for stdio transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, auth_value=auth_value, timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), @@ -4281,6 +4294,7 @@ class MCPServerManager: MCPClient( server_url=server_url, transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, timeout=( resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT @@ -4324,6 +4338,7 @@ class MCPServerManager: MCPClient( server_url=server_url, transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, auth_value=auth_value, auth_header_name=auth_header_name, diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index ebd26e4bf87..a19246b6e90 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -14,6 +14,8 @@ from mcp.types import ( CallToolRequest, CallToolRequestParams, CallToolResult, + DiscoverRequest, + DiscoverResult, GetPromptRequest, GetPromptRequestParams, GetPromptResult, @@ -28,10 +30,14 @@ from mcp.types import ( ListToolsResult, PaginatedRequestParams, Prompt, + PromptsCapability, ReadResourceRequest, ReadResourceRequestParams, + ResourcesCapability, ResourceTemplate, + ServerCapabilities, TextContent, + ToolsCapability, ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter @@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( cache_byok_credential, get_cached_byok_credential, ) +from litellm.proxy._experimental.mcp_server.capabilities import ( + GATEWAY_OPERATIONS, + build_discovery, + configured_versions, +) from litellm.proxy._experimental.mcp_server.contracts import ( AuthorizedToolCall, OperationContext, @@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import ( ) from litellm.types.mcp import ( DEFAULT_CREDENTIAL_HEADER, + MCP_LEGACY_VERSIONS, MCPAuth, + MCPTransport, without_header, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer @@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict): async def _execute_handle_list_tools( - context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None + context: OperationContext, + params: PaginatedRequestParams, + host_progress_callback: ProgressCallback | None = None, + *, + log_list_tools_to_spendlogs: bool = True, ) -> ListToolsResult: try: ( @@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, - log_list_tools_to_spendlogs=True, + log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, list_tools_log_source="mcp_protocol", client_ip=_client_ip, ) @@ -3065,6 +3082,7 @@ def prepare_context( client_ip: str | None = None, mcp_proxy_mode: bool = False, wire_compat: WireCompat = WireCompat.LEGACY, + protocol_version: str | None = None, ) -> OperationContext: return OperationContext( _caller=user_api_key_auth, @@ -3076,11 +3094,13 @@ def prepare_context( client_ip=client_ip, mcp_proxy_mode=mcp_proxy_mode, wire_compat=wire_compat, + protocol_version=protocol_version, ) GatewayOperation: TypeAlias = ( AuthorizedToolCall + | DiscoverRequest | ListToolsRequest | CallToolRequest | ListPromptsRequest @@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = ( | ReadResourceRequest ) GatewayResult: TypeAlias = ( - ListToolsResult + DiscoverResult + | ListToolsResult | CallToolResult | InputRequiredResult | ListPromptsResult @@ -3105,6 +3126,9 @@ class GatewayOperations: def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None: self._host_progress_callback = host_progress_callback + @overload + async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ... + @overload async def execute( self, operation: AuthorizedToolCall, context: OperationContext @@ -3137,6 +3161,51 @@ class GatewayOperations: async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult: match operation: + case DiscoverRequest(): + listings: Final = ( + () + if context.mcp_proxy_mode + else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest()) + ) + tasks: Final = ( + asyncio.create_task( + _execute_handle_list_tools( + context, + PaginatedRequestParams(), + self._host_progress_callback, + log_list_tools_to_spendlogs=False, + ) + ), + *(asyncio.create_task(self.execute(listing, context)) for listing in listings), + ) + try: + results: Final = await asyncio.gather(*tasks) + finally: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + return build_discovery( + configured=configured_versions(), + revision=context.protocol_version or "2025-11-25", + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=frozenset(MCP_LEGACY_VERSIONS), + capabilities=ServerCapabilities( + tools=ToolsCapability() + if any(isinstance(result, ListToolsResult) and result.tools for result in results) + else None, + prompts=PromptsCapability() + if any(isinstance(result, ListPromptsResult) and result.prompts for result in results) + else None, + resources=ResourcesCapability() + if any( + (isinstance(result, ListResourcesResult) and result.resources) + or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates) + for result in results + ) + else None, + ), + ) case AuthorizedToolCall(): auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth() return await _execute_mcp_tool( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7f519e2c0d9..02694f110b1 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1375,7 +1375,16 @@ if MCP_AVAILABLE: and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) else None ) - return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers) + preview_request: Final = ( + request.model_copy( + update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}} + ) + if saved_server is not None and "protocol_version" not in (request.mcp_info or {}) + else request + ) + return _StagedServerTest( + request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers + ) async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None: with anyio.move_on_after(deadline): @@ -1512,6 +1521,7 @@ if MCP_AVAILABLE: extra_headers=merged_headers, stdio_env=stdio_env, cred_provider=preview_cred_provider, + protocol_version_override=server_model.protocol_version, ) return await operation(client) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 1bd31d971b0..555aebc7434 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None: ``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which bypasses litellm's session/auth model, so the ASGI entry rejects it. """ + from litellm.proxy._experimental.mcp_server.capabilities import configured_versions + headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or () values: Final = tuple( raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER ) for value in values: - if value and value not in HANDSHAKE_PROTOCOL_VERSIONS: + if value and value not in configured_versions(): return value return None @@ -149,7 +151,10 @@ try: from mcp.server.session import ServerSession as _McpServerSession from mcp.types import ( BlobResourceContents, + DiscoverRequest, + DiscoverResult, GetPromptResult, + RequestParams, ResourceTemplate, TextResourceContents, ) @@ -526,11 +531,11 @@ if MCP_AVAILABLE: PaginatedRequestParams, ReadResourceRequestParams, ) - from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import ( MCPAuthenticatedUser, ) + from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, global_mcp_server_manager, @@ -585,6 +590,7 @@ if MCP_AVAILABLE: name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, ) + server.middleware.append(GatewayVersionPolicy()) server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server) sse: Final[SseServerTransport] = SseServerTransport("/sse/messages") @@ -830,6 +836,7 @@ if MCP_AVAILABLE: client_ip, _mcp_proxy_mode.get(), wire_compat_for(ctx.protocol_version), + ctx.protocol_version, ) async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: @@ -948,6 +955,11 @@ if MCP_AVAILABLE: ReadResourceRequest(params=params), context ) + async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult: + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context) + + server.add_request_handler("server/discover", RequestParams, discover) server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools) server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call) server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts) @@ -1954,7 +1966,7 @@ if MCP_AVAILABLE: reject_disallowed_mcp_origin(StarletteRequest(scope)) bad_version: Final = unsupported_protocol_version(scope) if bad_version is not None: - supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS)) + supported: Final = ", ".join(configured_versions()) await JSONResponse( status_code=400, content={ # mutable-ok: JSON-RPC error payload @@ -2299,7 +2311,7 @@ if MCP_AVAILABLE: reject_disallowed_mcp_origin(StarletteRequest(scope)) bad_version: Final = unsupported_protocol_version(scope) if bad_version is not None: - supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS)) + supported: Final = ", ".join(configured_versions()) await JSONResponse( status_code=400, content={ # mutable-ok: JSON-RPC error payload diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3da30070b9e..89fa0058644 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -37,6 +37,7 @@ from litellm.types.llms.openai import ( ResponsesAPIResponse, ) from litellm.types.mcp import ( + MCPAdvertisedVersions, MCPAllowedClient, MCPAuth, MCPAuthType, @@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).", ) + mcp_advertised_versions: MCPAdvertisedVersions | None = Field( + None, + description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. " + "Modern protocol serving and Apps/Tasks remain disabled.", + ) mcp_allowed_clients: list[MCPAllowedClient] | None = Field( None, description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ed4ea347c2e..61304f0d919 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6342,6 +6342,11 @@ class ProxyConfig: if general_settings is None: general_settings = {} + if general_settings.get("mcp_advertised_versions") is not None: + from litellm.types.mcp import MCPAdvertisedVersions + + TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"]) + if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) if declared_proxy_ranges(general_settings) is None: diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 83e719810d5..fec5e84c8df 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -4,7 +4,7 @@ import enum import re from collections.abc import Awaitable, Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal from urllib.parse import urlsplit import httpx @@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum): nov_2024 = "2024-11-05" mar_2025 = "2025-03-26" jun_2025 = "2025-06-18" + nov_2025 = "2025-11-25" + jul_2026 = "2026-07-28" class MCPAuth(str, enum.Enum): @@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok # MCP Literals MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio] -MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025] +MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"] +MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25") +MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"] +MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)] +MCPSpecVersionType = Literal[ + MCPSpecVersion.nov_2024, + MCPSpecVersion.mar_2025, + MCPSpecVersion.jun_2025, + MCPSpecVersion.nov_2025, + MCPSpecVersion.jul_2026, +] MCPAuthType = ( Literal[ MCPAuth.none, diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cb32299b143..c3b106c11d5 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,7 +1,7 @@ from datetime import datetime -from typing import Any, Final, Literal +from typing import Annotated, Any, Final, Literal -from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator from typing_extensions import Self from litellm.types.mcp import ( @@ -10,11 +10,19 @@ from litellm.types.mcp import ( MCPAuthType, MCPTokenEndpointAuthMethod, MCPTransportType, + MCPUpstreamProtocol, normalize_upstream_header_name, ) + # MCPInfo now allows arbitrary additional fields for custom metadata -MCPInfo = dict[str, Any] +def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]: + if "protocol_version" in value: + TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"]) + return value + + +MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)] class MCPOAuthMetadata(BaseModel): @@ -66,6 +74,7 @@ class MCPServer(BaseModel): server_name: str | None = None url: str | None = None transport: MCPTransportType + protocol_version: MCPUpstreamProtocol = "auto" spec_path: str | None = None auth_type: MCPAuthType | None = None authentication_token: str | None = None @@ -246,6 +255,14 @@ class MCPServer(BaseModel): """ return self.oauth2_flow == "client_credentials" + @model_validator(mode="after") + def resolve_protocol_version(self) -> Self: + if "protocol_version" not in self.model_fields_set and self.mcp_info is not None: + self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python( + self.mcp_info.get("protocol_version", "auto") + ) + return self + @model_validator(mode="after") def validate_identity_binding_mode(self) -> Self: binding: Final = self.oauth_identity_binding diff --git a/tests/integration/mcp/test_mcp_protocol_errors.py b/tests/integration/mcp/test_mcp_protocol_errors.py index bb06d8c6068..257e90e9d0e 100644 --- a/tests/integration/mcp/test_mcp_protocol_errors.py +++ b/tests/integration/mcp/test_mcp_protocol_errors.py @@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway) control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5}) assert control.status_code == 200 and control.json()["isError"] is False, control.text assert control.json()["content"][0]["text"] == "8" + + +@pytest.mark.parametrize("ingress", ("http", "sse")) +def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control( + gateway: Gateway, tmp_path, ingress: str +) -> None: + import asyncio + from pathlib import Path + + import yaml + from integration._support.mcp import mcp_peer + from integration._support.process import owned_proxy + from litellm.experimental_mcp_client.client import MCPClient + from litellm.types.mcp import MCPTransport + from mcp import MCPError + from mcp.types import CallToolRequestParams + + with mcp_peer() as upstream, gateway.scenario() as scenario: + alias: Final = "restricted" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"] + config_path: Final = tmp_path / "restricted.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted: + endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp") + headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity} + denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers) + allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers) + + async def exercise() -> None: + with pytest.raises(MCPError, match="Unsupported MCP protocol version"): + await denied.list_tools(raise_on_error=True) + assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True)) + result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5})) + assert result.is_error is False and result.content[0].text == "7" + + asyncio.run(exercise()) diff --git a/tests/integration/mcp/test_mcp_transports.py b/tests/integration/mcp/test_mcp_transports.py index 1862f11d07e..16004cdd501 100644 --- a/tests/integration/mcp/test_mcp_transports.py +++ b/tests/integration/mcp/test_mcp_transports.py @@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success assert outcome.error is not None, outcome.raw assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw assert len(tool_calls(peer.drain())) == 1 + + +@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio")) +@pytest.mark.parametrize("ingress", ("http", "sse")) +def test_pinned_revision_pairs_list_and_call_through_gateway( + gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str +) -> None: + import asyncio + + from mcp.types import CallToolRequestParams + + from litellm.experimental_mcp_client.client import MCPClient + from litellm.types.mcp import MCPTransport + + with peer_of(peer_kind) as peer, gateway.scenario() as scenario: + alias: Final = "versions" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream}) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp") + client: Final = MCPClient( + server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream, + extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15, + ) + + async def exercise() -> None: + tools: Final = await client.list_tools(raise_on_error=True) + assert f"{alias}-add" in tuple(tool.name for tool in tools) + result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4})) + assert result.is_error is False + assert result.content[0].text == "7" + + peer.drain() + asyncio.run(exercise()) + observed: Final = peer.drain() + negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize") + assert negotiations, "The operation must reach the upstream negotiation" + assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations + assert len(tool_calls(observed)) == 1 diff --git a/tests/test_litellm/integrations/levo/__init__.py b/tests/test_litellm/integrations/levo/__init__.py deleted file mode 100644 index 1560e78b7b9..00000000000 --- a/tests/test_litellm/integrations/levo/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Levo integration tests diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_capabilities.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_capabilities.py new file mode 100644 index 00000000000..f104e655704 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_capabilities.py @@ -0,0 +1,106 @@ +from typing import Final + +import pytest +from mcp import Client +from mcp.server import Server +from mcp.shared.exceptions import MCPError +from mcp.types import PromptsCapability, ResourcesCapability, ServerCapabilities, ToolsCapability +from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS + +from litellm.proxy._experimental.mcp_server.capabilities import ( + GATEWAY_OPERATIONS, + REVISION_SUPPORT, + TRANSLATION_PAIRS, + GatewayVersionPolicy, + build_discovery, +) +from litellm.types.mcp import MCPTransport + + +@pytest.mark.parametrize("revision", HANDSHAKE_PROTOCOL_VERSIONS) +@pytest.mark.parametrize("transport", tuple(MCPTransport)) +def test_discovery_only_exposes_authorized_completed_support(revision, transport): + result = build_discovery( + configured=(revision, "2026-07-28", "unknown"), + revision=revision, + transport=transport, + authorized_operations=frozenset({"tools/list", "tools/call"}), + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=ServerCapabilities( + tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability(), + extensions={"io.modelcontextprotocol/ui": {}}, + ), + client_extensions=frozenset({"io.modelcontextprotocol/ui"}), + upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}), + ) + assert result.supported_versions == [revision] + assert result.capabilities.tools is not None + assert result.capabilities.prompts is None + assert result.capabilities.resources is None + assert result.capabilities.extensions is None + assert result.capabilities.tasks is None + assert result.cache_scope == "private" + assert result.ttl_ms == 0 + + +@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})]) +def test_unproven_translation_never_advertises_operations(upstream): + result = build_discovery( + configured=HANDSHAKE_PROTOCOL_VERSIONS, + revision="2025-11-25", + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=upstream, + capabilities=ServerCapabilities(tools=ToolsCapability()), + ) + assert result.capabilities.tools is None + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "2024-11-05"]) +def test_unadvertised_revision_never_gains_capabilities(revision): + result = build_discovery( + configured=("2025-11-25",), revision=revision, transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=ServerCapabilities(tools=ToolsCapability()), + ) + assert result.capabilities.tools is None + + +def test_discovery_results_do_not_share_mutable_capabilities(): + capabilities = ServerCapabilities(tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability()) + args = dict( + configured=HANDSHAKE_PROTOCOL_VERSIONS, revision="2025-11-25", transport=MCPTransport.http, + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), capabilities=capabilities, + ) + allowed = build_discovery(**args, authorized_operations=GATEWAY_OPERATIONS) + denied = build_discovery(**args, authorized_operations=frozenset()) + assert allowed.capabilities.prompts is not None + assert allowed.capabilities.resources is not None + assert denied.capabilities.model_dump(exclude_none=True) == {} + assert allowed.capabilities.tools is not None + allowed.capabilities.tools.list_changed = True + assert capabilities.tools.list_changed is not True + + +def test_modern_candidates_do_not_enable_public_serving(): + modern = REVISION_SUPPORT["2026-07-28"] + assert modern.completed is False + assert "input_required" in modern.results + assert MCPTransport.sse not in modern.transports + assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("versions,accepted", [(("2025-11-25",), True), (("2025-06-18",), False)]) +async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted): + server: Final = Server("test-gateway", version="1") + server.middleware.append(GatewayVersionPolicy(lambda: versions)) + if accepted: + async with Client(server, mode="legacy") as client: + assert client.protocol_version == "2025-11-25" + result = await client.session.send_ping() + assert result is not None + else: + with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True): + async with Client(server, mode="legacy"): + pytest.fail("The excluded revision must not initialize") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 97b242831a2..ab00ec4da1e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10444,11 +10444,12 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque ) @pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp")) async def test_streamable_http_rejects_modern_protocol_version( - header_value: str, expected_rejected: bool, handler: str + header_value: str, expected_rejected: bool, handler: str, monkeypatch: pytest.MonkeyPatch ) -> None: from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) scope: Scope = { "type": "http", "method": "POST", @@ -10659,3 +10660,32 @@ async def test_legacy_sse_mount_emits_message_endpoint( await incoming.put({"type": "http.disconnect"}) await asyncio.wait_for(task, 2) assert await post(initialization) == 404 + + +@pytest.mark.parametrize("revision,rejected", [("2024-11-05", False), ("2025-11-25", True), ("2026-07-28", True)]) +def test_protocol_header_respects_configured_advertisement(revision, rejected): + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version + + with patch.dict(proxy_server.general_settings, {"mcp_advertised_versions": ["2024-11-05"]}): + result = unsupported_protocol_version({"headers": [(b"mcp-protocol-version", revision.encode())]}) + assert result == (revision if rejected else None) + + +@pytest.mark.asyncio +async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx): + from mcp.types import DiscoverResult, RequestParams, ServerCapabilities + from litellm.proxy._experimental.mcp_server import server + + expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities()) + 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.operations.GatewayOperations, "execute", dispatched), + ): + result = await server.discover(_mcp_request_ctx(), RequestParams()) + assert result is expected + context = dispatched.await_args.args[1] + assert context.user_api_key_auth.user_id == "discover-caller" + assert context.mcp_servers == ("allowed",) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 64d94065674..16bffa1a356 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -63,7 +63,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.types.llms.custom_http import httpxSpecialProvider -from litellm.types.mcp import MCPAuth, MCPAuthType +from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer from litellm.caching.caching import DualCache from litellm.caching.llm_caching_handler import LLMClientCache @@ -14637,3 +14637,27 @@ class TestSharedIdentifierPrefixWarning: assert "srv-b" in shared_warnings[0] assert "srv-c" not in shared_warnings[0] assert "'shared'" in shared_warnings[0] + + +@pytest.mark.asyncio +@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}}) + server = next(iter(manager.config_mcp_servers.values())) + client = await manager._create_mcp_client(server) + assert server.protocol_version == revision + assert client.protocol_version == revision + + +@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18")) +@pytest.mark.parametrize("explicit", (None, "auto", "2025-11-25")) +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 {}), + }) + assert server.protocol_version == (explicit if explicit is not None else revision) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py index 81f81045740..bb900de4f98 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -542,3 +542,126 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c assert [block.text for block in result.content] == [body] assert result.structured_content == (["a", "b"] if compat == "modern" else None) + + +@pytest.mark.asyncio +async def test_discovery_preserves_caller_scope_and_proxy_restrictions(): + from mcp.types import DiscoverRequest, ListToolsResult, Tool + + listed = AsyncMock(return_value=ListToolsResult(tools=[Tool(name="allowed", input_schema={"type": "object"})])) + context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["only-this"], mcp_proxy_mode=True, protocol_version="2025-06-18") + with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", listed): + result = await GatewayOperations().execute(DiscoverRequest(), context) + assert result.capabilities.tools is not None + assert result.capabilities.resources is None + assert result.capabilities.prompts is None + assert listed.await_args.args[0] is context + assert listed.await_args.args[0].user_api_key_auth.user_id == "scoped" + assert listed.await_args.args[0].mcp_servers == ("only-this",) + + +@pytest.mark.asyncio +async def test_discovery_denial_cannot_advertise_tools(): + from mcp.types import DiscoverRequest + from fastapi import HTTPException + + denied = AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden")) + with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", denied): + with pytest.raises(HTTPException) as error: + await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="denied"))) + assert error.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("available", ["none", "resources", "templates", "prompts"]) +async def test_discovery_lists_each_capability_with_the_same_caller(available): + from mcp.types import ( + DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, + ListResourceTemplatesResult, Prompt, Resource, ResourceTemplate, + ) + from litellm.proxy._experimental.mcp_server import operations + + context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["authorized"]) + tools = AsyncMock(return_value=ListToolsResult(tools=[])) + prompts = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="allowed")] if available == "prompts" else [])) + resources = AsyncMock(return_value=ListResourcesResult(resources=[Resource(name="allowed", uri="test://allowed")] if available == "resources" else [])) + templates = AsyncMock(return_value=ListResourceTemplatesResult(resource_templates=[ResourceTemplate(name="allowed", uri_template="test://{id}")] if available == "templates" else [])) + with ( + patch.object(operations, "_execute_handle_list_tools", tools), + patch.object(operations, "_execute_list_prompts", prompts), + patch.object(operations, "_execute_list_resources", resources), + patch.object(operations, "_execute_list_resource_templates", templates), + ): + result = await GatewayOperations().execute(DiscoverRequest(), context) + assert result.capabilities.tools is None + assert (result.capabilities.prompts is not None) == (available == "prompts") + assert (result.capabilities.resources is not None) == (available in {"resources", "templates"}) + for listing in (tools, prompts, resources, templates): + assert listing.await_args.args[0] is context + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"]) +async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome): + from mcp.types import DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult + from litellm.proxy._experimental.mcp_server import operations + + ready = [asyncio.Event() for _ in range(4)] + closed = [asyncio.Event() for _ in range(4)] + release = asyncio.Event() + responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]), ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[])) + + def listing(index): + async def run(*args, **kwargs): + ready[index].set() + try: + await release.wait() + if index == 0 and outcome == "failure": + raise ValueError("discovery failed") + if outcome != "success": + await asyncio.Event().wait() + return responses[index] + finally: + closed[index].set() + return run + + with ( + patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)) as tools, + patch.object(operations, "_execute_list_prompts", side_effect=listing(1)), + patch.object(operations, "_execute_list_resources", side_effect=listing(2)), + patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)), + ): + task = asyncio.create_task(GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped")))) + try: + await asyncio.wait_for(asyncio.gather(*(event.wait() for event in ready)), 1) + if outcome == "cancel": + task.cancel() + else: + release.set() + if outcome == "success": + result = await asyncio.wait_for(task, 1) + assert result.capabilities.model_dump(exclude_none=True) == {} + else: + with pytest.raises(asyncio.CancelledError if outcome == "cancel" else ValueError): + await asyncio.wait_for(task, 1) + assert all(event.is_set() for event in closed) + assert tools.call_args.kwargs["log_list_tools_to_spendlogs"] is False + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("log_enabled", [False, True]) +async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled): + from mcp.types import PaginatedRequestParams + from litellm.proxy._experimental.mcp_server import operations + + listing = AsyncMock(return_value=operations.AggregateToolListing(tools=[], outcomes={})) + with patch.object(operations, "_list_mcp_tools", listing): + result = await operations._execute_handle_list_tools( + prepare_context(UserAPIKeyAuth(user_id="caller")), PaginatedRequestParams(), + log_list_tools_to_spendlogs=log_enabled, + ) + assert result.tools == [] + assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 4da120cb26f..e82ab28bb4c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -28,7 +28,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.mcp import MCPAuth, MCPTransport +from litellm.types.mcp import MCPAuth, MCPTransport, MCPUpstreamProtocol from litellm.types.mcp_server.mcp_server_manager import MCPServer _OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False) @@ -1476,7 +1476,7 @@ class TestListToolsRestAPI: monkeypatch, ): """The REST tools/list path should include tools beyond the upstream first page.""" - from mcp.types import ListToolsResult, PaginatedRequestParams + from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities from mcp.types import Tool as MCPTool import litellm.experimental_mcp_client.client as mcp_client_module @@ -1512,7 +1512,11 @@ class TestListToolsRestAPI: mock_session_ctx = AsyncMock() mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock(return_value=None) + mock_session_instance.initialize = AsyncMock(return_value=InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="stub", version="1"), + )) mock_session_instance.list_tools.side_effect = [ ListToolsResult( tools=[ @@ -4628,3 +4632,74 @@ class TestClientAllowlistOnRestRoutes: assert denied.value.detail["error"] == "Forbidden" assert "'claude-code'" in denied.value.detail["details"] acting.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18")) +async def test_preview_client_honors_protocol_metadata(revision: MCPUpstreamProtocol) -> None: + from litellm.experimental_mcp_client.client import MCPClient + + payload: Final = NewMCPServerRequest( + server_name="preview", url="http://127.0.0.1:9/mcp", transport="http", + auth_type=MCPAuth.none, mcp_info={"protocol_version": revision}, + ) + + async def inspect_client(client: MCPClient) -> dict[str, str]: + return {"protocol_version": client.protocol_version} + + result: Final = await rest_endpoints._execute_with_mcp_client(payload, inspect_client) + assert result == {"protocol_version": revision} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type", (MCPAuth.none, MCPAuth.bearer_token, MCPAuth.oauth2)) +@pytest.mark.parametrize( + ("metadata", "expected"), + ( + (None, "2025-11-25"), + ({}, "2025-11-25"), + ({"description": "edited"}, "2025-11-25"), + ({"protocol_version": "auto"}, "auto"), + ({"protocol_version": "2024-11-05"}, "2024-11-05"), + ({"protocol_version": "2025-06-18"}, "2025-06-18"), + ), +) +async def test_saved_preview_protocol_omission_and_explicit_edits( + monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuth, + metadata: dict[str, str] | None, expected: MCPUpstreamProtocol, +) -> None: + from starlette.datastructures import Headers + + from litellm.experimental_mcp_client.client import MCPClient + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy.management_endpoints import mcp_management_endpoints + + saved: Final = MCPServer( + server_id="saved-preview", name="preview", url="https://example.com/mcp", + transport="http", auth_type=auth_type, protocol_version="2025-11-25", + authentication_token="stored-token", + authorization_url="https://example.com/authorize", token_url="https://example.com/token", + ) + manager: Final = MCPServerManager() + manager.registry = {saved.server_id: saved} + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager) + payload: Final = NewMCPServerRequest( + server_id=saved.server_id, server_name=saved.name, url=saved.url, transport="http", + auth_type=auth_type, mcp_info=metadata, + authorization_url=saved.authorization_url, token_url=saved.token_url, + ) + staged: Final = rest_endpoints._stage_server_test( + payload, Headers({"x-litellm-api-key": "sk-admin", "authorization": "Bearer preview-token"}) + ) + + async def inspect_client(client: MCPClient) -> dict[str, str]: + return {"protocol_version": client.protocol_version} + + result: Final = await rest_endpoints._execute_with_mcp_client( + staged.request, inspect_client, + mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, + ) + assert result == {"protocol_version": expected} + assert saved.protocol_version == "2025-11-25" + assert payload.mcp_info == metadata diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 7e198bc9131..b2ef327f50e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -4903,3 +4903,19 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f assert await pc._get_models_from_db(client) == [] assert pc.auto_router_db_catalog == () assert find_many.await_count == 3 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("versions", [None, ["2024-11-05"], [], ["2026-07-28"], ["unknown"]]) +async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions): + config = tmp_path / "mcp-versions.yaml" + config.write_text(json.dumps({"model_list": [], "general_settings": {"mcp_advertised_versions": versions}})) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + if versions is None or versions == ["2024-11-05"]: + _, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config)) + assert settings["mcp_advertised_versions"] == versions + return + with pytest.raises(ValidationError): + await ProxyConfig().load_config(router=None, config_file_path=str(config)) diff --git a/tests/test_litellm/proxy/test__types.py b/tests/test_litellm/proxy/test__types.py index f5abe0561db..b43a75d3323 100644 --- a/tests/test_litellm/proxy/test__types.py +++ b/tests/test_litellm/proxy/test__types.py @@ -377,3 +377,21 @@ def test_change_password_request_passwords_hidden_from_repr(): for rendered in (repr(request), str(request)): assert "hunter2hunter2" not in rendered assert "NewP@ssw0rd-2026" not in rendered +@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]]) +def test_mcp_advertised_versions_reject_unavailable_revisions(versions): + from pydantic import ValidationError + + from litellm.proxy._types import ConfigGeneralSettings + + with pytest.raises(ValidationError): + ConfigGeneralSettings(mcp_advertised_versions=versions) + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None]) +def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision): + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}} + for model in (NewMCPServerRequest, UpdateMCPServerRequest): + with pytest.raises(ValidationError): + model.model_validate(payload) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 368e34c455d..1a56227b008 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -58,6 +58,15 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer _JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage) +def _initialized(instructions: str | None = None) -> InitializeResult: + return InitializeResult( + protocol_version=LATEST_HANDSHAKE_VERSION, + capabilities=ServerCapabilities(), + server_info=Implementation(name="test", version="1"), + instructions=instructions, + ) + + class _MockTransportClient(MCPClient): """An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport.""" @@ -125,7 +134,7 @@ class TestMCPClient: mock_stdio_client.return_value = mock_stdio_ctx mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -168,7 +177,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -214,7 +223,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -266,7 +275,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -413,8 +422,7 @@ class TestMCPClientInstructionsCapture: ) mock_session = AsyncMock() - init_result = MagicMock() - init_result.instructions = " upstream says hello " + init_result = _initialized(" upstream says hello ") mock_session.initialize = AsyncMock(return_value=init_result) session_ctx = MagicMock() @@ -442,8 +450,7 @@ class TestMCPClientInstructionsCapture: ) mock_session = AsyncMock() - init_result = MagicMock() - init_result.instructions = None + init_result = _initialized() mock_session.initialize = AsyncMock(return_value=init_result) session_ctx = MagicMock() @@ -600,8 +607,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls): client = MCPClient(server_url="http://example.com/mcp", transport_type="http") - init_result = MagicMock() - init_result.instructions = None + init_result = _initialized() self._make_session(mock_session_cls, AsyncMock(return_value=init_result)) transport_ctx = self._make_transport(_FakeExceptionGroup("late", [httpx2.ConnectError("late cleanup error")])) @@ -634,7 +640,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.parametrize("original_error", (False, True)) @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_session_exit_cancellation_preserves_original_failure(self, session_class, original_error): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) cancelled: Final = asyncio.CancelledError("cancelled while closing session") session_class.return_value.__aexit__ = AsyncMock(side_effect=cancelled) original: Final = RuntimeError("operation failed") @@ -656,7 +662,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.parametrize("phase", ("session", "transport")) @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_cleanup_preserves_process_exit(self, session_class, phase, signal_type): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) signal: Final = signal_type("process stopping") if phase == "session": session_class.return_value.__aexit__ = AsyncMock(side_effect=signal) @@ -670,7 +676,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.asyncio @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_session_and_termination_share_one_cleanup_deadline(self, session_class): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) deleting: Final = asyncio.Event() async def close_session(*args): @@ -1883,16 +1889,17 @@ async def test_sse_read_failure_is_preserved() -> None: @pytest.mark.asyncio +@pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"]) @pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio]) @pytest.mark.parametrize("mode", ["ok", "closed", "silent"]) -async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None: +async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None: from mcp import ClientSession from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message logging_callback: Final = AsyncMock() read_timeout: Final = 0.2 if mode == "silent" else 30 client: Final = MCPClient( - server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback + server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version ) async def operation(session: ClientSession) -> CallToolResult: @@ -2754,12 +2761,13 @@ async def test_http_close_cancellation_cannot_turn_into_success(original_error: @pytest.mark.asyncio +@pytest.mark.parametrize("protocol_version", ("auto", "2025-06-18")) @pytest.mark.parametrize("cancel_mode", ("scope", "task", "wait_for", "read_timeout")) @pytest.mark.parametrize("concurrency", (1, 5)) @pytest.mark.parametrize("termination", ("ok", "hang", "hang_body")) @pytest.mark.parametrize("raise_on_error", (False, True)) async def test_cancellation_delivers_termination_over_tcp( - cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool + cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool, protocol_version: str ) -> None: started: Final = asyncio.Event() scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future() @@ -2813,6 +2821,8 @@ async def test_cancellation_delivers_termination_over_tcp( await stop.wait() return if payload["method"] == "initialize": + if cancel_mode != "read_timeout": + await asyncio.sleep(0.75) response: Final = json.dumps( { "jsonrpc": "2.0", @@ -2839,7 +2849,7 @@ async def test_cancellation_delivers_termination_over_tcp( listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0) port: Final = listener.sockets[0].getsockname()[1] client: Final = MCPClient( - server_url=f"http://127.0.0.1:{port}/mcp", timeout=2 if cancel_mode == "read_timeout" else 0.5 if termination != "ok" else 30 + server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30 ) async def calls(): @@ -2866,7 +2876,7 @@ async def test_cancellation_delivers_termination_over_tcp( try: task: Final = asyncio.create_task(invoke()) - await asyncio.wait_for(started.wait(), 3) + await asyncio.wait_for(started.wait(), 30) if cancel_mode == "scope": (await scope_ready).deadline = anyio.current_time() + 0.2 if cancel_mode == "task": @@ -2901,3 +2911,52 @@ async def test_cancellation_delivers_termination_over_tcp( closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2) assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed await asyncio.wait_for(listener.wait_closed(), 2) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revision", ["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "auto"]) +@pytest.mark.parametrize("accepted", [True, False]) +@pytest.mark.parametrize("callbacks", [False, True]) +async def test_configured_upstream_revision_is_offered_and_checked(revision, accepted, callbacks): + from mcp.types import JSONRPCRequest + from mcp_types.version import LATEST_HANDSHAKE_VERSION + + offered = LATEST_HANDSHAKE_VERSION if revision == "auto" else revision + + def respond(request): + if request.method == "DELETE": + return httpx2.Response(200) + payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + if not isinstance(payload, JSONRPCRequest): + return httpx2.Response(202) + if payload.method == "initialize": + assert payload.params["protocolVersion"] == offered + assert ("sampling" in payload.params["capabilities"]) == callbacks + assert ("elicitation" in payload.params["capabilities"]) == callbacks + return httpx2.Response(200, json={ + "jsonrpc": "2.0", "id": payload.id, + "result": {"protocolVersion": offered if accepted else "unsupported", + "capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}}, + }) + assert accepted, "No operation may execute after failed version negotiation" + return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}}) + + client = _MockTransportClient( + respond, server_url="https://example.com/mcp", protocol_version=revision, + sampling_callback=AsyncMock() if callbacks else None, + elicitation_callback=AsyncMock() if callbacks else None, + ) + if accepted: + result = await client.list_tools(raise_on_error=True) + assert [tool.name for tool in result] == ["echo"] + else: + with pytest.raises((MCPError, RuntimeError), match="protocol version"): + await client.list_tools(raise_on_error=True) + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None]) +def test_upstream_protocol_configuration_rejects_unavailable_modes(revision): + from pydantic import ValidationError + + with pytest.raises(ValidationError): + MCPClient(protocol_version=revision) diff --git a/tests/test_litellm/integrations/code_interpreter_interception/__init__.py b/tests/unit/integrations/SlackAlerting/__init__.py similarity index 100% rename from tests/test_litellm/integrations/code_interpreter_interception/__init__.py rename to tests/unit/integrations/SlackAlerting/__init__.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py rename to tests/unit/integrations/SlackAlerting/test_budget_alert_types.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_hanging_request_check.py b/tests/unit/integrations/SlackAlerting/test_hanging_request_check.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_hanging_request_check.py rename to tests/unit/integrations/SlackAlerting/test_hanging_request_check.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_model_deprecation_alert.py b/tests/unit/integrations/SlackAlerting/test_model_deprecation_alert.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_model_deprecation_alert.py rename to tests/unit/integrations/SlackAlerting/test_model_deprecation_alert.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_ms_teams.py b/tests/unit/integrations/SlackAlerting/test_ms_teams.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_ms_teams.py rename to tests/unit/integrations/SlackAlerting/test_ms_teams.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_digest.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_digest.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_digest.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting_digest.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py b/tests/unit/integrations/SlackAlerting/test_user_spend_alerts.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py rename to tests/unit/integrations/SlackAlerting/test_user_spend_alerts.py diff --git a/tests/test_litellm/integrations/gitlab/__init__.py b/tests/unit/integrations/arize/__init__.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/__init__.py rename to tests/unit/integrations/arize/__init__.py diff --git a/tests/test_litellm/integrations/arize/test_arize.py b/tests/unit/integrations/arize/test_arize.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize.py rename to tests/unit/integrations/arize/test_arize.py diff --git a/tests/test_litellm/integrations/arize/test_arize_health_check.py b/tests/unit/integrations/arize/test_arize_health_check.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_health_check.py rename to tests/unit/integrations/arize/test_arize_health_check.py diff --git a/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py b/tests/unit/integrations/arize/test_arize_otel_coexistence.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py rename to tests/unit/integrations/arize/test_arize_otel_coexistence.py diff --git a/tests/test_litellm/integrations/arize/test_arize_phoenix.py b/tests/unit/integrations/arize/test_arize_phoenix.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_phoenix.py rename to tests/unit/integrations/arize/test_arize_phoenix.py diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/unit/integrations/arize/test_arize_utils.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_utils.py rename to tests/unit/integrations/arize/test_arize_utils.py diff --git a/tests/test_litellm/integrations/open_telemetry/__init__.py b/tests/unit/integrations/azure_storage/__init__.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/__init__.py rename to tests/unit/integrations/azure_storage/__init__.py diff --git a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py b/tests/unit/integrations/azure_storage/test_azure_storage.py similarity index 100% rename from tests/test_litellm/integrations/azure_storage/test_azure_storage.py rename to tests/unit/integrations/azure_storage/test_azure_storage.py diff --git a/tests/unit/integrations/bitbucket/__init__.py b/tests/unit/integrations/bitbucket/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py b/tests/unit/integrations/bitbucket/test_bitbucket_integration.py similarity index 100% rename from tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py rename to tests/unit/integrations/bitbucket/test_bitbucket_integration.py diff --git a/tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py b/tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py similarity index 93% rename from tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py rename to tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py index d6668bf9ad8..a1a88653ee6 100644 --- a/tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py +++ b/tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py @@ -307,37 +307,6 @@ def test_bitbucket_prompt_manager_render_template_not_found(): manager.prompt_manager.render_template("nonexistent", {"some": "variable"}) -@patch("litellm.integrations.bitbucket.bitbucket_prompt_manager.BitBucketClient") -def test_bitbucket_prompt_manager_integration(mock_client_class): - """Test BitBucketPromptManager integration with BitBucketClient.""" - # Mock the BitBucket client - mock_client = MagicMock() - mock_client.get_file_content.return_value = """--- -model: gpt-4 -temperature: 0.7 ---- -Hello {{name}}!""" - mock_client_class.return_value = mock_client - - config = { - "workspace": "test-workspace", - "repository": "test-repo", - "access_token": "test-token", - } - - manager = BitBucketPromptManager(config, prompt_id="test_prompt") - - # Should have loaded the prompt - assert "test_prompt" in manager.prompt_manager.prompts - template = manager.prompt_manager.prompts["test_prompt"] - assert template.model == "gpt-4" - assert template.temperature == 0.7 - - # Test rendering - rendered = manager.prompt_manager.render_template("test_prompt", {"name": "World"}) - assert rendered == "Hello World!" - - def test_bitbucket_prompt_manager_parse_prompt_to_messages(): """Test parsing prompt content into messages.""" config = { diff --git a/tests/unit/integrations/cloudzero/__init__.py b/tests/unit/integrations/cloudzero/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/unit/integrations/cloudzero/test_cloudzero.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cloudzero.py rename to tests/unit/integrations/cloudzero/test_cloudzero.py diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py b/tests/unit/integrations/cloudzero/test_cloudzero_database.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py rename to tests/unit/integrations/cloudzero/test_cloudzero_database.py diff --git a/tests/test_litellm/integrations/cloudzero/test_cz_stream_api.py b/tests/unit/integrations/cloudzero/test_cz_stream_api.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cz_stream_api.py rename to tests/unit/integrations/cloudzero/test_cz_stream_api.py diff --git a/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py b/tests/unit/integrations/cloudzero/test_dry_run_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py rename to tests/unit/integrations/cloudzero/test_dry_run_endpoint.py diff --git a/tests/test_litellm/integrations/cloudzero/test_transform.py b/tests/unit/integrations/cloudzero/test_transform.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_transform.py rename to tests/unit/integrations/cloudzero/test_transform.py diff --git a/tests/unit/integrations/code_interpreter_interception/__init__.py b/tests/unit/integrations/code_interpreter_interception/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py b/tests/unit/integrations/code_interpreter_interception/test_handler.py similarity index 100% rename from tests/test_litellm/integrations/code_interpreter_interception/test_handler.py rename to tests/unit/integrations/code_interpreter_interception/test_handler.py diff --git a/tests/test_litellm/integrations/conftest.py b/tests/unit/integrations/conftest.py similarity index 94% rename from tests/test_litellm/integrations/conftest.py rename to tests/unit/integrations/conftest.py index adc8e36e0af..48a01ed8d48 100644 --- a/tests/test_litellm/integrations/conftest.py +++ b/tests/unit/integrations/conftest.py @@ -1,6 +1,7 @@ import functools import http.server import ipaddress +import os import queue import ssl import threading @@ -74,6 +75,14 @@ def write_self_signed_cert(directory: Path, stem: str) -> tuple[Path, Path]: return certificate_path, key_path +@pytest.fixture(autouse=True) +def restore_process_environment() -> Iterator[None]: + original: Final = dict(os.environ) + yield + os.environ.clear() + os.environ.update(original) + + @pytest.fixture def tls_sink(tmp_path: Path) -> Iterator[TlsSink]: certificate_path, key_path = write_self_signed_cert(tmp_path, "sink") diff --git a/tests/unit/integrations/datadog/__init__.py b/tests/unit/integrations/datadog/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/datadog/test_datadog_cost_management.py b/tests/unit/integrations/datadog/test_datadog_cost_management.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_cost_management.py rename to tests/unit/integrations/datadog/test_datadog_cost_management.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py b/tests/unit/integrations/datadog/test_datadog_llm_obs.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py rename to tests/unit/integrations/datadog/test_datadog_llm_obs.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py b/tests/unit/integrations/datadog/test_datadog_llm_obs_agent.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py rename to tests/unit/integrations/datadog/test_datadog_llm_obs_agent.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py b/tests/unit/integrations/datadog/test_datadog_logger_batching.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py rename to tests/unit/integrations/datadog/test_datadog_logger_batching.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_metrics.py b/tests/unit/integrations/datadog/test_datadog_metrics.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_metrics.py rename to tests/unit/integrations/datadog/test_datadog_metrics.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py b/tests/unit/integrations/datadog/test_datadog_tags_regression.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py rename to tests/unit/integrations/datadog/test_datadog_tags_regression.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_team_handler.py b/tests/unit/integrations/datadog/test_datadog_team_handler.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_team_handler.py rename to tests/unit/integrations/datadog/test_datadog_team_handler.py diff --git a/tests/unit/integrations/dotprompt/__init__.py b/tests/unit/integrations/dotprompt/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.prompt b/tests/unit/integrations/dotprompt/chat_prompt.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.prompt diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.v1.prompt b/tests/unit/integrations/dotprompt/chat_prompt.v1.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.v1.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.v1.prompt diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.v2.prompt b/tests/unit/integrations/dotprompt/chat_prompt.v2.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.v2.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.v2.prompt diff --git a/tests/test_litellm/integrations/dotprompt/coding_assistant.prompt b/tests/unit/integrations/dotprompt/coding_assistant.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/coding_assistant.prompt rename to tests/unit/integrations/dotprompt/coding_assistant.prompt diff --git a/tests/test_litellm/integrations/dotprompt/sample_prompt.prompt b/tests/unit/integrations/dotprompt/sample_prompt.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/sample_prompt.prompt rename to tests/unit/integrations/dotprompt/sample_prompt.prompt diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/unit/integrations/dotprompt/test_prompt_manager.py similarity index 100% rename from tests/test_litellm/integrations/dotprompt/test_prompt_manager.py rename to tests/unit/integrations/dotprompt/test_prompt_manager.py diff --git a/tests/unit/integrations/focus/__init__.py b/tests/unit/integrations/focus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/focus/test_csv_serializer.py b/tests/unit/integrations/focus/test_csv_serializer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_csv_serializer.py rename to tests/unit/integrations/focus/test_csv_serializer.py diff --git a/tests/test_litellm/integrations/focus/test_destination_factory.py b/tests/unit/integrations/focus/test_destination_factory.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_destination_factory.py rename to tests/unit/integrations/focus/test_destination_factory.py diff --git a/tests/test_litellm/integrations/focus/test_focus_database.py b/tests/unit/integrations/focus/test_focus_database.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_database.py rename to tests/unit/integrations/focus/test_focus_database.py diff --git a/tests/test_litellm/integrations/focus/test_focus_gcs_destination.py b/tests/unit/integrations/focus/test_focus_gcs_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_gcs_destination.py rename to tests/unit/integrations/focus/test_focus_gcs_destination.py diff --git a/tests/test_litellm/integrations/focus/test_focus_transformer.py b/tests/unit/integrations/focus/test_focus_transformer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_transformer.py rename to tests/unit/integrations/focus/test_focus_transformer.py diff --git a/tests/test_litellm/integrations/focus/test_mavvrik_destination.py b/tests/unit/integrations/focus/test_mavvrik_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_mavvrik_destination.py rename to tests/unit/integrations/focus/test_mavvrik_destination.py diff --git a/tests/test_litellm/integrations/focus/test_s3_destination.py b/tests/unit/integrations/focus/test_s3_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_s3_destination.py rename to tests/unit/integrations/focus/test_s3_destination.py diff --git a/tests/test_litellm/integrations/focus/test_transformer.py b/tests/unit/integrations/focus/test_transformer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_transformer.py rename to tests/unit/integrations/focus/test_transformer.py diff --git a/tests/test_litellm/integrations/focus/test_vantage_destination.py b/tests/unit/integrations/focus/test_vantage_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_vantage_destination.py rename to tests/unit/integrations/focus/test_vantage_destination.py diff --git a/tests/unit/integrations/gitlab/__init__.py b/tests/unit/integrations/gitlab/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_client.py b/tests/unit/integrations/gitlab/test_gitlab_client.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_client.py rename to tests/unit/integrations/gitlab/test_gitlab_client.py diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py b/tests/unit/integrations/gitlab/test_gitlab_integration.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_integration.py rename to tests/unit/integrations/gitlab/test_gitlab_integration.py diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py b/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py rename to tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py diff --git a/tests/unit/integrations/langfuse/__init__.py b/tests/unit/integrations/langfuse/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py b/tests/unit/integrations/langfuse/test_gemini_cached_tokens.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py rename to tests/unit/integrations/langfuse/test_gemini_cached_tokens.py diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py b/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py rename to tests/unit/integrations/langfuse/test_langfuse_prompt_management.py diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py b/tests/unit/integrations/langfuse/test_langfuse_sdk.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py rename to tests/unit/integrations/langfuse/test_langfuse_sdk.py diff --git a/tests/unit/integrations/newrelic/__init__.py b/tests/unit/integrations/newrelic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic.py b/tests/unit/integrations/newrelic/test_newrelic.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic.py rename to tests/unit/integrations/newrelic/test_newrelic.py diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py b/tests/unit/integrations/newrelic/test_newrelic_metrics.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py rename to tests/unit/integrations/newrelic/test_newrelic_metrics.py diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic_team_handler.py b/tests/unit/integrations/newrelic/test_newrelic_team_handler.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic_team_handler.py rename to tests/unit/integrations/newrelic/test_newrelic_team_handler.py diff --git a/tests/unit/integrations/open_telemetry/__init__.py b/tests/unit/integrations/open_telemetry/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/open_telemetry/_helpers.py b/tests/unit/integrations/open_telemetry/_helpers.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/_helpers.py rename to tests/unit/integrations/open_telemetry/_helpers.py diff --git a/tests/test_litellm/integrations/open_telemetry/conftest.py b/tests/unit/integrations/open_telemetry/conftest.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/conftest.py rename to tests/unit/integrations/open_telemetry/conftest.py diff --git a/tests/unit/integrations/open_telemetry/data/__init__.py b/tests/unit/integrations/open_telemetry/data/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/open_telemetry/data/captured_kwargs.json b/tests/unit/integrations/open_telemetry/data/captured_kwargs.json similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/data/captured_kwargs.json rename to tests/unit/integrations/open_telemetry/data/captured_kwargs.json diff --git a/tests/test_litellm/integrations/open_telemetry/data/captured_response.json b/tests/unit/integrations/open_telemetry/data/captured_response.json similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/data/captured_response.json rename to tests/unit/integrations/open_telemetry/data/captured_response.json diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_admin_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_admin_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py rename to tests/unit/integrations/open_telemetry/test_otel_exception_handler.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_passthrough_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_passthrough_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_passthrough_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_passthrough_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_unified_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_unified_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_unified_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_unified_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_passthrough_parent_span.py b/tests/unit/integrations/open_telemetry/test_passthrough_parent_span.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_passthrough_parent_span.py rename to tests/unit/integrations/open_telemetry/test_passthrough_parent_span.py diff --git a/tests/unit/integrations/otel/__init__.py b/tests/unit/integrations/otel/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/otel/test_db_endpoint.py b/tests/unit/integrations/otel/test_db_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_db_endpoint.py rename to tests/unit/integrations/otel/test_db_endpoint.py diff --git a/tests/test_litellm/integrations/otel/test_langfuse_logger.py b/tests/unit/integrations/otel/test_langfuse_logger.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_langfuse_logger.py rename to tests/unit/integrations/otel/test_langfuse_logger.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_baggage.py b/tests/unit/integrations/otel/test_otel_v2_baggage.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_baggage.py rename to tests/unit/integrations/otel/test_otel_v2_baggage.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/unit/integrations/otel/test_otel_v2_components.py similarity index 99% rename from tests/test_litellm/integrations/otel/test_otel_v2_components.py rename to tests/unit/integrations/otel/test_otel_v2_components.py index 07705e17d9a..fd10210c5ba 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/unit/integrations/otel/test_otel_v2_components.py @@ -42,7 +42,7 @@ from opentelemetry.trace.propagation.tracecontext import ( # noqa: E402 ) import litellm # noqa: E402 -from conftest import TlsSink # noqa: E402 +from tests.unit.integrations.conftest import TlsSink # noqa: E402 from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 from litellm.integrations.otel.plumbing import providers # noqa: E402 from litellm.integrations.otel.model.config import OpenTelemetryV2Config # noqa: E402 diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py rename to tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_destinations.py rename to tests/unit/integrations/otel/test_otel_v2_destinations.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_dynamic.py b/tests/unit/integrations/otel/test_otel_v2_dynamic.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_dynamic.py rename to tests/unit/integrations/otel/test_otel_v2_dynamic.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py b/tests/unit/integrations/otel/test_otel_v2_emitter.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_emitter.py rename to tests/unit/integrations/otel/test_otel_v2_emitter.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py similarity index 99% rename from tests/test_litellm/integrations/otel/test_otel_v2_logger.py rename to tests/unit/integrations/otel/test_otel_v2_logger.py index 00c1343f72e..d478c670e58 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -2945,14 +2945,6 @@ def test_success_without_pre_call_emits_deferred_span(): assert spans[0].end_time == 101_500_000_000 -def test_no_carrier_and_no_payload_is_noop(): - logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event({"litellm_params": {}}, None, None, None) - ) - assert exporter.get_finished_spans() == () - - def test_second_close_for_same_call_does_not_duplicate_span(): """Success and failure can both fire on one logging object for the same call id. The first close pops the carrier and finishes the boundary span; the diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/unit/integrations/otel/test_otel_v2_metrics.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_metrics.py rename to tests/unit/integrations/otel/test_otel_v2_metrics.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_mount.py b/tests/unit/integrations/otel/test_otel_v2_mount.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_mount.py rename to tests/unit/integrations/otel/test_otel_v2_mount.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_multibackend.py b/tests/unit/integrations/otel/test_otel_v2_multibackend.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_multibackend.py rename to tests/unit/integrations/otel/test_otel_v2_multibackend.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_presets.py b/tests/unit/integrations/otel/test_otel_v2_presets.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_presets.py rename to tests/unit/integrations/otel/test_otel_v2_presets.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py rename to tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_vendor_mappers.py b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_vendor_mappers.py rename to tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py diff --git a/tests/test_litellm/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_runtime.py rename to tests/unit/integrations/otel/test_runtime.py diff --git a/tests/test_litellm/integrations/rubrik_test_helpers.py b/tests/unit/integrations/rubrik_test_helpers.py similarity index 100% rename from tests/test_litellm/integrations/rubrik_test_helpers.py rename to tests/unit/integrations/rubrik_test_helpers.py diff --git a/tests/test_litellm/integrations/test_agentops.py b/tests/unit/integrations/test_agentops.py similarity index 100% rename from tests/test_litellm/integrations/test_agentops.py rename to tests/unit/integrations/test_agentops.py diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py similarity index 100% rename from tests/test_litellm/integrations/test_anthropic_cache_control_hook.py rename to tests/unit/integrations/test_anthropic_cache_control_hook.py diff --git a/tests/test_litellm/integrations/test_athina.py b/tests/unit/integrations/test_athina.py similarity index 100% rename from tests/test_litellm/integrations/test_athina.py rename to tests/unit/integrations/test_athina.py diff --git a/tests/test_litellm/integrations/test_azure_sentinel.py b/tests/unit/integrations/test_azure_sentinel.py similarity index 100% rename from tests/test_litellm/integrations/test_azure_sentinel.py rename to tests/unit/integrations/test_azure_sentinel.py diff --git a/tests/test_litellm/integrations/test_braintrust_logging.py b/tests/unit/integrations/test_braintrust_logging.py similarity index 100% rename from tests/test_litellm/integrations/test_braintrust_logging.py rename to tests/unit/integrations/test_braintrust_logging.py diff --git a/tests/test_litellm/integrations/test_braintrust_span_name.py b/tests/unit/integrations/test_braintrust_span_name.py similarity index 100% rename from tests/test_litellm/integrations/test_braintrust_span_name.py rename to tests/unit/integrations/test_braintrust_span_name.py diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_guardrail.py rename to tests/unit/integrations/test_custom_guardrail.py diff --git a/tests/test_litellm/integrations/test_custom_guardrail_recursion.py b/tests/unit/integrations/test_custom_guardrail_recursion.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_guardrail_recursion.py rename to tests/unit/integrations/test_custom_guardrail_recursion.py diff --git a/tests/test_litellm/integrations/test_custom_prompt_management.py b/tests/unit/integrations/test_custom_prompt_management.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_prompt_management.py rename to tests/unit/integrations/test_custom_prompt_management.py diff --git a/tests/test_litellm/integrations/test_deepeval.py b/tests/unit/integrations/test_deepeval.py similarity index 100% rename from tests/test_litellm/integrations/test_deepeval.py rename to tests/unit/integrations/test_deepeval.py diff --git a/tests/test_litellm/integrations/test_galileo.py b/tests/unit/integrations/test_galileo.py similarity index 100% rename from tests/test_litellm/integrations/test_galileo.py rename to tests/unit/integrations/test_galileo.py diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/unit/integrations/test_guardrail_logging_sync.py similarity index 100% rename from tests/test_litellm/integrations/test_guardrail_logging_sync.py rename to tests/unit/integrations/test_guardrail_logging_sync.py diff --git a/tests/test_litellm/integrations/test_helicone.py b/tests/unit/integrations/test_helicone.py similarity index 100% rename from tests/test_litellm/integrations/test_helicone.py rename to tests/unit/integrations/test_helicone.py diff --git a/tests/test_litellm/integrations/test_langfuse.py b/tests/unit/integrations/test_langfuse.py similarity index 100% rename from tests/test_litellm/integrations/test_langfuse.py rename to tests/unit/integrations/test_langfuse.py diff --git a/tests/test_litellm/integrations/test_langfuse_otel.py b/tests/unit/integrations/test_langfuse_otel.py similarity index 100% rename from tests/test_litellm/integrations/test_langfuse_otel.py rename to tests/unit/integrations/test_langfuse_otel.py diff --git a/tests/test_litellm/integrations/test_langsmith_init.py b/tests/unit/integrations/test_langsmith_init.py similarity index 100% rename from tests/test_litellm/integrations/test_langsmith_init.py rename to tests/unit/integrations/test_langsmith_init.py diff --git a/tests/test_litellm/integrations/test_lunary.py b/tests/unit/integrations/test_lunary.py similarity index 100% rename from tests/test_litellm/integrations/test_lunary.py rename to tests/unit/integrations/test_lunary.py diff --git a/tests/test_litellm/integrations/test_mlflow.py b/tests/unit/integrations/test_mlflow.py similarity index 100% rename from tests/test_litellm/integrations/test_mlflow.py rename to tests/unit/integrations/test_mlflow.py diff --git a/tests/test_litellm/integrations/test_openmeter.py b/tests/unit/integrations/test_openmeter.py similarity index 100% rename from tests/test_litellm/integrations/test_openmeter.py rename to tests/unit/integrations/test_openmeter.py diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py similarity index 99% rename from tests/test_litellm/integrations/test_opentelemetry.py rename to tests/unit/integrations/test_opentelemetry.py index 974961f2eb5..52eeec31e71 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -33,7 +33,7 @@ from parameterized import parameterized import requests -from conftest import TlsSink, write_self_signed_cert +from tests.unit.integrations.conftest import TlsSink, write_self_signed_cert import litellm from litellm.integrations import opentelemetry as otel_module from litellm.integrations.opentelemetry import ( @@ -1244,64 +1244,6 @@ class TestOpenTelemetry(unittest.TestCase): time.sleep(self.POLL_INTERVAL) return [] - @patch("litellm.integrations.opentelemetry.datetime") - def test_create_guardrail_span_with_valid_info(self, mock_datetime): - # Setup - otel = OpenTelemetry() - otel.tracer = MagicMock() - mock_span = MagicMock() - otel.tracer.start_span.return_value = mock_span - - # Create guardrail information - guardrail_info = { - "guardrail_name": "test_guardrail", - "guardrail_mode": "input", - "masked_entity_count": {"CREDIT_CARD": 2}, - "guardrail_response": "filtered_content", - "start_time": 1609459200.0, - "end_time": 1609459201.0, - } - - # Create a kwargs dict with standard_logging_object containing guardrail information - kwargs = { - "standard_logging_object": {"guardrail_information": [guardrail_info]} - } - - # Call the method - otel._create_guardrail_span(kwargs=kwargs, context=None) - - # Assertions - otel.tracer.start_span.assert_called_once() - - # print all calls to mock_span.set_attribute - print("Calls to mock_span.set_attribute:") - for call in mock_span.set_attribute.call_args_list: - print(call) - - # Check that the span has the correct attributes set - mock_span.set_attribute.assert_any_call("guardrail_name", "test_guardrail") - mock_span.set_attribute.assert_any_call("guardrail_mode", "input") - mock_span.set_attribute.assert_any_call( - "guardrail_response", safe_dumps("filtered_content") - ) - mock_span.set_attribute.assert_any_call( - "masked_entity_count", safe_dumps({"CREDIT_CARD": 2}) - ) - - # Verify that the span was ended - mock_span.end.assert_called_once() - - def test_create_guardrail_span_with_no_info(self): - # Setup - otel = OpenTelemetry() - otel.tracer = MagicMock() - - # Test with no guardrail information - kwargs = {"standard_logging_object": {}} - otel._create_guardrail_span(kwargs=kwargs, context=None) - - # Verify that start_span was never called - otel.tracer.start_span.assert_not_called() def test_get_tracer_to_use_for_request_with_dynamic_headers(self): """Test that get_tracer_to_use_for_request returns a dynamic tracer when dynamic headers are present.""" @@ -5461,10 +5403,6 @@ class TestOpenTelemetryPreprocessingDuration(unittest.TestCase): ) assert "litellm.preprocessing.duration_ms" not in self._attr(span, exp) - def test_none_span_is_noop(self): - OpenTelemetry().set_preprocessing_duration_attribute( - None, {"first_api_call_start_time": datetime(2026, 1, 1)} - ) def test_non_dict_container_is_noop(self): otel = OpenTelemetry() diff --git a/tests/test_litellm/integrations/test_opentelemetry_dynamic_imports.py b/tests/unit/integrations/test_opentelemetry_dynamic_imports.py similarity index 100% rename from tests/test_litellm/integrations/test_opentelemetry_dynamic_imports.py rename to tests/unit/integrations/test_opentelemetry_dynamic_imports.py diff --git a/tests/test_litellm/integrations/test_opik_utils.py b/tests/unit/integrations/test_opik_utils.py similarity index 100% rename from tests/test_litellm/integrations/test_opik_utils.py rename to tests/unit/integrations/test_opik_utils.py diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/unit/integrations/test_otel_guardrail_violation_spans.py similarity index 100% rename from tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py rename to tests/unit/integrations/test_otel_guardrail_violation_spans.py diff --git a/tests/test_litellm/integrations/test_otel_team_attributes_matrix.py b/tests/unit/integrations/test_otel_team_attributes_matrix.py similarity index 100% rename from tests/test_litellm/integrations/test_otel_team_attributes_matrix.py rename to tests/unit/integrations/test_otel_team_attributes_matrix.py diff --git a/tests/test_litellm/integrations/test_prometheus_api_promql_escape.py b/tests/unit/integrations/test_prometheus_api_promql_escape.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_api_promql_escape.py rename to tests/unit/integrations/test_prometheus_api_promql_escape.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py b/tests/unit/integrations/test_prometheus_budget_metric_guard.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py rename to tests/unit/integrations/test_prometheus_budget_metric_guard.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py b/tests/unit/integrations/test_prometheus_budget_metrics_db_lookups.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py rename to tests/unit/integrations/test_prometheus_budget_metrics_db_lookups.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py b/tests/unit/integrations/test_prometheus_budget_metrics_timeout.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py rename to tests/unit/integrations/test_prometheus_budget_metrics_timeout.py diff --git a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py b/tests/unit/integrations/test_prometheus_cache_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_cache_metrics.py rename to tests/unit/integrations/test_prometheus_cache_metrics.py index aa031bb813b..21f13f0ec5c 100644 --- a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py +++ b/tests/unit/integrations/test_prometheus_cache_metrics.py @@ -1,7 +1,7 @@ """ Unit tests for cache Prometheus metrics. -Run with: uv run pytest tests/test_litellm/integrations/test_prometheus_cache_metrics.py -v +Run with: uv run pytest tests/unit/integrations/test_prometheus_cache_metrics.py -v """ import pytest diff --git a/tests/test_litellm/integrations/test_prometheus_caller_identity.py b/tests/unit/integrations/test_prometheus_caller_identity.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_caller_identity.py rename to tests/unit/integrations/test_prometheus_caller_identity.py diff --git a/tests/test_litellm/integrations/test_prometheus_carried_budget_state.py b/tests/unit/integrations/test_prometheus_carried_budget_state.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_carried_budget_state.py rename to tests/unit/integrations/test_prometheus_carried_budget_state.py diff --git a/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py b/tests/unit/integrations/test_prometheus_client_ip_user_agent.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py rename to tests/unit/integrations/test_prometheus_client_ip_user_agent.py diff --git a/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py b/tests/unit/integrations/test_prometheus_custom_metadata_label_counts.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py rename to tests/unit/integrations/test_prometheus_custom_metadata_label_counts.py diff --git a/tests/test_litellm/integrations/test_prometheus_deployment_state_proxy_rejects.py b/tests/unit/integrations/test_prometheus_deployment_state_proxy_rejects.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_deployment_state_proxy_rejects.py rename to tests/unit/integrations/test_prometheus_deployment_state_proxy_rejects.py diff --git a/tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py b/tests/unit/integrations/test_prometheus_end_user_cardinality.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py rename to tests/unit/integrations/test_prometheus_end_user_cardinality.py diff --git a/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py b/tests/unit/integrations/test_prometheus_input_sequence_length_label.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py rename to tests/unit/integrations/test_prometheus_input_sequence_length_label.py diff --git a/tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py b/tests/unit/integrations/test_prometheus_invalid_key_filtering.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py rename to tests/unit/integrations/test_prometheus_invalid_key_filtering.py diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/unit/integrations/test_prometheus_labels.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_labels.py rename to tests/unit/integrations/test_prometheus_labels.py diff --git a/tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py b/tests/unit/integrations/test_prometheus_mcp_tool_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py rename to tests/unit/integrations/test_prometheus_mcp_tool_metrics.py index 22c36f00ca9..da5a0b35e9d 100644 --- a/tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py +++ b/tests/unit/integrations/test_prometheus_mcp_tool_metrics.py @@ -5,7 +5,7 @@ These metrics expose ``mcp_tool_call_metadata`` in Prometheus so Grafana dashboards can break down MCP usage by server and tool name. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py -v + uv run pytest tests/unit/integrations/test_prometheus_mcp_tool_metrics.py -v """ from typing import get_args diff --git a/tests/test_litellm/integrations/test_prometheus_media_generation_metrics.py b/tests/unit/integrations/test_prometheus_media_generation_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_media_generation_metrics.py rename to tests/unit/integrations/test_prometheus_media_generation_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py b/tests/unit/integrations/test_prometheus_metric_name_consistency.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py rename to tests/unit/integrations/test_prometheus_metric_name_consistency.py diff --git a/tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py b/tests/unit/integrations/test_prometheus_metrics_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py rename to tests/unit/integrations/test_prometheus_metrics_endpoint.py diff --git a/tests/test_litellm/integrations/test_prometheus_missing_metrics.py b/tests/unit/integrations/test_prometheus_missing_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_missing_metrics.py rename to tests/unit/integrations/test_prometheus_missing_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_none_metadata.py b/tests/unit/integrations/test_prometheus_none_metadata.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_none_metadata.py rename to tests/unit/integrations/test_prometheus_none_metadata.py diff --git a/tests/test_litellm/integrations/test_prometheus_overhead_with_guardrails.py b/tests/unit/integrations/test_prometheus_overhead_with_guardrails.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_overhead_with_guardrails.py rename to tests/unit/integrations/test_prometheus_overhead_with_guardrails.py diff --git a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py b/tests/unit/integrations/test_prometheus_queue_guardrail_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py rename to tests/unit/integrations/test_prometheus_queue_guardrail_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py b/tests/unit/integrations/test_prometheus_rate_limit_labels.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py rename to tests/unit/integrations/test_prometheus_rate_limit_labels.py diff --git a/tests/test_litellm/integrations/test_prometheus_remaining_tokens_router_fallback.py b/tests/unit/integrations/test_prometheus_remaining_tokens_router_fallback.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_remaining_tokens_router_fallback.py rename to tests/unit/integrations/test_prometheus_remaining_tokens_router_fallback.py diff --git a/tests/test_litellm/integrations/test_prometheus_requested_model_cardinality.py b/tests/unit/integrations/test_prometheus_requested_model_cardinality.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_requested_model_cardinality.py rename to tests/unit/integrations/test_prometheus_requested_model_cardinality.py diff --git a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py b/tests/unit/integrations/test_prometheus_service_tier_label.py similarity index 98% rename from tests/test_litellm/integrations/test_prometheus_service_tier_label.py rename to tests/unit/integrations/test_prometheus_service_tier_label.py index b2212c4ff41..8b8131b5af2 100644 --- a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py +++ b/tests/unit/integrations/test_prometheus_service_tier_label.py @@ -6,7 +6,7 @@ between the tier a provider served and the tier a caller requested, and the end-to-end emit wiring through async_log_success_event. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_service_tier_label.py -v + uv run pytest tests/unit/integrations/test_prometheus_service_tier_label.py -v """ import datetime diff --git a/tests/test_litellm/integrations/test_prometheus_services.py b/tests/unit/integrations/test_prometheus_services.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_services.py rename to tests/unit/integrations/test_prometheus_services.py diff --git a/tests/test_litellm/integrations/test_prometheus_spend_capture_rate.py b/tests/unit/integrations/test_prometheus_spend_capture_rate.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_spend_capture_rate.py rename to tests/unit/integrations/test_prometheus_spend_capture_rate.py diff --git a/tests/test_litellm/integrations/test_prometheus_spend_logs_metadata.py b/tests/unit/integrations/test_prometheus_spend_logs_metadata.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_spend_logs_metadata.py rename to tests/unit/integrations/test_prometheus_spend_logs_metadata.py diff --git a/tests/test_litellm/integrations/test_prometheus_stream_label.py b/tests/unit/integrations/test_prometheus_stream_label.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_stream_label.py rename to tests/unit/integrations/test_prometheus_stream_label.py diff --git a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py b/tests/unit/integrations/test_prometheus_token_detail_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py rename to tests/unit/integrations/test_prometheus_token_detail_metrics.py index 5e3846d6fa2..03d67316080 100644 --- a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py +++ b/tests/unit/integrations/test_prometheus_token_detail_metrics.py @@ -6,7 +6,7 @@ from the Usage object that providers report. They are sparse — only incremented when the underlying detail is populated and > 0. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py -v + uv run pytest tests/unit/integrations/test_prometheus_token_detail_metrics.py -v """ from typing import get_args diff --git a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py b/tests/unit/integrations/test_prometheus_user_team_metrics.py similarity index 98% rename from tests/test_litellm/integrations/test_prometheus_user_team_metrics.py rename to tests/unit/integrations/test_prometheus_user_team_metrics.py index 0fc91748af2..ab0fb67b52f 100644 --- a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py +++ b/tests/unit/integrations/test_prometheus_user_team_metrics.py @@ -102,27 +102,6 @@ class TestPrometheusUserTeamCountMetrics: f"litellm_teams_count_metric should accept value {value}: {e}" ) - def test_user_count_metric_with_zero(self, prometheus_logger): - """Test that user count metric handles zero users""" - metric = prometheus_logger.litellm_total_users_metric - - # Should handle zero gracefully - try: - metric.set(0) - assert True - except Exception as e: - pytest.fail(f"litellm_total_users_metric should handle zero: {e}") - - def test_team_count_metric_with_zero(self, prometheus_logger): - """Test that team count metric handles zero teams""" - metric = prometheus_logger.litellm_teams_count_metric - - # Should handle zero gracefully - try: - metric.set(0) - assert True - except Exception as e: - pytest.fail(f"litellm_teams_count_metric should handle zero: {e}") def test_metrics_can_be_updated_multiple_times(self, prometheus_logger): """Test that metrics can be updated multiple times (simulating refresh cycle)""" diff --git a/tests/test_litellm/integrations/test_prometheus_zero_cost_metric.py b/tests/unit/integrations/test_prometheus_zero_cost_metric.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_zero_cost_metric.py rename to tests/unit/integrations/test_prometheus_zero_cost_metric.py diff --git a/tests/test_litellm/integrations/test_prompt_manager_ssti.py b/tests/unit/integrations/test_prompt_manager_ssti.py similarity index 100% rename from tests/test_litellm/integrations/test_prompt_manager_ssti.py rename to tests/unit/integrations/test_prompt_manager_ssti.py diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/unit/integrations/test_responses_background_cost.py similarity index 100% rename from tests/test_litellm/integrations/test_responses_background_cost.py rename to tests/unit/integrations/test_responses_background_cost.py diff --git a/tests/test_litellm/integrations/test_rubrik.py b/tests/unit/integrations/test_rubrik.py similarity index 99% rename from tests/test_litellm/integrations/test_rubrik.py rename to tests/unit/integrations/test_rubrik.py index 4a2ee487c65..f3fea292bde 100644 --- a/tests/test_litellm/integrations/test_rubrik.py +++ b/tests/unit/integrations/test_rubrik.py @@ -19,7 +19,7 @@ from litellm.integrations.rubrik import ( ) from litellm.proxy._types import UserAPIKeyAuth -from tests.test_litellm.integrations.rubrik_test_helpers import ( +from tests.unit.integrations.rubrik_test_helpers import ( make_inputs_with_tools, make_tool_call_dict, ) diff --git a/tests/test_litellm/integrations/test_s3.py b/tests/unit/integrations/test_s3.py similarity index 100% rename from tests/test_litellm/integrations/test_s3.py rename to tests/unit/integrations/test_s3.py diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/unit/integrations/test_s3_v2.py similarity index 100% rename from tests/test_litellm/integrations/test_s3_v2.py rename to tests/unit/integrations/test_s3_v2.py diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/unit/integrations/test_shadow_eval_logger.py similarity index 100% rename from tests/test_litellm/integrations/test_shadow_eval_logger.py rename to tests/unit/integrations/test_shadow_eval_logger.py diff --git a/tests/test_litellm/integrations/test_weave_otel.py b/tests/unit/integrations/test_weave_otel.py similarity index 100% rename from tests/test_litellm/integrations/test_weave_otel.py rename to tests/unit/integrations/test_weave_otel.py diff --git a/tests/unit/integrations/websearch_interception/__init__.py b/tests/unit/integrations/websearch_interception/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py b/tests/unit/integrations/websearch_interception/test_websearch_agentic_loop_cap.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py rename to tests/unit/integrations/websearch_interception/test_websearch_agentic_loop_cap.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_chat_completion.py b/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_chat_completion.py rename to tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py rename to tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_thinking.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py rename to tests/unit/integrations/websearch_interception/test_websearch_interception_thinking.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_native_blocks.py b/tests/unit/integrations/websearch_interception/test_websearch_native_blocks.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_native_blocks.py rename to tests/unit/integrations/websearch_interception/test_websearch_native_blocks.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_responses.py b/tests/unit/integrations/websearch_interception/test_websearch_responses.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_responses.py rename to tests/unit/integrations/websearch_interception/test_websearch_responses.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_rich_query_shape.py b/tests/unit/integrations/websearch_interception/test_websearch_rich_query_shape.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_rich_query_shape.py rename to tests/unit/integrations/websearch_interception/test_websearch_rich_query_shape.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py rename to tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_streaming_wrap.py b/tests/unit/integrations/websearch_interception/test_websearch_streaming_wrap.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_streaming_wrap.py rename to tests/unit/integrations/websearch_interception/test_websearch_streaming_wrap.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py b/tests/unit/integrations/websearch_interception/test_websearch_thinking_constraint.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py rename to tests/unit/integrations/websearch_interception/test_websearch_thinking_constraint.py diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py index 6438525706a..1a592aa1c9a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py @@ -12,7 +12,7 @@ import litellm.experimental_mcp_client.client as mcp_client_module from litellm.experimental_mcp_client.client import MCPClient from litellm.types.mcp import MCPAuth, MCPTransport from mcp.types import CallToolResult as MCPCallToolResult -from mcp.types import ListToolsResult, PaginatedRequestParams +from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities from mcp.types import Tool as MCPTool @@ -128,6 +128,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) client = MCPClient( @@ -163,6 +168,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_tools = [ @@ -204,6 +214,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) first_page_tools = [ @@ -245,6 +260,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_session_instance.list_tools.side_effect = [ @@ -277,6 +297,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_result = MCPCallToolResult(content=[]) @@ -289,7 +314,7 @@ class TestMCPClientUnitTests: assert result == mock_result mock_session_instance.initialize.assert_called_once() mock_session_instance.call_tool.assert_called_once_with( - name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY + name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY, allow_input_required=False ) 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 dfd250338ff..26674373b4e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock(): ) with ( - patch("litellm.proxy._experimental.mcp_server.server.server.run", run), + patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -511,7 +511,7 @@ async def test_sse_mcp_handler_mock(): # Call the handler await handle_sse_mcp(mock_scope, mock_receive, mock_send) - assert run.await_args.args[:2] == (read_stream, write_stream) + assert run.await_args.args[1:3] == (read_stream, write_stream) assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse" diff --git a/tests/unit/secret_managers/__init__.py b/tests/unit/secret_managers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/secret_managers/hashicorp_vault_parity.json b/tests/unit/secret_managers/hashicorp_vault_parity.json similarity index 100% rename from tests/test_litellm/secret_managers/hashicorp_vault_parity.json rename to tests/unit/secret_managers/hashicorp_vault_parity.json diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py b/tests/unit/secret_managers/test_aws_secret_manager_replication.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py rename to tests/unit/secret_managers/test_aws_secret_manager_replication.py diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py b/tests/unit/secret_managers/test_aws_secret_manager_rotation.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py rename to tests/unit/secret_managers/test_aws_secret_manager_rotation.py diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py b/tests/unit/secret_managers/test_aws_secret_manager_v2.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py rename to tests/unit/secret_managers/test_aws_secret_manager_v2.py diff --git a/tests/test_litellm/secret_managers/test_base_secret_manager.py b/tests/unit/secret_managers/test_base_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_base_secret_manager.py rename to tests/unit/secret_managers/test_base_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_custom_secret_manager.py b/tests/unit/secret_managers/test_custom_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_custom_secret_manager.py rename to tests/unit/secret_managers/test_custom_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_cyberark_secret_manager.py b/tests/unit/secret_managers/test_cyberark_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_cyberark_secret_manager.py rename to tests/unit/secret_managers/test_cyberark_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/unit/secret_managers/test_get_azure_ad_token_provider.py similarity index 100% rename from tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py rename to tests/unit/secret_managers/test_get_azure_ad_token_provider.py diff --git a/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py b/tests/unit/secret_managers/test_hashicorp_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py rename to tests/unit/secret_managers/test_hashicorp_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_secret_manager_handler.py b/tests/unit/secret_managers/test_secret_manager_handler.py similarity index 100% rename from tests/test_litellm/secret_managers/test_secret_manager_handler.py rename to tests/unit/secret_managers/test_secret_manager_handler.py diff --git a/tests/test_litellm/secret_managers/test_secret_managers_main.py b/tests/unit/secret_managers/test_secret_managers_main.py similarity index 100% rename from tests/test_litellm/secret_managers/test_secret_managers_main.py rename to tests/unit/secret_managers/test_secret_managers_main.py diff --git a/tests/unit/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index e464402c9d8..b528a75d0df 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -38,6 +38,7 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl "PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}", "GITHUB_OUTPUT": str(tmp_path / "github_output"), "TEST_PATH": test_path, + "UNIT_FLAG": "", "WORKERS": workers, "UNIT_FLAG": "", }, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 95975b68ac2..3c82d4f4be9 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28350,6 +28350,11 @@ export interface components { * @description Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted. */ maximum_spend_logs_retention_period?: string | null; + /** + * Mcp Advertised Versions + * @description MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. Modern protocol serving and Apps/Tasks remain disabled. + */ + mcp_advertised_versions?: ("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25")[] | null; /** * Mcp Allowed Clients * @description MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.