mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): preserve token pricing and legacy request handling
This commit is contained in:
parent
bc5ec23354
commit
cbec06c183
5 changed files with 51 additions and 7 deletions
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue