diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 5b7cc7ebeeb..96fd49056c6 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4752,6 +4752,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): tpm_limited_tags=stash.tpm_limited_tags if stash is not None else frozenset(), model_group=reconcile_model.group if reconcile_model is not None else None, ) + targets.extend( + scope + for scope in sorted(reserved_scopes) + if scope[0] in ("agent", "agent_session") and scope not in targets + ) charged_targets: Final = ( [target for target in targets if target[0] != "model_per_team"] if self._key_owns_model_tpm_limit_from_request_metadata(request_metadata, reconcile_model) 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 4f19ed5647d..0ded3da917a 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 @@ -7562,24 +7562,36 @@ async def test_success_tpm_accounting_keeps_the_admission_target_after_an_alias_ @pytest.mark.parametrize("self_call", [False, True]) -def test_managed_invocations_enforce_actor_and_target_rate_policies( +async 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) + actor: Final = AgentResponse( + agent_id="actor", agent_name="Actor", agent_card_params={}, rpm_limit=10, tpm_limit=1000 + ) target: Final = AgentResponse( - agent_id="target", agent_name="Target", agent_card_params={}, rpm_limit=1, session_rpm_limit=1 + agent_id="target", + agent_name="Target", + agent_card_params={}, + rpm_limit=1, + tpm_limit=1000, + session_rpm_limit=1, + session_tpm_limit=1000, ) 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())) + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) 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, + 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 == ( @@ -7588,3 +7600,26 @@ def test_managed_invocations_enforce_actor_and_target_rate_policies( else {("agent", "actor"): 10, ("agent", "target"): 1, ("agent_session", "target:session"): 1} ) assert len(descriptors) == len(limits) + await handler.async_pre_call_hook( + user_api_key_dict=auth, + cache=cache, + data={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 20, + "litellm_session_id": "session", + }, + call_type="acompletion", + ) + stash: Final = get_request_stash() + assert stash is not None and stash.reserved_tokens > 3 + response: Final = ModelResponse(usage=Usage(prompt_tokens=2, completion_tokens=1, total_tokens=3)) + operations: Final = handler._build_success_event_pipeline_operations( + kwargs={"standard_logging_object": {"metadata": {"agent_id": auth.invoked_agent_id, "session_id": "session"}}}, + response_obj=response, + rate_limit_type="total", + ) + increments: Final = {op["key"]: op["increment_value"] for op in operations} + for scope in stash.reserved_scopes: + if scope[0] in ("agent", "agent_session"): + assert increments[handler.create_rate_limit_keys(*scope, "tokens")] == 3 - stash.reserved_tokens