From 2a7c4187c6f6d675d7b2ad3e0a0f540bbf4828d5 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sun, 27 Sep 2026 09:32:37 -0700 Subject: [PATCH] fix(agents): exclude discovery from billing and repair budget fixtures --- .../auth/managed_authorization.py | 6 +- litellm/proxy/auth/user_api_key_auth.py | 3 +- .../auth/test_managed_authorization.py | 13 +++- .../proxy/auth/test_user_api_key_auth.py | 76 +++++++++++++++++++ .../test_access_group_management.py | 3 + .../test_organization_endpoints.py | 1 + .../test_tag_management_endpoints.py | 8 +- .../test_team_endpoints.py | 1 + .../test_management_helpers_utils.py | 1 + 9 files changed, 104 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 0218e05cf44..9aa6e80f4af 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -242,7 +242,11 @@ async def prepare_agent_invocation( raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent")) auth.invoked_agent_id = effective.agent_id auth.invoked_agent_policy = effective - if auth.agent_id is None and (effective.identity_managed or effective.litellm_budget_table is not None): + if ( + billable + and auth.agent_id is None + and (effective.identity_managed or effective.litellm_budget_table is not None) + ): auth.billing_agent_policy = effective raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0 try: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2829bc95b42..014c7f4fcdf 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3244,7 +3244,8 @@ async def _authorize_authenticated_request( user_api_key_auth_obj, target_name, store, - billable=request_data.get("method") + billable=request.method == "POST" + and request_data.get("method") in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), ) await _run_centralized_common_checks( 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 7dfcf890583..03285d360c2 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 @@ -129,9 +129,11 @@ async def test_agent_budget_accumulates_across_credentials_and_denies_the_next_a @pytest.mark.asyncio @pytest.mark.parametrize("autonomous", (True, False)) +@pytest.mark.parametrize("billable", (True, False)) async def test_invocation_prepares_target_fee_for_the_correct_agent( monkeypatch: pytest.MonkeyPatch, autonomous: bool, + billable: bool, ) -> None: from unittest.mock import AsyncMock, MagicMock @@ -158,11 +160,14 @@ async def test_invocation_prepares_target_fee_for_the_correct_agent( caller: Final = agent(agent_id="caller", object_permission=permission.model_dump()) auth.managed_agent_policy = caller auth.billing_agent_policy = caller - await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) - assert auth.agent_invocation_cost == pytest.approx(0.25) + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert auth.agent_invocation_cost == pytest.approx(0.25 if billable else 0.0) assert auth.invoked_agent_id == "agent" - assert auth.billing_agent_policy is not None - assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent") + if autonomous or billable: + assert auth.billing_agent_policy is not None + assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent") + else: + assert auth.billing_agent_policy is None @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 268b70c9596..767da6d9fbf 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -9490,3 +9490,79 @@ async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeyp assert failure.value.code == "403" assert "without virtual-key mapping" in failure.value.message client.writer_db.litellm_agentstable.find_unique.assert_awaited_once() + + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "http_method,route,body,billed", + [ + ("GET", "/a2a/agent/.well-known/agent-card.json", {}, False), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "tasks/get", "params": {"id": "t"}}, False), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/send", "params": {}}, True), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True), + ], +) +async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do( + monkeypatch: pytest.MonkeyPatch, http_method: str, route: str, body: dict, billed: bool +) -> None: + from typing import Final + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + target: Final = AgentResponse( + agent_id="agent", + agent_name="Agent", + agent_card_params={}, + identity_managed=True, + execution_mode="both", + litellm_params={"cost_per_query": 0.25}, + litellm_budget_table={"budget_id": "agent-budget", "max_budget": 10.0}, + identity=AgentIdentityBinding( + agent_id="agent", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + issuer="issuer", + revision="current", + ), + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + request = _alias_request(route, body) + request.scope["method"] = http_method + auth: Final = UserAPIKeyAuth( + api_key="human-key", + user_id="human", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["agent"]), + ) + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new_callable=AsyncMock, + return_value=None, + ) as reserve: + assert await _authorize_authenticated_request(auth, request, body, route, "human-key") is None + assert auth.invoked_agent_id == "agent" + assert auth.agent_invocation_cost == pytest.approx(0.25 if billed else 0.0), (http_method, body.get("method")) + if billed: + assert auth.billing_agent_policy is not None and auth.billing_agent_policy.agent_id == "agent" + else: + assert auth.billing_agent_policy is None, (http_method, body.get("method")) + reserve.assert_awaited_once() + reserved: Final = reserve.call_args.kwargs["valid_token"] + assert reserved is auth and (reserved.billing_agent_policy is not None) is billed, (http_method, body.get("method")) diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 59c2921e0d0..2843bb60965 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -714,6 +714,9 @@ class _FakePrismaClient: litellm_modelaccessgroupbudgettable=self.access_group_budget_table, litellm_proxymodeltable=self.model_table, ) + self.writer_db = SimpleNamespace( + litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None)), + ) def jsonify_object(self, data): return dict(data) diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index 3c6afa86c45..2bdc81756c7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -1176,6 +1176,7 @@ async def _run_legacy_update_organization( mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org) mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock()) mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) diff --git a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py index 3cfdd345a45..830ffdcad92 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py @@ -35,6 +35,10 @@ class _BudgetState: return SimpleNamespace(**self._values) +def _writer_db_without_agent_budgets() -> SimpleNamespace: + return SimpleNamespace(litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None))) + + class FakeVerificationTokenTable: """Stand-in for ``prisma_client.db.litellm_verificationtoken``. @@ -316,7 +320,7 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets()) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) @@ -370,7 +374,7 @@ async def test_update_tag_explicit_null_clears_budget_duration(): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets()) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 54827f094f0..4c2ecf270e3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -2998,6 +2998,7 @@ async def test_upsert_team_member_budget_table_clears_duration_kept_budget(mock_ mock_db_client.db.litellm_budgettable.update = AsyncMock( side_effect=lambda where, data: SimpleNamespace(**data) ) + mock_db_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) result = await TeamMemberBudgetHandler.upsert_team_member_budget_table( team_table=team_table, diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 922504ecc58..46bac64db19 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -845,6 +845,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): db: Final = _FakeDb() prisma_client: Final = MagicMock() prisma_client.db = db + prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN) team_id: Final = "team-shared-default" default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0})