fix(agents): reconcile every reserved agent token counter

This commit is contained in:
Joshua Valluru 2026-09-27 14:15:22 -07:00
parent 4ce23d0726
commit 91c331ae73
2 changed files with 46 additions and 6 deletions

View file

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

View file

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