From 723ca75b8aac0f2d2f9978f21195186d2ffbc505 Mon Sep 17 00:00:00 2001 From: sinksilk <785976238@qq.com> Date: Fri, 25 Sep 2026 01:09:56 +0800 Subject: [PATCH] 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> --- .../context_management/editors/compact.py | 31 +- .../hooks/parallel_request_limiter_v3.py | 101 +++++-- tests/unit/proxy/hooks/test_tpm_concurrent.py | 265 ++++++++---------- 3 files changed, 234 insertions(+), 163 deletions(-) diff --git a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py index 59d85528d9e..4be645b8f72 100644 --- a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py @@ -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) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 1e08403083a..f8ad7fd2f7e 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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( diff --git a/tests/unit/proxy/hooks/test_tpm_concurrent.py b/tests/unit/proxy/hooks/test_tpm_concurrent.py index 91aad818f52..83f00211895 100644 --- a/tests/unit/proxy/hooks/test_tpm_concurrent.py +++ b/tests/unit/proxy/hooks/test_tpm_concurrent.py @@ -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__":