fix(agents): preserve token pricing and legacy request handling

This commit is contained in:
Joshua Valluru 2026-09-30 18:41:55 -07:00
parent bc5ec23354
commit cbec06c183
5 changed files with 51 additions and 7 deletions

View file

@ -738,7 +738,9 @@ async def invoke_agent_a2a(
str, object
] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict
**(agent.litellm_params or MappingProxyType({})),
"cost_per_query": user_api_key_dict.agent_invocation_cost,
"cost_per_query": user_api_key_dict.agent_invocation_cost
if (agent.litellm_params or MappingProxyType({})).get("cost_per_query") is not None
else None,
}
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")

View file

@ -82,9 +82,11 @@ async def route_a2a_agent_request(
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
# Inject API base and route to litellm
data.pop("litellm_params", None)
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
api_base: Final = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base)
invocation_fee: Final = user_api_key_dict.agent_invocation_cost if user_api_key_dict is not None else None
return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee}))
provider_data: Final = MappingProxyType({key: value for key, value in data.items() if key != "litellm_params"})
return getattr(litellm, f"{route_type}")(
**MappingProxyType({**provider_data, "api_base": api_base, "cost_per_query": invocation_fee})
)

View file

@ -3264,7 +3264,7 @@ async def _authorize_authenticated_request(
general_settings,
user_model,
request.path_params.get("model") or request.path_params.get("model_name"),
request.query_params.get("model"),
_safe_get_request_query_params(request).get("model"),
model_group_alias=router_settings.get("model_group_alias")
if isinstance(router_settings, Mapping)
else None,

View file

@ -27,6 +27,7 @@ class CapturedAgentCall:
agent_extra_headers: dict[str, str] | None
cost_per_query: object
api_base: object
pricing: Mapping[str, object]
@pytest.mark.asyncio
@ -500,6 +501,7 @@ async def _invoke_message_method(
agent_extra_headers=kwargs.get("agent_extra_headers"),
cost_per_query=kwargs["litellm_params"].get("cost_per_query"),
api_base=kwargs["api_base"],
pricing=kwargs["litellm_params"],
)
@ -2753,3 +2755,41 @@ async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch
)
assert response.status_code == 404
assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("message/send", "message/stream"))
@pytest.mark.parametrize("fixed_fee", (None, 0.0, 0.25))
async def test_unbudgeted_managed_agent_keeps_token_pricing_without_a_fixed_fee(
monkeypatch: pytest.MonkeyPatch, method: str, fixed_fee: float | None,
) -> None:
from litellm import Usage
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
policy: Final = AgentResponse(
agent_id="test-agent", agent_name="test-agent", identity_managed=True,
agent_card_params={"url": "https://agent.test/"},
litellm_params={"input_cost_per_token": 0.02, "output_cost_per_token": 0.03,
**({"cost_per_query": fixed_fee} if fixed_fee is not None else {})},
)
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(policy)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
with ExitStack() as stack:
for context in _base_patches(policy):
stack.enter_context(context)
await prepare_agent_invocation(auth, policy.agent_id, AgentIdentityStore.from_client(database))
captured: Final = await _invoke_message_method(
method, _make_request_mock(method, _HELLO_MESSAGE_PARAMS), auth, agent=policy,
)
logging: Final = MagicMock(model_call_details={
"litellm_params": captured.pricing,
"usage": Usage(prompt_tokens=10, completion_tokens=2, total_tokens=12),
})
assert A2ACostCalculator.calculate_a2a_cost(logging) == pytest.approx(0.26 if fixed_fee is None else fixed_fee)

View file

@ -2270,7 +2270,7 @@ class TestGigachatProxyRoute:
mock_request.headers = {"content-type": "application/json"}
mock_request.query_params = {}
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict = UserAPIKeyAuth()
mock_llm_router.allm_passthrough_route = AsyncMock(
return_value=httpx.Response(200, json={"response": "success"})
)