diff --git a/messages/en-US.json b/messages/en-US.json index 137440d42..a3bfe73b8 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -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.", diff --git a/server/lib/aiModelCatalog.ts b/server/lib/aiModelCatalog.ts index a74378e6b..5f2f7fb94 100644 --- a/server/lib/aiModelCatalog.ts +++ b/server/lib/aiModelCatalog.ts @@ -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(); + 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. diff --git a/server/routers/aiProvider/index.ts b/server/routers/aiProvider/index.ts index 835991fbd..32378d3d9 100644 --- a/server/routers/aiProvider/index.ts +++ b/server/routers/aiProvider/index.ts @@ -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"; diff --git a/server/routers/aiProvider/listCatalogModels.ts b/server/routers/aiProvider/listCatalogModels.ts index 9a3241276..5320ada8a 100644 --- a/server/routers/aiProvider/listCatalogModels.ts +++ b/server/routers/aiProvider/listCatalogModels.ts @@ -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(); - 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(res, { data: { models }, success: true, diff --git a/server/routers/aiProvider/listCatalogModelsByType.ts b/server/routers/aiProvider/listCatalogModelsByType.ts new file mode 100644 index 000000000..4fbca54e9 --- /dev/null +++ b/server/routers/aiProvider/listCatalogModelsByType.ts @@ -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 { + 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(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") + ); + } +} diff --git a/server/routers/external.ts b/server/routers/external.ts index c391665b9..17cf0d921 100644 --- a/server/routers/external.ts +++ b/server/routers/external.ts @@ -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, diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx index f8ace5cb2..b85c0076c 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx @@ -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() { -
- - -

- {t("aiProviderModelsAllowDescription")} -

-
- -
- - -

- {t("aiProviderModelsBlockDescription")} -

-
+
diff --git a/src/app/[orgId]/settings/ai-providers/create/page.tsx b/src/app/[orgId]/settings/ai-providers/create/page.tsx index 722874028..7684b2d48 100644 --- a/src/app/[orgId]/settings/ai-providers/create/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/create/page.tsx @@ -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([]); + const [blockItems, setBlockItems] = useState([]); const targetsRef = useRef([]); 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() { + + + + + {t("aiProviderModels")} + + + {t("aiProviderCreateModelsDescription")} + + + + + + + + +
diff --git a/src/components/AiProviderModelsLists.tsx b/src/components/AiProviderModelsLists.tsx new file mode 100644 index 000000000..8ed7f6fef --- /dev/null +++ b/src/components/AiProviderModelsLists.tsx @@ -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 ( +
+
+ + +

+ {t("aiProviderModelsAllowDescription")} +

+
+ +
+ + +

+ {t("aiProviderModelsBlockDescription")} +

+
+
+ ); +} diff --git a/src/lib/queries.ts b/src/lib/queries.ts index 56a41cee8..c58c6c567 100644 --- a/src/lib/queries.ts +++ b/src/lib/queries.ts @@ -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 + >(`/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,