add model picker to create provider wizard

This commit is contained in:
miloschwartz
2026-08-13 15:08:10 -04:00
parent d9a9ae14fd
commit 9b10292e02
10 changed files with 341 additions and 81 deletions
+1
View File
@@ -1864,6 +1864,7 @@
"aiProviderErrorNoUpdate": "AI provider is not available to update",
"aiProviderModels": "Models",
"aiProviderModelsDescription": "Define allow and block lists for this provider. Requests must match an allow entry and must not match a block entry.",
"aiProviderCreateModelsDescription": "Choose which models this provider can serve. An empty allow list denies all traffic. You can add or change models later.",
"aiProviderModelsPlaceholder": "Search models or type a custom key",
"aiProviderModelsAllow": "Allow List",
"aiProviderModelsAllowDescription": "Models that may be used through this provider. Empty means deny all.",
+30
View File
@@ -284,6 +284,36 @@ export class AiModelCatalog {
export const aiModelCatalog = new AiModelCatalog();
export function listCatalogModelsForType(
type: AiProviderType,
query?: string
): { model: string }[] {
const catalogProvider = getCatalogProviderForType(type);
let models = catalogProvider
? aiModelCatalog.list(catalogProvider).map((entry) => ({
model: entry.model
}))
: [];
if (query) {
const q = query.toLowerCase();
models = models.filter((m) => m.model.toLowerCase().includes(q));
}
const seen = new Set<string>();
models = models.filter((m) => {
if (seen.has(m.model)) {
return false;
}
seen.add(m.model);
return true;
});
models.sort((a, b) => a.model.localeCompare(b.model));
return models;
}
/**
* Loads the AI model pricing catalog into memory and schedules periodic
* background refreshes. Call once at server startup.
+1
View File
@@ -6,6 +6,7 @@ export * from "./deleteAiProvider";
export * from "./createAiModel";
export * from "./listAiModels";
export * from "./listCatalogModels";
export * from "./listCatalogModelsByType";
export * from "./getAiModel";
export * from "./updateAiModel";
export * from "./deleteAiModel";
+4 -31
View File
@@ -8,10 +8,7 @@ import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { eq } from "drizzle-orm";
import {
aiModelCatalog,
getCatalogProviderForType
} from "@server/lib/aiModelCatalog";
import { listCatalogModelsForType } from "@server/lib/aiModelCatalog";
import type { AiProviderType } from "@server/lib/aiProviderDefaults";
import type { ListCatalogModelsResponse } from "@server/routers/aiProvider/types";
@@ -86,35 +83,11 @@ export async function listCatalogModels(
);
}
const catalogProvider = getCatalogProviderForType(
provider.type as AiProviderType
const models = listCatalogModelsForType(
provider.type as AiProviderType,
parsedQuery.data.query
);
let models = catalogProvider
? aiModelCatalog.list(catalogProvider).map((entry) => ({
model: entry.model
}))
: [];
const { query } = parsedQuery.data;
if (query) {
const q = query.toLowerCase();
models = models.filter((m) => m.model.toLowerCase().includes(q));
}
// Deduplicate model keys (catalog may have duplicates after provider
// normalization, e.g. bedrock + bedrock_converse).
const seen = new Set<string>();
models = models.filter((m) => {
if (seen.has(m.model)) {
return false;
}
seen.add(m.model);
return true;
});
models.sort((a, b) => a.model.localeCompare(b.model));
return response<ListCatalogModelsResponse>(res, {
data: { models },
success: true,
@@ -0,0 +1,81 @@
import { Request, Response, NextFunction } from "express";
import { z } from "zod";
import response from "@server/lib/response";
import HttpCode from "@server/types/HttpCode";
import createHttpError from "http-errors";
import logger from "@server/logger";
import { fromError } from "zod-validation-error";
import { OpenAPITags, registry } from "@server/openApi";
import { listCatalogModelsForType } from "@server/lib/aiModelCatalog";
import type { ListCatalogModelsResponse } from "@server/routers/aiProvider/types";
import { aiProviderTypeSchema } from "@server/routers/aiProvider/validation";
const paramsSchema = z.strictObject({
orgId: z.string().nonempty()
});
const querySchema = z.strictObject({
type: aiProviderTypeSchema,
query: z.string().optional()
});
registry.registerPath({
method: "get",
path: "/org/{orgId}/ai-catalog-models",
description:
"List known catalog models for an AI provider type. Used for model key suggestions before a provider exists.",
tags: [OpenAPITags.AiModel],
request: {
params: paramsSchema,
query: querySchema
},
responses: {
200: {
description: "Successful response"
}
}
});
export async function listCatalogModelsByType(
req: Request,
res: Response,
next: NextFunction
): Promise<any> {
try {
const parsedParams = paramsSchema.safeParse(req.params);
if (!parsedParams.success) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
fromError(parsedParams.error).toString()
)
);
}
const parsedQuery = querySchema.safeParse(req.query);
if (!parsedQuery.success) {
return next(
createHttpError(
HttpCode.BAD_REQUEST,
fromError(parsedQuery.error).toString()
)
);
}
const { type, query } = parsedQuery.data;
const models = listCatalogModelsForType(type, query);
return response<ListCatalogModelsResponse>(res, {
data: { models },
success: true,
error: false,
message: "Catalog models retrieved successfully",
status: HttpCode.OK
});
} catch (error) {
logger.error(error);
return next(
createHttpError(HttpCode.INTERNAL_SERVER_ERROR, "An error occurred")
);
}
}
+7
View File
@@ -1645,6 +1645,13 @@ authenticated.get(
aiProvider.listCatalogModels
);
authenticated.get(
"/org/:orgId/ai-catalog-models",
verifyOrgAccess,
verifyUserHasAction(ActionsEnum.listAiModels),
aiProvider.listCatalogModelsByType
);
authenticated.get(
"/ai-model/:modelId",
verifyAiModelAccess,
@@ -11,12 +11,11 @@ import {
SettingsSectionTitle
} from "@app/components/Settings";
import {
AiProviderModelListEditor,
type AiProviderModelListItem,
type ModelListType
} from "@app/components/AiProviderModelListEditor";
import { AiProviderModelsLists } from "@app/components/AiProviderModelsLists";
import { Button } from "@app/components/ui/button";
import { Label } from "@app/components/ui/label";
import { useAiProviderContext } from "@app/hooks/useAiProviderContext";
import { useEnvContext } from "@app/hooks/useEnvContext";
import { toast } from "@app/hooks/useToast";
@@ -65,15 +64,6 @@ export default function AiProviderModelsPage() {
[catalogQuery.data]
);
const allowExcludeKeys = useMemo(
() => new Set(blockItems.map((item) => item.modelKey)),
[blockItems]
);
const blockExcludeKeys = useMemo(
() => new Set(allowItems.map((item) => item.modelKey)),
[allowItems]
);
useEffect(() => {
if (!modelsQuery.data) return;
setAllowItems(
@@ -232,45 +222,15 @@ export default function AiProviderModelsPage() {
<SettingsSectionBody>
<SettingsSectionForm>
<div className="space-y-2">
<Label>{t("aiProviderModelsAllow")}</Label>
<AiProviderModelListEditor
orgId={provider.orgId}
listType="allow"
items={allowItems}
onChange={setAllowItems}
catalogModels={catalogModels}
excludeKeys={allowExcludeKeys}
disabled={modelsQuery.isLoading}
emptyMessage={t("aiProviderModelsAllowEmpty")}
addPlaceholder={t(
"aiProviderModelsAllowPlaceholder"
)}
/>
<p className="text-sm text-muted-foreground">
{t("aiProviderModelsAllowDescription")}
</p>
</div>
<div className="space-y-2">
<Label>{t("aiProviderModelsBlock")}</Label>
<AiProviderModelListEditor
orgId={provider.orgId}
listType="block"
items={blockItems}
onChange={setBlockItems}
catalogModels={catalogModels}
excludeKeys={blockExcludeKeys}
disabled={modelsQuery.isLoading}
emptyMessage={t("aiProviderModelsBlockEmpty")}
addPlaceholder={t(
"aiProviderModelsBlockPlaceholder"
)}
/>
<p className="text-sm text-muted-foreground">
{t("aiProviderModelsBlockDescription")}
</p>
</div>
<AiProviderModelsLists
orgId={provider.orgId}
allowItems={allowItems}
onAllowChange={setAllowItems}
blockItems={blockItems}
onBlockChange={setBlockItems}
catalogModels={catalogModels}
disabled={modelsQuery.isLoading}
/>
</SettingsSectionForm>
</SettingsSectionBody>
@@ -21,6 +21,11 @@ import {
import HeaderTitle from "@app/components/SettingsSectionTitle";
import { AiProviderAuthTypeSelect } from "@app/components/AiProviderAuthTypeSelect";
import { AiProviderCapabilitiesSelect } from "@app/components/AiProviderCapabilitiesSelect";
import {
type AiProviderModelListItem,
type ModelListType
} from "@app/components/AiProviderModelListEditor";
import { AiProviderModelsLists } from "@app/components/AiProviderModelsLists";
import {
AiProviderTypeSelect,
aiProviderTypeLabelMap
@@ -53,7 +58,9 @@ import {
} from "@app/lib/aiProviderFormSchema";
import { zodResolver } from "@hookform/resolvers/zod";
import { authTypeRequiresApiKey } from "@app/lib/aiProviderDefaults";
import { aiProviderQueries } from "@app/lib/queries";
import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/types";
import { useQuery } from "@tanstack/react-query";
import type { AxiosResponse } from "axios";
import { useTranslations } from "next-intl";
import { useParams, useRouter } from "next/navigation";
@@ -69,6 +76,8 @@ export default function CreateAiProviderPage() {
const t = useTranslations();
const [loading, setLoading] = useState(false);
const [headersValid, setHeadersValid] = useState(true);
const [allowItems, setAllowItems] = useState<AiProviderModelListItem[]>([]);
const [blockItems, setBlockItems] = useState<AiProviderModelListItem[]>([]);
const targetsRef = useRef<LocalTarget[]>([]);
const formSchema = useMemo(() => createAiProviderCreateFormSchema(t), [t]);
@@ -98,6 +107,17 @@ export default function CreateAiProviderPage() {
const showTargets = providerType === "custom" && routingMode === "target";
const showApiKey = authTypeRequiresApiKey(authType ?? "bearer");
const catalogQuery = useQuery(
aiProviderQueries.catalogModelsByType({
orgId,
type: providerType
})
);
const catalogModels = useMemo(
() => (catalogQuery.data ?? []).map((entry) => entry.model),
[catalogQuery.data]
);
async function createTargets(
providerId: number,
localTargets: LocalTarget[]
@@ -135,6 +155,19 @@ export default function CreateAiProviderPage() {
}
}
async function createModels(
providerId: number,
items: { modelKey: string; listType: ModelListType }[]
) {
for (const item of items) {
await api.put(`/ai-provider/${providerId}/model`, {
modelKey: item.modelKey,
name: item.modelKey,
listType: item.listType
});
}
}
async function onSubmit(values: AiProviderFormValues) {
const targets = targetsRef.current;
@@ -157,6 +190,31 @@ export default function CreateAiProviderPage() {
}
}
const nextAllow = new Set(
allowItems.map((item) => item.modelKey.trim()).filter(Boolean)
);
const nextBlock = new Set(
blockItems.map((item) => item.modelKey.trim()).filter(Boolean)
);
const overlap = [...nextAllow].filter((key) => nextBlock.has(key));
if (overlap.length > 0) {
toast({
variant: "destructive",
title: t("aiProviderErrorCreate"),
description: t("aiProviderModelsOverlapError", {
keys: overlap.join(", ")
})
});
return;
}
const modelItems = [...allowItems, ...blockItems]
.map((item) => ({
modelKey: item.modelKey.trim(),
listType: item.listType
}))
.filter((item) => item.modelKey);
setLoading(true);
try {
const res = await api.put<
@@ -184,6 +242,25 @@ export default function CreateAiProviderPage() {
}
}
if (modelItems.length > 0) {
try {
await createModels(providerId, modelItems);
} catch (e) {
toast({
variant: "destructive",
title: t("aiProviderErrorCreate"),
description: formatAxiosError(
e,
t("aiProviderErrorCreate")
)
});
router.push(
`/${orgId}/settings/ai-providers/${providerId}/models`
);
return;
}
}
toast({
title: t("success"),
description: t("aiProviderCreated")
@@ -253,6 +330,12 @@ export default function CreateAiProviderPage() {
field.onChange(
value
);
setAllowItems(
[]
);
setBlockItems(
[]
);
form.setValue(
"upstreamUrl",
emptyUpstreamForType(
@@ -689,6 +772,30 @@ export default function CreateAiProviderPage() {
</SettingsSectionForm>
</SettingsSectionBody>
</SettingsSection>
<SettingsSection>
<SettingsSectionHeader>
<SettingsSectionTitle>
{t("aiProviderModels")}
</SettingsSectionTitle>
<SettingsSectionDescription>
{t("aiProviderCreateModelsDescription")}
</SettingsSectionDescription>
</SettingsSectionHeader>
<SettingsSectionBody>
<SettingsSectionForm>
<AiProviderModelsLists
orgId={orgId}
allowItems={allowItems}
onAllowChange={setAllowItems}
blockItems={blockItems}
onBlockChange={setBlockItems}
catalogModels={catalogModels}
/>
</SettingsSectionForm>
</SettingsSectionBody>
</SettingsSection>
</SettingsContainer>
<div className="flex justify-end space-x-2 mt-8">
+80
View File
@@ -0,0 +1,80 @@
"use client";
import {
AiProviderModelListEditor,
type AiProviderModelListItem
} from "@app/components/AiProviderModelListEditor";
import { Label } from "@app/components/ui/label";
import { useTranslations } from "next-intl";
import { useMemo } from "react";
export type AiProviderModelsListsProps = {
orgId: string;
allowItems: AiProviderModelListItem[];
onAllowChange: (items: AiProviderModelListItem[]) => void;
blockItems: AiProviderModelListItem[];
onBlockChange: (items: AiProviderModelListItem[]) => void;
catalogModels: string[];
disabled?: boolean;
};
export function AiProviderModelsLists({
orgId,
allowItems,
onAllowChange,
blockItems,
onBlockChange,
catalogModels,
disabled
}: AiProviderModelsListsProps) {
const t = useTranslations();
const allowExcludeKeys = useMemo(
() => new Set(blockItems.map((item) => item.modelKey)),
[blockItems]
);
const blockExcludeKeys = useMemo(
() => new Set(allowItems.map((item) => item.modelKey)),
[allowItems]
);
return (
<div className="space-y-6">
<div className="space-y-2">
<Label>{t("aiProviderModelsAllow")}</Label>
<AiProviderModelListEditor
orgId={orgId}
listType="allow"
items={allowItems}
onChange={onAllowChange}
catalogModels={catalogModels}
excludeKeys={allowExcludeKeys}
disabled={disabled}
emptyMessage={t("aiProviderModelsAllowEmpty")}
addPlaceholder={t("aiProviderModelsAllowPlaceholder")}
/>
<p className="text-sm text-muted-foreground">
{t("aiProviderModelsAllowDescription")}
</p>
</div>
<div className="space-y-2">
<Label>{t("aiProviderModelsBlock")}</Label>
<AiProviderModelListEditor
orgId={orgId}
listType="block"
items={blockItems}
onChange={onBlockChange}
catalogModels={catalogModels}
excludeKeys={blockExcludeKeys}
disabled={disabled}
emptyMessage={t("aiProviderModelsBlockEmpty")}
addPlaceholder={t("aiProviderModelsBlockPlaceholder")}
/>
<p className="text-sm text-muted-foreground">
{t("aiProviderModelsBlockDescription")}
</p>
</div>
</div>
);
}
+20
View File
@@ -75,6 +75,7 @@ import type {
ListAiProvidersResponse,
ListCatalogModelsResponse
} from "@server/routers/aiProvider/types";
import type { AiProviderType } from "@app/lib/aiProviderDefaults";
import type { ListAiBudgetsByScopeResponse } from "@server/routers/aiBudget/types";
import {
getAiBudgetScopeListPath,
@@ -1495,6 +1496,25 @@ export const aiProviderQueries = {
return res.data.data.models;
}
}),
catalogModelsByType: ({
orgId,
type
}: {
orgId: string;
type: AiProviderType;
}) =>
queryOptions({
queryKey: ["AI_PROVIDERS", orgId, "CATALOG_MODELS", type] as const,
queryFn: async ({ signal, meta }) => {
const res = await meta!.api.get<
AxiosResponse<ListCatalogModelsResponse>
>(`/org/${orgId}/ai-catalog-models`, {
params: { type },
signal
});
return res.data.data.models;
}
}),
orgProviders: ({ orgId, query }: { orgId: string; query?: string }) =>
queryOptions({
queryKey: ["AI_PROVIDERS", orgId, "LIST", query ?? ""] as const,