mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test(proxy): type the model search fakes precisely
The fake Prisma table and proxy config behind the deployment id search tests took bare dicts, poked attributes onto MagicMocks and left every local unannotated. The rows are now a frozen slotted dataclass, the deployments are ReadOnly TypedDicts, the where evaluator is typed against the production _ModelSearchWhere shape, the mocks are built through constructor kwargs and every local is Final. Both tests still fail when either the router side or the DB side id match is removed Claude-Session: https://claude.ai/code/session_018yW93iDaEMhoQUXcYjus7D
This commit is contained in:
parent
d83076c0ca
commit
222ae2c00d
1 changed files with 91 additions and 83 deletions
|
|
@ -7,8 +7,11 @@ import re
|
|||
import socket
|
||||
import subprocess
|
||||
import types
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest import mock
|
||||
from unittest.mock import AsyncMock, MagicMock, mock_open, patch
|
||||
|
||||
|
|
@ -19,6 +22,7 @@ import yaml
|
|||
from fastapi import FastAPI
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.testclient import TestClient
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
||||
import litellm
|
||||
|
|
@ -29,7 +33,7 @@ from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
|||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.proxy_server import app, initialize
|
||||
from litellm.proxy.proxy_server import _ModelSearchWhere, app, initialize
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
|
||||
example_embedding_result = {
|
||||
|
|
@ -2239,59 +2243,73 @@ async def test_apply_search_filter_bounds_db_fetch_by_page_and_cap():
|
|||
assert take < 10_000, "sorted search must cap below the full match set"
|
||||
|
||||
|
||||
def _db_model_row(model_id: str, model_name: str) -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.model_id = model_id
|
||||
row.model_name = model_name
|
||||
row.model_info = {"id": model_id, "db_model": True}
|
||||
return row
|
||||
class _ModelInfo(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
db_model: ReadOnly[bool]
|
||||
|
||||
|
||||
def _prisma_client_with_model_rows(rows: list[MagicMock]) -> MagicMock:
|
||||
class _Deployment(TypedDict):
|
||||
model_name: ReadOnly[str]
|
||||
litellm_params: ReadOnly[Mapping[str, str]]
|
||||
model_info: ReadOnly[_ModelInfo]
|
||||
|
||||
|
||||
def _deployment(model_name: str, litellm_model: str, model_id: str, db_model: bool = False) -> _Deployment:
|
||||
deployment: Final[_Deployment] = {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {"model": litellm_model},
|
||||
"model_info": {"id": model_id, "db_model": db_model},
|
||||
}
|
||||
return deployment
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _DbModelRow:
|
||||
model_id: str
|
||||
model_name: str
|
||||
|
||||
@property
|
||||
def model_info(self) -> _ModelInfo:
|
||||
model_info: Final[_ModelInfo] = {"id": self.model_id, "db_model": True}
|
||||
return model_info
|
||||
|
||||
@property
|
||||
def columns(self) -> Mapping[str, str]:
|
||||
return types.MappingProxyType({"model_id": self.model_id, "model_name": self.model_name})
|
||||
|
||||
|
||||
def _db_row_matches(row: _DbModelRow, where: _ModelSearchWhere) -> bool:
|
||||
"""
|
||||
Fake `LiteLLM_ProxyModelTable` that evaluates the Prisma `where` the
|
||||
search sends: top-level fields AND together, `OR` takes any clause,
|
||||
`contains` is a case-insensitive substring and `not.in` excludes ids.
|
||||
Evaluate the Prisma `where` the search sends: `OR` takes any clause,
|
||||
`contains` is a case-insensitive substring, `model_id.not.in` drops
|
||||
the rows the router already serves and `model_name` is an exact match.
|
||||
"""
|
||||
search_hits: Final = any(
|
||||
predicate["contains"].lower() in row.columns[column].lower()
|
||||
for clause in where["OR"]
|
||||
for column, predicate in clause.items()
|
||||
)
|
||||
not_in_router: Final = row.model_id not in where["model_id"]["not"]["in"]
|
||||
in_model_group: Final = where.get("model_name", row.model_name) == row.model_name
|
||||
return search_hits and not_in_router and in_model_group
|
||||
|
||||
def _column_matches(row: MagicMock, column: str, predicate: object) -> bool:
|
||||
value = getattr(row, column)
|
||||
if isinstance(predicate, str):
|
||||
return value == predicate
|
||||
assert isinstance(predicate, dict)
|
||||
if "contains" in predicate:
|
||||
return predicate["contains"].lower() in value.lower()
|
||||
return value not in predicate["not"]["in"]
|
||||
|
||||
def _row_matches(row: MagicMock, where: dict) -> bool:
|
||||
columns_match = all(
|
||||
_column_matches(row, column, predicate) for column, predicate in where.items() if column != "OR"
|
||||
)
|
||||
any_clause_matches = any(
|
||||
all(_column_matches(row, column, predicate) for column, predicate in clause.items())
|
||||
for clause in where.get("OR", ())
|
||||
)
|
||||
return columns_match and ("OR" not in where or any_clause_matches)
|
||||
def _prisma_client_with_model_rows(rows: Sequence[_DbModelRow]) -> MagicMock:
|
||||
async def _count(where: _ModelSearchWhere) -> int:
|
||||
return sum(_db_row_matches(row, where) for row in rows)
|
||||
|
||||
async def _count(where: dict) -> int:
|
||||
return sum(_row_matches(row, where) for row in rows)
|
||||
async def _find_many(where: _ModelSearchWhere, take: int) -> tuple[_DbModelRow, ...]:
|
||||
return tuple(row for row in rows if _db_row_matches(row, where))[:take]
|
||||
|
||||
async def _find_many(where: dict, take: int) -> list[MagicMock]:
|
||||
return [row for row in rows if _row_matches(row, where)][:take]
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(side_effect=_count)
|
||||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=_find_many)
|
||||
return prisma_client
|
||||
table: Final = MagicMock(count=AsyncMock(side_effect=_count), find_many=AsyncMock(side_effect=_find_many))
|
||||
return MagicMock(db=MagicMock(litellm_proxymodeltable=table))
|
||||
|
||||
|
||||
def _proxy_config_decrypting_rows() -> MagicMock:
|
||||
proxy_config = MagicMock()
|
||||
proxy_config.decrypt_model_list_from_db = lambda rows: [
|
||||
{"model_name": r.model_name, "model_info": r.model_info, "litellm_params": {"model": r.model_name}}
|
||||
for r in rows
|
||||
]
|
||||
return proxy_config
|
||||
def _decrypt(rows: Sequence[_DbModelRow]) -> tuple[_Deployment, ...]:
|
||||
return tuple(_deployment(row.model_name, row.model_name, row.model_id, db_model=True) for row in rows)
|
||||
|
||||
return MagicMock(decrypt_model_list_from_db=_decrypt)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2304,14 +2322,14 @@ async def test_apply_search_filter_honours_exact_model_name_in_db_query():
|
|||
"""
|
||||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||||
|
||||
prisma_client = _prisma_client_with_model_rows(
|
||||
[
|
||||
_db_model_row("aaa-in-sonnet-group", "anthropic-sonnet-5"),
|
||||
_db_model_row("bbb-sonnet-in-id-only", "anthropic-opus-5"),
|
||||
_db_model_row("ccc-opus-in-id-only", "anthropic-sonnet-5"),
|
||||
]
|
||||
prisma_client: Final = _prisma_client_with_model_rows(
|
||||
(
|
||||
_DbModelRow("aaa-in-sonnet-group", "anthropic-sonnet-5"),
|
||||
_DbModelRow("bbb-sonnet-in-id-only", "anthropic-opus-5"),
|
||||
_DbModelRow("ccc-opus-in-id-only", "anthropic-sonnet-5"),
|
||||
)
|
||||
)
|
||||
proxy_config = _proxy_config_decrypting_rows()
|
||||
proxy_config: Final = _proxy_config_decrypting_rows()
|
||||
|
||||
sonnet_group, sonnet_group_count = await _apply_search_filter_to_models(
|
||||
all_models=[],
|
||||
|
|
@ -2320,7 +2338,7 @@ async def test_apply_search_filter_honours_exact_model_name_in_db_query():
|
|||
proxy_config=proxy_config,
|
||||
model_name="anthropic-sonnet-5",
|
||||
)
|
||||
assert [m["model_info"]["id"] for m in sonnet_group] == ["aaa-in-sonnet-group", "ccc-opus-in-id-only"]
|
||||
assert tuple(m["model_info"]["id"] for m in sonnet_group) == ("aaa-in-sonnet-group", "ccc-opus-in-id-only")
|
||||
assert sonnet_group_count == 2
|
||||
|
||||
opus_id_in_sonnet_group, opus_id_in_sonnet_group_count = await _apply_search_filter_to_models(
|
||||
|
|
@ -2330,7 +2348,7 @@ async def test_apply_search_filter_honours_exact_model_name_in_db_query():
|
|||
proxy_config=proxy_config,
|
||||
model_name="anthropic-sonnet-5",
|
||||
)
|
||||
assert [m["model_info"]["id"] for m in opus_id_in_sonnet_group] == ["ccc-opus-in-id-only"]
|
||||
assert tuple(m["model_info"]["id"] for m in opus_id_in_sonnet_group) == ("ccc-opus-in-id-only",)
|
||||
assert opus_id_in_sonnet_group_count == 1
|
||||
|
||||
every_sonnet, every_sonnet_count = await _apply_search_filter_to_models(
|
||||
|
|
@ -2353,70 +2371,60 @@ async def test_apply_search_filter_matches_deployment_id():
|
|||
"""
|
||||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||||
|
||||
config_model_id = "39a9705e6b3a94a207bfbd3200b39a8ed886d2c6daa0deab123d15b721c8c898"
|
||||
config_model = {
|
||||
"model_name": "gpt-4.1",
|
||||
"litellm_params": {"model": "openai/gpt-4.1"},
|
||||
"model_info": {"id": config_model_id, "db_model": False},
|
||||
}
|
||||
sibling_config_model = {
|
||||
"model_name": "gpt-4.1",
|
||||
"litellm_params": {"model": "azure/gpt-4.1"},
|
||||
"model_info": {"id": "0f0f0f0f-sibling", "db_model": False},
|
||||
}
|
||||
router_db_row = _db_model_row("db-row-loaded-in-router", "claude-sonnet-4-5")
|
||||
unloaded_db_row = _db_model_row("lit4738-db-row-not-in-router", "gemini-2.5-pro")
|
||||
router_db_model = {
|
||||
"model_name": router_db_row.model_name,
|
||||
"litellm_params": {"model": "anthropic/claude-sonnet-4-5"},
|
||||
"model_info": router_db_row.model_info,
|
||||
}
|
||||
all_models = [config_model, sibling_config_model, router_db_model]
|
||||
prisma_client = _prisma_client_with_model_rows([router_db_row, unloaded_db_row])
|
||||
proxy_config = _proxy_config_decrypting_rows()
|
||||
config_model_id: Final = "39a9705e6b3a94a207bfbd3200b39a8ed886d2c6daa0deab123d15b721c8c898"
|
||||
config_model: Final = _deployment("gpt-4.1", "openai/gpt-4.1", config_model_id)
|
||||
sibling_config_model: Final = _deployment("gpt-4.1", "azure/gpt-4.1", "0f0f0f0f-sibling")
|
||||
router_db_row: Final = _DbModelRow("db-row-loaded-in-router", "claude-sonnet-4-5")
|
||||
unloaded_db_row: Final = _DbModelRow("lit4738-db-row-not-in-router", "gemini-2.5-pro")
|
||||
router_db_model: Final = _deployment(
|
||||
router_db_row.model_name, "anthropic/claude-sonnet-4-5", router_db_row.model_id, db_model=True
|
||||
)
|
||||
router_models: Final = list(dict(m) for m in (config_model, sibling_config_model, router_db_model))
|
||||
prisma_client: Final = _prisma_client_with_model_rows((router_db_row, unloaded_db_row))
|
||||
proxy_config: Final = _proxy_config_decrypting_rows()
|
||||
|
||||
by_config_id, by_config_id_count = await _apply_search_filter_to_models(
|
||||
all_models=all_models,
|
||||
all_models=router_models,
|
||||
search=config_model_id,
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
assert by_config_id == [config_model]
|
||||
assert tuple(by_config_id) == (config_model,)
|
||||
assert by_config_id_count == 1
|
||||
|
||||
by_unloaded_db_id, by_unloaded_db_id_count = await _apply_search_filter_to_models(
|
||||
all_models=all_models,
|
||||
all_models=router_models,
|
||||
search=unloaded_db_row.model_id,
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
assert [m["model_info"]["id"] for m in by_unloaded_db_id] == [unloaded_db_row.model_id]
|
||||
assert tuple(m["model_info"]["id"] for m in by_unloaded_db_id) == (unloaded_db_row.model_id,)
|
||||
assert by_unloaded_db_id_count == 1
|
||||
|
||||
by_router_db_id, by_router_db_id_count = await _apply_search_filter_to_models(
|
||||
all_models=all_models,
|
||||
all_models=router_models,
|
||||
search=router_db_row.model_id,
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
assert by_router_db_id == [router_db_model], "a DB row already in the router must not be returned twice"
|
||||
assert tuple(by_router_db_id) == (router_db_model,), "a DB row already in the router must not be returned twice"
|
||||
assert by_router_db_id_count == 1
|
||||
|
||||
by_id_prefix_upper, _ = await _apply_search_filter_to_models(
|
||||
all_models=all_models,
|
||||
all_models=router_models,
|
||||
search=config_model_id[:12].upper(),
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
assert by_id_prefix_upper == [config_model]
|
||||
assert tuple(by_id_prefix_upper) == (config_model,)
|
||||
|
||||
by_name, by_name_count = await _apply_search_filter_to_models(
|
||||
all_models=all_models,
|
||||
all_models=router_models,
|
||||
search="gpt-4",
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
assert by_name == [config_model, sibling_config_model]
|
||||
assert tuple(by_name) == (config_model, sibling_config_model)
|
||||
assert by_name_count == 2
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue