mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(agents): exclude discovery from billing and repair budget fixtures
This commit is contained in:
parent
d51659e817
commit
2a7c4187c6
9 changed files with 104 additions and 8 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue