mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: address A2A review feedback
This commit is contained in:
parent
eba758223d
commit
a5997bb908
10 changed files with 81 additions and 33 deletions
|
|
@ -89,11 +89,12 @@ class A2ACompletionBridgeHandler:
|
|||
"api_base": api_base,
|
||||
"stream": stream,
|
||||
}
|
||||
configured_headers: Final[object] = litellm_params.get("extra_headers") or litellm_params.get("headers")
|
||||
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
|
||||
litellm_params_to_add: Final = {
|
||||
k: v
|
||||
for k, v in litellm_params.items()
|
||||
if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS
|
||||
if k not in ("model", "custom_llm_provider", "extra_headers", "headers") and k not in _AGENT_ONLY_PARAMS
|
||||
}
|
||||
completion_params.update(litellm_params_to_add)
|
||||
# Apply forward metadata AFTER the litellm_params merge so the helper
|
||||
|
|
@ -105,10 +106,10 @@ class A2ACompletionBridgeHandler:
|
|||
params=params,
|
||||
)
|
||||
|
||||
if agent_extra_headers:
|
||||
if agent_extra_headers or configured_headers:
|
||||
completion_params["extra_headers"] = merge_agent_headers(
|
||||
dynamic_headers=agent_extra_headers,
|
||||
static_headers=completion_params.get("extra_headers"),
|
||||
static_headers=configured_headers if isinstance(configured_headers, Mapping) else None,
|
||||
)
|
||||
|
||||
return completion_params
|
||||
|
|
@ -361,6 +362,7 @@ class A2ACompletionBridgeHandler:
|
|||
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text=content,
|
||||
index=choice_index,
|
||||
)
|
||||
yield artifact_event
|
||||
finally:
|
||||
|
|
@ -380,13 +382,10 @@ class A2ACompletionBridgeHandler:
|
|||
completed_event["result"]["finish_reason"] = stream_finish_reason
|
||||
if stream_usage is not None:
|
||||
completed_event["usage"] = stream_usage
|
||||
if choice_delta_fields.get(0):
|
||||
completed_event["result"].update(choice_delta_fields[0])
|
||||
if 0 in choice_logprobs:
|
||||
completed_event["result"]["logprobs"] = choice_logprobs[0]
|
||||
if len(choice_texts) > 1:
|
||||
completed_event["result"]["choices"] = [
|
||||
{
|
||||
choice_payloads: list[dict[str, object]] = []
|
||||
for choice_index in sorted(choice_texts):
|
||||
choice_payload: dict[str, object] = {
|
||||
"index": choice_index,
|
||||
"message": {
|
||||
"kind": "message",
|
||||
|
|
@ -397,7 +396,6 @@ class A2ACompletionBridgeHandler:
|
|||
if choice_tool_calls.get(choice_index)
|
||||
else {}
|
||||
),
|
||||
**choice_delta_fields.get(choice_index, {}),
|
||||
},
|
||||
**(
|
||||
{"finish_reason": choice_finish_reasons[choice_index]}
|
||||
|
|
@ -406,8 +404,29 @@ class A2ACompletionBridgeHandler:
|
|||
),
|
||||
**({"logprobs": choice_logprobs[choice_index]} if choice_index in choice_logprobs else {}),
|
||||
}
|
||||
for choice_index in sorted(choice_texts)
|
||||
]
|
||||
if choice_delta_fields.get(choice_index):
|
||||
choice_payload["delta"] = choice_delta_fields[choice_index]
|
||||
choice_payloads.append(choice_payload)
|
||||
completed_event["result"]["choices"] = choice_payloads
|
||||
else:
|
||||
metadata_indices = sorted(set(choice_delta_fields) | set(choice_logprobs))
|
||||
if metadata_indices:
|
||||
completed_event["result"]["choices"] = [
|
||||
{
|
||||
"index": choice_index,
|
||||
**(
|
||||
{"delta": choice_delta_fields[choice_index]}
|
||||
if choice_delta_fields.get(choice_index)
|
||||
else {}
|
||||
),
|
||||
**(
|
||||
{"logprobs": choice_logprobs[choice_index]}
|
||||
if choice_index in choice_logprobs
|
||||
else {}
|
||||
),
|
||||
}
|
||||
for choice_index in metadata_indices
|
||||
]
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
|
|
|
|||
|
|
@ -262,6 +262,10 @@ class A2ACompletionBridgeTransformation:
|
|||
}
|
||||
if usage is not None:
|
||||
a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage
|
||||
for field in ("system_fingerprint", "service_tier"):
|
||||
value = getattr(response, field, None)
|
||||
if value is not None:
|
||||
a2a_response[field] = value
|
||||
if len(serialized_choices) > 1:
|
||||
a2a_response["choices"] = serialized_choices
|
||||
|
||||
|
|
@ -354,6 +358,7 @@ class A2ACompletionBridgeTransformation:
|
|||
def create_artifact_update_event(
|
||||
ctx: A2AStreamingContext,
|
||||
text: str,
|
||||
index: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Create an artifact update event with content.
|
||||
|
|
@ -362,15 +367,18 @@ class A2ACompletionBridgeTransformation:
|
|||
ctx: Streaming context
|
||||
text: The text content for the artifact
|
||||
"""
|
||||
artifact: Final[dict[str, Any]] = {
|
||||
"artifactId": str(uuid4()),
|
||||
"name": "response",
|
||||
"parts": [{"kind": "text", "text": text}],
|
||||
}
|
||||
if index is not None:
|
||||
artifact["index"] = index
|
||||
return {
|
||||
"id": ctx.request_id,
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"artifact": {
|
||||
"artifactId": str(uuid4()),
|
||||
"name": "response",
|
||||
"parts": [{"kind": "text", "text": text}],
|
||||
},
|
||||
"artifact": artifact,
|
||||
"contextId": ctx.context_id,
|
||||
"kind": "artifact-update",
|
||||
"taskId": ctx.task_id,
|
||||
|
|
|
|||
|
|
@ -71,6 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
try:
|
||||
# Extract text from A2A response
|
||||
result: Final = chunk.get("result", {})
|
||||
chunk_index = 0
|
||||
if isinstance(result, Mapping):
|
||||
artifact = result.get("artifact")
|
||||
if isinstance(artifact, Mapping) and isinstance(artifact.get("index"), int):
|
||||
chunk_index = artifact["index"]
|
||||
choices = result.get("choices")
|
||||
if isinstance(choices, list) and choices and isinstance(choices[0], Mapping):
|
||||
raw_index = choices[0].get("index")
|
||||
if isinstance(raw_index, int):
|
||||
chunk_index = raw_index
|
||||
status: Final = result.get("status", {}) if isinstance(result, Mapping) else {}
|
||||
is_working_status: Final = (
|
||||
isinstance(result, Mapping)
|
||||
|
|
@ -132,7 +142,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
is_finished=bool(finish_reason or tool_calls),
|
||||
finish_reason=finish_reason or ("tool_calls" if tool_calls else ""),
|
||||
usage=usage,
|
||||
index=0,
|
||||
index=chunk_index,
|
||||
tool_use=tool_calls,
|
||||
provider_specific_fields=provider_fields or None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -860,7 +860,7 @@ async def invoke_agent_a2a(
|
|||
_enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None)
|
||||
if _enqueue_fn is not None:
|
||||
logging_obj._enqueue_deferred_logging = None
|
||||
_enqueue_fn()
|
||||
_enqueue_fn(response)
|
||||
|
||||
response_dict: Final[dict[str, Any]] = (
|
||||
response.model_dump(mode="json", exclude_none=True)
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
|
|||
"frequency_penalty",
|
||||
"functions",
|
||||
"function_call",
|
||||
"guided_json",
|
||||
"include_server_side_tool_invocations",
|
||||
"logit_bias",
|
||||
"logprobs",
|
||||
|
|
@ -305,6 +306,10 @@ async def _route_registered_provider(
|
|||
id=str(response.get("id") or request_id),
|
||||
model=model_name,
|
||||
choices=model_choices,
|
||||
system_fingerprint=response.get("system_fingerprint")
|
||||
if isinstance(response.get("system_fingerprint"), str)
|
||||
else None,
|
||||
service_tier=response.get("service_tier") if isinstance(response.get("service_tier"), str) else None,
|
||||
)
|
||||
raw_usage: Final = response.get("usage")
|
||||
usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage
|
||||
|
|
@ -315,10 +320,10 @@ async def _route_registered_provider(
|
|||
|
||||
if isinstance(logging_obj, Logging):
|
||||
|
||||
def _enqueue_logging() -> None:
|
||||
def _enqueue_logging(final_response: ModelResponse | None = None) -> None:
|
||||
asyncio.create_task(
|
||||
logging_obj.dispatch_success_handlers(
|
||||
model_response,
|
||||
final_response if final_response is not None else model_response,
|
||||
cache_hit=False,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
|
|
@ -544,7 +549,7 @@ async def route_a2a_agent_request(
|
|||
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
|
||||
api_base: Final = configured_api_base if isinstance(configured_api_base, str) and configured_api_base else agent_url
|
||||
registered_model: Final = registered_params_value.get("model") if registered_params_value else None
|
||||
cardless_provider: Final = registered_provider == "watsonx_orchestrate" or (
|
||||
registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model
|
||||
|
|
|
|||
|
|
@ -2544,6 +2544,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
|
||||
logging_obj=logging_obj,
|
||||
exception_raised=_exception_raised,
|
||||
response=response,
|
||||
)
|
||||
|
||||
# Streaming cleanup: if an exception occurred AND the deferred
|
||||
|
|
@ -3027,6 +3028,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
def _flush_deferred_async_logging(
|
||||
logging_obj: Any,
|
||||
exception_raised: bool,
|
||||
response: Any | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Fire the deferred async-success closure stored by wrapper_async, then
|
||||
|
|
@ -3057,7 +3059,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if exception_raised:
|
||||
return
|
||||
try:
|
||||
_enqueue_fn()
|
||||
_enqueue_fn(response) if response is not None else _enqueue_fn()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error firing deferred logging: %s", e)
|
||||
|
||||
|
|
|
|||
|
|
@ -1818,11 +1818,11 @@ def client(original_function):
|
|||
if not _is_litellm_internal_call:
|
||||
if getattr(logging_obj, "_defer_async_logging", False):
|
||||
|
||||
def _enqueue_deferred_logging() -> None:
|
||||
def _enqueue_deferred_logging(final_response=None) -> None:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
result=final_response if final_response is not None else result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
|
|
|
|||
|
|
@ -10,10 +10,9 @@ Verifies that:
|
|||
"""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent"
|
||||
SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}"
|
||||
|
|
@ -482,6 +481,7 @@ class TestHandlerIntegration:
|
|||
api_base=None,
|
||||
litellm_params=SAMPLE_LITELLM_PARAMS,
|
||||
agent_extra_headers=None,
|
||||
agent_static_headers=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -324,10 +324,13 @@ async def test_handle_streaming_preserves_non_text_delta_fields():
|
|||
]
|
||||
|
||||
result = events[-1]["result"]
|
||||
assert result["audio"] == {"data": "abc"}
|
||||
assert result["reasoning_content"] == "thinking"
|
||||
assert result["provider_specific_fields"] == {"trace_id": "trace-1"}
|
||||
assert result["logprobs"] == {"content": []}
|
||||
choice_result = result["choices"][0]
|
||||
assert choice_result["delta"] == {
|
||||
"audio": {"data": "abc"},
|
||||
"reasoning_content": "thinking",
|
||||
"provider_specific_fields": {"trace_id": "trace-1"},
|
||||
}
|
||||
assert choice_result["logprobs"] == {"content": []}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ async def test_route_a2a_model_bypasses_router():
|
|||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
|
||||
mock_registry,
|
||||
):
|
||||
result = await route_request(
|
||||
await route_request(
|
||||
data=data,
|
||||
llm_router=mock_router,
|
||||
user_model=None,
|
||||
|
|
@ -593,6 +593,7 @@ async def test_route_a2a_stream_uses_registered_provider():
|
|||
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
|
||||
)
|
||||
logging_obj = Mock(spec=Logging)
|
||||
logging_obj.model_call_details = {}
|
||||
data = {
|
||||
"model": "a2a/test-agent",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
|
|
@ -745,7 +746,7 @@ def _router_without_models():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_replica(monkeypatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
agent_name = "a2a-sibling-replica-agent"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue