mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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.
This commit is contained in:
parent
7eb6dc4c74
commit
00d656c341
2 changed files with 116 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue