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 62826865894..69089d20cba 100644 --- a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py @@ -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( @@ -996,6 +1037,39 @@ def _build_summary_messages( return summary_messages +def _count_summary_message_tokens( + model: str, + messages: Sequence[Mapping[str, object]], +) -> int: + """Count summary-call tokens through a fully annotated wrapper. + + ``litellm.token_counter``'s own signature is partially unknown (bare + ``Sequence`` / ``dict`` parameters). Passing that function into + ``asyncify`` is a new ``reportUnknownArgumentType``. This wrapper's + signature is fully known, so the asyncify boundary stays typed. + """ + return litellm.token_counter(model=model, messages=messages) + + +async def _estimate_summary_input_tokens( + *, + summary_model: str, + summary_messages: Sequence[Mapping[str, object]], + fallback_tokens: int, +) -> int: + try: + return await asyncify(_count_summary_message_tokens)( + model=summary_model, + messages=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" @@ -1290,9 +1364,19 @@ 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=estimated_summary_input, + estimated_output_tokens=_read_summary_max_tokens_setting(), ): verbose_logger.warning( "compact_20260112: caller over rate limit for summary_model=%s; skipping summary call", @@ -1305,8 +1389,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 2bfaf57f0fc..3e681c68a3d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -24,6 +24,7 @@ from typing import ( Protocol, TypeAlias, TypedDict, + cast, ) from fastapi import HTTPException @@ -555,6 +556,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): @@ -740,6 +742,19 @@ def _call_id_from_callback_kwargs(kwargs: object) -> str | None: return call_id if isinstance(call_id, str) else None +def _as_str_object_dict(value: object) -> dict[str, object] | None: # mutable-ok: model-group helper requires a dict + """Return a ``dict[str, object]`` view of an untyped callback payload. + + ``isinstance(..., dict)`` narrows to ``dict[Unknown, Unknown]``. Keep an + ``object`` alias from before that narrowing and cast that alias, so the + cast argument stays a known type. + """ + raw: Final[object] = value + if not isinstance(value, dict): + return None + return cast("dict[str, object]", raw) # cast-ok: success-callback payload is an untyped dict + + def _parse_output_cap_value(raw_value: object) -> int | None: if isinstance(raw_value, bool) or not isinstance(raw_value, (int, float, str)): return None @@ -4510,18 +4525,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( @@ -4529,6 +4542,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: @@ -4542,6 +4556,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 @@ -4647,6 +4712,111 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return 0, 0, False + def _collect_project_io_scope_targets( + self, + standard_logging_metadata: Mapping[str, object], + model_group: str | None, + ) -> Sequence[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 ( + (PROJECT_ITPM_DESCRIPTOR_KEY, descriptor_value), + (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: object, + response_obj: object, + ) -> 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, + ) + + callback_kwargs: Final = _as_str_object_dict(kwargs) + if callback_kwargs is None: + return () + logging_map: Final = _as_str_object_dict(callback_kwargs.get("standard_logging_object")) + if logging_map is None: + return () + standard_logging_metadata: Final = _as_str_object_dict(logging_map.get("metadata")) + if standard_logging_metadata is None: + return () + + logged_group: Final = logging_map.get("model_group") + model_group: Final = get_model_group_from_litellm_kwargs(callback_kwargs) or ( + logged_group if isinstance(logged_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_object: Final = callback_kwargs.get("combined_usage_object") + combined_usage: Final = self._resolve_io_token_reconcile_usage(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(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 = 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_unreserved_scoped_token_ops( + targets=itpm_targets, + actual_tokens=billable_input, + ) + self._build_unreserved_scoped_token_ops( + targets=otpm_targets, + actual_tokens=completion_tokens, + ) + def _build_io_token_reservation_ops( self, kwargs: object, @@ -4659,12 +4829,17 @@ 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. """ + callback_kwargs: Final[object] = kwargs 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(callback_kwargs, response_obj) itpm_reserved: Final = stash.itpm_reserved_tokens otpm_reserved: Final = stash.otpm_reserved_tokens diff --git a/tests/unit/proxy/hooks/test_tpm_concurrent.py b/tests/unit/proxy/hooks/test_tpm_concurrent.py index 42c1f489bdd..f4da0d00223 100644 --- a/tests/unit/proxy/hooks/test_tpm_concurrent.py +++ b/tests/unit/proxy/hooks/test_tpm_concurrent.py @@ -30,6 +30,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _as_str_object_dict, _call_id_from_callback_kwargs, _request_stash, get_or_create_request_stash, @@ -3708,5 +3709,325 @@ 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: summary subrequests gate and charge project ITPM/OTPM.""" + from typing import Final + + from litellm.llms.anthropic.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) + previous_hook: Final = proxy_server.proxy_logging_obj.proxy_hook_mapping.get( + "parallel_request_limiter" + ) + + def install_limiter(limiter: RateLimitHandler) -> None: + # Summary gate resolves the limiter via get_proxy_hook(), not the + # legacy max_parallel_request_limiter attribute alone. + proxy_server.proxy_logging_obj.max_parallel_request_limiter = limiter + proxy_server.proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + + install_limiter(handler) + try: + model: Final = "gpt-4o-mini" + project: Final = "proj-summary-io" + + 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, + }, + ) + + def request_data() -> dict[str, object]: + return { + "model": model, + "messages": [{"role": "user", "content": "x " * 300}], + "litellm_call_id": "parent-call-id", + } + + 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 + ) + + 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 + ) + + rpm_handler: Final = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache())) + install_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())) + install_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, + } + 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(), + ) + + itpm_key: Final = charging.create_rate_limit_keys( + PROJECT_ITPM_DESCRIPTOR_KEY, f"{project}:{model}", "tokens" + ) + otpm_key: Final = charging.create_rate_limit_keys( + PROJECT_OTPM_DESCRIPTOR_KEY, f"{project}:{model}", "tokens" + ) + 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 + + 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 previous_hook is None: + proxy_server.proxy_logging_obj.proxy_hook_mapping.pop( + "parallel_request_limiter", None + ) + else: + proxy_server.proxy_logging_obj.proxy_hook_mapping[ + "parallel_request_limiter" + ] = previous_hook + + +@pytest.mark.asyncio +async def test_unreserved_project_io_covers_empty_and_fallback_usage(rate_limiter): + """Cover the unreserved ITPM/OTPM branches the happy-path summary test skips.""" + from typing import Final + + handler, _cache = rate_limiter + assert _as_str_object_dict("nope") is None + echoed: Final = _as_str_object_dict({"a": 1}) + assert echoed is not None + assert echoed["a"] == 1 + + assert handler._build_unreserved_project_io_token_ops("nope", {}) == () + assert handler._build_unreserved_project_io_token_ops({}, {}) == () + assert ( + handler._build_unreserved_project_io_token_ops( + {"standard_logging_object": "x"}, + {}, + ) + == () + ) + assert ( + handler._build_unreserved_project_io_token_ops( + {"standard_logging_object": {"metadata": "x"}}, + {}, + ) + == () + ) + assert ( + handler._build_unreserved_project_io_token_ops( + { + "standard_logging_object": { + "metadata": {"user_api_key_project_id": "proj"}, + "model_group": 1, + } + }, + {}, + ) + == () + ) + + combined_ops: Final = handler._build_unreserved_project_io_token_ops( + { + "standard_logging_object": { + "metadata": {"user_api_key_project_id": "proj"}, + "model_group": "gpt-4o-mini", + }, + "combined_usage_object": {"prompt_tokens": 5, "completion_tokens": 0}, + }, + {}, + ) + assert len(combined_ops) == 1 + assert combined_ops[0]["increment_value"] == 5 + + aggregate_ops: Final = handler._build_unreserved_project_io_token_ops( + { + "standard_logging_object": { + "metadata": {"user_api_key_project_id": "proj"}, + "model_group": "gpt-4o-mini", + }, + }, + {"total_tokens": 9}, + ) + assert len(aggregate_ops) == 2 + assert {op["increment_value"] for op in aggregate_ops} == {9} + + await handler._seed_rate_limit_window_if_absent(window_key="already-open", ttl=None) + await handler._seed_rate_limit_window_if_absent(window_key="already-open", ttl=None) + await handler._apply_one_reservation_aware_token_increment( + operation={ + "key": "plain-counter", + "increment_value": 3, + "ttl": 60, + }, + parent_otel_span=None, + ) + assert int(await handler.internal_usage_cache.async_get_cache(key="plain-counter", litellm_parent_otel_span=None) or 0) == 3 + + +@pytest.mark.asyncio +async def test_summary_token_estimate_uses_counter_or_falls_back(monkeypatch): + from typing import Final + + from litellm.llms.anthropic.pass_through.context_management.editors.compact import ( + _check_summary_model_rate_limit, + _estimate_summary_input_tokens, + ) + from litellm.proxy import proxy_server + + messages: Final = ({"role": "user", "content": "hi"},) + + def fail_counter(**_kwargs: object) -> int: + raise RuntimeError("counter down") + + monkeypatch.setattr("litellm.token_counter", fail_counter) + assert ( + await _estimate_summary_input_tokens( + summary_model="gpt-4o-mini", + summary_messages=messages, + fallback_tokens=42, + ) + == 42 + ) + + def fixed_counter(**_kwargs: object) -> int: + return 17 + + monkeypatch.setattr("litellm.token_counter", fixed_counter) + assert ( + await _estimate_summary_input_tokens( + summary_model="gpt-4o-mini", + summary_messages=messages, + fallback_tokens=42, + ) + == 17 + ) + + handler: Final = RateLimitHandler(internal_usage_cache=InternalUsageCache(DualCache())) + + async def junk_statuses(self, descriptors, parent_otel_span=None, read_only=False, **_kwargs): + return { + "overall_code": "OK", + "statuses": [ + "not-a-status", + {"descriptor_key": "model_per_project_itpm", "limit_remaining": "lots"}, + {"descriptor_key": "model_per_project_itpm", "limit_remaining": 1000}, + ], + } + + monkeypatch.setattr(RateLimitHandler, "should_rate_limit", junk_statuses) + previous_hook: Final = proxy_server.proxy_logging_obj.proxy_hook_mapping.get("parallel_request_limiter") + proxy_server.proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = handler + try: + assert ( + await _check_summary_model_rate_limit( + user_api_key_auth=UserAPIKeyAuth( + api_key="sk-proj-key", + project_id="proj-summary-io", + project_metadata={"model_itpm_limit": {"gpt-4o-mini": 2000}}, + ), + summary_model="gpt-4o-mini", + estimated_input_tokens=10, + estimated_output_tokens=1, + ) + is True + ) + finally: + if previous_hook is None: + proxy_server.proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None) + else: + proxy_server.proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = previous_hook + + if __name__ == "__main__": pytest.main([__file__, "-v", "-s"])