fix(agents): enforce and settle chat adapter invocation fees

This commit is contained in:
Joshua Valluru 2026-09-30 16:54:32 -07:00
parent 8f6d78d652
commit b33fda1dd3
6 changed files with 96 additions and 4 deletions

View file

@ -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)

View file

@ -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,

View file

@ -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}))

View file

@ -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,

View file

@ -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(

View file

@ -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()