fix(proxy): enforce project ITPM/OTPM on context-management summaries

Summary subrequests skipped project model_itpm_limit/model_otpm_limit on
both the read-only gate and post-call charge path. Include those descriptors
with estimate headroom checks, and charge unreserved summary call ids from
logging metadata like combined TPM.

Signed-off-by: sinksilk <785976238@qq.com>
This commit is contained in:
sinksilk 2026-09-22 23:24:08 +08:00
parent 615ed7900f
commit 791d38d87b
3 changed files with 310 additions and 6 deletions

View file

@ -521,6 +521,9 @@ def _without_parallel_request_gauge(descriptor: "RateLimitDescriptor") -> "RateL
async def _check_summary_model_rate_limit(
user_api_key_auth: Optional["UserAPIKeyAuth"],
summary_model: str,
*,
estimated_input_tokens: int = 1,
estimated_output_tokens: int = 1,
) -> bool:
"""Return True when the caller is within their configured RPM/TPM limits
for ``summary_model``.
@ -530,14 +533,22 @@ async def _check_summary_model_rate_limit(
user RPM or TPM could still drive an extra summary-model completion per
allowed ``/v1/messages`` request. This mirrors the read side of
``_PROXY_MaxParallelRequestsHandler_v3.async_pre_call_hook`` for the
summary model: it builds the same descriptor set and runs the check in
``read_only`` mode so no counter is reserved or incremented — the summary
call's actual usage is still charged exactly once by the limiter's
post-call success hook (via the propagated ``litellm_metadata``).
summary model: it builds the same descriptor set (including project
ITPM/OTPM) and runs the check in ``read_only`` mode so no counter is
reserved or incremented — the summary call's actual usage is still
charged exactly once by the limiter's post-call success hook (via the
propagated ``litellm_metadata``).
``max_parallel_requests`` gauges are left out of the check: the summary
call runs inside the caller's already admitted request, whose own slot
would otherwise count against it.
Project ITPM/OTPM are reservation-style quotas (pre-call reserves an
estimate). A read-only ``OVER_LIMIT`` only fires once the counter is
already at the cap, so this gate also compares ``limit_remaining`` to
``estimated_input_tokens`` / ``estimated_output_tokens`` for those
descriptors — matching how an ordinary request would be refused when the
next reservation cannot fit.
Returns True (allow) outside the proxy, when the active limiter does not
expose the read-only descriptor check (legacy limiter), or when the
descriptor set cannot be built — the deny signals are a definitive
@ -566,6 +577,9 @@ async def _check_summary_model_rate_limit(
add_project_descriptor: Final[_AddModelRateLimitDescriptor | None] = getattr(
limiter, "_add_project_model_rate_limit_descriptor_from_metadata", None
)
add_project_io_descriptor: Final[_AddModelRateLimitDescriptor | None] = getattr(
limiter, "add_project_io_token_rate_limit_descriptors_from_metadata", None
)
create_org_descriptors: Final[_CreateOrgRateLimitDescriptors | None] = getattr(
limiter, "create_organization_rate_limit_descriptor", None
)
@ -599,6 +613,15 @@ async def _check_summary_model_rate_limit(
requested_model=summary_model,
descriptors=base_descriptors,
)
# Project ITPM/OTPM are reserved (not merely read) on the main pre-call
# path, so the summary gate must add those descriptors explicitly —
# otherwise an exhausted project IO quota still allows compaction.
if add_project_io_descriptor is not None:
add_project_io_descriptor(
user_api_key_dict=user_api_key_auth,
requested_model=summary_model,
descriptors=base_descriptors,
)
descriptors: Final = _without_parallel_request_gauges(
(*base_descriptors, *create_org_descriptors(user_api_key_auth, summary_model))
)
@ -624,7 +647,25 @@ async def _check_summary_model_rate_limit(
e,
)
return True
return response.get("overall_code") != "OVER_LIMIT"
if response.get("overall_code") == "OVER_LIMIT":
return False
# Reservation-style project IO quotas: deny when the estimated summary
# cannot fit in remaining headroom (ordinary traffic fails the same way).
input_estimate: Final = max(1, estimated_input_tokens)
output_estimate: Final = max(1, estimated_output_tokens)
for status in response.get("statuses") or ():
if not isinstance(status, Mapping):
continue
descriptor_key = status.get("descriptor_key")
remaining = status.get("limit_remaining")
if not isinstance(remaining, int):
continue
if descriptor_key == "model_per_project_itpm" and remaining < input_estimate:
return False
if descriptor_key == "model_per_project_otpm" and remaining < output_estimate:
return False
return True
def _find_latest_compaction_index(
@ -1293,6 +1334,8 @@ async def apply_compact_20260112(
if not await _check_summary_model_rate_limit(
user_api_key_auth=user_api_key_auth,
summary_model=summary_model,
estimated_input_tokens=current_tokens,
estimated_output_tokens=_read_summary_max_tokens_setting(),
):
verbose_logger.warning(
"compact_20260112: caller over rate limit for summary_model=%s; skipping summary call",

View file

@ -4637,6 +4637,90 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return 0, 0, False
def _collect_project_io_scope_targets(
self,
standard_logging_metadata: Mapping[str, Any],
model_group: str | None,
) -> list[tuple[str, str]]:
"""Rebuild project ITPM/OTPM scopes from logging metadata.
Combined TPM already charges ``model_per_project`` from metadata when
no reservation owns the call. Summary/compaction subrequests carry a
distinct ``litellm_call_id`` and never own the parent stash, so their
IO quotas must use the same metadata rebuild.
"""
user_api_key_project_id: Final = standard_logging_metadata.get("user_api_key_project_id")
if not user_api_key_project_id or not model_group:
return []
descriptor_value: Final = f"{user_api_key_project_id}:{model_group}"
return [ # mutable-ok: caller may filter ITPM vs OTPM scopes
(PROJECT_ITPM_DESCRIPTOR_KEY, descriptor_value),
(PROJECT_OTPM_DESCRIPTOR_KEY, descriptor_value),
]
def _build_unreserved_project_io_token_ops(
self,
kwargs: dict[str, Any],
response_obj: object,
) -> Sequence[RedisPipelineIncrementOperation]:
"""Charge full actual ITPM/OTPM when no pre-call reservation owns this call.
Summary subrequests never claim the parent stash (``owner_litellm_call_id``
pins it), so without this path their input/output tokens never hit the
project IO counters even though combined TPM still charges them.
"""
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
standard_logging_object: Final = kwargs.get("standard_logging_object") or {}
if not isinstance(standard_logging_object, dict):
return ()
standard_logging_metadata: Final = standard_logging_object.get("metadata") or {}
if not isinstance(standard_logging_metadata, Mapping):
return ()
model_group: Final = get_model_group_from_litellm_kwargs(kwargs) or (
standard_logging_object.get("model_group")
if isinstance(standard_logging_object.get("model_group"), str)
else None
)
targets: Final = self._collect_project_io_scope_targets(
standard_logging_metadata=standard_logging_metadata,
model_group=model_group if isinstance(model_group, str) else None,
)
if not targets:
return ()
response_usage: Final = self._resolve_io_token_reconcile_usage(response_obj)
combined_usage: Final = self._resolve_io_token_reconcile_usage(kwargs.get("combined_usage_object"))
aggregate_total: Final = self._aggregate_only_total_tokens(
self._response_usage(response_obj)
) or self._aggregate_only_total_tokens(self._response_usage(kwargs.get("combined_usage_object")))
if not response_usage[2] and not combined_usage[2] and aggregate_total <= 0:
return ()
resolved_usage: Final = (
response_usage
if response_usage[2]
else combined_usage
if combined_usage[2]
else (aggregate_total, aggregate_total, True)
)
billable_input, completion_tokens, _ = resolved_usage
itpm_targets: Final = [t for t in targets if t[0] == PROJECT_ITPM_DESCRIPTOR_KEY]
otpm_targets: Final = [t for t in targets if t[0] == PROJECT_OTPM_DESCRIPTOR_KEY]
return self._build_reservation_aware_tpm_ops(
targets=itpm_targets,
reserved_scopes=frozenset(),
actual_tokens=billable_input,
reserved_tokens=0,
) + self._build_reservation_aware_tpm_ops(
targets=otpm_targets,
reserved_scopes=frozenset(),
actual_tokens=completion_tokens,
reserved_tokens=0,
)
def _build_io_token_reservation_ops(
self,
kwargs: object,
@ -4649,12 +4733,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
are stored in the same ":tokens" cache bucket as combined TPM, just
under distinct scope keys, so the reservation-aware increment math is
identical; only the usage fields being reconciled against differ.
When this call id does not own the request stash (summary/compaction
subrequest), falls through to the unreserved metadata rebuild so
project IO quotas still receive the summary's actual usage.
"""
if not isinstance(kwargs, dict):
return ()
stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
if stash is None:
return ()
return self._build_unreserved_project_io_token_ops(kwargs, response_obj)
itpm_reserved: Final = stash.itpm_reserved_tokens
otpm_reserved: Final = stash.otpm_reserved_tokens

View file

@ -3708,5 +3708,178 @@ async def test_the_project_itpm_reservation_counts_the_request_off_the_event_loo
assert_loop_stayed_free(took, lags)
@pytest.mark.asyncio
async def test_summary_subrequest_honors_project_itpm_otpm(rate_limiter):
"""Regression for #41395: context-management summary subrequests must be
gated on project ITPM/OTPM and charged against those buckets afterwards.
The summary call carries its own litellm_call_id, so it never owns the
parent stash. Without the unreserved IO charge path, combined TPM still
increments while model_per_project_itpm/otpm stay untouched.
Project IO quotas are reservation-style: after ordinary traffic is refused
there may still be residual headroom, so the summary gate passes the
estimated summary size (as apply_compact does) to compare against
limit_remaining.
"""
from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import (
_check_summary_model_rate_limit,
)
from litellm.proxy import proxy_server
handler, _cache = rate_limiter
proxy_server.proxy_logging_obj.max_parallel_request_limiter = handler
model = "gpt-4o-mini"
project = "proj-summary-io"
itpm_limit = 2000
otpm_limit = 10**6
def make_auth(**extra_project_metadata) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-proj-key",
project_id=project,
project_metadata={
"model_itpm_limit": {model: itpm_limit},
"model_otpm_limit": {model: otpm_limit},
**extra_project_metadata,
},
)
def request_data() -> dict:
return {
"model": model,
"messages": [{"role": "user", "content": "x " * 300}],
"litellm_call_id": "parent-call-id",
}
async def drive_until_refused(make_auth_fn) -> tuple[int, str | None]:
allowed = 0
for _ in range(30):
try:
await handler.async_pre_call_hook(
user_api_key_dict=make_auth_fn(),
cache=DualCache(),
data=request_data(),
call_type="completion",
)
allowed += 1
except Exception as e:
return allowed, str(e)
return allowed, None
allowed, refusal = await drive_until_refused(make_auth)
assert allowed >= 1
assert refusal is not None
assert "model_per_project_itpm" in refusal
# Same residual headroom that refused the next ordinary reservation must
# refuse a summary whose estimated input cannot fit.
assert (
await _check_summary_model_rate_limit(
user_api_key_auth=make_auth(),
summary_model=model,
estimated_input_tokens=200,
estimated_output_tokens=1,
)
is False
)
# CONTROL: RPM exhaustion still denies via the same gate.
rpm_handler = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache()))
proxy_server.proxy_logging_obj.max_parallel_request_limiter = rpm_handler
allowed_rpm, refusal_rpm = 0, None
for _ in range(30):
try:
await rpm_handler.async_pre_call_hook(
user_api_key_dict=make_auth(model_rpm_limit={model: 4}),
cache=DualCache(),
data=request_data(),
call_type="completion",
)
allowed_rpm += 1
except Exception as e:
refusal_rpm = str(e)
break
assert allowed_rpm == 4
assert refusal_rpm is not None
assert (
await _check_summary_model_rate_limit(
user_api_key_auth=make_auth(model_rpm_limit={model: 4}),
summary_model=model,
)
is False
)
# Post-call: summary call id must charge ITPM/OTPM like combined TPM.
charging = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache()))
proxy_server.proxy_logging_obj.max_parallel_request_limiter = charging
async def in_one_request_context():
await charging.async_pre_call_hook(
user_api_key_dict=make_auth(),
cache=DualCache(),
data=request_data(),
call_type="completion",
)
response = ModelResponse(
usage=Usage(prompt_tokens=5000, completion_tokens=5000, total_tokens=10000)
)
metadata = {
"user_api_key_project_id": project,
"user_api_key_hash": "sk-proj-key",
"model_group": model,
}
def kwargs_for(call_id: str) -> dict:
return {
"litellm_call_id": call_id,
"model": model,
"litellm_params": {"metadata": metadata},
"standard_logging_object": {"metadata": metadata, "model_group": model},
}
parent_ops = list(
charging._build_io_token_reservation_ops(
kwargs=kwargs_for("parent-call-id"),
response_obj=response,
)
)
summary_ops = list(
charging._build_io_token_reservation_ops(
kwargs=kwargs_for("summary-call-id"),
response_obj=response,
)
)
summary_tpm = charging._build_success_event_pipeline_operations(
kwargs=kwargs_for("summary-call-id"),
response_obj=response,
rate_limit_type=charging.get_rate_limit_type(),
)
return parent_ops, summary_ops, summary_tpm
parent_ops, summary_ops, summary_tpm = await asyncio.create_task(
in_one_request_context()
)
assert parent_ops, "parent call should reconcile reserved ITPM/OTPM"
assert summary_ops, "summary call must charge project ITPM/OTPM without owning the stash"
summary_keys = {op["key"] for op in summary_ops}
assert any("model_per_project_itpm" in key for key in summary_keys)
assert any("model_per_project_otpm" in key for key in summary_keys)
assert any(
"model_per_project:" in op["key"] and op["increment_value"] == 10000
for op in summary_tpm
)
# Unreserved summary path charges full actual usage (no reservation delta).
assert any(
"model_per_project_itpm" in op["key"] and op["increment_value"] == 5000
for op in summary_ops
)
assert any(
"model_per_project_otpm" in op["key"] and op["increment_value"] == 5000
for op in summary_ops
)
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])