mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +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(),
|
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,
|
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 = (
|
charged_targets: Final = (
|
||||||
[target for target in targets if target[0] != "model_per_team"]
|
[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)
|
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])
|
@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
|
monkeypatch: pytest.MonkeyPatch, self_call: bool
|
||||||
) -> None:
|
) -> None:
|
||||||
from litellm.types.agents import AgentResponse
|
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(
|
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: Final = UserAPIKeyAuth(agent_id="actor")
|
||||||
auth.managed_agent_policy = actor
|
auth.managed_agent_policy = actor
|
||||||
auth.invoked_agent_id = "actor" if self_call else "target"
|
auth.invoked_agent_id = "actor" if self_call else "target"
|
||||||
auth.invoked_agent_policy = 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)
|
monkeypatch.setattr(handler, "_get_agent_from_registry", lambda _: None)
|
||||||
descriptors: Final = handler._create_rate_limit_descriptors(
|
descriptors: Final = handler._create_rate_limit_descriptors(
|
||||||
user_api_key_dict=auth, data={"model": "a2a/target", "litellm_session_id": "session"},
|
user_api_key_dict=auth,
|
||||||
rpm_limit_type=None, tpm_limit_type=None, model_has_failures=False,
|
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}
|
limits: Final = {(item["key"], item["value"]): item["rate_limit"]["requests_per_unit"] for item in descriptors}
|
||||||
assert limits == (
|
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}
|
else {("agent", "actor"): 10, ("agent", "target"): 1, ("agent_session", "target:session"): 1}
|
||||||
)
|
)
|
||||||
assert len(descriptors) == len(limits)
|
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