fix(proxy): seed windows for unreserved summary ITPM/OTPM charges

Unreserved summary charges now open the TPM window with the increment so a later reservation cannot wipe them. Admission estimates use the built summary messages and summary_model

Signed-off-by: sinksilk <785976238@qq.com>
This commit is contained in:
sinksilk 2026-09-25 01:09:56 +08:00
parent a77f233961
commit 723ca75b8a
3 changed files with 234 additions and 163 deletions

View file

@ -1037,6 +1037,25 @@ def _build_summary_messages(
return summary_messages
async def _estimate_summary_input_tokens(
*,
summary_model: str,
summary_messages: Sequence[Mapping[str, object]],
fallback_tokens: int,
) -> int:
try:
return await asyncify(litellm.token_counter)(
model=summary_model,
messages=list(summary_messages),
)
except Exception as e:
verbose_logger.warning(
"compact_20260112: summary token estimate failed; falling back to parent current_tokens: %s",
e,
)
return fallback_tokens
def _is_user_message(msg: object) -> bool:
return isinstance(msg, dict) and msg.get("role") == "user"
@ -1331,10 +1350,18 @@ async def apply_compact_20260112(
applied_edits=[applied],
)
prompt: Final = _build_summary_prompt(edit_spec, tools)
summary_messages: Final = _build_summary_messages(effective_messages, prompt, system=augmented_system)
estimated_summary_input: Final = await _estimate_summary_input_tokens(
summary_model=summary_model,
summary_messages=summary_messages,
fallback_tokens=current_tokens,
)
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_input_tokens=estimated_summary_input,
estimated_output_tokens=_read_summary_max_tokens_setting(),
):
verbose_logger.warning(
@ -1348,8 +1375,6 @@ async def apply_compact_20260112(
applied_edits=[applied],
)
prompt: Final = _build_summary_prompt(edit_spec, tools)
summary_messages: Final = _build_summary_messages(effective_messages, prompt, system=augmented_system)
propagated_metadata: Final = _propagate_metadata(litellm_metadata)
allowed_model_region: Final = getattr(user_api_key_auth, "allowed_model_region", None)

View file

@ -554,6 +554,7 @@ class ReservationAwareIncrementOperation(RedisPipelineIncrementOperation):
window_key: NotRequired[str]
expected_window_start: NotRequired[str]
reservation_backend: NotRequired[Literal["redis", "local"]]
seed_window_if_absent: NotRequired[bool]
class RateLimitResponseWithDescriptors(TypedDict):
@ -4500,18 +4501,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
parent_otel_span: Span | None = None,
) -> None:
for operation in pipeline_operations:
if operation.get("window_key") is None or operation.get("expected_window_start") is None:
await self.internal_usage_cache.async_increment_cache(
key=operation["key"],
value=operation["increment_value"],
litellm_parent_otel_span=parent_otel_span,
ttl=operation["ttl"],
)
await self._apply_one_reservation_aware_token_increment(
operation=operation,
parent_otel_span=parent_otel_span,
)
local_guarded_operations: Final = tuple(
operation
for operation in pipeline_operations
if operation.get("window_key") is not None
and operation.get("expected_window_start") is not None
and not operation.get("seed_window_if_absent")
and operation.get("reservation_backend") == "local"
)
redis_guarded_operations: Final = tuple(
@ -4519,6 +4518,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
for operation in pipeline_operations
if operation.get("window_key") is not None
and operation.get("expected_window_start") is not None
and not operation.get("seed_window_if_absent")
and operation.get("reservation_backend") != "local"
)
if local_guarded_operations:
@ -4532,6 +4532,57 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
parent_otel_span=parent_otel_span,
)
async def _apply_one_reservation_aware_token_increment(
self,
*,
operation: ReservationAwareIncrementOperation,
parent_otel_span: Span | None,
) -> None:
if operation.get("seed_window_if_absent"):
window_key: Final = operation.get("window_key")
if window_key is not None:
await self._seed_rate_limit_window_if_absent(
window_key=window_key,
ttl=operation["ttl"],
parent_otel_span=parent_otel_span,
)
await self.internal_usage_cache.async_increment_cache(
key=operation["key"],
value=operation["increment_value"],
litellm_parent_otel_span=parent_otel_span,
ttl=operation["ttl"],
)
return
if operation.get("window_key") is None or operation.get("expected_window_start") is None:
await self.internal_usage_cache.async_increment_cache(
key=operation["key"],
value=operation["increment_value"],
litellm_parent_otel_span=parent_otel_span,
ttl=operation["ttl"],
)
async def _seed_rate_limit_window_if_absent(
self,
*,
window_key: str,
ttl: int | None,
parent_otel_span: Span | None = None,
) -> None:
"""Open a TPM window around an unreserved charge so a later reservation does not wipe it."""
active_window: Final = await self.internal_usage_cache.async_get_cache(
key=window_key,
litellm_parent_otel_span=parent_otel_span,
)
if active_window is not None:
return
window_ttl: Final = ttl if ttl is not None else self.window_size
await self.internal_usage_cache.async_set_cache(
key=window_key,
value=str(int(self._get_current_time().timestamp())),
ttl=window_ttl,
litellm_parent_otel_span=parent_otel_span,
)
def get_rate_limit_type(self) -> Literal["output", "input", "total"]:
from litellm.proxy.proxy_server import general_settings
@ -4639,7 +4690,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def _collect_project_io_scope_targets(
self,
standard_logging_metadata: Mapping[str, Any],
standard_logging_metadata: Mapping[str, object],
model_group: str | None,
) -> Sequence[tuple[str, str]]:
"""Rebuild project ITPM/OTPM scopes from logging metadata.
@ -4658,16 +4709,38 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
(PROJECT_OTPM_DESCRIPTOR_KEY, descriptor_value),
)
def _build_unreserved_scoped_token_ops(
self,
targets: Sequence[tuple[str, str]],
actual_tokens: int,
) -> tuple[ReservationAwareIncrementOperation, ...]:
if actual_tokens == 0:
return ()
return tuple(
ReservationAwareIncrementOperation(
key=self.create_rate_limit_keys(scope_key, scope_value, "tokens"),
increment_value=actual_tokens,
ttl=self.window_size,
window_key=f"{{{scope_key}:{scope_value}}}:window",
seed_window_if_absent=True,
)
for scope_key, scope_value in targets
)
def _build_unreserved_project_io_token_ops(
self,
kwargs: dict[str, Any],
kwargs: Mapping[str, object],
response_obj: object,
) -> Sequence[RedisPipelineIncrementOperation]:
) -> tuple[ReservationAwareIncrementOperation, ...]:
"""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.
Unreserved charges also seed the TPM window key. A plain counter increment
without a window is wiped when the next ordinary reservation treats a
missing window as expired and resets sibling counters.
"""
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
@ -4709,16 +4782,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
billable_input, completion_tokens, _ = resolved_usage
itpm_targets: Final = tuple(t for t in targets if t[0] == PROJECT_ITPM_DESCRIPTOR_KEY)
otpm_targets: Final = tuple(t for t in targets if t[0] == PROJECT_OTPM_DESCRIPTOR_KEY)
return self._build_reservation_aware_tpm_ops(
return self._build_unreserved_scoped_token_ops(
targets=itpm_targets,
reserved_scopes=frozenset(),
actual_tokens=billable_input,
reserved_tokens=0,
) + self._build_reservation_aware_tpm_ops(
) + self._build_unreserved_scoped_token_ops(
targets=otpm_targets,
reserved_scopes=frozenset(),
actual_tokens=completion_tokens,
reserved_tokens=0,
)
def _build_io_token_reservation_ops(

View file

@ -3710,175 +3710,152 @@ async def test_the_project_itpm_reservation_counts_the_request_off_the_event_loo
@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.
"""Regression for #41395: summary subrequests gate and charge project ITPM/OTPM."""
from typing import Final
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
previous_limiter: Final = getattr(proxy_server.proxy_logging_obj, "max_parallel_request_limiter", None)
proxy_server.proxy_logging_obj.max_parallel_request_limiter = handler
try:
model: Final = "gpt-4o-mini"
project: Final = "proj-summary-io"
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",
def make_auth(**extra_project_metadata) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-proj-key",
project_id=project,
project_metadata={
"model_itpm_limit": {model: 2000},
"model_otpm_limit": {model: 10**6},
**extra_project_metadata,
},
)
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
def request_data() -> dict[str, object]:
return {
"model": model,
"messages": [{"role": "user", "content": "x " * 300}],
"litellm_call_id": "parent-call-id",
}
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",
async def drive_until_refused(
limiter: RateLimitHandler, auth: UserAPIKeyAuth
) -> tuple[int, str | None]:
successes: Final[list[bool]] = []
for _ in range(30):
try:
await limiter.async_pre_call_hook(
user_api_key_dict=auth,
cache=DualCache(),
data=request_data(),
call_type="completion",
)
successes.append(True)
except Exception as e:
return len(successes), str(e)
return len(successes), None
allowed, refusal = await drive_until_refused(handler, make_auth())
assert allowed >= 1
assert refusal is not None
assert "model_per_project_itpm" in refusal
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
)
response = ModelResponse(
usage=Usage(prompt_tokens=5000, completion_tokens=5000, total_tokens=10000)
assert (
await _check_summary_model_rate_limit(
user_api_key_auth=make_auth(
model_itpm_limit={model: 10**6},
model_otpm_limit={model: 5},
),
summary_model=model,
estimated_input_tokens=1,
estimated_output_tokens=20,
)
is False
)
metadata = {
rpm_handler: Final = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache()))
proxy_server.proxy_logging_obj.max_parallel_request_limiter = rpm_handler
allowed_rpm, refusal_rpm = await drive_until_refused(
rpm_handler, make_auth(model_rpm_limit={model: 4})
)
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
)
charging: Final = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache()))
proxy_server.proxy_logging_obj.max_parallel_request_limiter = charging
summary_response: Final = ModelResponse(
usage=Usage(prompt_tokens=60, completion_tokens=40, total_tokens=100)
)
metadata: Final = {
"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,
await charging.async_log_success_event(
kwargs={
"litellm_call_id": "summary-call-id",
"model": model,
"litellm_params": {"metadata": metadata},
"standard_logging_object": {"metadata": metadata, "model_group": model},
}
},
response_obj=summary_response,
start_time=datetime.now(),
end_time=datetime.now(),
)
parent_ops = list(
charging._build_io_token_reservation_ops(
kwargs=kwargs_for("parent-call-id"),
response_obj=response,
)
itpm_key: Final = charging.create_rate_limit_keys(
PROJECT_ITPM_DESCRIPTOR_KEY, f"{project}:{model}", "tokens"
)
summary_ops = list(
charging._build_io_token_reservation_ops(
kwargs=kwargs_for("summary-call-id"),
response_obj=response,
)
otpm_key: Final = charging.create_rate_limit_keys(
PROJECT_OTPM_DESCRIPTOR_KEY, f"{project}:{model}", "tokens"
)
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
itpm_window: Final = f"{{{PROJECT_ITPM_DESCRIPTOR_KEY}:{project}:{model}}}:window"
otpm_window: Final = f"{{{PROJECT_OTPM_DESCRIPTOR_KEY}:{project}:{model}}}:window"
dual: Final = charging.internal_usage_cache.dual_cache
assert int(await dual.async_get_cache(key=itpm_key) or 0) == 60
assert int(await dual.async_get_cache(key=otpm_key) or 0) == 40
assert await dual.async_get_cache(key=itpm_window) is not None
assert await dual.async_get_cache(key=otpm_window) is not None
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
)
await charging.async_pre_call_hook(
user_api_key_dict=make_auth(
model_itpm_limit={model: 10**6},
model_otpm_limit={model: 10**6},
),
cache=DualCache(),
data={
"model": model,
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 10,
"litellm_call_id": "follow-up-call",
},
call_type="completion",
)
assert int(await dual.async_get_cache(key=itpm_key) or 0) >= 60
finally:
proxy_server.proxy_logging_obj.max_parallel_request_limiter = previous_limiter
if __name__ == "__main__":