diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 641be7078b7..d40af60cbb3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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: diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 9cffa415a9a..41a02b10e4c 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index e4b782d5ff3..5b7cc7ebeeb 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 9aff2636c42..4f19ed5647d 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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)