diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index 03285d360c2..a5691d49a64 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -586,3 +586,57 @@ async def test_new_budget_window_isolated_from_inflight_previous_window_charge(m await check_agent_budget(auth) assert await cache.async_get_cache(old_key) == 10.5 assert await cache.async_get_cache(new_key) == 1.1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("autonomous", [False, True]) +@pytest.mark.parametrize( + "pricing,billable,bounded,rejected,fee", + [ + ({"input_cost_per_token": 0.01}, True, True, True, None), + ({"output_cost_per_token": 0.01}, True, True, True, None), + ({"input_cost_per_token": 0.01}, False, True, False, 0.0), + ({"input_cost_per_token": 0.01}, True, False, False, 0.0), + ({"input_cost_per_token": 0.0, "output_cost_per_token": 0.0}, True, True, False, 0.0), + ({"cost_per_query": 0.25, "output_cost_per_token": 0.01}, True, True, False, 0.25), + ({"cost_per_query": 0.0, "output_cost_per_token": 0.01}, True, True, False, 0.0), + ], +) +async def test_budgeted_invocation_requires_a_bounded_price( + monkeypatch: pytest.MonkeyPatch, + autonomous: bool, + pricing: dict[str, float], + billable: bool, + bounded: bool, + rejected: bool, + fee: float | None, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + budget: Final = {"budget_id": "budget", "max_budget": 1.0} if bounded else None + target: Final = agent(litellm_params=pricing, litellm_budget_table=budget) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"]) + auth: Final = UserAPIKeyAuth(agent_id="caller" if autonomous else None, object_permission=permission) + if autonomous: + caller: Final = agent( + agent_id="caller", object_permission=permission.model_dump(), litellm_budget_table=budget + ) + auth.managed_agent_policy = caller + auth.billing_agent_policy = caller + if rejected: + with pytest.raises(HTTPException) as exc: + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert exc.value.status_code == 503 + assert "cost_per_query" in str(exc.value.detail) + else: + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert auth.agent_invocation_cost == fee