mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(agents): reconcile every reserved agent token counter
This commit is contained in:
parent
4ce23d0726
commit
91c331ae73
2 changed files with 46 additions and 6 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue