From 00d656c3414f8aa6cd2039d23719e0120c797b17 Mon Sep 17 00:00:00 2001 From: aayushbaluni <73417844+aayushbaluni@users.noreply.github.com> Date: Tue, 18 Aug 2026 19:31:45 +0530 Subject: [PATCH] fix(proxy): keep strict token counting when deployment selection fails Strictness was read only from the selected deployment's model_info, so when async_get_available_deployment raised (all deployments cooling down, rate limited) a model marked strict_token_count silently returned a local estimate - exactly the value the flag exists to refuse. Resolve the policy from the router's configuration as well, and treat a model as strict if any of its configured deployments asks for it, since the caller cannot choose which deployment serves them. Raised in review by veria-ai. --- litellm/proxy/proxy_server.py | 52 +++++++++++++- .../test_strict_token_count.py | 67 +++++++++++++++++++ 2 files changed, 116 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5cd59e8fd5a..88f99664be8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11685,6 +11685,48 @@ def _get_provider_token_counter( return None, None, None +def _deployment_wants_strict_count(deployment: Mapping[str, Any]) -> bool: + """Whether one configured deployment asks for exact token counts.""" + info: Final = deployment.get("model_info") + if info is None: + return False + return bool(info.get("strict_token_count", False)) + + +def _is_strict_token_count_model( + llm_router: Router | None, + model_name: str | None, + model_info: ModelMapInfo | None, +) -> bool: + """Whether this model requires an exact token count. + + Prefers the selected deployment's `model_info`, then falls back to the + router's configuration for the requested model. The fallback matters + because deployment selection can fail for reasons unrelated to the + policy, and a strict model must not quietly return an estimate then. + """ + if model_info is not None and bool(model_info.get("strict_token_count", False)): + return True + + if llm_router is None or model_name is None: + return False + + # Strict if any configured deployment for this model asks for it: the + # caller cannot choose which deployment serves them, so the safe reading + # of a mixed configuration is the strict one. + try: + deployments: Final = llm_router.get_model_list(model_name=model_name) + if deployments is None: + return False + return any(_deployment_wants_strict_count(d) for d in deployments) + except (KeyError, AttributeError, TypeError, ValueError): + verbose_proxy_logger.debug( + "litellm.proxy.proxy_server._is_strict_token_count_model(): could not list deployments for %s", + model_name, + ) + return False + + async def _try_provider_token_count( provider_counter: "BaseTokenCounter", custom_llm_provider: str | None, @@ -11806,9 +11848,13 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) # that same behaviour, so a deployment can require an exact count for one # model without giving up the local estimate for every other model. ######################################################### - strict_token_count: bool = litellm.disable_token_counter is True - if strict_token_count is False and model_info is not None: - strict_token_count = bool(model_info.get("strict_token_count", False)) + # Resolved from the router's configuration rather than the selected + # deployment, so the policy still holds when deployment selection fails + # (every deployment cooling down, rate limited, ...). Failing open there + # would hand back the estimate the flag exists to refuse. + strict_token_count: Final = litellm.disable_token_counter is True or _is_strict_token_count_model( + llm_router=llm_router, model_name=request.model, model_info=model_info + ) # Try provider-specific token counting first - only for non-direct requests (from provider endpoints) provider_counter: BaseTokenCounter | None = None diff --git a/tests/proxy_unit_tests/test_strict_token_count.py b/tests/proxy_unit_tests/test_strict_token_count.py index 71e5502b19a..66cfc3d67ae 100644 --- a/tests/proxy_unit_tests/test_strict_token_count.py +++ b/tests/proxy_unit_tests/test_strict_token_count.py @@ -168,6 +168,73 @@ async def test_strict_token_count_does_not_affect_other_models(): setattr(proxy_server, "llm_router", original_router) +@pytest.mark.asyncio +async def test_strict_read_from_the_selected_deployment(): + """Strictness on the *selected deployment* is honoured on its own. + + The router's configured list here does not carry the flag, so only the + deployment returned by selection does. This pins the `model_info` branch + independently of the router-config fallback. + """ + router = _router() # configured without strict_token_count + + original = Router.async_get_available_deployment + + async def _strict_deployment(self, *args, **kwargs): + deployment = await original(self, *args, **kwargs) + deployment = dict(deployment) + deployment["model_info"] = { + **(deployment.get("model_info") or {}), + "strict_token_count": True, + } + return deployment + + with patch.object(Router, "async_get_available_deployment", new=_strict_deployment): + with _unsupported_count_tokens(): + with pytest.raises(ProxyException) as exc_info: + await _count_tokens(router) + + assert exc_info.value.type == "token_counting_error" + assert UNSUPPORTED_MODEL_ERROR in exc_info.value.message + + +@pytest.mark.asyncio +async def test_strict_survives_deployment_selection_failure(): + """A strict model must not fall back to an estimate when routing fails. + + Deployment selection can fail for reasons unrelated to the policy - every + deployment cooling down, rate limited, unhealthy. The selected deployment's + `model_info` is unavailable then, so resolving strictness only from it would + hand back exactly the estimate the flag exists to refuse. + """ + router = _router(model_info={"strict_token_count": True}) + + with patch.object( + Router, + "async_get_available_deployment", + new=AsyncMock(side_effect=Exception("No deployments available - cooldown")), + ): + with pytest.raises(ProxyException) as exc_info: + await _count_tokens(router) + + assert exc_info.value.type == "token_counting_disabled" + + +@pytest.mark.asyncio +async def test_non_strict_model_still_estimates_when_selection_fails(): + """The failure path stays permissive for models that never opted in.""" + router = _router() + + with patch.object( + Router, + "async_get_available_deployment", + new=AsyncMock(side_effect=Exception("No deployments available - cooldown")), + ): + response = await _count_tokens(router) + + assert response.total_tokens > 0 + + @pytest.mark.asyncio async def test_disable_token_counter_still_applies_proxy_wide(): """The existing proxy-wide flag must keep working for unmarked models."""