mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(agents): enforce actor and target invocation rate limits
This commit is contained in:
parent
1b0162bbfb
commit
4ce23d0726
4 changed files with 43 additions and 8 deletions
|
|
@ -3342,6 +3342,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
|
|||
),
|
||||
)
|
||||
invoked_agent_id: str | None = Field(default=None, exclude=True)
|
||||
invoked_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
agent_invocation_cost: float | None = Field(default=None, exclude=True)
|
||||
billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
|
|
@ -3393,6 +3394,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
|
|||
values.pop("managed_agent_context", None)
|
||||
values.pop("managed_agent_policy", None)
|
||||
values.pop("invoked_agent_id", None)
|
||||
values.pop("invoked_agent_policy", None)
|
||||
values.pop("agent_invocation_cost", None)
|
||||
values.pop("billing_agent_policy", None)
|
||||
if values.get("api_key") is not None:
|
||||
|
|
|
|||
|
|
@ -221,6 +221,7 @@ async def prepare_agent_invocation(
|
|||
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
|
||||
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:
|
||||
auth.billing_agent_policy = effective
|
||||
raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0
|
||||
|
|
|
|||
|
|
@ -2881,6 +2881,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self,
|
||||
agent_id: str,
|
||||
data: dict,
|
||||
policy: "AgentResponse | None" = None,
|
||||
) -> list[RateLimitDescriptor]:
|
||||
"""
|
||||
Create rate limit descriptors for agent-level and session-level limits.
|
||||
|
|
@ -2890,7 +2891,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"""
|
||||
descriptors: Final[list[RateLimitDescriptor]] = []
|
||||
|
||||
agent: Final = self._get_agent_from_registry(agent_id)
|
||||
agent: Final = policy if policy is not None else self._get_agent_from_registry(agent_id)
|
||||
if agent is None:
|
||||
return descriptors
|
||||
|
||||
|
|
@ -3085,15 +3086,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
descriptors=descriptors,
|
||||
)
|
||||
|
||||
# Agent-level and session-level rate limits
|
||||
resolved_agent_id: Final = self._get_resolved_agent_id(user_api_key_dict, data)
|
||||
|
||||
if resolved_agent_id:
|
||||
for agent_id in dict.fromkeys((resolved_agent_id, user_api_key_dict.invoked_agent_id)):
|
||||
if agent_id is None:
|
||||
continue
|
||||
policy: Final = (
|
||||
user_api_key_dict.managed_agent_policy
|
||||
if agent_id == user_api_key_dict.agent_id
|
||||
else user_api_key_dict.invoked_agent_policy
|
||||
)
|
||||
descriptors.extend(
|
||||
self._create_agent_rate_limit_descriptors(
|
||||
agent_id=resolved_agent_id,
|
||||
data=data,
|
||||
)
|
||||
self._create_agent_rate_limit_descriptors(agent_id=agent_id, data=data, policy=policy)
|
||||
)
|
||||
|
||||
return descriptors
|
||||
|
|
|
|||
|
|
@ -7559,3 +7559,32 @@ async def test_success_tpm_accounting_keeps_the_admission_target_after_an_alias_
|
|||
charged: Final = {op["key"]: op["increment_value"] for op in ops}
|
||||
assert charged[admission_bucket] == 150 - stash.reserved_tokens
|
||||
assert not any(":target-b" in key for key in charged)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("self_call", [False, True])
|
||||
def test_managed_invocations_enforce_actor_and_target_rate_policies(
|
||||
monkeypatch: pytest.MonkeyPatch, self_call: bool
|
||||
) -> None:
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
actor: Final = AgentResponse(agent_id="actor", agent_name="Actor", agent_card_params={}, rpm_limit=10)
|
||||
target: Final = AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, rpm_limit=1, session_rpm_limit=1
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth.managed_agent_policy = actor
|
||||
auth.invoked_agent_id = "actor" if self_call else "target"
|
||||
auth.invoked_agent_policy = actor if self_call else target
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
monkeypatch.setattr(handler, "_get_agent_from_registry", lambda _: None)
|
||||
descriptors: Final = handler._create_rate_limit_descriptors(
|
||||
user_api_key_dict=auth, data={"model": "a2a/target", "litellm_session_id": "session"},
|
||||
rpm_limit_type=None, tpm_limit_type=None, model_has_failures=False,
|
||||
)
|
||||
limits: Final = {(item["key"], item["value"]): item["rate_limit"]["requests_per_unit"] for item in descriptors}
|
||||
assert limits == (
|
||||
{("agent", "actor"): 10}
|
||||
if self_call
|
||||
else {("agent", "actor"): 10, ("agent", "target"): 1, ("agent_session", "target:session"): 1}
|
||||
)
|
||||
assert len(descriptors) == len(limits)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue