Merge pull request #41508 from BerriAI/litellm_1789600151_discover_context_limits

feat(router): discover token limits for hosted OpenAI-compatible models
This commit is contained in:
tin-berri 2026-09-16 20:29:57 -07:00 committed by GitHub
commit d18e06f736
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 856 additions and 17 deletions

View file

@ -0,0 +1,92 @@
import hashlib
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import Annotated, Final, TypeAlias
import httpx
from pydantic import BaseModel, BeforeValidator, ConfigDict
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.utils import _add_path_to_api_base # pyright: ignore[reportPrivateUsage] # shared provider URL helper
MODEL_INFO_REFRESH_SECONDS: Final = 300
MODEL_INFO_REFRESH_CONCURRENCY: Final = 8
MODEL_INFO_DISCOVERY_PROVIDERS: Final = frozenset({"hosted_vllm", "openai", "text-completion-openai", "openai_like"})
_EMPTY_LIMITS: Final[Mapping[str, int]] = MappingProxyType({})
def _positive_limit(value: object) -> int | None:
return value if isinstance(value, int) and not isinstance(value, bool) and value > 0 else None
_TokenLimit: TypeAlias = Annotated[int | None, BeforeValidator(_positive_limit)]
class _ModelCard(BaseModel):
model_config = ConfigDict(frozen=True)
id: str
max_model_len: _TokenLimit = None
context_length: _TokenLimit = None
max_input_tokens: _TokenLimit = None
max_output_tokens: _TokenLimit = None
def token_limits(self) -> Mapping[str, int]:
context: Final = self.max_model_len or self.context_length
input_limit: Final = self.max_input_tokens or context
output_limit: Final = self.max_output_tokens or context
return MappingProxyType(
{
key: value
for key, value in (
("max_tokens", context),
("max_input_tokens", min(input_limit, context) if input_limit and context else input_limit),
("max_output_tokens", min(output_limit, context) if output_limit and context else output_limit),
)
if value is not None
}
)
class _ModelList(BaseModel):
model_config = ConfigDict(frozen=True)
data: tuple[_ModelCard, ...] = ()
async def get_openai_compatible_model_info(
*,
model: str,
api_base: str,
headers: Mapping[str, str],
client: AsyncHTTPHandler,
cache: InMemoryCache,
) -> Mapping[str, int]:
url: Final = _add_path_to_api_base(api_base, "/v1/models")
cache_key: Final = (
"upstream_model_info:" + hashlib.sha256(json.dumps((url, sorted(headers.items()))).encode()).hexdigest()
)
cached: Final[object] = cache.get_cache(cache_key)
if isinstance(cached, _ModelList):
return next((card.token_limits() for card in cached.data if card.id == model), _EMPTY_LIMITS)
try:
response: Final = await client.get(
url=url,
headers=dict(headers), # mutable-ok: AsyncHTTPHandler requires a concrete dict
timeout=httpx.Timeout(5.0),
follow_redirects=False,
max_response_bytes=2 * 1024 * 1024,
)
response.raise_for_status()
models: Final = _ModelList.model_validate_json(response.content)
except Exception: # noqa: BLE001 # optional upstream metadata must not interrupt proxy refresh
verbose_logger.debug("Could not discover upstream model token limits")
cache.set_cache(cache_key, _ModelList(), ttl=60)
return _EMPTY_LIMITS
cache.set_cache(cache_key, models, ttl=MODEL_INFO_REFRESH_SECONDS)
return next((card.token_limits() for card in models.data if card.id == model), _EMPTY_LIMITS)

View file

@ -305,6 +305,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
mask_sensitive_keys,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
from litellm.proxy._types import *
@ -1384,9 +1385,27 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
## Initialize shared aiohttp session for connection reuse
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
model_info_scheduler: Final = scheduler if scheduler is not None else AsyncIOScheduler()
model_info_scheduler.add_job(
ProxyStartupEvent.refresh_model_info,
"interval",
seconds=MODEL_INFO_REFRESH_SECONDS,
id="refresh_model_info",
next_run_time=datetime.now(timezone.utc),
max_instances=1,
replace_existing=True,
)
if not model_info_scheduler.running:
model_info_scheduler.start()
# End of startup event
yield
if model_info_scheduler.running:
model_info_scheduler.remove_job("refresh_model_info")
if model_info_scheduler is not scheduler:
model_info_scheduler.shutdown(wait=False)
# Shutdown event - drain in-flight requests before tearing down dependencies
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
GracefulShutdownManager.start_shutdown()
@ -9337,6 +9356,11 @@ def giveup(e):
class ProxyStartupEvent:
@staticmethod
async def refresh_model_info() -> None:
if llm_router is not None:
await llm_router.arefresh_model_info()
@staticmethod
def _warn_budget_without_db(max_budget: float | None, prisma_client: PrismaClient | None) -> None:
if prisma_client is not None or not max_budget or max_budget <= 0:
@ -13593,8 +13617,11 @@ def _enrich_model_info_with_litellm_data(
litellm_model_info = litellm.get_model_info(model=litellm_model, custom_llm_provider=split_model[0])
except Exception:
litellm_model_info = {}
for k, v in litellm_model_info.items():
if k not in model_info:
discovered_model_info: Final = (
llm_router.get_discovered_model_info(model_info.get("id")) if llm_router is not None else MappingProxyType({})
)
for k, v in MappingProxyType({**litellm_model_info, **discovered_model_info}).items():
if k not in model_info or (model_info[k] is None and k in discovered_model_info):
model_info[k] = v
model["model_info"] = model_info
# don't return the api key / vertex credentials
@ -15059,8 +15086,11 @@ def _get_proxy_model_info(model: dict) -> dict:
litellm_model_info = litellm.get_model_info(model=litellm_model, custom_llm_provider=split_model[0])
except Exception:
litellm_model_info = {}
for k, v in litellm_model_info.items():
if k not in model_info:
discovered_model_info: Final = (
llm_router.get_discovered_model_info(model_info.get("id")) if llm_router is not None else MappingProxyType({})
)
for k, v in MappingProxyType({**litellm_model_info, **discovered_model_info}).items():
if k not in model_info or (model_info[k] is None and k in discovered_model_info):
model_info[k] = v
model["model_info"] = model_info
# don't return the llm credentials

View file

@ -109,7 +109,14 @@ from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
vector_store_request_metadata,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.llms.openai_like.model_info import (
MODEL_INFO_DISCOVERY_PROVIDERS,
MODEL_INFO_REFRESH_CONCURRENCY,
MODEL_INFO_REFRESH_SECONDS,
get_openai_compatible_model_info,
)
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
@ -242,6 +249,7 @@ from litellm.types.router import (
Deployment,
DeploymentModelListingInfo,
DeploymentTypedDict,
DiscoveredDeploymentModelInfo,
FallbackAccessCheck,
FallbackBudgetCheck,
GuardrailTypedDict,
@ -973,6 +981,10 @@ class Router:
self.cached_deployment_model_info = lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE)(
self.get_deployment_model_info
)
self._discovered_model_info_cache: InMemoryCache = InMemoryCache(
max_size_in_memory=max(len(model_list or ()), 1),
default_ttl=2 * MODEL_INFO_REFRESH_SECONDS,
)
self._routing_group_rows: tuple[DeploymentTypedDict, ...] | None = None
self._init_routing_groups(None)
self._provider_unresolved_deployments: tuple[Callable[[], Deployment | None], ...] = ()
@ -9492,6 +9504,7 @@ class Router:
def set_model_list(self, model_list: list):
original_model_list: Final = copy.deepcopy(model_list)
self._discovered_model_info_cache.flush_cache()
self.model_list = []
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
@ -9786,6 +9799,7 @@ class Router:
- model_id: str - the id of the deployment that was removed
- removal_idx: int - the index where the deployment was removed from model_list
"""
self._discovered_model_info_cache.delete_cache(model_id)
# Update indices for all models after the removed one
for deployment_id, idx in self.model_id_to_deployment_index_map.items():
if idx > removal_idx:
@ -10316,11 +10330,85 @@ class Router:
return None
return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable
async def arefresh_model_info(self, *, client: AsyncHTTPHandler | None = None) -> None:
"""Refresh token limits advertised by configured OpenAI-compatible deployments."""
deployments: Final = iter(tuple(self.model_list))
async def refresh_worker() -> None:
for raw_deployment in deployments:
try:
await self._arefresh_deployment_model_info(raw_deployment, client=client)
except Exception: # noqa: BLE001 # one invalid deployment must not prevent refreshing the others
verbose_router_logger.debug("Could not refresh deployment model info")
await asyncio.gather(*(refresh_worker() for _ in range(MODEL_INFO_REFRESH_CONCURRENCY)))
self._invalidate_model_group_info_cache()
async def _arefresh_deployment_model_info(
self, raw_deployment: Mapping[str, object], *, client: AsyncHTTPHandler | None
) -> None:
deployment: Final = Deployment.model_validate(raw_deployment)
params: Final = LiteLLM_Params.model_validate(
MappingProxyType(
{
**deployment.litellm_params.model_dump(exclude_none=True),
**(
self.get_deployment_credentials_with_provider(deployment.model_info.id or "")
or MappingProxyType({})
),
}
)
)
model, provider, dynamic_api_key, api_base = litellm.get_llm_provider(model=params.model, litellm_params=params)
if provider not in MODEL_INFO_DISCOVERY_PROVIDERS:
return
if api_base is None or "*" in model or params.get("use_clientside_credentials"):
return
api_key: Final = params.api_key or dynamic_api_key
headers: Final = TypeAdapter(Mapping[str, str]).validate_python(
params.get("extra_headers") or params.get("headers") or MappingProxyType({})
)
auth_headers: Final = (
MappingProxyType({"authorization": f"Bearer {api_key}"}) if api_key else MappingProxyType({})
)
limits: Final = await get_openai_compatible_model_info(
model=model,
api_base=api_base,
headers=MappingProxyType(
{
**auth_headers,
**MappingProxyType({key.lower(): value for key, value in headers.items()}),
}
),
client=client or get_async_httpx_client(llm_provider=LlmProviders.OPENAI),
cache=self.cache.in_memory_cache,
)
model_id: Final = deployment.model_info.id
if not limits or model_id is None or self.get_model_info(model_id) is not raw_deployment:
return
self._discovered_model_info_cache.max_size_in_memory = max(len(self.model_list), 1)
self._discovered_model_info_cache.delete_cache(model_id)
self._discovered_model_info_cache.set_cache(
model_id, DiscoveredDeploymentModelInfo(deployment=raw_deployment, limits=limits)
)
self._invalidate_model_group_info_cache()
def get_discovered_model_info(self, model_id: str | None) -> Mapping[str, int]:
cached: Final[object] = self._discovered_model_info_cache.get_cache(model_id)
if (
model_id is not None
and isinstance(cached, DiscoveredDeploymentModelInfo)
and cached.deployment is self.get_model_info(model_id)
):
configured: Final = TypeAdapter(Mapping[str, object]).validate_python(cached.deployment["model_info"])
return MappingProxyType({key: value for key, value in cached.limits.items() if configured.get(key) is None})
return MappingProxyType({})
def get_model_listing_info(self, model_name: str) -> DeploymentModelListingInfo | None:
"""
Return what the concrete deployments behind model_name contribute to its
/v1/models entry: the cost-map keys for their underlying models, plus the widest
token limits explicitly configured in their model_info. Resolved via O(1) index
configured or discovered token limits. Resolved via O(1) index
lookup.
Returns None for wildcard-expanded or unknown names, where the listed name is the
@ -10340,7 +10428,21 @@ class Router:
return None
deployments: Final = tuple(self.model_list[index] for index in indices)
model_infos: Final = tuple(deployment.get("model_info") or MappingProxyType({}) for deployment in deployments)
model_infos: Final = tuple(
MappingProxyType(
{
**self.get_discovered_model_info((deployment.get("model_info") or MappingProxyType({})).get("id")),
**MappingProxyType(
{
k: v
for k, v in (deployment.get("model_info") or MappingProxyType({})).items()
if v is not None
}
),
}
)
for deployment in deployments
)
params: Final = tuple(deployment.get("litellm_params") or MappingProxyType({}) for deployment in deployments)
# base_model resolution mirrors get_router_model_info: unset or blank means the
# deployment's own model name is the cost-map key.
@ -10372,8 +10474,8 @@ class Router:
def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]":
"""
Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete
deployment's model_info for model_name, via O(1) index lookup.
Return (max_input_tokens, max_output_tokens) configured or discovered for a concrete
deployment of model_name, via O(1) index lookup.
Returns (None, None) for wildcard-expanded or unknown names, and treats a
malformed configured value as absent rather than failing the caller.
@ -10386,7 +10488,12 @@ class Router:
if deployment is None:
return (None, None)
model_info: Final = deployment.model_info
model_info: Final = MappingProxyType(
{
**self.get_discovered_model_info(deployment.model_info.id),
**deployment.model_info.model_dump(exclude_none=True),
}
)
return (
coerce_token_limit(model_info.get("max_input_tokens")),
coerce_token_limit(model_info.get("max_output_tokens")),
@ -10651,11 +10758,13 @@ class Router:
# get_model_info() hands back an lru_cache'd dict, so merge into a copy; unset
# values are skipped or Deployment's None pricing defaults would erase the map's
merged_model_info: Final = copy.deepcopy(model_info)
if user_model_info:
for key, value in user_model_info.items():
if value is not None:
merged_model_info[key] = value
merged_model_info: Final[ModelMapInfo] = {
**copy.deepcopy(model_info),
**self.get_discovered_model_info((deployment.get("model_info") or {}).get("id")),
**MappingProxyType(
{key: value for key, value in (user_model_info or MappingProxyType({})).items() if value is not None}
),
}
return merged_model_info
@ -10702,7 +10811,14 @@ class Router:
litellm_model_name_model_info: ModelInfo | None = None
try:
custom_model_info = copy.deepcopy(litellm.model_cost.get(model_id))
custom_model_info = (
{ # mutable-ok: the legacy model-info merge updates this private copy
**copy.deepcopy(litellm.model_cost.get(model_id) or MappingProxyType({})),
**self.get_discovered_model_info(model_id),
}
if model_id in litellm.model_cost
else None
)
except Exception:
pass

View file

@ -623,6 +623,12 @@ class Deployment(BaseModel):
setattr(self, key, value)
@dataclass(frozen=True, slots=True)
class DiscoveredDeploymentModelInfo:
deployment: Mapping[str, object]
limits: Mapping[str, int]
@dataclass(frozen=True, slots=True)
class DeploymentModelListingInfo:
"""What the deployments behind a model name contribute to its OpenAI-compatible listing entry.

View file

@ -0,0 +1,126 @@
from collections.abc import Mapping
from typing import Final
from unittest.mock import Mock
import httpx
import pytest
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.openai_like.model_info import (
MODEL_INFO_REFRESH_SECONDS,
get_openai_compatible_model_info,
)
@pytest.mark.parametrize(
("card", "expected"),
(
({"max_model_len": 8192}, {"max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192}),
(
{"context_length": 4096, "max_output_tokens": 1024},
{"max_tokens": 4096, "max_input_tokens": 4096, "max_output_tokens": 1024},
),
(
{"max_model_len": 4096, "max_input_tokens": 2048, "max_output_tokens": 8192},
{"max_tokens": 4096, "max_input_tokens": 2048, "max_output_tokens": 4096},
),
({"max_input_tokens": 2048}, {"max_input_tokens": 2048}),
({"max_output_tokens": 1024}, {"max_output_tokens": 1024}),
({"max_model_len": True, "max_output_tokens": -1}, {}),
({"max_model_len": "8192", "max_input_tokens": 0, "max_output_tokens": 1.5}, {}),
({}, {}),
),
)
async def test_discovers_only_valid_advertised_limits(card: Mapping[str, object], expected: Mapping[str, int]) -> None:
def respond(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/tenant/v1/models"
assert request.headers["authorization"] == "Bearer local-key"
return httpx.Response(200, json={"data": [{"id": "org/model", **card}]})
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler.client = client
cache: Final = InMemoryCache()
result: Final = await get_openai_compatible_model_info(
model="org/model",
api_base="https://backend.test/tenant/v1/",
headers={"Authorization": "Bearer local-key"},
client=handler,
cache=cache,
)
assert result == expected
assert (
await get_openai_compatible_model_info(
model="missing",
api_base="https://backend.test/tenant/v1/",
headers={"Authorization": "Bearer local-key"},
client=handler,
cache=cache,
)
== {}
)
async def test_cache_is_scoped_to_endpoint_and_authentication_and_expires() -> None:
clock: Final = Mock(return_value=0)
responder: Final = Mock(
side_effect=(
httpx.Response(
200, json={"data": [{"id": "first", "max_model_len": 1024}, {"id": "second", "max_model_len": 2048}]}
),
httpx.Response(200, json={"data": [{"id": "first", "max_model_len": 4096}]}),
httpx.Response(200, json={"data": [{"id": "first", "max_model_len": 8192}]}),
httpx.Response(200, json={"data": [{"id": "first", "max_model_len": 16384}]}),
)
)
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(responder)) as client:
handler.client = client
cache: Final = InMemoryCache(clock=clock)
async def lookup(model: str = "first", host: str = "one.test", key: str = "one") -> Mapping[str, int]:
return await get_openai_compatible_model_info(
model=model, api_base=f"https://{host}", headers={"Authorization": key}, client=handler, cache=cache
)
assert (await lookup())["max_input_tokens"] == 1024
assert (await lookup("second"))["max_input_tokens"] == 2048
assert responder.call_count == 1
assert (await lookup(key="two"))["max_input_tokens"] == 4096
assert (await lookup(host="two.test"))["max_input_tokens"] == 8192
clock.return_value = MODEL_INFO_REFRESH_SECONDS + 1
assert (await lookup())["max_input_tokens"] == 16384
assert responder.call_count == 4
@pytest.mark.parametrize(
"response",
(
httpx.Response(404),
httpx.Response(401),
httpx.Response(302, headers={"location": "https://elsewhere.test"}),
httpx.Response(200, content=b"not json"),
httpx.Response(200, json={"data": None}),
httpx.ReadTimeout("backend unavailable"),
),
)
async def test_unavailable_metadata_is_best_effort_and_negative_cached(
response: httpx.Response | Exception,
) -> None:
responder: Final = Mock(side_effect=response if isinstance(response, Exception) else None, return_value=response)
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(responder), follow_redirects=True) as client:
handler.client = client
cache: Final = InMemoryCache()
for _ in range(2):
assert (
await get_openai_compatible_model_info(
model="model", api_base="https://backend.test", headers={}, client=handler, cache=cache
)
== {}
)
assert responder.call_count == 1

View file

@ -9,14 +9,159 @@ Pins (PR2):
from __future__ import annotations
import copy
from collections.abc import Callable
from contextlib import AbstractContextManager
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi.testclient import TestClient
import litellm
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.utils import _invalidate_model_cost_lowercase_map
from .conftest import normalize # type: ignore[import-not-found]
@pytest.mark.parametrize(
("backend_model", "base_model"),
(
("azure/hosted-model", "fallback-model"),
("openai/org/fallback-model", None),
("openai/hosted-model", "fallback-model"),
("openai/fallback-model", "unknown-base-model"),
),
)
@pytest.mark.parametrize("advertised_limit", (None, 2048))
async def test_discovery_preserves_model_info_fallbacks(
backend_model: str, base_model: str | None, advertised_limit: int | None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
router: Final = litellm.Router(
model_list=[
{
"model_name": "local",
"litellm_params": {
"model": backend_model,
"api_base": "https://fallback.test/v1",
"api_key": "local-key",
},
"model_info": {"id": "fallback-deployment", "base_model": base_model, "max_output_tokens": 333},
}
]
)
builtin: Final = {
"litellm_provider": "openai",
"mode": "chat",
"max_input_tokens": 7000,
"max_output_tokens": 2000,
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
}
monkeypatch.setattr(
litellm,
"model_cost",
{
"fallback-model": builtin,
"openai/fallback-model": builtin,
"fallback-deployment": {"litellm_provider": "openai", "mode": "chat"},
},
)
_invalidate_model_cost_lowercase_map()
monkeypatch.setattr(proxy_server, "llm_router", router)
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(
transport=httpx.MockTransport(
lambda request: httpx.Response(
200,
json={
"data": [
{
"id": backend_model.split("/", 1)[1],
"max_model_len": advertised_limit,
}
]
},
)
)
) as client:
handler.client = client
await router.arefresh_model_info(client=handler)
deployment: Final = {
**router.model_list[0],
"model_info": {**router.model_list[0]["model_info"], "mode": None},
}
enriched_models: Final = (
proxy_server._get_proxy_model_info(copy.deepcopy(deployment)),
proxy_server._enrich_model_info_with_litellm_data(copy.deepcopy(deployment), llm_router=router),
)
expected_input: Final = (
advertised_limit
if advertised_limit is not None and backend_model.startswith("openai/")
else builtin["max_input_tokens"]
)
for enriched in enriched_models:
info: Final = enriched["model_info"]
assert info.get("max_input_tokens") == expected_input
assert info["max_output_tokens"] == 333
assert info["input_cost_per_token"] == builtin["input_cost_per_token"]
assert info["output_cost_per_token"] == builtin["output_cost_per_token"]
assert info["mode"] is None
_invalidate_model_cost_lowercase_map()
async def test_upstream_limits_reach_model_info_routes(
client: TestClient,
auth_as: Callable[[], AbstractContextManager[object]],
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
router: Final = litellm.Router(
model_list=[
{
"model_name": "local",
"litellm_params": {
"model": "hosted_vllm/org/local-model",
"api_base": "https://backend.test/v1",
"api_key": "local-key",
},
"model_info": {"id": "local-deployment", "max_output_tokens": 512, "max_input_tokens": None},
}
]
)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "llm_model_list", router.get_model_list())
monkeypatch.setattr(proxy_server, "user_model", None)
def respond(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/models"
return httpx.Response(200, json={"data": [{"id": "org/local-model", "max_model_len": 4096}]})
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as upstream:
handler.client = upstream
litellm.in_memory_llm_clients_cache.set_cache("async_httpx_clientopenai", handler)
await proxy_server.ProxyStartupEvent.refresh_model_info()
with auth_as():
for path in ("/v1/model/info", "/model/info"):
response: Final = client.get(path)
assert response.status_code == 200, response.text
info: Final = response.json()["data"][0]["model_info"]
assert (info["max_input_tokens"], info["max_output_tokens"]) == (4096, 512)
group_response: Final = client.get("/model_group/info")
assert group_response.status_code == 200, group_response.text
assert group_response.json()["data"][0]["max_input_tokens"] == 4096
_invalidate_model_cost_lowercase_map()
# ---------------------------------------------------------------------------
# GET /v2/model/info
# ---------------------------------------------------------------------------

View file

@ -7,18 +7,24 @@ and one has explicit zero-cost pricing in model_info, the other deployment
should still use the built-in pricing.
"""
import asyncio
import copy
import logging
import os
import re
from unittest.mock import patch
from typing import Final
from unittest.mock import Mock, patch
import httpx
import pytest
import litellm
from litellm import Router
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import DEFAULT_MAX_LRU_CACHE_SIZE
from litellm.litellm_core_utils.ptu_pricing import ptu_config_error
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
@ -60,6 +66,324 @@ def _restore_model_cost_entries(original_entries):
_invalidate_model_cost_lowercase_map()
@pytest.mark.parametrize("initial_count", (1, DEFAULT_MAX_LRU_CACHE_SIZE + 1))
async def test_discovered_limits_survive_deployment_growth_and_removal(
initial_count: int, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
deployments: Final = tuple(
Deployment(
model_name=f"local-{index}",
litellm_params=LiteLLM_Params(
model="hosted_vllm/local-model", api_base="https://capacity.test/v1", api_key="local-key"
),
model_info=ModelInfo(id=f"capacity-{index}"),
)
for index in range(DEFAULT_MAX_LRU_CACHE_SIZE + 2)
)
router: Final = Router(model_list=[deployment.to_json() for deployment in deployments[:initial_count]])
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(
transport=httpx.MockTransport(
lambda request: httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 4096}]})
)
) as client:
handler.client = client
await router.arefresh_model_info(client=handler)
assert all(
router.get_configured_token_limits(deployment.model_name) == (4096, 4096)
for deployment in deployments[:initial_count]
)
for deployment in deployments[initial_count:]:
router.add_deployment(deployment)
await router._arefresh_deployment_model_info(router.model_list[-1], client=handler)
assert all(
router.get_configured_token_limits(deployment.model_name) == (4096, 4096) for deployment in deployments
)
for deployment in deployments[-2:]:
router.delete_deployment(deployment.model_info.id or "")
await router._arefresh_deployment_model_info(router.model_list[0], client=handler)
assert all(
router.get_configured_token_limits(deployment.model_name) == (4096, 4096) for deployment in deployments[:-2]
)
_invalidate_model_cost_lowercase_map()
async def test_discovery_discards_metadata_for_a_replaced_deployment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
router: Final = Router(model_list=[{
"model_name": "local",
"litellm_params": {
"model": "hosted_vllm/local-model",
"api_base": "https://original.test/v1",
"api_key": "local-key",
},
"model_info": {"id": "replaced-deployment"},
}])
def respond(request: httpx.Request) -> httpx.Response:
if request.url.host == "original.test":
router.upsert_deployment(Deployment(
model_name="local",
litellm_params=LiteLLM_Params(
model="hosted_vllm/local-model",
api_base="https://replacement.test/v1",
api_key="local-key",
),
model_info=ModelInfo(id="replaced-deployment"),
))
return httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 8192}]})
assert request.url.host == "replacement.test"
return httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 2048}]})
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler.client = client
await router._arefresh_deployment_model_info(router.model_list[0], client=handler)
assert router.get_configured_token_limits("local") == (None, None)
await router.arefresh_model_info(client=handler)
assert router.get_configured_token_limits("local") == (2048, 2048)
_invalidate_model_cost_lowercase_map()
async def test_discovery_is_isolated_across_routers_and_reused_ids(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
first, second = tuple(
Router(model_list=[{
"model_name": "local",
"litellm_params": {
"model": "hosted_vllm/local-model",
"api_base": f"https://{host}.test/v1",
"api_key": "local-key",
},
"model_info": {"id": "shared-discovery-id"},
}])
for host in ("first", "second")
)
def respond(request: httpx.Request) -> httpx.Response:
if request.url.host == "unavailable.test":
return httpx.Response(503)
limit: Final = 8192 if request.url.host == "first.test" else 2048
return httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": limit}]})
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler.client = client
await first.arefresh_model_info(client=handler)
assert second.get_configured_token_limits("local") == (None, None)
await second.arefresh_model_info(client=handler)
assert first.get_discovered_model_info("shared-discovery-id")["max_input_tokens"] == 8192
assert first.get_configured_token_limits("local") == (8192, 8192)
assert second.get_configured_token_limits("local") == (2048, 2048)
assert litellm.model_cost["shared-discovery-id"].get("max_input_tokens") is None
first.upsert_deployment(Deployment(
model_name="local",
litellm_params=LiteLLM_Params(
model="hosted_vllm/local-model",
api_base="https://unavailable.test/v1",
api_key="local-key",
),
model_info=ModelInfo(id="shared-discovery-id"),
))
assert first.get_configured_token_limits("local") == (None, None)
await first.arefresh_model_info(client=handler)
assert first.get_configured_token_limits("local") == (None, None)
assert second.get_configured_token_limits("local") == (2048, 2048)
_invalidate_model_cost_lowercase_map()
async def test_discovery_refreshes_other_endpoints_while_one_is_pending(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
second_started: Final = asyncio.Event()
router: Final = Router(model_list=[
{
"model_name": host,
"litellm_params": {
"model": "hosted_vllm/local-model",
"api_base": f"https://{host}.test/v1",
"api_key": "local-key",
},
}
for host in ("first", "second", "third")
])
async def respond(request: httpx.Request) -> httpx.Response:
if request.url.host == "first.test":
await second_started.wait()
if request.url.host == "second.test":
second_started.set()
return httpx.Response(503)
return httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 2048}]})
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler.client = client
await asyncio.wait_for(router.arefresh_model_info(client=handler), timeout=2)
assert router.get_configured_token_limits("first") == (2048, 2048)
assert router.get_configured_token_limits("second") == (None, None)
assert router.get_configured_token_limits("third") == (2048, 2048)
_invalidate_model_cost_lowercase_map()
async def test_discovered_limits_expire_after_the_last_successful_refresh(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
clock: Final = Mock(return_value=0.0)
router: Final = Router(model_list=[{
"model_name": "local",
"litellm_params": {
"model": "hosted_vllm/local-model",
"api_base": "https://expiry.test/v1",
"api_key": "local-key",
},
"model_info": {"id": "expiring-discovery"},
}])
router._discovered_model_info_cache = InMemoryCache(clock=clock, default_ttl=2 * MODEL_INFO_REFRESH_SECONDS)
responses: Final = iter((
httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 4096}]}),
httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 8192}]}),
httpx.Response(503),
))
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(lambda request: next(responses))) as client:
handler.client = client
await router.arefresh_model_info(client=handler)
clock.return_value = MODEL_INFO_REFRESH_SECONDS
router.cache.in_memory_cache.flush_cache()
await router.arefresh_model_info(client=handler)
clock.return_value = 2 * MODEL_INFO_REFRESH_SECONDS + 1
router.cache.in_memory_cache.flush_cache()
await router.arefresh_model_info(client=handler)
assert router.get_configured_token_limits("local") == (8192, 8192)
group: Final = router.get_model_group_info("local")
assert group is not None
assert group.max_input_tokens == 8192
clock.return_value = 3 * MODEL_INFO_REFRESH_SECONDS + 1
await router.arefresh_model_info(client=handler)
assert router.get_configured_token_limits("local") == (None, None)
expired_group: Final = router.get_model_group_info("local")
assert expired_group is not None
assert expired_group.max_input_tokens is None
_invalidate_model_cost_lowercase_map()
@pytest.mark.parametrize("provider", ("hosted_vllm", "openai", "openai_like", "text-completion-openai"))
async def test_discovered_limits_are_isolated_overridable_and_refreshable(
provider: str, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
upstream_limit: Final = iter((8192, 4096, 16384, 2048))
def respond(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/models"
assert request.headers["authorization"] == "Bearer local-key"
return httpx.Response(200, json={"data": [{"id": "org/local-model", "max_model_len": next(upstream_limit)}]})
router: Final = Router(
model_list=[
{
"model_name": "local",
"litellm_params": {
"model": f"{provider}/org/local-model",
"api_base": f"https://{host}.test/v1",
"api_key": "local-key",
},
"model_info": {"id": host, **overrides},
}
for host, overrides in (("one", {}), ("two", {"max_output_tokens": 512}))
],
enable_pre_call_checks=True,
)
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler.client = client
await router.arefresh_model_info(client=handler)
first: Final = router.get_router_model_info(id="one", deployment=None, received_model_name="local")
second: Final = router.get_router_model_info(id="two", deployment=None, received_model_name="local")
assert (first["max_input_tokens"], first["max_output_tokens"]) == (8192, 8192)
assert (second["max_input_tokens"], second["max_output_tokens"]) == (4096, 512)
group: Final = router.get_model_group_info("local")
assert group is not None
assert group.max_input_tokens == 8192
listing: Final = router.get_model_listing_info("local")
assert listing is not None
assert listing.max_input_tokens == 8192
assert router.get_configured_token_limits("local") == (8192, 8192)
assert router._deployment_max_input_tokens("local", router.model_list[1]) == 4096
allowed: Final = router._pre_call_checks(
model="local", healthy_deployments=router.model_list, input="prompt", input_token_count=5000
)
assert [deployment["model_info"]["id"] for deployment in allowed] == ["one"]
assert router.model_list[0]["model_info"].get("max_input_tokens") is None
assert litellm.model_cost[f"{provider}/org/local-model"].get("max_input_tokens") is None
router.cache.in_memory_cache.flush_cache()
await router.arefresh_model_info(client=handler)
refreshed: Final = router.get_model_group_info("local")
assert refreshed is not None
assert refreshed.max_input_tokens == 16384
assert (
router.get_router_model_info(id="two", deployment=None, received_model_name="local")["max_output_tokens"]
== 512
)
_invalidate_model_cost_lowercase_map()
async def test_discovery_preserves_input_overrides_and_survives_outages(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost))
responses: Final = iter((
httpx.Response(200, json={"data": [{"id": "local-model", "max_model_len": 4096}]}),
httpx.Response(503),
))
def respond(request: httpx.Request) -> httpx.Response:
assert request.url.host == "backend.test"
assert request.headers["authorization"] == "Bearer local-key"
assert request.headers["x-tenant"] == "tenant"
return next(responses)
router: Final = Router(model_list=[
{
"model_name": "configured",
"litellm_params": {
"model": "hosted_vllm/local-model",
"api_base": "https://backend.test/v1",
"api_key": "unused-key",
"extra_headers": {"authorization": "Bearer local-key", "X-Tenant": "tenant"},
},
"model_info": {"id": "configured", "max_input_tokens": 1024},
},
{
"model_name": "byok",
"litellm_params": {
"model": "openai/local-model",
"api_base": "https://caller.test/v1",
"use_clientside_credentials": True,
},
},
{"model_name": "default-openai", "litellm_params": {"model": "openai/local-model", "api_key": "unused"}},
])
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
responder: Final = Mock(side_effect=respond)
async with httpx.AsyncClient(transport=httpx.MockTransport(responder)) as client:
handler.client = client
await router.arefresh_model_info(client=handler)
assert router.get_configured_token_limits("configured") == (1024, 4096)
router.cache.in_memory_cache.flush_cache()
await router.arefresh_model_info(client=handler)
assert router.get_configured_token_limits("configured") == (1024, 4096)
assert router.get_configured_token_limits("byok") == (None, None)
assert next(responses, None) is None
assert responder.call_count == 2
_invalidate_model_cost_lowercase_map()
def test_should_not_pollute_shared_key_with_zero_cost_pricing():
"""
When deployment A has input_cost_per_token=0 and deployment B has no