diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 0067af3c7db..45e3bbd30df 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -530,15 +530,11 @@ async def asend_message_streaming( "Either a2a_client or api_base is required for standard A2A flow" ) # Mirror the non-streaming path: always include trace and agent-id headers - streaming_extra_headers: Dict[str, str] = { - "X-LiteLLM-Trace-Id": str(request.id), - } - if agent_id: - streaming_extra_headers["X-LiteLLM-Agent-Id"] = agent_id - if agent_extra_headers: - streaming_extra_headers.update(agent_extra_headers) - a2a_client = await create_a2a_client( - base_url=api_base, extra_headers=streaming_extra_headers + a2a_client = await _create_streaming_a2a_client( + api_base=api_base, + request_id=request.id, + agent_id=agent_id, + agent_extra_headers=agent_extra_headers, ) # Type assertion: a2a_client is guaranteed to be non-None here @@ -614,6 +610,21 @@ async def asend_message_streaming( raise +async def _create_streaming_a2a_client( + api_base: str, + request_id: Any, + agent_id: Optional[str], + agent_extra_headers: Optional[Dict[str, str]], +) -> "A2AClientType": + """Build trace/agent-id headers and create an A2A streaming client.""" + extra_headers: Dict[str, str] = {"X-LiteLLM-Trace-Id": str(request_id)} + if agent_id: + extra_headers["X-LiteLLM-Agent-Id"] = agent_id + if agent_extra_headers: + extra_headers.update(agent_extra_headers) + return await create_a2a_client(base_url=api_base, extra_headers=extra_headers) + + async def create_a2a_client( base_url: str, timeout: float = 60.0, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5c063839304..b2062ee136f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1644,6 +1644,18 @@ if MCP_AVAILABLE: }, ) + def _format_mcp_auth_header( + mcp_auth_header: str, + mcp_server: Optional["MCPServer"], + ) -> str: + """Format the Authorization header value based on the server's auth_type.""" + server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None + if server_auth_type == MCPAuth.api_key: + return f"ApiKey {mcp_auth_header}" + if server_auth_type == MCPAuth.basic: + return f"Basic {mcp_auth_header}" + return f"Bearer {mcp_auth_header}" + async def execute_mcp_tool( name: str, arguments: Dict[str, Any], @@ -1768,15 +1780,11 @@ if MCP_AVAILABLE: # because the tool function has headers baked into its closure. # Pre-format the full Authorization header value using the server's # configured auth_type so the generator doesn't need to know the prefix. - auth_header_value: Optional[str] = None - if mcp_auth_header: - server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None - if server_auth_type == MCPAuth.api_key: - auth_header_value = f"ApiKey {mcp_auth_header}" - elif server_auth_type == MCPAuth.basic: - auth_header_value = f"Basic {mcp_auth_header}" - else: - auth_header_value = f"Bearer {mcp_auth_header}" + auth_header_value: Optional[str] = ( + _format_mcp_auth_header(mcp_auth_header, mcp_server) + if mcp_auth_header + else None + ) _auth_token = _request_auth_header.set(auth_header_value) try: local_content = await _handle_local_mcp_tool(name, arguments) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 344070d17fc..07f185f0188 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -20,6 +20,33 @@ from litellm.types.utils import all_litellm_params router = APIRouter() +def _build_agent_request_headers( + agent: Any, + request: Request, +) -> Optional[Dict[str, str]]: + """Build merged extra headers for forwarding to the backend agent.""" + static_headers: Dict[str, str] = dict(agent.static_headers or {}) + raw_headers = dict(request.headers) + normalized = {k.lower(): v for k, v in raw_headers.items()} + dynamic_headers: Dict[str, str] = {} + if agent.extra_headers: + for header_name in agent.extra_headers: + val = normalized.get(header_name.lower()) + if val is not None: + dynamic_headers[header_name] = val + for alias in (agent.agent_id.lower(), agent.agent_name.lower()): + prefix = f"x-a2a-{alias}-" + for key, val in normalized.items(): + if key.startswith(prefix): + header_name = key[len(prefix):] + if header_name: + dynamic_headers[header_name] = val + return merge_agent_headers( + dynamic_headers=dynamic_headers or None, + static_headers=static_headers or None, + ) + + def _jsonrpc_error( request_id: Optional[str], code: int, @@ -389,34 +416,7 @@ async def invoke_agent_a2a( ) # Build merged headers for the backend agent - static_headers: Dict[str, str] = dict(agent.static_headers or {}) - - raw_headers = dict(request.headers) - normalized = {k.lower(): v for k, v in raw_headers.items()} - - dynamic_headers: Dict[str, str] = {} - - # 1. Admin-configured extra_headers: forward named headers from client request - if agent.extra_headers: - for header_name in agent.extra_headers: - val = normalized.get(header_name.lower()) - if val is not None: - dynamic_headers[header_name] = val - - # 2. Convention-based forwarding: x-a2a-{agent_id_or_name}-{header_name} - # Matches both agent_id (UUID) and agent_name (alias), case-insensitive. - for alias in (agent.agent_id.lower(), agent.agent_name.lower()): - prefix = f"x-a2a-{alias}-" - for key, val in normalized.items(): - if key.startswith(prefix): - header_name = key[len(prefix) :] - if header_name: - dynamic_headers[header_name] = val - - agent_extra_headers = merge_agent_headers( - dynamic_headers=dynamic_headers or None, - static_headers=static_headers or None, - ) + agent_extra_headers = _build_agent_request_headers(agent=agent, request=request) # Route through SDK functions if method == "message/send": diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index b48db72a536..f7a4cec301b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -81,7 +81,6 @@ if MCP_AVAILABLE: delete_user_credential, get_all_mcp_servers_for_user, get_mcp_server, - get_user_credential, store_user_credential, update_mcp_server, )