mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
fix(router): preserve discovered limits and model info fallbacks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
d0ed8145b0
commit
bab273ea0f
4 changed files with 153 additions and 17 deletions
|
|
@ -9324,12 +9324,6 @@ def get_litellm_model_info(model: dict = {}):
|
|||
model_info: Final = model.get("model_info", {})
|
||||
model_to_lookup = model.get("litellm_params", {}).get("model", None)
|
||||
try:
|
||||
if llm_router is not None and model_info.get("id") is not None:
|
||||
deployment_info: Final = llm_router.get_deployment_model_info(
|
||||
model_id=model_info["id"], model_name=model_to_lookup
|
||||
)
|
||||
if deployment_info is not None:
|
||||
return deployment_info
|
||||
if "azure" in model_to_lookup or model_info.get("base_model"):
|
||||
model_to_lookup = model_info.get("base_model", None)
|
||||
litellm_model_info: Final = litellm.get_model_info(model_to_lookup)
|
||||
|
|
@ -13623,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 model_info.get(k) is None:
|
||||
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
|
||||
|
|
@ -15089,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
|
||||
|
|
|
|||
|
|
@ -982,7 +982,7 @@ class Router:
|
|||
self.get_deployment_model_info
|
||||
)
|
||||
self._discovered_model_info_cache: InMemoryCache = InMemoryCache(
|
||||
max_size_in_memory=DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
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
|
||||
|
|
@ -9504,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
|
||||
|
|
@ -9798,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:
|
||||
|
|
@ -10384,13 +10386,14 @@ class Router:
|
|||
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]:
|
||||
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
|
||||
|
|
@ -10428,7 +10431,7 @@ class Router:
|
|||
model_infos: Final = tuple(
|
||||
MappingProxyType(
|
||||
{
|
||||
**self._get_discovered_model_info((deployment.get("model_info") or MappingProxyType({})).get("id")),
|
||||
**self.get_discovered_model_info((deployment.get("model_info") or MappingProxyType({})).get("id")),
|
||||
**MappingProxyType(
|
||||
{
|
||||
k: v
|
||||
|
|
@ -10487,7 +10490,7 @@ class Router:
|
|||
|
||||
model_info: Final = MappingProxyType(
|
||||
{
|
||||
**self._get_discovered_model_info(deployment.model_info.id),
|
||||
**self.get_discovered_model_info(deployment.model_info.id),
|
||||
**deployment.model_info.model_dump(exclude_none=True),
|
||||
}
|
||||
)
|
||||
|
|
@ -10757,7 +10760,7 @@ class Router:
|
|||
# values are skipped or Deployment's None pricing defaults would erase the map's
|
||||
merged_model_info: Final[ModelMapInfo] = {
|
||||
**copy.deepcopy(model_info),
|
||||
**self._get_discovered_model_info((deployment.get("model_info") or {}).get("id")),
|
||||
**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}
|
||||
),
|
||||
|
|
@ -10811,7 +10814,7 @@ class Router:
|
|||
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),
|
||||
**self.get_discovered_model_info(model_id),
|
||||
}
|
||||
if model_id in litellm.model_cost
|
||||
else None
|
||||
|
|
|
|||
|
|
@ -28,6 +28,94 @@ 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]],
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ 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
|
||||
|
|
@ -65,6 +66,50 @@ 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=[{
|
||||
|
|
@ -131,7 +176,7 @@ async def test_discovery_is_isolated_across_routers_and_reused_ids(monkeypatch:
|
|||
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_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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue