mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): enforce and settle chat adapter invocation fees
This commit is contained in:
parent
8f6d78d652
commit
b33fda1dd3
6 changed files with 96 additions and 4 deletions
|
|
@ -1529,7 +1529,14 @@ def completion_cost(
|
|||
completion_tokens = token_counter(model=model, text=completion)
|
||||
|
||||
# Handle A2A calls before model check - A2A doesn't require a model
|
||||
if call_type in _A2A_CALL_TYPES:
|
||||
if call_type in _A2A_CALL_TYPES or (
|
||||
custom_llm_provider == "a2a"
|
||||
and litellm_logging_obj is not None
|
||||
and (litellm_logging_obj.model_call_details.get("litellm_params") or MappingProxyType({})).get(
|
||||
"cost_per_query"
|
||||
)
|
||||
is not None
|
||||
):
|
||||
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
|
||||
|
||||
return A2ACostCalculator.calculate_a2a_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
|
|
|
|||
|
|
@ -5662,6 +5662,7 @@ def completion(
|
|||
preset_cache_key=preset_cache_key,
|
||||
no_log=no_log,
|
||||
cost_per_second=cost_per_second,
|
||||
cost_per_query=kwargs.get("cost_per_query"),
|
||||
input_cost_per_second=input_cost_per_second,
|
||||
input_cost_per_token=input_cost_per_token,
|
||||
output_cost_per_second=output_cost_per_second,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Handles routing for A2A agents (models with "a2a/<agent-name>" prefix).
|
|||
Looks up agents in the registry and injects their API base URL.
|
||||
"""
|
||||
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -78,4 +79,13 @@ async def route_a2a_agent_request(
|
|||
data["api_base"] = agent.agent_card_params["url"]
|
||||
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
|
||||
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
invocation_pricing: Final = (
|
||||
MappingProxyType({"cost_per_query": user_api_key_dict.agent_invocation_cost})
|
||||
if user_api_key_dict is not None
|
||||
and user_api_key_dict.agent_invocation_cost is not None
|
||||
and user_api_key_dict.invoked_agent_policy is not None
|
||||
and (user_api_key_dict.invoked_agent_policy.litellm_params or MappingProxyType({})).get("cost_per_query")
|
||||
is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, **invocation_pricing}))
|
||||
|
|
|
|||
|
|
@ -3261,8 +3261,11 @@ async def _authorize_authenticated_request(
|
|||
target_name,
|
||||
store,
|
||||
billable=request.method == "POST"
|
||||
and request_data.get("method")
|
||||
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
|
||||
and (
|
||||
not RouteChecks.check_route_access(route, ("/a2a/{agent_id}", "/v1/a2a/{agent_id}"))
|
||||
or request_data.get("method")
|
||||
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage")
|
||||
),
|
||||
)
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
|
|
|
|||
|
|
@ -9701,6 +9701,9 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon
|
|||
("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True),
|
||||
("POST", "/a2a/agent", {"method": "message/send", "model": "free-model", "params": {}}, True),
|
||||
("POST", "/a2a/agent", {"method": "message/stream", "model": "free-model", "params": {}}, True),
|
||||
("POST", "/v1/chat/completions", {"model": "a2a/agent", "method": "tasks/get"}, True),
|
||||
("POST", "/chat/completions", {"model": "a2a/agent", "method": "tasks/cancel", "stream": True}, True),
|
||||
("POST", "/v1/a2a/agent/message/send", {"method": "tasks/get", "params": {}}, True),
|
||||
],
|
||||
)
|
||||
async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do(
|
||||
|
|
|
|||
|
|
@ -117,6 +117,7 @@ class CostLogger(CustomLogger):
|
|||
|
||||
def __init__(self):
|
||||
self.response_cost: Optional[float] = None
|
||||
self.logged: asyncio.Event = asyncio.Event()
|
||||
super().__init__()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -125,6 +126,7 @@ class CostLogger(CustomLogger):
|
|||
self.response_cost = (
|
||||
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
||||
)
|
||||
self.logged.set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -449,3 +451,69 @@ async def test_asend_message_streaming_triggers_callbacks():
|
|||
assert callback_logger.agent_id == test_agent_id, (
|
||||
f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stream", (False, True))
|
||||
@pytest.mark.parametrize("claimed_fee", (None, 0.0, 99.0))
|
||||
async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
|
||||
monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None
|
||||
) -> None:
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
await _reset_callbacks_and_settle_pending_logs()
|
||||
logger: Final = CostLogger()
|
||||
monkeypatch.setattr(litellm, "callbacks", [logger])
|
||||
target: Final = AgentResponse(
|
||||
agent_id="fee-target", agent_name="fee-target",
|
||||
agent_card_params={"url": "https://agent.test/", "capabilities": {"streaming": True}},
|
||||
litellm_params={"cost_per_query": 0.25},
|
||||
)
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
registry.register_agent(target)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
|
||||
auth.invoked_agent_policy = target
|
||||
auth.invoked_agent_id = target.agent_id
|
||||
auth.agent_invocation_cost = 0.25
|
||||
|
||||
def reply(request: httpx.Request) -> httpx.Response:
|
||||
body: Final = json.loads(request.content)
|
||||
assert body["method"] == ("message/stream" if stream else "message/send")
|
||||
result: Final = {
|
||||
"jsonrpc": "2.0", "id": body["id"],
|
||||
"result": {"kind": "message", "messageId": "reply", "role": "agent",
|
||||
"parts": [{"kind": "text", "text": "Paid reply"}]},
|
||||
}
|
||||
if stream:
|
||||
return httpx.Response(200, text=f"data: {json.dumps(result)}\n\n", headers={"content-type": "text/event-stream"})
|
||||
return httpx.Response(200, json=result)
|
||||
|
||||
client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(reply))
|
||||
try:
|
||||
pending: Final = await route_a2a_agent_request(
|
||||
data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}],
|
||||
"stream": stream, "client": client,
|
||||
**({"cost_per_query": claimed_fee} if claimed_fee is not None else {})},
|
||||
route_type="acompletion", user_api_key_dict=auth,
|
||||
)
|
||||
response: Final = await pending
|
||||
if stream:
|
||||
chunks: Final = tuple([chunk async for chunk in response])
|
||||
assert any(chunk.choices[0].delta.content == "Paid reply" for chunk in chunks)
|
||||
else:
|
||||
assert response.choices[0].message.content == "Paid reply"
|
||||
await asyncio.wait_for(logger.logged.wait(), timeout=10.0)
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
|
||||
assert logger.response_cost == pytest.approx(0.25)
|
||||
finally:
|
||||
await client.close()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue