mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(proxy): skip team model tpm accounting when key owns model tpm limit
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
460a0128c9
commit
db0f06d153
2 changed files with 71 additions and 3 deletions
|
|
@ -2893,6 +2893,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
return batch_limiter
|
||||
return None
|
||||
|
||||
def _key_owns_model_limit(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: str,
|
||||
rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
|
||||
) -> bool:
|
||||
key_own_limits: Final = get_key_own_model_rate_limit(user_api_key_dict, rate_limit_key)
|
||||
return key_own_limits is not None and key_own_limits.get(requested_model) is not None
|
||||
|
||||
def _inherited_team_model_limit(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -2903,11 +2912,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
team_limit: Final = team_limits.get(requested_model) if team_limits else None
|
||||
if team_limit is None:
|
||||
return None
|
||||
key_own_limits: Final = get_key_own_model_rate_limit(user_api_key_dict, rate_limit_key)
|
||||
if key_own_limits and key_own_limits.get(requested_model) is not None:
|
||||
if self._key_owns_model_limit(user_api_key_dict, requested_model, rate_limit_key):
|
||||
return None
|
||||
return team_limit
|
||||
|
||||
def _key_owns_model_tpm_limit_from_request_metadata(
|
||||
self,
|
||||
request_metadata: dict[str, Any],
|
||||
model_group: str | None,
|
||||
) -> bool:
|
||||
if model_group is None:
|
||||
return False
|
||||
key_view: Final = UserAPIKeyAuth(
|
||||
metadata=request_metadata.get("user_api_key_metadata") or {},
|
||||
model_max_budget=request_metadata.get("user_api_key_model_max_budget") or {},
|
||||
)
|
||||
return self._key_owns_model_limit(key_view, model_group, "model_tpm_limit")
|
||||
|
||||
def _add_team_model_rate_limit_descriptor_from_metadata(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -4463,6 +4484,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
kwargs=kwargs,
|
||||
model_group=reconcile_model,
|
||||
)
|
||||
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)
|
||||
else targets
|
||||
)
|
||||
if reserved_tokens > 0 and total_tokens < reserved_tokens:
|
||||
verbose_proxy_logger.debug(
|
||||
"Releasing unused TPM budget on success: reserved=%s, actual=%s, release=%s",
|
||||
|
|
@ -4472,7 +4498,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
pipeline_operations.extend(
|
||||
self._build_reservation_aware_tpm_ops(
|
||||
targets=targets,
|
||||
targets=charged_targets,
|
||||
reserved_scopes=reserved_scopes,
|
||||
actual_tokens=total_tokens,
|
||||
reserved_tokens=reserved_tokens,
|
||||
|
|
|
|||
|
|
@ -6597,3 +6597,45 @@ async def test_key_model_rpm_override_keeps_team_model_tpm_limit(key_limits, ove
|
|||
assert exc.value.status_code == 429
|
||||
assert "model_per_team" in str(exc.value.detail)
|
||||
assert exc.value.headers["rate_limit_type"] == "tokens"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_metadata, charges_team_model_pool",
|
||||
[
|
||||
({}, True),
|
||||
({"model_rpm_limit": {"test-model": 10}}, True),
|
||||
({"model_tpm_limit": {"test-model": 5000}}, False),
|
||||
({"model_tpm_limit": {"other-model": 5000}}, True),
|
||||
],
|
||||
ids=["no_override", "rpm_only_override", "tpm_override", "tpm_override_on_other_model"],
|
||||
)
|
||||
def test_success_tpm_accounting_skips_team_model_pool_when_key_owns_model_tpm_limit(
|
||||
key_metadata, charges_team_model_pool
|
||||
):
|
||||
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
response = ModelResponse(
|
||||
id="team-pool-tpm",
|
||||
object="chat.completion",
|
||||
created=int(datetime.now().timestamp()),
|
||||
model="test-model",
|
||||
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
|
||||
choices=[],
|
||||
)
|
||||
kwargs = {
|
||||
"standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-pool"), "user_api_key_team_id": "t"}},
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"model_group": "test-model",
|
||||
"user_api_key_metadata": key_metadata,
|
||||
"user_api_key_team_metadata": {"model_tpm_limit": {"test-model": 500}},
|
||||
}
|
||||
},
|
||||
"model": "test-model",
|
||||
}
|
||||
|
||||
ops = handler._build_success_event_pipeline_operations(kwargs=kwargs, response_obj=response, rate_limit_type="output")
|
||||
|
||||
charged_keys = {op["key"] for op in ops}
|
||||
assert handler.create_rate_limit_keys("model_per_key", f"{hash_token('sk-pool')}:test-model", "tokens") in charged_keys
|
||||
team_pool_key = handler.create_rate_limit_keys("model_per_team", "t:test-model", "tokens")
|
||||
assert (team_pool_key in charged_keys) is charges_team_model_pool
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue