All Models Backend Search

This commit is contained in:
yuneng-jiang
2026-01-22 22:00:22 -08:00
parent 76fdaa2039
commit 3ee7aab5f2
5 changed files with 409 additions and 16 deletions
+95 -2
View File
@@ -7727,6 +7727,9 @@ async def model_info_v2(
debug: Optional[bool] = False,
page: int = Query(1, description="Page number", ge=1),
size: int = Query(50, description="Page size", ge=1),
search: Optional[str] = fastapi.Query(
None, description="Search model names (case-insensitive partial match)"
),
):
"""
BETA ENDPOINT. Might change unexpectedly. Use `/v1/model/info` for now.
@@ -7760,6 +7763,95 @@ async def model_info_v2(
if model is not None:
all_models = [m for m in all_models if m["model_name"] == model]
# Track total count for search (will be calculated if searching)
search_total_count = None
# Apply search filter if provided
if search is not None and search.strip():
search_lower = search.lower().strip()
# First, filter ALL models in router by search term (both config and db models)
filtered_router_models = [
m for m in all_models
if search_lower in m.get("model_name", "").lower()
]
# Separate filtered models into config vs db models, and track db model IDs
filtered_config_models = []
db_model_ids_in_router = set()
for m in filtered_router_models:
model_info = m.get("model_info", {})
is_db_model = model_info.get("db_model", False)
model_id = model_info.get("id")
if is_db_model and model_id:
db_model_ids_in_router.add(model_id)
else:
filtered_config_models.append(m)
config_models_count = len(filtered_config_models)
db_models_in_router_count = len(db_model_ids_in_router)
router_models_count = config_models_count + db_models_in_router_count
# Query database for additional models with search term (not already in router)
# We need enough models to fill the current page (size * page total models)
db_models = []
db_models_total_count = 0
models_needed_for_page = size * page # Total models needed up to current page
try:
# Build where condition for database query
db_where_condition: Dict[str, Any] = {
"model_name": {
"contains": search_lower,
"mode": "insensitive",
}
}
# Exclude models already in router if we have any
if db_model_ids_in_router:
db_where_condition["model_id"] = {
"not": {"in": list(db_model_ids_in_router)}
}
# Get total count of matching database models (excluding those already in router)
db_models_total_count = await prisma_client.db.litellm_proxymodeltable.count(
where=db_where_condition
)
# Calculate total count for search results
search_total_count = router_models_count + db_models_total_count
# Fetch database models if we need more for the current page
if router_models_count < models_needed_for_page:
models_to_fetch = min(
models_needed_for_page - router_models_count,
db_models_total_count
)
if models_to_fetch > 0:
db_models_raw = await prisma_client.db.litellm_proxymodeltable.find_many(
where=db_where_condition,
take=models_to_fetch,
)
# Convert database models to router format
for db_model in db_models_raw:
decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
if decrypted_models:
db_models.extend(decrypted_models)
except Exception as e:
verbose_proxy_logger.exception(
f"Error querying database models with search: {str(e)}"
)
# If error, use router models count as fallback
search_total_count = router_models_count
# Combine all models: config models first, then db models from router, then db models from database
# filtered_router_models already contains both config and db models from router, so we can use it directly
all_models = filtered_router_models + db_models
# else: No search - models are already in all_models from llm_router.model_list
if user_models_only:
all_models = await non_admin_all_models(
all_models=all_models,
@@ -7783,7 +7875,8 @@ async def model_info_v2(
verbose_proxy_logger.debug("all_models: %s", all_models)
total_count = len(all_models)
# Use search_total_count if searching, otherwise use len(all_models)
total_count = search_total_count if search_total_count is not None else len(all_models)
skip = (page - 1) * size
@@ -7792,7 +7885,7 @@ async def model_info_v2(
paginated_models = all_models[skip : skip + size]
verbose_proxy_logger.debug(
f"Pagination: skip={skip}, take={size}, total_count={total_count}, total_pages={total_pages}"
f"Pagination: skip={skip}, take={size}, total_count={total_count}, total_pages={total_pages}, search={search}"
)
return {
@@ -3443,6 +3443,283 @@ async def test_model_info_v2_pagination_edge_cases(monkeypatch):
app.dependency_overrides = original_overrides
@pytest.mark.asyncio
async def test_model_info_v2_search_config_models(monkeypatch):
"""
Test search parameter for config models (models from config.yaml).
Config models don't have db_model=True in model_info.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
# Create mock config models (no db_model flag or db_model=False)
mock_config_models = [
{
"model_name": "gpt-4-turbo",
"litellm_params": {"model": "gpt-4-turbo"},
"model_info": {"id": "gpt-4-turbo"}, # No db_model flag = config model
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo"},
"model_info": {"id": "gpt-3.5-turbo", "db_model": False}, # Explicitly config model
},
{
"model_name": "claude-3-opus",
"litellm_params": {"model": "claude-3-opus"},
"model_info": {"id": "claude-3-opus"}, # No db_model flag = config model
},
{
"model_name": "gemini-pro",
"litellm_params": {"model": "gemini-pro"},
"model_info": {"id": "gemini-pro"}, # No db_model flag = config model
},
]
# Mock llm_router
mock_router = MagicMock()
mock_router.model_list = mock_config_models
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock proxy_config.get_config
mock_get_config = AsyncMock(return_value={})
# Mock user authentication
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.user_id = "test-user"
mock_user_api_key_dict.api_key = "test-key"
mock_user_api_key_dict.team_models = []
mock_user_api_key_dict.models = []
# Apply monkeypatches
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
# Override auth dependency
original_overrides = app.dependency_overrides.copy()
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict
client = TestClient(app)
try:
# Test search for "gpt" - should return gpt-4-turbo and gpt-3.5-turbo
response = client.get("/v2/model/info", params={"search": "gpt"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 2 # Only config models matching search
assert len(data["data"]) == 2
model_names = [m["model_name"] for m in data["data"]]
assert "gpt-4-turbo" in model_names
assert "gpt-3.5-turbo" in model_names
assert "claude-3-opus" not in model_names
assert "gemini-pro" not in model_names
# Test search for "claude" - should return claude-3-opus
response = client.get("/v2/model/info", params={"search": "claude"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 1
assert len(data["data"]) == 1
assert data["data"][0]["model_name"] == "claude-3-opus"
# Test case-insensitive search
response = client.get("/v2/model/info", params={"search": "GPT"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 2
assert len(data["data"]) == 2
# Test partial match
response = client.get("/v2/model/info", params={"search": "turbo"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 2
assert len(data["data"]) == 2
model_names = [m["model_name"] for m in data["data"]]
assert "gpt-4-turbo" in model_names
assert "gpt-3.5-turbo" in model_names
# Test search with no matches
response = client.get("/v2/model/info", params={"search": "nonexistent"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 0
assert len(data["data"]) == 0
finally:
app.dependency_overrides = original_overrides
@pytest.mark.asyncio
async def test_model_info_v2_search_db_models(monkeypatch):
"""
Test search parameter for db models (models from database).
DB models have db_model=True and id in model_info.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
# Create mock db models (db_model=True with id)
mock_db_models_in_router = [
{
"model_name": "db-gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "db-model-1", "db_model": True}, # DB model
},
{
"model_name": "db-claude-3",
"litellm_params": {"model": "claude-3"},
"model_info": {"id": "db-model-2", "db_model": True}, # DB model
},
]
# Mock llm_router
mock_router = MagicMock()
mock_router.model_list = mock_db_models_in_router
# Mock prisma_client with database query methods
mock_db_models_from_db = [
MagicMock(
model_id="db-model-3",
model_name="db-gemini-pro",
litellm_params='{"model": "gemini-pro"}',
model_info='{"id": "db-model-3", "db_model": true}',
),
MagicMock(
model_id="db-model-4",
model_name="db-gpt-3.5",
litellm_params='{"model": "gpt-3.5-turbo"}',
model_info='{"id": "db-model-4", "db_model": true}',
),
]
# Mock the database count and find_many methods dynamically based on search
async def mock_db_count_func(*args, **kwargs):
where_condition = kwargs.get("where", {})
search_term = where_condition.get("model_name", {}).get("contains", "")
excluded_ids = where_condition.get("model_id", {}).get("not", {}).get("in", [])
# Count models matching search term but not in excluded_ids
count = 0
for model in mock_db_models_from_db:
if search_term.lower() in model.model_name.lower():
if model.model_id not in excluded_ids:
count += 1
return count
async def mock_db_find_many_func(*args, **kwargs):
where_condition = kwargs.get("where", {})
search_term = where_condition.get("model_name", {}).get("contains", "")
excluded_ids = where_condition.get("model_id", {}).get("not", {}).get("in", [])
take = kwargs.get("take", 10)
# Return models matching search term but not in excluded_ids
result = []
for model in mock_db_models_from_db:
if search_term.lower() in model.model_name.lower():
if model.model_id not in excluded_ids:
result.append(model)
if len(result) >= take:
break
return result
mock_db_count = AsyncMock(side_effect=mock_db_count_func)
mock_db_find_many = AsyncMock(side_effect=mock_db_find_many_func)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_proxymodeltable.count = mock_db_count
mock_prisma_client.db.litellm_proxymodeltable.find_many = mock_db_find_many
# Mock proxy_config.decrypt_model_list_from_db to return router-format models
def mock_decrypt_models(db_models_list):
result = []
for db_model in db_models_list:
result.append(
{
"model_name": db_model.model_name,
"litellm_params": {"model": db_model.model_name.replace("db-", "")},
"model_info": {"id": db_model.model_id, "db_model": True},
}
)
return result
# Mock proxy_config.get_config
mock_get_config = AsyncMock(return_value={})
# Mock user authentication
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.user_id = "test-user"
mock_user_api_key_dict.api_key = "test-key"
mock_user_api_key_dict.team_models = []
mock_user_api_key_dict.models = []
# Apply monkeypatches
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
monkeypatch.setattr(proxy_config, "decrypt_model_list_from_db", mock_decrypt_models)
# Override auth dependency
original_overrides = app.dependency_overrides.copy()
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict
client = TestClient(app)
try:
# Test search for "gpt" - should return db-gpt-4 from router and db-gpt-3.5 from db
response = client.get("/v2/model/info", params={"search": "gpt"})
assert response.status_code == 200
data = response.json()
# Should have db-gpt-4 from router + db-gpt-3.5 from db = 2 total
assert data["total_count"] == 2
assert len(data["data"]) == 2
model_names = [m["model_name"] for m in data["data"]]
assert "db-gpt-4" in model_names
assert "db-gpt-3.5" in model_names
# Verify database was queried
mock_db_count.assert_called()
# Verify the where condition excludes models already in router
call_args = mock_db_count.call_args
assert call_args is not None
where_condition = call_args[1]["where"]
assert "model_name" in where_condition
assert where_condition["model_name"]["contains"] == "gpt"
assert where_condition["model_name"]["mode"] == "insensitive"
# Test search for "claude" - should return db-claude-3 from router only
response = client.get("/v2/model/info", params={"search": "claude"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 1
assert len(data["data"]) == 1
assert data["data"][0]["model_name"] == "db-claude-3"
# Test search for "gemini" - should return db-gemini-pro from db only
response = client.get("/v2/model/info", params={"search": "gemini"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 1
assert len(data["data"]) == 1
assert data["data"][0]["model_name"] == "db-gemini-pro"
# Test case-insensitive search
response = client.get("/v2/model/info", params={"search": "GPT"})
assert response.status_code == 200
data = response.json()
assert data["total_count"] == 2
finally:
app.dependency_overrides = original_overrides
def test_enrich_model_info_with_litellm_data():
"""
Test the _enrich_model_info_with_litellm_data helper function.
@@ -27,7 +27,7 @@ const modelHubKeys = createQueryKeys("modelHub");
const allProxyModelsKeys = createQueryKeys("allProxyModels");
const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels");
export const useModelsInfo = (page: number = 1, size: number = 50) => {
export const useModelsInfo = (page: number = 1, size: number = 50, search?: string) => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<PaginatedModelInfoResponse>({
queryKey: modelKeys.list({
@@ -36,9 +36,10 @@ export const useModelsInfo = (page: number = 1, size: number = 50) => {
...(userRole && { userRole }),
page,
size,
...(search && { search }),
},
}),
queryFn: async () => await modelInfoCall(accessToken!, userId!, userRole!, page, size),
queryFn: async () => await modelInfoCall(accessToken!, userId!, userRole!, page, size, search),
enabled: Boolean(accessToken && userId && userRole),
});
};
@@ -8,10 +8,11 @@ import { getDisplayModelName } from "@/components/view_model/model_name_display"
import { InfoCircleOutlined } from "@ant-design/icons";
import { PaginationState } from "@tanstack/react-table";
import { Grid, Select, SelectItem, TabPanel, Text } from "@tremor/react";
import { Skeleton } from "antd";
import debounce from "lodash/debounce";
import { useEffect, useMemo, useState } from "react";
import { useModelsInfo } from "../../hooks/models/useModels";
import { transformModelData } from "../utils/modelDataTransformer";
import { Skeleton } from "antd";
type ModelViewMode = "all" | "current_team";
interface AllModelsTabProps {
@@ -36,6 +37,7 @@ const AllModelsTab = ({
const { data: teams } = useTeams();
const [modelNameSearch, setModelNameSearch] = useState<string>("");
const [debouncedSearch, setDebouncedSearch] = useState<string>("");
const [modelViewMode, setModelViewMode] = useState<ModelViewMode>("current_team");
const [currentTeam, setCurrentTeam] = useState<Team | "personal">("personal");
const [showFilters, setShowFilters] = useState<boolean>(false);
@@ -48,7 +50,26 @@ const AllModelsTab = ({
pageSize: 50,
});
const { data: rawModelData, isLoading: isLoadingModelsInfo } = useModelsInfo(currentPage, pageSize);
// Debounce search input
const debouncedUpdateSearch = useMemo(
() =>
debounce((value: string) => {
setDebouncedSearch(value);
// Reset to page 1 when search changes
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}, 200),
[]
);
useEffect(() => {
debouncedUpdateSearch(modelNameSearch);
return () => {
debouncedUpdateSearch.cancel();
};
}, [modelNameSearch, debouncedUpdateSearch]);
const { data: rawModelData, isLoading: isLoadingModelsInfo } = useModelsInfo(currentPage, pageSize, debouncedSearch || undefined);
const isLoading = isLoadingModelsInfo || isLoadingModelCostMap;
const getProviderFromModel = (model: string) => {
@@ -88,10 +109,8 @@ const AllModelsTab = ({
return [];
}
// Server-side search is now handled by the API, so we only filter by other criteria
return modelData.data.filter((model: any) => {
const searchMatch =
modelNameSearch === "" || model.model_name.toLowerCase().includes(modelNameSearch.toLowerCase());
const modelNameMatch =
selectedModelGroup === "all" ||
model.model_name === selectedModelGroup ||
@@ -120,13 +139,13 @@ const AllModelsTab = ({
}
}
return searchMatch && modelNameMatch && accessGroupMatch && teamAccessMatch;
return modelNameMatch && accessGroupMatch && teamAccessMatch;
});
}, [modelData, modelNameSearch, selectedModelGroup, selectedModelAccessGroupFilter, currentTeam, modelViewMode]);
}, [modelData, selectedModelGroup, selectedModelAccessGroupFilter, currentTeam, modelViewMode]);
useEffect(() => {
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}, [modelNameSearch, selectedModelGroup, selectedModelAccessGroupFilter, currentTeam, modelViewMode]);
}, [selectedModelGroup, selectedModelAccessGroupFilter, currentTeam, modelViewMode]);
const resetFilters = () => {
setModelNameSearch("");
@@ -354,8 +373,8 @@ const AllModelsTab = ({
<Skeleton.Input active style={{ width: 184, height: 20 }} />
) : (
<span className="text-sm text-gray-700">
{filteredData.length > 0
? `Showing 1 - ${filteredData.length} of ${filteredData.length} results`
{paginationMeta.total_count > 0
? `Showing ${((currentPage - 1) * pageSize) + 1} - ${Math.min(currentPage * pageSize, paginationMeta.total_count)} of ${paginationMeta.total_count} results`
: "Showing 0 results"}
</span>
)}
@@ -2007,17 +2007,20 @@ export const regenerateKeyCall = async (accessToken: string, keyToRegenerate: st
let ModelListerrorShown = false;
let errorTimer: NodeJS.Timeout | null = null;
export const modelInfoCall = async (accessToken: string, userID: string, userRole: string, page: number = 1, size: number = 50) => {
export const modelInfoCall = async (accessToken: string, userID: string, userRole: string, page: number = 1, size: number = 50, search?: string) => {
/**
* Get all models on proxy
*/
try {
console.log("modelInfoCall:", accessToken, userID, userRole, page, size);
console.log("modelInfoCall:", accessToken, userID, userRole, page, size, search);
let url = proxyBaseUrl ? `${proxyBaseUrl}/v2/model/info` : `/v2/model/info`;
const params = new URLSearchParams();
params.append("include_team_models", "true");
params.append("page", page.toString());
params.append("size", size.toString());
if (search && search.trim()) {
params.append("search", search.trim());
}
if (params.toString()) {
url += `?${params.toString()}`;
}