Merge remote-tracking branch 'origin/litellm_internal_staging' into devin_ai_fix_model_new_read_replica_lag_38556

This commit is contained in:
mateo-berri 2026-08-29 11:29:07 -07:00
commit a8c36e8307
22 changed files with 504 additions and 32 deletions

View file

@ -9,6 +9,7 @@ from typing import Final
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL
from litellm.proxy._types import ProxyErrorTypes, ProxyException
SUGGEST_TOOL: Final = {
"type": "function",
@ -60,6 +61,18 @@ class AiPolicySuggester:
system_prompt: Final = self._build_system_prompt(templates)
user_prompt: Final = self._build_user_prompt(attack_examples, description)
model = model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
custom_llm_provider: Final = model.split("/", 1)[0] if "/" in model else None
supported_params: Final = litellm.get_supported_openai_params(
model=model,
custom_llm_provider=custom_llm_provider,
)
if supported_params is not None and "tools" not in supported_params:
raise ProxyException(
message=(f"AI policy suggestion requires tool calling; model '{model}' does not support it"),
type=ProxyErrorTypes.validation_error.value,
param="model",
code=400,
)
try:
response: Final = await litellm.acompletion(
@ -74,6 +87,7 @@ class AiPolicySuggester:
"function": {"name": "select_policy_templates"},
},
temperature=0.2,
drop_params=True,
)
tool_calls: Final = response.choices[0].message.tool_calls

View file

@ -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 (
@ -12801,6 +12801,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
@ -12817,7 +12818,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)}}
@ -12864,6 +12867,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.
@ -12884,6 +12888,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.
@ -12941,7 +12950,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,
@ -12953,6 +12963,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:
@ -13506,7 +13517,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(
@ -13518,6 +13529,7 @@ async def model_info_v2(
page=page,
size=size,
sort_by=sortBy,
model_name=model,
)
if user_models_only:
@ -14032,7 +14044,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

View file

@ -7,6 +7,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.policy_endpoints.ai_policy_suggester import (
SUGGEST_TOOL,
AiPolicySuggester,
@ -234,6 +237,7 @@ class TestAiPolicySuggester:
call_kwargs = mock_acompletion.call_args.kwargs
assert call_kwargs["model"] == "gpt-4o-mini"
assert call_kwargs["temperature"] == 0.2
assert call_kwargs["drop_params"] is True
assert len(call_kwargs["tools"]) == 1
assert call_kwargs["tools"][0]["function"]["name"] == "select_policy_templates"
assert (
@ -242,3 +246,76 @@ class TestAiPolicySuggester:
assert len(call_kwargs["messages"]) == 2
assert call_kwargs["messages"][0]["role"] == "system"
assert call_kwargs["messages"][1]["role"] == "user"
class TestSuggesterRejectsModelsWithoutToolCalling:
@pytest.mark.asyncio
async def test_a_tools_less_model_is_rejected(self, local_model_cost_map):
with pytest.raises(ProxyException) as exc:
await AiPolicySuggester().suggest(
templates=SAMPLE_TEMPLATES,
attack_examples=["Ignore all previous instructions"],
description="Block prompt injection attempts",
model="perplexity/sonar",
)
assert int(exc.value.code) == 400
assert exc.value.param == "model"
assert "tool calling" in exc.value.message
def test_a_model_without_forced_tool_choice_support_remains_eligible(self, local_model_cost_map):
supported_params = litellm.get_supported_openai_params(
model="amazon.nova-pro-v1:0",
custom_llm_provider="bedrock",
)
assert supported_params is not None
assert "tools" in supported_params
assert "tool_choice" not in supported_params
class TestSuggesterToleratesAModelThatRefusesItsSamplingParams:
"""The model is operator-supplied, so it can be a reasoning model whose only accepted
temperature is 1. This call pins temperature=0.2 for tool-selection determinism, which such
a model rejects outright: without drop_params litellm raises UnsupportedParamsError and the
whole suggestion fails rather than degrading. Every other internal LLM call in the proxy
already opts in through judge_acompletion; this one was the exception.
"""
@pytest.mark.asyncio
async def test_a_reasoning_model_gets_past_param_mapping(self, monkeypatch, local_model_cost_map):
"""Drives the real entry point with no patching and no network. Which exception escapes is
the discriminator: param mapping runs before any credential check, so UnsupportedParamsError
means the call died on the pinned temperature, while AuthenticationError means it survived
that and got as far as needing a key. Asserting the latter is what the caller observes.
"""
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
with pytest.raises(litellm.AuthenticationError):
await AiPolicySuggester().suggest(
templates=SAMPLE_TEMPLATES,
attack_examples=["My SSN is 123-45-6789"],
description="",
model="gpt-5.6-terra",
)
def test_the_pinned_temperature_is_what_such_a_model_refuses(self, local_model_cost_map):
"""The other half of the discriminator above: the same temperature this call pins is
exactly what the model rejects, and drop_params is what removes it."""
from litellm.utils import get_optional_params
optional_params = get_optional_params(
model="gpt-5.6-terra",
custom_llm_provider="openai",
temperature=0.2,
tools=[SUGGEST_TOOL],
tool_choice={"type": "function", "function": {"name": "select_policy_templates"}},
drop_params=True,
)
assert "temperature" not in optional_params
assert optional_params["tools"] == [SUGGEST_TOOL]
assert optional_params["tool_choice"] == {
"type": "function",
"function": {"name": "select_policy_templates"},
}

View file

@ -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

View file

@ -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():
"""

View file

@ -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,
);
});

View file

@ -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),
});

View file

@ -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" />);

View file

@ -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;

View file

@ -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));
});
});

View file

@ -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 };
}

View file

@ -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}

View file

@ -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");

View file

@ -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());
}

View file

@ -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();
});
});

View file

@ -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>
);
}

View file

@ -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(() => {}));

View file

@ -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>

View file

@ -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

View file

@ -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>

View 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();
},
);
});

View file

@ -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)}`;
}