mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: close registered A2A review gaps
This commit is contained in:
parent
07cea813c6
commit
32cb4c076b
8 changed files with 200 additions and 33 deletions
|
|
@ -59,7 +59,12 @@ class A2ACompletionBridgeHandler:
|
|||
message: Final = params.get("message", {})
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
supplied_messages: Final = params.get("messages")
|
||||
openai_messages: Final = (
|
||||
supplied_messages
|
||||
if isinstance(supplied_messages, list)
|
||||
else A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
|
|
@ -149,10 +154,12 @@ class A2ACompletionBridgeHandler:
|
|||
if a2a_provider_config is not None:
|
||||
verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider)
|
||||
|
||||
provider_params: Final = {key: value for key, value in params.items() if key != "messages"}
|
||||
return await a2a_provider_config.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
params=provider_params,
|
||||
api_base=api_base,
|
||||
timeout=litellm_params.get("timeout") or 60.0,
|
||||
litellm_params=litellm_params,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
|
@ -218,10 +225,12 @@ class A2ACompletionBridgeHandler:
|
|||
if a2a_provider_config is not None:
|
||||
verbose_logger.info("A2A: Using provider config for %s (streaming)", custom_llm_provider)
|
||||
|
||||
provider_params: Final = {key: value for key, value in params.items() if key != "messages"}
|
||||
async for chunk in a2a_provider_config.handle_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
params=provider_params,
|
||||
api_base=api_base,
|
||||
timeout=litellm_params.get("timeout") or 60.0,
|
||||
litellm_params=litellm_params,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
):
|
||||
|
|
@ -261,6 +270,7 @@ class A2ACompletionBridgeHandler:
|
|||
|
||||
# 3. Accumulate content and emit artifact update
|
||||
accumulated_text = ""
|
||||
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
|
||||
chunk_count = 0
|
||||
async for chunk in response:
|
||||
chunk_count += 1
|
||||
|
|
@ -271,6 +281,9 @@ class A2ACompletionBridgeHandler:
|
|||
choice = chunk.choices[0]
|
||||
if hasattr(choice, "delta") and choice.delta:
|
||||
content = choice.delta.content or ""
|
||||
tool_calls = getattr(choice.delta, "tool_calls", None)
|
||||
if isinstance(tool_calls, (list, tuple)):
|
||||
accumulated_tool_calls.extend(tool_calls)
|
||||
|
||||
if content:
|
||||
accumulated_text += content
|
||||
|
|
@ -289,6 +302,8 @@ class A2ACompletionBridgeHandler:
|
|||
state="completed",
|
||||
final=True,
|
||||
)
|
||||
if accumulated_tool_calls:
|
||||
completed_event["result"]["tool_calls"] = accumulated_tool_calls
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
|
|
|
|||
|
|
@ -202,12 +202,16 @@ class A2ACompletionBridgeTransformation:
|
|||
if finish_reason:
|
||||
a2a_message["finish_reason"] = finish_reason
|
||||
|
||||
usage: Final = getattr(response, "usage", None)
|
||||
|
||||
# Build A2A response
|
||||
a2a_response: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": a2a_message,
|
||||
}
|
||||
if usage is not None:
|
||||
a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage
|
||||
|
||||
verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(content))
|
||||
|
||||
|
|
|
|||
|
|
@ -71,15 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
|
||||
# Determine finish reason
|
||||
finish_reason: Final = self._get_finish_reason(chunk)
|
||||
tool_calls: Final = self._get_tool_calls(chunk)
|
||||
|
||||
# Return generic streaming chunk
|
||||
return GenericStreamingChunk(
|
||||
text=text,
|
||||
is_finished=bool(finish_reason),
|
||||
finish_reason=finish_reason or "",
|
||||
is_finished=bool(finish_reason or tool_calls),
|
||||
finish_reason=finish_reason or ("tool_calls" if tool_calls else ""),
|
||||
usage=None,
|
||||
index=0,
|
||||
tool_use=None,
|
||||
tool_use=tool_calls,
|
||||
)
|
||||
except Exception:
|
||||
# Return empty chunk on parse error
|
||||
|
|
@ -92,9 +93,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
tool_use=None,
|
||||
)
|
||||
|
||||
def _handle_string_chunk(
|
||||
self, str_line: str | dict
|
||||
) -> GenericStreamingChunk | ModelResponseStream:
|
||||
def _handle_string_chunk(self, str_line: str | dict) -> GenericStreamingChunk | ModelResponseStream:
|
||||
if isinstance(str_line, dict):
|
||||
return self.chunk_parser(chunk=str_line)
|
||||
return super()._handle_string_chunk(str_line=str_line)
|
||||
|
|
@ -118,3 +117,15 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
return "stop"
|
||||
|
||||
return None
|
||||
|
||||
def _get_tool_calls(self, chunk: dict) -> list[dict] | None:
|
||||
result: Final = chunk.get("result", {})
|
||||
if not isinstance(result, dict):
|
||||
return None
|
||||
tool_calls = result.get("tool_calls")
|
||||
if isinstance(tool_calls, list):
|
||||
return tool_calls
|
||||
message = result.get("message")
|
||||
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list):
|
||||
return message["tool_calls"]
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -80,9 +80,7 @@ def _get_agent_request_headers(data: Mapping[str, object]) -> dict[str, str]:
|
|||
metadata = data.get("litellm_metadata")
|
||||
raw_headers = metadata.get("headers") if isinstance(metadata, Mapping) else None
|
||||
return (
|
||||
{str(key).lower(): str(value) for key, value in raw_headers.items()}
|
||||
if isinstance(raw_headers, Mapping)
|
||||
else {}
|
||||
{str(key).lower(): str(value) for key, value in raw_headers.items()} if isinstance(raw_headers, Mapping) else {}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -95,10 +93,12 @@ class _A2AMessage(TypedDict):
|
|||
role: ReadOnly[str]
|
||||
parts: ReadOnly[tuple[_A2ATextPart, ...]]
|
||||
messageId: ReadOnly[str]
|
||||
contextId: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _A2AParams(TypedDict):
|
||||
message: ReadOnly[_A2AMessage]
|
||||
messages: ReadOnly[list[AllMessageValues]]
|
||||
|
||||
|
||||
async def _route_registered_provider(
|
||||
|
|
@ -120,20 +120,27 @@ async def _route_registered_provider(
|
|||
messages: Final = _MESSAGES_ADAPTER.validate_python(raw_messages)
|
||||
stream: Final = data.get("stream") is True
|
||||
request_id: Final = str(uuid4())
|
||||
raw_session_id: Final = data.get("litellm_session_id")
|
||||
metadata: Final = data.get("metadata")
|
||||
session_id: Final = (
|
||||
raw_session_id
|
||||
if isinstance(raw_session_id, str)
|
||||
else metadata.get("session_id")
|
||||
if isinstance(metadata, Mapping) and isinstance(metadata.get("session_id"), str)
|
||||
else None
|
||||
)
|
||||
params: Final[_A2AParams] = {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": ({"kind": "text", "text": convert_messages_to_prompt(messages)},),
|
||||
"messageId": str(uuid4()),
|
||||
}
|
||||
"contextId": session_id,
|
||||
},
|
||||
"messages": messages,
|
||||
}
|
||||
provider_params: Final = {
|
||||
**_OBJECT_DICT_ADAPTER.validate_python(litellm_params),
|
||||
**{
|
||||
key: data[key]
|
||||
for key in _FORWARDED_REQUEST_PARAMS
|
||||
if key in data and data[key] is not None
|
||||
},
|
||||
**{key: data[key] for key in _FORWARDED_REQUEST_PARAMS if key in data and data[key] is not None},
|
||||
}
|
||||
bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python(params)
|
||||
configured_headers: Final = litellm_params.get("extra_headers") or litellm_params.get("headers")
|
||||
|
|
@ -161,6 +168,7 @@ async def _route_registered_provider(
|
|||
logging_obj.litellm_params.update(pricing_params)
|
||||
logging_obj.model_call_details["litellm_params"].update(pricing_params)
|
||||
logging_obj.custom_pricing = True
|
||||
provider_params["no-log"] = True
|
||||
|
||||
if stream:
|
||||
streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming(
|
||||
|
|
@ -226,9 +234,12 @@ async def _route_registered_provider(
|
|||
)
|
||||
],
|
||||
)
|
||||
usage: Final = response.get("usage")
|
||||
raw_usage: Final = response.get("usage")
|
||||
usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage
|
||||
if usage is not None:
|
||||
setattr(model_response, "usage", usage)
|
||||
if isinstance(logging_obj, Logging):
|
||||
logging_obj.model_call_details["usage"] = usage
|
||||
|
||||
if isinstance(logging_obj, Logging):
|
||||
|
||||
|
|
@ -253,9 +264,7 @@ def _merge_agent_guardrails(
|
|||
if not agent_guardrails:
|
||||
return data
|
||||
|
||||
configured_guardrails: list[object] = (
|
||||
agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails]
|
||||
)
|
||||
configured_guardrails: list[object] = agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails]
|
||||
metadata_key: Final = "litellm_metadata" if "litellm_metadata" in data else "metadata"
|
||||
metadata = data.get(metadata_key)
|
||||
metadata_guardrails = metadata.get("guardrails") if isinstance(metadata, dict) else None
|
||||
|
|
@ -296,6 +305,49 @@ async def merge_a2a_agent_guardrails_before_hooks(data: Mapping[str, object]) ->
|
|||
return _merge_agent_guardrails(data, agent.litellm_params.get("guardrails"))
|
||||
|
||||
|
||||
async def authorize_a2a_agent_before_hooks(
|
||||
data: Mapping[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth | None,
|
||||
) -> Mapping[str, object]:
|
||||
model_name: Final = data.get("model")
|
||||
if not isinstance(model_name, str) or not model_name.startswith("a2a/"):
|
||||
return data
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
|
||||
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
|
||||
|
||||
agent = await get_agent_with_read_through(model_name[4:])
|
||||
if agent is None:
|
||||
return data
|
||||
|
||||
is_admin: Final = user_api_key_dict is not None and (
|
||||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
)
|
||||
if not is_admin:
|
||||
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
|
||||
agent_id=agent.agent_id,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
if not is_allowed:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"Agent '{agent.agent_name}' is not allowed for your key/team. Contact proxy admin for access.",
|
||||
)
|
||||
|
||||
if (agent.litellm_params or {}).get("require_trace_id_on_calls_to_agent"):
|
||||
_enforce_inbound_trace_id(data, agent.agent_id)
|
||||
|
||||
if isinstance(data, dict):
|
||||
data["agent_id"] = agent.agent_id
|
||||
metadata = data.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
metadata = {}
|
||||
data["metadata"] = metadata
|
||||
metadata["agent_id"] = agent.agent_id
|
||||
return data
|
||||
|
||||
|
||||
def _get_agent_dynamic_headers(
|
||||
data: Mapping[str, object],
|
||||
agent_id: str,
|
||||
|
|
@ -409,6 +461,16 @@ async def route_a2a_agent_request(
|
|||
registered_params_value.get("custom_llm_provider") if registered_params_value else None
|
||||
)
|
||||
registered_provider: Final = registered_provider_value if isinstance(registered_provider_value, str) else None
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import A2A_USER_API_KEY_HASH_PARAM
|
||||
|
||||
registered_params_for_route: Final[Mapping[str, object]] = (
|
||||
{
|
||||
**registered_params_value,
|
||||
A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key,
|
||||
}
|
||||
if registered_params_value and user_api_key_dict is not None and user_api_key_dict.api_key
|
||||
else registered_params_value or {}
|
||||
)
|
||||
configured_api_base: Final = registered_params_value.get("api_base") if registered_params_value else None
|
||||
api_base: Final = configured_api_base if isinstance(configured_api_base, str) else agent_url
|
||||
registered_model: Final = registered_params_value.get("model") if registered_params_value else None
|
||||
|
|
@ -454,7 +516,7 @@ async def route_a2a_agent_request(
|
|||
data=routed_data,
|
||||
model_name=model_name,
|
||||
api_base=api_base,
|
||||
litellm_params=registered_params_value,
|
||||
litellm_params=registered_params_for_route,
|
||||
static_headers=registered_static_headers,
|
||||
dynamic_headers=registered_dynamic_headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -37,9 +37,7 @@ async def append_agents_to_model_group(
|
|||
if agent is not None:
|
||||
agent_params: Final = agent.litellm_params
|
||||
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
|
||||
custom_llm_provider: Final = (
|
||||
provider_value if isinstance(provider_value, str) else "a2a"
|
||||
)
|
||||
custom_llm_provider: Final = provider_value if isinstance(provider_value, str) else "a2a"
|
||||
model_groups.append(
|
||||
ModelGroupInfoProxy(
|
||||
model_group=f"a2a/{agent.agent_name}",
|
||||
|
|
@ -79,9 +77,7 @@ async def append_agents_to_model_info(
|
|||
if agent is not None:
|
||||
agent_params: Final = agent.litellm_params
|
||||
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
|
||||
custom_llm_provider: Final = (
|
||||
provider_value if isinstance(provider_value, str) else "a2a"
|
||||
)
|
||||
custom_llm_provider: Final = provider_value if isinstance(provider_value, str) else "a2a"
|
||||
models.append(
|
||||
{
|
||||
"model_name": f"a2a/{agent.agent_name}",
|
||||
|
|
|
|||
|
|
@ -1836,6 +1836,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call
|
||||
## IMPORTANT Note: - initialize this before running pre-call checks. Ensures we log rejected requests to langfuse.
|
||||
from litellm.proxy.agent_endpoints.a2a_routing import (
|
||||
authorize_a2a_agent_before_hooks,
|
||||
merge_a2a_agent_guardrails_before_hooks,
|
||||
)
|
||||
|
||||
self.data = await authorize_a2a_agent_before_hooks(
|
||||
data=self.data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
logging_obj, self.data = litellm.utils.function_setup(
|
||||
original_function=route_type,
|
||||
rules_obj=litellm.utils.Rules(),
|
||||
|
|
@ -1845,10 +1855,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
self.data["litellm_logging_obj"] = logging_obj
|
||||
|
||||
from litellm.proxy.agent_endpoints.a2a_routing import (
|
||||
merge_a2a_agent_guardrails_before_hooks,
|
||||
)
|
||||
|
||||
self.data = await merge_a2a_agent_guardrails_before_hooks(self.data)
|
||||
|
||||
# Merge model-level guardrails before pre_call_hook so DB/UI-configured
|
||||
|
|
|
|||
|
|
@ -27,6 +27,27 @@ async def test_async_iterator_accepts_decoded_a2a_events():
|
|||
assert chunk["text"] == "Hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_iterator_preserves_tool_calls():
|
||||
tool_calls = [
|
||||
{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": {"name": "lookup", "arguments": "{}"},
|
||||
}
|
||||
]
|
||||
|
||||
async def _events():
|
||||
yield {"jsonrpc": "2.0", "result": {"tool_calls": tool_calls}}
|
||||
|
||||
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
|
||||
|
||||
chunk = await iterator.__aiter__().__anext__()
|
||||
|
||||
assert chunk["tool_use"] == tool_calls
|
||||
assert chunk["finish_reason"] == "tool_calls"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_iterator_propagates_jsonrpc_errors():
|
||||
async def _events():
|
||||
|
|
|
|||
|
|
@ -290,6 +290,58 @@ async def test_route_a2a_registered_provider_preserves_identity_headers():
|
|||
assert "x-litellm-user-id" not in {key.lower() for key in headers if key != "X-LiteLLM-User-Id"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_a2a_registered_provider_preserves_messages_and_session():
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import A2A_USER_API_KEY_HASH_PARAM
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
agent = AgentResponse(
|
||||
agent_id="test-agent-id",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://agent.example.com"},
|
||||
litellm_params={"custom_llm_provider": "langflow", "model": "flow"},
|
||||
)
|
||||
data = {
|
||||
"model": "a2a/test-agent",
|
||||
"messages": [
|
||||
{"role": "system", "content": "Be concise"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
"litellm_session_id": "session-1",
|
||||
}
|
||||
bridge_response = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "request-id",
|
||||
"result": {"kind": "message", "parts": [{"kind": "text", "text": "Hello back"}]},
|
||||
}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
|
||||
AsyncMock(return_value=agent),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
patch(
|
||||
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
|
||||
AsyncMock(return_value=bridge_response),
|
||||
) as bridge,
|
||||
):
|
||||
call = await route_a2a_agent_request(
|
||||
data,
|
||||
"acompletion",
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
)
|
||||
await call
|
||||
|
||||
bridge_kwargs = bridge.await_args.kwargs
|
||||
assert bridge_kwargs["params"]["messages"] == data["messages"]
|
||||
assert bridge_kwargs["params"]["message"]["contextId"] == "session-1"
|
||||
assert bridge_kwargs["litellm_params"][A2A_USER_API_KEY_HASH_PARAM] == "hashed-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_a2a_requires_inbound_trace_id():
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue