mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: forward A2A headers and errors
This commit is contained in:
parent
02c5c74b37
commit
41bec1c2d0
4 changed files with 80 additions and 3 deletions
|
|
@ -7,7 +7,7 @@ from typing import Final
|
|||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
||||
|
||||
from ..common_utils import extract_text_from_a2a_response
|
||||
from ..common_utils import A2AError, extract_text_from_a2a_response
|
||||
|
||||
|
||||
class A2AModelResponseIterator(BaseModelResponseIterator):
|
||||
|
|
@ -56,6 +56,15 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
}
|
||||
}
|
||||
"""
|
||||
if "error" in chunk:
|
||||
error_value: Final = chunk["error"]
|
||||
error_message: Final = (
|
||||
error_value.get("message")
|
||||
if isinstance(error_value, dict) and isinstance(error_value.get("message"), str)
|
||||
else str(error_value)
|
||||
)
|
||||
raise A2AError(status_code=500, message=f"A2A error: {error_message}")
|
||||
|
||||
try:
|
||||
# Extract text from A2A response
|
||||
text: Final = extract_text_from_a2a_response(chunk)
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ async def _route_registered_provider(
|
|||
api_base: str | None,
|
||||
litellm_params: Mapping[str, object],
|
||||
static_headers: Mapping[str, str] | None,
|
||||
dynamic_headers: Mapping[str, str] | None = None,
|
||||
) -> ModelResponse | CustomStreamWrapper:
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
|
|
@ -122,7 +123,10 @@ async def _route_registered_provider(
|
|||
_HEADERS_ADAPTER.validate_python(configured_headers) if isinstance(configured_headers, dict) else None
|
||||
)
|
||||
agent_extra_headers: Final = merge_agent_headers(
|
||||
dynamic_headers=configured_headers_dict,
|
||||
dynamic_headers=merge_agent_headers(
|
||||
dynamic_headers=dynamic_headers,
|
||||
static_headers=configured_headers_dict,
|
||||
),
|
||||
static_headers=static_headers,
|
||||
)
|
||||
if agent_extra_headers:
|
||||
|
|
@ -233,6 +237,40 @@ def _merge_agent_guardrails(
|
|||
return merged_data
|
||||
|
||||
|
||||
def _get_agent_dynamic_headers(
|
||||
data: Mapping[str, object],
|
||||
agent_id: str,
|
||||
agent_name: str,
|
||||
extra_headers: list[str] | None,
|
||||
) -> dict[str, str]:
|
||||
proxy_request: Final = data.get("proxy_server_request")
|
||||
raw_headers: object = proxy_request.get("headers") if isinstance(proxy_request, Mapping) else None
|
||||
if not isinstance(raw_headers, Mapping):
|
||||
metadata: Final = data.get("metadata")
|
||||
raw_headers = metadata.get("headers") if isinstance(metadata, Mapping) else None
|
||||
normalized_headers: Final = (
|
||||
{str(key).lower(): str(value) for key, value in raw_headers.items()}
|
||||
if isinstance(raw_headers, Mapping)
|
||||
else {}
|
||||
)
|
||||
|
||||
dynamic_headers: dict[str, str] = {}
|
||||
for header_name in extra_headers or []:
|
||||
header_name_str: Final = str(header_name)
|
||||
value: Final = normalized_headers.get(header_name_str.lower())
|
||||
if value is not None:
|
||||
dynamic_headers[header_name_str] = value
|
||||
|
||||
for alias in (agent_id.lower(), agent_name.lower()):
|
||||
prefix: Final = f"x-a2a-{alias}-"
|
||||
for key, value in normalized_headers.items():
|
||||
if key.startswith(prefix):
|
||||
header_name: Final = key[len(prefix) :]
|
||||
if header_name:
|
||||
dynamic_headers[header_name] = value
|
||||
return dynamic_headers
|
||||
|
||||
|
||||
async def route_a2a_agent_request(
|
||||
data: Mapping[str, object],
|
||||
route_type: str,
|
||||
|
|
@ -310,6 +348,12 @@ async def route_a2a_agent_request(
|
|||
data=data,
|
||||
agent_guardrails=registered_params_value.get("guardrails") if registered_params_value else None,
|
||||
)
|
||||
registered_dynamic_headers: Final = _get_agent_dynamic_headers(
|
||||
data=routed_data,
|
||||
agent_id=agent.agent_id,
|
||||
agent_name=agent.agent_name,
|
||||
extra_headers=agent.extra_headers,
|
||||
)
|
||||
if (
|
||||
registered_provider
|
||||
and registered_provider != "a2a"
|
||||
|
|
@ -323,6 +367,7 @@ async def route_a2a_agent_request(
|
|||
api_base=api_base,
|
||||
litellm_params=registered_params_value,
|
||||
static_headers=agent.static_headers,
|
||||
dynamic_headers=registered_dynamic_headers,
|
||||
)
|
||||
|
||||
completion_data: Final = MappingProxyType({**routed_data, "api_base": api_base})
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
import pytest
|
||||
|
||||
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
|
||||
from litellm.llms.a2a.common_utils import A2AError
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -24,3 +25,14 @@ async def test_async_iterator_accepts_decoded_a2a_events():
|
|||
chunk = await iterator.__aiter__().__anext__()
|
||||
|
||||
assert chunk["text"] == "Hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_iterator_propagates_jsonrpc_errors():
|
||||
async def _events():
|
||||
yield {"jsonrpc": "2.0", "error": {"code": -32000, "message": "agent failed"}}
|
||||
|
||||
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
|
||||
|
||||
with pytest.raises(A2AError, match="agent failed"):
|
||||
await iterator.__aiter__().__anext__()
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ async def test_route_a2a_model_uses_registered_provider():
|
|||
"guardrails": ["agent-guardrail"],
|
||||
},
|
||||
static_headers={"Authorization": "Bearer static"},
|
||||
extra_headers=["X-Tenant"],
|
||||
)
|
||||
data = {
|
||||
"model": "a2a/test-agent",
|
||||
|
|
@ -92,6 +93,12 @@ async def test_route_a2a_model_uses_registered_provider():
|
|||
"temperature": 0.2,
|
||||
"timeout": 12.0,
|
||||
"tools": [{"type": "function", "function": {"name": "lookup"}}],
|
||||
"proxy_server_request": {
|
||||
"headers": {
|
||||
"x-tenant": "tenant-1",
|
||||
"x-a2a-test-agent-x-run": "run-1",
|
||||
}
|
||||
},
|
||||
}
|
||||
bridge_response = {
|
||||
"jsonrpc": "2.0",
|
||||
|
|
@ -133,7 +140,11 @@ async def test_route_a2a_model_uses_registered_provider():
|
|||
assert bridge_kwargs["litellm_params"]["timeout"] == 12.0
|
||||
assert bridge_kwargs["litellm_params"]["tools"] == data["tools"]
|
||||
assert bridge_kwargs["litellm_params"]["guardrails"] == ["request-guardrail", "agent-guardrail"]
|
||||
assert bridge_kwargs["litellm_params"]["extra_headers"] == {"Authorization": "Bearer static"}
|
||||
assert bridge_kwargs["litellm_params"]["extra_headers"] == {
|
||||
"X-Tenant": "tenant-1",
|
||||
"x-run": "run-1",
|
||||
"Authorization": "Bearer static",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue