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

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 

7 

8from litellm.proxy._types import UserAPIKeyAuth 

9from litellm.types.mcp_server.mcp_server_manager import MCPServer 

10 

11 

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 

17 

18 

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 

29 

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 ) 

50 

51 @property 

52 def user_api_key_auth(self) -> UserAPIKeyAuth | None: 

53 return copy_caller(self._caller) 

54 

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 ) 

77 

78 

79class ProgressCallback(Protocol): 

80 async def __call__(self, progress: float, total: float | None, /) -> None: ... 80 ↛ exitline 80 didn't return from function '__call__' because

81 

82 

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]