mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): account for completion bridge invocation fees
This commit is contained in:
parent
69b8ada9ac
commit
92e13728f5
3 changed files with 81 additions and 13 deletions
|
|
@ -471,13 +471,20 @@ async def asend_message(
|
|||
if custom_llm_provider:
|
||||
if request is None:
|
||||
raise ValueError("request is required for completion bridge")
|
||||
return await _send_message_via_completion_bridge(
|
||||
bridge_response: Final = await _send_message_via_completion_bridge(
|
||||
request=request,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
bridge_prompt_tokens, bridge_completion_tokens, _ = await asyncify(
|
||||
A2ARequestUtils.calculate_usage_from_request_response
|
||||
)(request=request, response_dict=bridge_response.model_dump(mode="json", exclude_none=True))
|
||||
_set_usage_on_logging_obj(kwargs, bridge_prompt_tokens, bridge_completion_tokens)
|
||||
_set_litellm_params_on_logging_obj(kwargs, litellm_params)
|
||||
_set_agent_id_on_logging_obj(kwargs, agent_id)
|
||||
return bridge_response
|
||||
|
||||
# Standard A2A client flow
|
||||
if request is None:
|
||||
|
|
@ -692,12 +699,35 @@ async def asend_message_streaming(
|
|||
request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
)
|
||||
|
||||
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
|
||||
bridge_name: Final = str(litellm_params.get("model") or agent_id or "agent")
|
||||
existing_logging: Final = kwargs.get("litellm_logging_obj")
|
||||
bridge_logging: Final = (
|
||||
existing_logging
|
||||
if isinstance(existing_logging, Logging)
|
||||
else _build_streaming_logging_obj(
|
||||
request=request,
|
||||
agent_name=bridge_name,
|
||||
agent_id=agent_id,
|
||||
litellm_params=litellm_params,
|
||||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
)
|
||||
)
|
||||
bridge_context: Final = {"litellm_logging_obj": bridge_logging}
|
||||
_set_litellm_params_on_logging_obj(bridge_context, litellm_params)
|
||||
_set_agent_id_on_logging_obj(bridge_context, agent_id)
|
||||
bridge_stream: Final = A2ACompletionBridgeHandler.handle_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
async for chunk in A2AStreamingIterator(
|
||||
stream=bridge_stream,
|
||||
request=request,
|
||||
logging_obj=bridge_logging,
|
||||
agent_name=bridge_name,
|
||||
):
|
||||
yield chunk
|
||||
return
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support.
|
|||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from typing import TYPE_CHECKING, Final, Generic, TypeVar
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -18,7 +18,10 @@ if TYPE_CHECKING:
|
|||
from a2a.compat.v0_3.types import SendStreamingMessageRequest, SendStreamingMessageResponse
|
||||
|
||||
|
||||
class A2AStreamingIterator:
|
||||
_StreamChunk = TypeVar("_StreamChunk", bound="SendStreamingMessageResponse | dict[str, object]")
|
||||
|
||||
|
||||
class A2AStreamingIterator(Generic[_StreamChunk]):
|
||||
"""
|
||||
Async iterator for A2A streaming responses with token tracking.
|
||||
|
||||
|
|
@ -27,7 +30,7 @@ class A2AStreamingIterator:
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
stream: AsyncIterator["SendStreamingMessageResponse"],
|
||||
stream: AsyncIterator[_StreamChunk],
|
||||
request: "SendStreamingMessageRequest",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
agent_name: str = "unknown",
|
||||
|
|
@ -39,14 +42,14 @@ class A2AStreamingIterator:
|
|||
self.start_time = datetime.now()
|
||||
|
||||
# Collect chunks for token counting
|
||||
self.chunks: list[SendStreamingMessageResponse] = []
|
||||
self.chunks: list[_StreamChunk] = []
|
||||
self.collected_text_parts: list[str] = []
|
||||
self.final_chunk: SendStreamingMessageResponse | None = None
|
||||
self.final_chunk: _StreamChunk | None = None
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> "SendStreamingMessageResponse":
|
||||
async def __anext__(self) -> _StreamChunk:
|
||||
try:
|
||||
chunk: Final = await self.stream.__anext__()
|
||||
|
||||
|
|
@ -69,20 +72,20 @@ class A2AStreamingIterator:
|
|||
await self._handle_stream_complete()
|
||||
raise
|
||||
|
||||
def _collect_text_from_chunk(self, chunk: "SendStreamingMessageResponse") -> None:
|
||||
def _collect_text_from_chunk(self, chunk: _StreamChunk) -> None:
|
||||
"""Extract text from a streaming chunk and add to collected parts."""
|
||||
try:
|
||||
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
|
||||
text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict)
|
||||
if text:
|
||||
self.collected_text_parts.append(text)
|
||||
except Exception:
|
||||
verbose_logger.debug("Failed to extract text from A2A streaming chunk")
|
||||
|
||||
def _is_completed_chunk(self, chunk: "SendStreamingMessageResponse") -> bool:
|
||||
def _is_completed_chunk(self, chunk: _StreamChunk) -> bool:
|
||||
"""Check if chunk indicates stream completion."""
|
||||
try:
|
||||
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
|
||||
result: Final = chunk_dict.get("result", {})
|
||||
if isinstance(result, dict):
|
||||
status: Final = result.get("status", {})
|
||||
|
|
@ -160,7 +163,11 @@ class A2AStreamingIterator:
|
|||
# Add final chunk result if available
|
||||
if self.final_chunk:
|
||||
try:
|
||||
chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True)
|
||||
chunk_dict: Final = (
|
||||
self.final_chunk
|
||||
if isinstance(self.final_chunk, dict)
|
||||
else self.final_chunk.model_dump(mode="json", exclude_none=True)
|
||||
)
|
||||
result["result"] = chunk_dict.get("result", {})
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -539,3 +539,34 @@ def test_streaming_logging_obj_keeps_agent_credentials_out_of_logging_params():
|
|||
assert logging_obj.litellm_params == expected
|
||||
assert logging_obj.optional_params == expected
|
||||
assert logging_obj.model_call_details["litellm_params"] == expected
|
||||
|
||||
|
||||
class _AgentFeeRecorder(CustomLogger):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.logged = asyncio.Event()
|
||||
self.fees = ()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if kwargs.get("call_type") in ("asend_message", "asend_message_streaming"):
|
||||
self.fees = (*self.fees, (kwargs.get("agent_id"), kwargs["standard_logging_object"]["response_cost"]))
|
||||
self.logged.set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("streaming", [False, True])
|
||||
async def test_completion_bridge_records_one_agent_fee(streaming, monkeypatch):
|
||||
from litellm.a2a_protocol.main import asend_message_streaming
|
||||
|
||||
recorder = _AgentFeeRecorder()
|
||||
monkeypatch.setattr(litellm, "callbacks", [recorder])
|
||||
params = {"custom_llm_provider": "openai", "model": "gpt-4o-mini", "mock_response": "hello back", "cost_per_query": 0.01}
|
||||
if streaming:
|
||||
request = SendStreamingMessageRequest(id="bridge-stream", params=_request().params)
|
||||
chunks = [chunk async for chunk in asend_message_streaming(request=request, litellm_params=params, agent_id="budgeted-agent")]
|
||||
assert chunks[-1]["result"]["final"] is True
|
||||
else:
|
||||
response = await asend_message(request=_request(), litellm_params=params, agent_id="budgeted-agent")
|
||||
assert response.id == "r1"
|
||||
await asyncio.wait_for(recorder.logged.wait(), timeout=2)
|
||||
assert recorder.fees == (("budgeted-agent", pytest.approx(0.01)),)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue