fix(agents): exclude discovery from billing and repair budget fixtures

This commit is contained in:
Joshua Valluru 2026-09-27 09:32:37 -07:00
parent d51659e817
commit 2a7c4187c6
9 changed files with 104 additions and 8 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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