mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'upstream/main' into litellm_forward_reasoning_content
This commit is contained in:
commit
eac8c95cb9
216 changed files with 982 additions and 189 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
4
.github/workflows/test-unit.yml
vendored
4
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
4
Makefile
4
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -118,7 +118,7 @@ pub(super) struct ParityCase {
|
|||
pub(super) fn parity_cases() -> Vec<ParityCase> {
|
||||
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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
| --- | --- |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
|
|
@ -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})
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
# Levo integration tests
|
||||
|
|
@ -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")
|
||||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
0
tests/unit/integrations/bitbucket/__init__.py
Normal file
0
tests/unit/integrations/bitbucket/__init__.py
Normal file
|
|
@ -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 = {
|
||||
0
tests/unit/integrations/cloudzero/__init__.py
Normal file
0
tests/unit/integrations/cloudzero/__init__.py
Normal file
|
|
@ -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")
|
||||
0
tests/unit/integrations/datadog/__init__.py
Normal file
0
tests/unit/integrations/datadog/__init__.py
Normal file
0
tests/unit/integrations/dotprompt/__init__.py
Normal file
0
tests/unit/integrations/dotprompt/__init__.py
Normal file
0
tests/unit/integrations/focus/__init__.py
Normal file
0
tests/unit/integrations/focus/__init__.py
Normal file
0
tests/unit/integrations/gitlab/__init__.py
Normal file
0
tests/unit/integrations/gitlab/__init__.py
Normal file
0
tests/unit/integrations/langfuse/__init__.py
Normal file
0
tests/unit/integrations/langfuse/__init__.py
Normal file
0
tests/unit/integrations/newrelic/__init__.py
Normal file
0
tests/unit/integrations/newrelic/__init__.py
Normal file
0
tests/unit/integrations/open_telemetry/__init__.py
Normal file
0
tests/unit/integrations/open_telemetry/__init__.py
Normal file
0
tests/unit/integrations/open_telemetry/data/__init__.py
Normal file
0
tests/unit/integrations/open_telemetry/data/__init__.py
Normal file
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue