Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/contracts.py: 71%
45 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from collections.abc import Mapping
2from copy import deepcopy
3from dataclasses import dataclass, field
4from datetime import datetime
5from types import MappingProxyType
6from typing import Final, Protocol
8from litellm.proxy._types import UserAPIKeyAuth
9from litellm.types.mcp_server.mcp_server_manager import MCPServer
12def copy_caller(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None:
13 if auth is None:
14 return None
15 span: Final = auth.parent_otel_span
16 return deepcopy(auth, {id(span): span} if span is not None else None) # mutable-ok: deepcopy mutates its memo
19@dataclass(frozen=True, slots=True)
20class OperationContext:
21 _caller: UserAPIKeyAuth | None = field(repr=False)
22 mcp_auth_header: str | None = field(default=None, repr=False)
23 mcp_servers: tuple[str, ...] | None = None
24 mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = field(default=None, repr=False)
25 oauth2_headers: Mapping[str, str] | None = field(default=None, repr=False)
26 raw_headers: Mapping[str, str] | None = field(default=None, repr=False)
27 client_ip: str | None = None
28 mcp_proxy_mode: bool = False
30 def __post_init__(self) -> None:
31 object.__setattr__(self, "_caller", copy_caller(self._caller))
32 object.__setattr__(self, "mcp_servers", tuple(self.mcp_servers) if self.mcp_servers is not None else None)
33 object.__setattr__(
34 self,
35 "oauth2_headers",
36 MappingProxyType(dict(self.oauth2_headers)) if self.oauth2_headers is not None else None,
37 )
38 object.__setattr__(
39 self, "raw_headers", MappingProxyType(dict(self.raw_headers)) if self.raw_headers is not None else None
40 )
41 object.__setattr__(
42 self,
43 "mcp_server_auth_headers",
44 MappingProxyType(
45 {key: MappingProxyType(dict(value)) for key, value in self.mcp_server_auth_headers.items()}
46 )
47 if self.mcp_server_auth_headers is not None
48 else None,
49 )
51 @property
52 def user_api_key_auth(self) -> UserAPIKeyAuth | None:
53 return copy_caller(self._caller)
55 def legacy_auth(
56 self,
57 ) -> tuple[
58 UserAPIKeyAuth | None,
59 str | None,
60 list[str] | None,
61 dict[str, dict[str, str]] | None,
62 dict[str, str] | None,
63 dict[str, str] | None,
64 str | None,
65 ]:
66 return (
67 self.user_api_key_auth,
68 self.mcp_auth_header,
69 list(self.mcp_servers) if self.mcp_servers is not None else None, # mutable-ok: legacy policy list input
70 {key: dict(value) for key, value in self.mcp_server_auth_headers.items()}
71 if self.mcp_server_auth_headers is not None
72 else None,
73 dict(self.oauth2_headers) if self.oauth2_headers is not None else None,
74 dict(self.raw_headers) if self.raw_headers is not None else None, # mutable-ok: legacy request header input
75 self.client_ip,
76 )
79class ProgressCallback(Protocol):
80 async def __call__(self, progress: float, total: float | None, /) -> None: ... 80 ↛ exitline 80 didn't return from function '__call__' because
83@dataclass(frozen=True, slots=True)
84class AuthorizedToolCall:
85 name: str
86 arguments: Mapping[str, object]
87 allowed_mcp_servers: tuple[MCPServer, ...]
88 start_time: datetime
89 host_progress_callback: ProgressCallback | None
90 guardrail_context: Mapping[str, object] | None
91 logging_data: Mapping[str, object]