fix(agents): enforce actor and target invocation rate limits

This commit is contained in:
Joshua Valluru 2026-09-27 14:01:18 -07:00
parent 1b0162bbfb
commit 4ce23d0726
4 changed files with 43 additions and 8 deletions

View file

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

View file

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

View file

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

View file

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