fix(lint): fix all ruff PLR0915 and F401 errors on main

- litellm/a2a_protocol/main.py: extract _create_streaming_a2a_client
  helper to reduce asend_message_streaming statement count (54→49)
- litellm/proxy/agent_endpoints/a2a_endpoints.py: extract
  _build_agent_request_headers helper to reduce invoke_agent_a2a
  statement count (61→45)
- litellm/proxy/_experimental/mcp_server/server.py: extract
  _format_mcp_auth_header helper to reduce execute_mcp_tool
  statement count (53→48)
- litellm/proxy/management_endpoints/mcp_management_endpoints.py:
  remove unused get_user_credential import (F401)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Harshit28j 2026-03-06 21:17:55 +05:30
parent 8b0375f99c
commit 4514c1a16d
4 changed files with 65 additions and 47 deletions

View file

@ -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,

View file

@ -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)

View file

@ -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":

View file

@ -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,
)