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:
yassin 2026-09-15 22:35:48 +00:00
parent 460a0128c9
commit db0f06d153
2 changed files with 71 additions and 3 deletions

View file

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

View file

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