mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge pull request #38626 from BerriAI/litellm_ui_model_links_team_key_info
feat(ui): link team and key model chips to the models page filtered to that group
This commit is contained in:
commit
fa25ff2a2e
20 changed files with 413 additions and 32 deletions
|
|
@ -16,7 +16,7 @@ import threading
|
|||
import time
|
||||
import traceback
|
||||
import warnings
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping, MutableMapping, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import (
|
||||
|
|
@ -12792,6 +12792,7 @@ async def _fetch_db_models_for_search(
|
|||
size: int,
|
||||
sort_by: str | None,
|
||||
is_byok_outside_caller_teams: Callable[[dict[str, JsonValue]], bool],
|
||||
model_name: str | None = None,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
"""
|
||||
Run the bounded DB query that backs `/v2/model/info?search=`. Returns
|
||||
|
|
@ -12808,7 +12809,9 @@ async def _fetch_db_models_for_search(
|
|||
filter for `team_public_model_name` instead and keep the DB cost
|
||||
bounded by `search`.
|
||||
"""
|
||||
db_where_condition: Final[dict[str, Any]] = {"model_name": {"contains": search_lower, "mode": "insensitive"}}
|
||||
db_where_condition: Final[dict[str, Any]] = {
|
||||
"model_name": {"contains": search_lower, "mode": "insensitive"} if model_name is None else model_name
|
||||
}
|
||||
if db_model_ids_in_router:
|
||||
db_where_condition["model_id"] = {"not": {"in": list(db_model_ids_in_router)}}
|
||||
|
||||
|
|
@ -12855,6 +12858,7 @@ async def _apply_search_filter_to_models(
|
|||
page: int = 1,
|
||||
size: int = 50,
|
||||
sort_by: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> tuple[list[dict[str, Any]], int | None]:
|
||||
"""
|
||||
Apply search filter to models, querying database for additional matching models.
|
||||
|
|
@ -12875,6 +12879,11 @@ async def _apply_search_filter_to_models(
|
|||
sort_by: Sort field. When set, results must be sorted across the
|
||||
full match set, so the DB fetch is capped at
|
||||
``_SORTED_SEARCH_DB_FETCH_CAP`` instead of one page.
|
||||
model_name: Exact ``model_name`` the caller already narrowed
|
||||
``all_models`` to (``?model=``). The DB query matches it
|
||||
exactly instead of the substring, and is skipped when the
|
||||
substring cannot occur in it, otherwise rows from other model
|
||||
groups leak into the result and the count.
|
||||
|
||||
Returns:
|
||||
Tuple of (filtered_models, total_count). total_count is None if not searching.
|
||||
|
|
@ -12932,7 +12941,8 @@ async def _apply_search_filter_to_models(
|
|||
|
||||
# Query database for additional models with search term
|
||||
db_models: list[dict[str, Any]] = []
|
||||
if prisma_client is not None:
|
||||
exact_name_can_match: Final = model_name is None or search_lower in model_name.lower()
|
||||
if prisma_client is not None and exact_name_can_match:
|
||||
try:
|
||||
db_models, db_models_total_count = await _fetch_db_models_for_search(
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -12944,6 +12954,7 @@ async def _apply_search_filter_to_models(
|
|||
size=size,
|
||||
sort_by=sort_by,
|
||||
is_byok_outside_caller_teams=_is_byok_outside_caller_teams,
|
||||
model_name=model_name,
|
||||
)
|
||||
search_total_count = router_models_count + db_models_total_count
|
||||
except Exception as e:
|
||||
|
|
@ -13497,7 +13508,7 @@ async def model_info_v2(
|
|||
all_models += [user_model]
|
||||
|
||||
if model is not None:
|
||||
all_models = [m for m in all_models if m["model_name"] == model]
|
||||
all_models = [m for m in all_models if _deployment_matches_allowed_model_names(m, frozenset((model,)))]
|
||||
|
||||
# Apply search filter if provided
|
||||
all_models, search_total_count = await _apply_search_filter_to_models(
|
||||
|
|
@ -13509,6 +13520,7 @@ async def model_info_v2(
|
|||
page=page,
|
||||
size=size,
|
||||
sort_by=sortBy,
|
||||
model_name=model,
|
||||
)
|
||||
|
||||
if user_models_only:
|
||||
|
|
@ -14023,7 +14035,7 @@ async def model_metrics_exceptions(
|
|||
return {"data": response, "exception_types": list(exception_types)}
|
||||
|
||||
|
||||
def _deployment_matches_allowed_model_names(model: dict[str, JsonValue], allowed_model_names: set[str]) -> bool:
|
||||
def _deployment_matches_allowed_model_names(model: dict[str, JsonValue], allowed_model_names: Collection[str]) -> bool:
|
||||
"""Match a router deployment against allowed public model names.
|
||||
|
||||
Team-scoped rows store an internal routing key in ``model_name``; callers
|
||||
|
|
|
|||
|
|
@ -154,6 +154,59 @@ async def test_model_info_v2_translates_team_model_name(monkeypatch):
|
|||
assert "model_name_team-abc-123_4a6b8" not in names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_v2_exact_model_filter_matches_team_public_name(monkeypatch):
|
||||
"""`/v2/model/info?model=<public name>` must keep the team-scoped row whose
|
||||
`model_name` is the internal routing key: the dashboard links team model
|
||||
chips with the public name, and the exact filter ran before translation."""
|
||||
global_row = {
|
||||
"model_name": "gpt-4o",
|
||||
"litellm_params": {"model": "gpt-4o"},
|
||||
"model_info": {"id": "normal-id-1", "db_model": False},
|
||||
}
|
||||
router = MagicMock()
|
||||
router.model_list = [_team_row(), global_row]
|
||||
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||||
monkeypatch.setattr(ps.proxy_config, "get_config", AsyncMock(return_value={}))
|
||||
monkeypatch.setattr(
|
||||
ps,
|
||||
"_apply_search_filter_to_models",
|
||||
AsyncMock(side_effect=lambda all_models, **kw: (all_models, len(all_models))),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
|
||||
)
|
||||
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
|
||||
|
||||
monkeypatch.setattr(
|
||||
mlh,
|
||||
"append_agents_to_model_info",
|
||||
AsyncMock(side_effect=lambda models, **kw: models),
|
||||
)
|
||||
|
||||
admin = UserAPIKeyAuth(user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
resp = await ps.model_info_v2(
|
||||
user_api_key_dict=admin,
|
||||
model="team-claude-sonnet",
|
||||
user_models_only=False,
|
||||
include_team_models=False,
|
||||
debug=False,
|
||||
page=1,
|
||||
size=50,
|
||||
search=None,
|
||||
modelId=None,
|
||||
teamId=None,
|
||||
sortBy=None,
|
||||
sortOrder="asc",
|
||||
)
|
||||
|
||||
assert [m["model_name"] for m in resp["data"]] == ["team-claude-sonnet"]
|
||||
assert resp["total_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_v1_list_path_translates_team_model_name(monkeypatch):
|
||||
"""/v1/model/info list path (no litellm_model_id) must include team-scoped
|
||||
|
|
|
|||
|
|
@ -2126,6 +2126,53 @@ 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"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_search_filter_honours_exact_model_name_in_db_query():
|
||||
"""
|
||||
`/v2/model/info?model=<group>&search=<term>`: the router list is already
|
||||
narrowed to the exact group, so the DB count and fetch must be too, or
|
||||
other groups' rows leak into the page and inflate total_count.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=0)
|
||||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
proxy_config = MagicMock()
|
||||
proxy_config.decrypt_model_list_from_db = lambda rows: []
|
||||
|
||||
await _apply_search_filter_to_models(
|
||||
all_models=[],
|
||||
search="sonnet",
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
model_name="anthropic-sonnet-5",
|
||||
)
|
||||
where = prisma_client.db.litellm_proxymodeltable.count.call_args.kwargs["where"]
|
||||
assert where["model_name"] == "anthropic-sonnet-5"
|
||||
assert prisma_client.db.litellm_proxymodeltable.find_many.call_args.kwargs["where"] == where
|
||||
|
||||
prisma_client.db.litellm_proxymodeltable.count.reset_mock()
|
||||
_, total_count = await _apply_search_filter_to_models(
|
||||
all_models=[],
|
||||
search="opus",
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
model_name="anthropic-sonnet-5",
|
||||
)
|
||||
prisma_client.db.litellm_proxymodeltable.count.assert_not_called()
|
||||
assert total_count == 0
|
||||
|
||||
await _apply_search_filter_to_models(
|
||||
all_models=[],
|
||||
search="sonnet",
|
||||
prisma_client=prisma_client,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
where = prisma_client.db.litellm_proxymodeltable.count.call_args.kwargs["where"]
|
||||
assert where["model_name"] == {"contains": "sonnet", "mode": "insensitive"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filter_models_by_team_id_excludes_viewer_direct_access():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -118,6 +118,7 @@ describe("useModelsInfo", () => {
|
|||
// exclude_auto_routers defaults off: only the Models + Endpoints table opts in, so
|
||||
// every other consumer of this hook keeps seeing auto-routers.
|
||||
false,
|
||||
undefined,
|
||||
);
|
||||
expect(modelInfoCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
|
@ -145,6 +146,7 @@ describe("useModelsInfo", () => {
|
|||
// exclude_auto_routers defaults off: only the Models + Endpoints table opts in, so
|
||||
// every other consumer of this hook keeps seeing auto-routers.
|
||||
false,
|
||||
undefined,
|
||||
);
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ export const useModelsInfo = (
|
|||
sortBy?: string,
|
||||
sortOrder?: string,
|
||||
excludeAutoRouters: boolean = false,
|
||||
modelName?: string,
|
||||
) => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
return useQuery<PaginatedModelInfoResponse>({
|
||||
|
|
@ -48,6 +49,7 @@ export const useModelsInfo = (
|
|||
page,
|
||||
size,
|
||||
...(search && { search }),
|
||||
...(modelName && { modelName }),
|
||||
...(modelId && { modelId }),
|
||||
...(teamId && { teamId }),
|
||||
...(sortBy && { sortBy }),
|
||||
|
|
@ -70,6 +72,7 @@ export const useModelsInfo = (
|
|||
sortBy,
|
||||
sortOrder,
|
||||
excludeAutoRouters,
|
||||
modelName,
|
||||
),
|
||||
enabled: Boolean(accessToken && userId && userRole),
|
||||
});
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ interface ModelsInfoArgs {
|
|||
teamId?: string;
|
||||
sortBy?: string;
|
||||
sortOrder?: string;
|
||||
modelName?: string;
|
||||
}
|
||||
|
||||
const modelsInfoCalls: ModelsInfoArgs[] = [];
|
||||
|
|
@ -47,12 +48,14 @@ type UseModelsInfoArgs = [
|
|||
teamId?: string,
|
||||
sortBy?: string,
|
||||
sortOrder?: string,
|
||||
excludeAutoRouters?: boolean,
|
||||
modelName?: string,
|
||||
];
|
||||
|
||||
vi.mock("../../hooks/models/useModels", () => ({
|
||||
useModelsInfo: (...args: UseModelsInfoArgs) => {
|
||||
const [page, size, search, , teamId, sortBy, sortOrder] = args;
|
||||
const call: ModelsInfoArgs = { page, size, search, teamId, sortBy, sortOrder };
|
||||
const [page, size, search, , teamId, sortBy, sortOrder, , modelName] = args;
|
||||
const call: ModelsInfoArgs = { page, size, search, teamId, sortBy, sortOrder, modelName };
|
||||
modelsInfoCalls.push(call);
|
||||
return { ...modelsInfoResult, refetch: mockRefetch };
|
||||
},
|
||||
|
|
@ -260,6 +263,28 @@ describe("AllModelsTab", () => {
|
|||
expect(within(table).queryByText("gpt-4")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("asks the server for the exact selected model group so deployments beyond the first page are found", () => {
|
||||
render(<AllModelsTab {...defaultProps} selectedModelGroup="claude-opus" />);
|
||||
|
||||
expect(lastModelsInfoCall().modelName).toBe("claude-opus");
|
||||
expect(lastModelsInfoCall().search).toBeUndefined();
|
||||
});
|
||||
|
||||
it.each(["all", "wildcard"])("sends no exact model name for the %s pseudo group", (group) => {
|
||||
render(<AllModelsTab {...defaultProps} selectedModelGroup={group} />);
|
||||
|
||||
expect(lastModelsInfoCall().modelName).toBeUndefined();
|
||||
});
|
||||
|
||||
it("keeps the exact model group alongside a typed search", async () => {
|
||||
render(<AllModelsTab {...defaultProps} selectedModelGroup="claude-opus" />);
|
||||
|
||||
fireEvent.change(screen.getByPlaceholderText("Search model names…"), { target: { value: "opus" } });
|
||||
|
||||
await waitFor(() => expect(lastModelsInfoCall().search).toBe("opus"));
|
||||
expect(lastModelsInfoCall().modelName).toBe("claude-opus");
|
||||
});
|
||||
|
||||
it("resets search, filters, team and sorting from the drawer reset button", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<AllModelsTab {...defaultProps} selectedModelGroup="gpt-4" />);
|
||||
|
|
|
|||
|
|
@ -81,6 +81,11 @@ const AllModelsTab = ({
|
|||
}, [modelNameSearch, debouncedUpdateSearch]);
|
||||
|
||||
const teamIdForQuery = selectedTeamValue === PERSONAL_TEAM_VALUE ? undefined : selectedTeamValue;
|
||||
const isConcreteModelGroup =
|
||||
Boolean(selectedModelGroup) &&
|
||||
selectedModelGroup !== ALL_MODEL_GROUPS_VALUE &&
|
||||
selectedModelGroup !== WILDCARD_MODEL_GROUP_VALUE;
|
||||
const modelNameForQuery = isConcreteModelGroup ? selectedModelGroup ?? undefined : undefined;
|
||||
|
||||
const sortBy = useMemo(() => {
|
||||
if (sorting.length === 0) return undefined;
|
||||
|
|
@ -108,6 +113,7 @@ const AllModelsTab = ({
|
|||
// Auto-routers are routing constructs, not deployments; the sibling Auto-Routers tab
|
||||
// lists and manages them. Excluded server-side so total_count stays honest.
|
||||
true,
|
||||
modelNameForQuery,
|
||||
);
|
||||
const isLoading = isLoadingModelsInfo || isLoadingModelCostMap;
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import { act, renderHook, waitFor } from "@testing-library/react";
|
||||
import { withNuqsTestingAdapter, type UrlUpdateEvent } from "nuqs/adapters/testing";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { useModelDetailRouting } from "./detailNavigation";
|
||||
import { useModelDetailRouting, useModelGroupFilterRouting } from "./detailNavigation";
|
||||
|
||||
describe("useModelDetailRouting", () => {
|
||||
it("openModel sets ?model= with a history push", async () => {
|
||||
|
|
@ -54,3 +54,29 @@ describe("useModelDetailRouting", () => {
|
|||
expect(result.current.teamId).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("useModelGroupFilterRouting", () => {
|
||||
it("reads the selected group from ?model_group=", () => {
|
||||
const { result } = renderHook(() => useModelGroupFilterRouting(), {
|
||||
wrapper: withNuqsTestingAdapter({ searchParams: "?model_group=gpt-4.1" }),
|
||||
});
|
||||
expect(result.current.modelGroup).toBe("gpt-4.1");
|
||||
});
|
||||
|
||||
it("writes the selected group to ?model_group= and clears it on null", async () => {
|
||||
const onUrlUpdate = vi.fn<(event: UrlUpdateEvent) => void>();
|
||||
const { result } = renderHook(() => useModelGroupFilterRouting(), {
|
||||
wrapper: withNuqsTestingAdapter({ onUrlUpdate }),
|
||||
});
|
||||
await act(async () => {
|
||||
result.current.setModelGroup("claude-sonnet-5");
|
||||
});
|
||||
await waitFor(() => expect(onUrlUpdate).toHaveBeenCalled());
|
||||
expect(onUrlUpdate.mock.calls.at(-1)?.[0].searchParams.get("model_group")).toBe("claude-sonnet-5");
|
||||
|
||||
await act(async () => {
|
||||
result.current.setModelGroup(null);
|
||||
});
|
||||
await waitFor(() => expect(onUrlUpdate.mock.calls.at(-1)?.[0].searchParams.has("model_group")).toBe(false));
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import { parseAsString, useQueryStates } from "nuqs";
|
||||
import { parseAsString, useQueryState, useQueryStates } from "nuqs";
|
||||
import { useCallback } from "react";
|
||||
|
||||
export interface ModelDetailRouting {
|
||||
|
|
@ -41,3 +41,21 @@ export function useModelDetailRouting(): ModelDetailRouting {
|
|||
close,
|
||||
};
|
||||
}
|
||||
|
||||
export interface ModelGroupFilterRouting {
|
||||
modelGroup: string | null;
|
||||
setModelGroup: (modelGroup: string | null) => void;
|
||||
}
|
||||
|
||||
export function useModelGroupFilterRouting(): ModelGroupFilterRouting {
|
||||
const [modelGroup, setParam] = useQueryState("model_group", parseAsString);
|
||||
|
||||
const setModelGroup = useCallback(
|
||||
(next: string | null) => {
|
||||
void setParam(next);
|
||||
},
|
||||
[setParam],
|
||||
);
|
||||
|
||||
return { modelGroup, setModelGroup };
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,22 @@
|
|||
"use client";
|
||||
|
||||
import { useState } from "react";
|
||||
import AllModelsTab from "@/app/(dashboard)/models-and-endpoints/components/AllModelsTab";
|
||||
import { ALL_MODEL_GROUPS_VALUE } from "@/app/(dashboard)/models-and-endpoints/components/AllModelsTable";
|
||||
import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/useModelDashboardData";
|
||||
import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation";
|
||||
import {
|
||||
useModelDetailRouting,
|
||||
useModelGroupFilterRouting,
|
||||
} from "@/app/(dashboard)/models-and-endpoints/detailNavigation";
|
||||
|
||||
export default function AllModelsPanel() {
|
||||
const [selectedModelGroup, setSelectedModelGroup] = useState<string | null>(null);
|
||||
const { modelGroup, setModelGroup } = useModelGroupFilterRouting();
|
||||
const { availableModelGroups, availableModelAccessGroups } = useModelDashboardData();
|
||||
const { openModel, openTeam } = useModelDetailRouting();
|
||||
|
||||
return (
|
||||
<AllModelsTab
|
||||
selectedModelGroup={selectedModelGroup}
|
||||
setSelectedModelGroup={setSelectedModelGroup}
|
||||
selectedModelGroup={modelGroup}
|
||||
setSelectedModelGroup={(group) => setModelGroup(group === ALL_MODEL_GROUPS_VALUE ? null : group)}
|
||||
availableModelGroups={availableModelGroups}
|
||||
availableModelAccessGroups={availableModelAccessGroups}
|
||||
setSelectedModelId={openModel}
|
||||
|
|
|
|||
|
|
@ -104,6 +104,45 @@ describe("loginCall - storeLoginToken integration", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("modelInfoCall", () => {
|
||||
let currentFetch: typeof global.fetch;
|
||||
|
||||
beforeEach(() => {
|
||||
currentFetch = global.fetch;
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
global.fetch = currentFetch;
|
||||
});
|
||||
|
||||
it("sends the exact model name as the model query param and leaves search alone", async () => {
|
||||
const mockFetch = vi.fn().mockResolvedValue({ ok: true, json: vi.fn().mockResolvedValue({ data: [] }) } as any);
|
||||
global.fetch = mockFetch as any;
|
||||
|
||||
await Networking.modelInfoCall(
|
||||
"token",
|
||||
"user",
|
||||
"Admin",
|
||||
2,
|
||||
25,
|
||||
undefined,
|
||||
undefined,
|
||||
undefined,
|
||||
undefined,
|
||||
undefined,
|
||||
true,
|
||||
"gpt-4",
|
||||
);
|
||||
|
||||
const parsed = new URL(mockFetch.mock.calls[0][0] as string, "http://example.com");
|
||||
expect(parsed.pathname).toBe("/v2/model/info");
|
||||
expect(parsed.searchParams.get("model")).toBe("gpt-4");
|
||||
expect(parsed.searchParams.has("search")).toBe(false);
|
||||
expect(parsed.searchParams.get("page")).toBe("2");
|
||||
expect(parsed.searchParams.get("exclude_auto_routers")).toBe("true");
|
||||
});
|
||||
});
|
||||
|
||||
describe("daily activity helpers", () => {
|
||||
const startTime = new Date("2025-02-12T00:00:00.000Z");
|
||||
const endTime = new Date("2025-02-19T00:00:00.000Z");
|
||||
|
|
|
|||
|
|
@ -1677,6 +1677,7 @@ export const modelInfoCall = async (
|
|||
sortBy?: string,
|
||||
sortOrder?: string,
|
||||
excludeAutoRouters?: boolean,
|
||||
modelName?: string,
|
||||
) => {
|
||||
/**
|
||||
* Get all models on proxy
|
||||
|
|
@ -1690,6 +1691,9 @@ export const modelInfoCall = async (
|
|||
if (search && search.trim()) {
|
||||
params.append("search", search.trim());
|
||||
}
|
||||
if (modelName && modelName.trim()) {
|
||||
params.append("model", modelName.trim());
|
||||
}
|
||||
if (modelId && modelId.trim()) {
|
||||
params.append("modelId", modelId.trim());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,12 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
import { StatusBadge, type StatusTone } from "./status_badge";
|
||||
|
||||
const push = vi.fn();
|
||||
vi.mock("next/navigation", () => ({ useRouter: () => ({ push }) }));
|
||||
|
||||
describe("StatusBadge", () => {
|
||||
const toneClasses: Record<StatusTone, string[]> = {
|
||||
success: ["border-success/20", "bg-success/10", "text-success"],
|
||||
|
|
@ -39,4 +42,19 @@ describe("StatusBadge", () => {
|
|||
await user.hover(screen.getByText("Blocked"));
|
||||
expect(await screen.findByText("This key was blocked by SCIM")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders a tinted anchor that navigates client-side when href is given", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<StatusBadge tone="info" label="gpt-4.1" href="/models-and-endpoints?model_group=gpt-4.1" />);
|
||||
const link = screen.getByRole("link", { name: "gpt-4.1" });
|
||||
expect(link).toHaveAttribute("href", "/models-and-endpoints?model_group=gpt-4.1");
|
||||
expect(link).toHaveClass("text-info");
|
||||
await user.click(link);
|
||||
expect(push).toHaveBeenCalledWith("/models-and-endpoints?model_group=gpt-4.1");
|
||||
});
|
||||
|
||||
it("renders no anchor without an href", () => {
|
||||
render(<StatusBadge tone="info" label="gpt-4.1" />);
|
||||
expect(screen.queryByRole("link")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import * as React from "react";
|
||||
|
||||
import { useEntityLinkClick } from "@/components/shared/EntityLink";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
|
|
@ -23,15 +24,17 @@ interface StatusBadgeProps {
|
|||
tooltip?: React.ReactNode;
|
||||
dataTestId?: string;
|
||||
className?: string;
|
||||
href?: string;
|
||||
}
|
||||
|
||||
export function StatusBadge({ tone, label, tooltip, dataTestId, className }: StatusBadgeProps) {
|
||||
const badge = (
|
||||
<Badge
|
||||
variant="outline"
|
||||
data-testid={dataTestId}
|
||||
className={cn("whitespace-nowrap font-normal", TONE_CLASS[tone], className)}
|
||||
>
|
||||
export function StatusBadge({ tone, label, tooltip, dataTestId, className, href }: StatusBadgeProps) {
|
||||
const badgeClassName = cn("whitespace-nowrap font-normal", TONE_CLASS[tone], className);
|
||||
const badge = href ? (
|
||||
<LinkedStatusBadge href={href} dataTestId={dataTestId} className={badgeClassName}>
|
||||
{label}
|
||||
</LinkedStatusBadge>
|
||||
) : (
|
||||
<Badge variant="outline" data-testid={dataTestId} className={badgeClassName}>
|
||||
{label}
|
||||
</Badge>
|
||||
);
|
||||
|
|
@ -41,3 +44,25 @@ export function StatusBadge({ tone, label, tooltip, dataTestId, className }: Sta
|
|||
}
|
||||
return <CellTooltip content={tooltip} trigger={badge} />;
|
||||
}
|
||||
|
||||
interface LinkedStatusBadgeProps {
|
||||
href: string;
|
||||
dataTestId?: string;
|
||||
className: string;
|
||||
children: string;
|
||||
}
|
||||
|
||||
function LinkedStatusBadge({ href, dataTestId, className, children }: LinkedStatusBadgeProps) {
|
||||
const handleClick = useEntityLinkClick(href);
|
||||
|
||||
return (
|
||||
<Badge
|
||||
variant="outline"
|
||||
data-testid={dataTestId}
|
||||
className={cn("cursor-pointer hover:underline", className)}
|
||||
render={<a href={href} onClick={handleClick} />}
|
||||
>
|
||||
{children}
|
||||
</Badge>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -21,7 +21,10 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
|||
}),
|
||||
}));
|
||||
|
||||
vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) }));
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
serverRootPath: "",
|
||||
teamInfoCall: vi.fn(),
|
||||
teamMemberDeleteCall: vi.fn(),
|
||||
teamMemberAddCall: vi.fn(),
|
||||
|
|
@ -278,6 +281,36 @@ describe("TeamInfoView", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("links direct and access-group model badges to the models page filtered to that group", async () => {
|
||||
vi.mocked(networking.teamInfoCall).mockResolvedValue(
|
||||
createMockTeamData({
|
||||
models: ["gpt-4.1"],
|
||||
access_group_models: ["claude-sonnet-5"],
|
||||
access_group_details: [{ access_group_id: "ag-1", access_group_name: "prod", models: ["claude-sonnet-5"] }],
|
||||
}),
|
||||
);
|
||||
|
||||
renderWithProviders(<TeamInfoView {...defaultProps} />);
|
||||
|
||||
expect(await screen.findByRole("link", { name: "gpt-4.1" })).toHaveAttribute(
|
||||
"href",
|
||||
expect.stringContaining("/models-and-endpoints?model_group=gpt-4.1"),
|
||||
);
|
||||
expect(screen.getByRole("link", { name: "claude-sonnet-5" })).toHaveAttribute(
|
||||
"href",
|
||||
expect.stringContaining("/models-and-endpoints?model_group=claude-sonnet-5"),
|
||||
);
|
||||
});
|
||||
|
||||
it("keeps the all-proxy-models badge non-clickable", async () => {
|
||||
vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData({ models: ["all-proxy-models"] }));
|
||||
|
||||
renderWithProviders(<TeamInfoView {...defaultProps} />);
|
||||
|
||||
expect(await screen.findByText("All proxy models")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("link", { name: "All proxy models" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display loading state while fetching team data", () => {
|
||||
vi.mocked(networking.teamInfoCall).mockImplementation(() => new Promise(() => {}));
|
||||
|
||||
|
|
|
|||
|
|
@ -22,7 +22,9 @@ import type { ObjectPermission } from "@/components/object_permission_types";
|
|||
import { isProxyAdminRole } from "@/utils/roles";
|
||||
import { ArrowLeftIcon } from "@heroicons/react/outline";
|
||||
import { StatusBadge, type StatusTone } from "@/components/shared/table_cells/status_badge";
|
||||
import { BadgeLink } from "@/components/shared/BadgeLink";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { modelGroupHref } from "@/utils/entityLinks";
|
||||
import { Card } from "@/components/ui/card";
|
||||
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
|
||||
import { Input as UIInput } from "@/components/ui/input";
|
||||
|
|
@ -53,6 +55,7 @@ import {
|
|||
computeTeamModelBadges,
|
||||
normalizeTeamModelSelection,
|
||||
TeamAccessGroupModelGrant,
|
||||
TeamModelBadge,
|
||||
TeamModelBadgeKind,
|
||||
} from "./teamModelAccess";
|
||||
import MetadataKeyValueFields, {
|
||||
|
|
@ -111,6 +114,9 @@ const TEAM_MODEL_BADGE_TONES: Record<TeamModelBadgeKind, StatusTone> = {
|
|||
"access-group": "success",
|
||||
};
|
||||
|
||||
const teamModelBadgeHref = (badge: TeamModelBadge): string | undefined =>
|
||||
badge.kind === "direct" || badge.kind === "access-group" ? modelGroupHref(badge.label) : undefined;
|
||||
|
||||
export interface TeamMembership {
|
||||
user_id: string;
|
||||
team_id: string;
|
||||
|
|
@ -1006,7 +1012,11 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
(badge, index) => (
|
||||
<SimpleTooltip key={`${badge.kind}-${badge.label}-${index}`} content={badge.tooltip}>
|
||||
<span>
|
||||
<StatusBadge tone={TEAM_MODEL_BADGE_TONES[badge.kind]} label={badge.label} />
|
||||
<StatusBadge
|
||||
tone={TEAM_MODEL_BADGE_TONES[badge.kind]}
|
||||
label={badge.label}
|
||||
href={teamModelBadgeHref(badge)}
|
||||
/>
|
||||
</span>
|
||||
</SimpleTooltip>
|
||||
),
|
||||
|
|
@ -1727,9 +1737,9 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
<p className="font-medium">Models</p>
|
||||
<div className="flex flex-wrap gap-2 mt-1">
|
||||
{info.models.map((model, index) => (
|
||||
<Badge key={index} variant="secondary">
|
||||
<BadgeLink key={index} href={modelGroupHref(model)}>
|
||||
{model}
|
||||
</Badge>
|
||||
</BadgeLink>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
|
@ -1738,9 +1748,9 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
<p className="font-medium">Default Member Models</p>
|
||||
<div className="flex flex-wrap gap-2 mt-1">
|
||||
{info.default_team_member_models.map((model, index) => (
|
||||
<Badge key={index} variant="secondary">
|
||||
<BadgeLink key={index} href={modelGroupHref(model)}>
|
||||
{model}
|
||||
</Badge>
|
||||
</BadgeLink>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -561,6 +561,32 @@ describe("KeyInfoView", () => {
|
|||
);
|
||||
});
|
||||
|
||||
it("links each model chip to the models page filtered to that model group", async () => {
|
||||
const keyData = { ...MOCK_KEY_DATA, models: ["gpt-4.1", "anthropic/*"] };
|
||||
renderWithProviders(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId="test-key-id" onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
);
|
||||
|
||||
expect(await screen.findByRole("link", { name: "gpt-4.1" })).toHaveAttribute(
|
||||
"href",
|
||||
expect.stringContaining("/models-and-endpoints?model_group=gpt-4.1"),
|
||||
);
|
||||
expect(screen.getByRole("link", { name: "anthropic/*" })).toHaveAttribute(
|
||||
"href",
|
||||
expect.stringContaining("/models-and-endpoints?model_group=anthropic%2F*"),
|
||||
);
|
||||
});
|
||||
|
||||
it("keeps the all-proxy-models grant chip non-clickable", async () => {
|
||||
const keyData = { ...MOCK_KEY_DATA, models: ["all-proxy-models"] };
|
||||
renderWithProviders(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId="test-key-id" onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
);
|
||||
|
||||
expect((await screen.findAllByText("all-proxy-models")).length).toBeGreaterThan(0);
|
||||
expect(screen.queryByRole("link", { name: "all-proxy-models" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders no team link when the key has no team", async () => {
|
||||
renderWithProviders(
|
||||
<KeyInfoView
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ import { Card } from "@/components/ui/card";
|
|||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { EntityLink } from "@/components/shared/EntityLink";
|
||||
import { teamDetailHref } from "@/utils/entityLinks";
|
||||
import { modelGroupHref, teamDetailHref } from "@/utils/entityLinks";
|
||||
import { BadgeLink } from "@/components/shared/BadgeLink";
|
||||
import { KeyInfoHeader } from "./KeyInfoHeader";
|
||||
import KeySavingsTab from "./KeySavingsTab";
|
||||
import { useEffect, useState } from "react";
|
||||
|
|
@ -660,9 +661,9 @@ export default function KeyInfoView({
|
|||
<div className="mt-2 flex flex-wrap gap-2">
|
||||
{currentKeyData.models && currentKeyData.models.length > 0 ? (
|
||||
currentKeyData.models.map((model, index) => (
|
||||
<Badge key={index} variant="secondary" className="min-w-0 break-words">
|
||||
<BadgeLink key={index} href={modelGroupHref(model)} className="min-w-0 break-words">
|
||||
{model}
|
||||
</Badge>
|
||||
</BadgeLink>
|
||||
))
|
||||
) : (
|
||||
<p className="text-sm">No models specified</p>
|
||||
|
|
@ -996,9 +997,9 @@ export default function KeyInfoView({
|
|||
<div className="flex flex-wrap gap-2 mt-1">
|
||||
{currentKeyData.models && currentKeyData.models.length > 0 ? (
|
||||
currentKeyData.models.map((model, index) => (
|
||||
<span key={index} className="px-2 py-1 bg-info/15 rounded-sm text-xs">
|
||||
<BadgeLink key={index} href={modelGroupHref(model)} className="min-w-0 break-words">
|
||||
{model}
|
||||
</span>
|
||||
</BadgeLink>
|
||||
))
|
||||
) : (
|
||||
<p className="text-sm">No models specified</p>
|
||||
|
|
|
|||
19
ui/litellm-dashboard/src/utils/entityLinks.test.ts
Normal file
19
ui/litellm-dashboard/src/utils/entityLinks.test.ts
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
vi.mock("@/components/networking", () => ({ serverRootPath: "" }));
|
||||
|
||||
import { modelGroupHref } from "./entityLinks";
|
||||
|
||||
describe("modelGroupHref", () => {
|
||||
it("targets the models page filtered to the encoded model group", () => {
|
||||
expect(modelGroupHref("gpt-4.1")).toMatch(/\/models-and-endpoints\?model_group=gpt-4\.1$/);
|
||||
expect(modelGroupHref("openai/*")).toMatch(/\?model_group=openai%2F\*$/);
|
||||
});
|
||||
|
||||
it.each(["all-proxy-models", "all-team-models", "no-default-models"])(
|
||||
"returns no href for the %s grant sentinel",
|
||||
(sentinel) => {
|
||||
expect(modelGroupHref(sentinel)).toBeUndefined();
|
||||
},
|
||||
);
|
||||
});
|
||||
|
|
@ -1,5 +1,11 @@
|
|||
import { migratedHref } from "@/utils/migratedPages";
|
||||
|
||||
const MODEL_GRANT_SENTINELS: ReadonlySet<string> = new Set([
|
||||
"all-proxy-models",
|
||||
"all-team-models",
|
||||
"no-default-models",
|
||||
]);
|
||||
|
||||
export function teamDetailHref(teamId: string): string {
|
||||
return `${migratedHref("teams")}?team=${encodeURIComponent(teamId)}`;
|
||||
}
|
||||
|
|
@ -15,3 +21,8 @@ export function userDetailHref(userId: string): string {
|
|||
export function orgDetailHref(orgId: string): string {
|
||||
return `${migratedHref("organizations")}?org=${encodeURIComponent(orgId)}`;
|
||||
}
|
||||
|
||||
export function modelGroupHref(modelGroup: string): string | undefined {
|
||||
if (MODEL_GRANT_SENTINELS.has(modelGroup)) return undefined;
|
||||
return `${migratedHref("models-and-endpoints")}?model_group=${encodeURIComponent(modelGroup)}`;
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue