From 7759d878352d09ab1cfde3d8cd7b89fc0e13fb39 Mon Sep 17 00:00:00 2001 From: miloschwartz Date: Tue, 4 Aug 2026 16:54:23 -0400 Subject: [PATCH] add basic ui for private inference resource --- messages/en-US.json | 14 ++ server/db/pg/schema/schema.ts | 8 +- server/db/sqlite/schema/schema.ts | 8 +- server/lib/aiInferenceResource.ts | 79 ++++-- server/routers/aiGateway/chatCompletions.ts | 32 --- server/routers/resource/createResource.ts | 6 +- .../resource/setResourceAiProviders.ts | 6 +- .../siteResource/createSiteResource.ts | 2 +- .../setSiteResourceAiProviders.ts | 6 +- .../ai-providers/[providerId]/layout.tsx | 4 + .../ai-providers/[providerId]/models/page.tsx | 150 +++++++++++ .../private/[niceId]/general/page.tsx | 30 ++- .../private/[niceId]/inference/page.tsx | 112 +-------- .../resources/private/[niceId]/layout.tsx | 45 ++-- .../private/[niceId]/providers/page.tsx | 232 ++++++++++++++++++ .../resources/private/create/page.tsx | 91 ++++++- src/components/AiProvidersSelector.tsx | 61 +++++ src/lib/privateResourceForm.ts | 92 +++++-- src/lib/queries.ts | 58 +++++ 19 files changed, 828 insertions(+), 208 deletions(-) create mode 100644 src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx create mode 100644 src/app/[orgId]/settings/resources/private/[niceId]/providers/page.tsx create mode 100644 src/components/AiProvidersSelector.tsx diff --git a/messages/en-US.json b/messages/en-US.json index 81a60f0a1..2fc4e778e 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1714,6 +1714,20 @@ "aiProviderQuestionRemove": "Are you sure you want to delete this AI provider?", "aiProviderMessageRemove": "This will permanently delete the provider and its models and targets. This cannot be undone.", "aiProviderErrorNoUpdate": "AI provider is not available to update", + "aiProviderModels": "Models", + "aiProviderModelsDescription": "Define model names available on this provider. Requests must use one of these model keys.", + "aiProviderModelsPlaceholder": "Type a model name and press Enter", + "aiProviderModelsUpdated": "Models updated", + "aiProviderModelsErrorUpdate": "Failed to update models", + "aiResourceProviders": "Providers", + "aiResourceProvidersDescription": "Choose which AI providers this inference resource can use", + "aiResourceProvidersHelp": "Models must be defined on each provider. Model names cannot overlap across selected providers.", + "aiResourceProvidersSelect": "Select providers", + "aiResourceProvidersEmpty": "No AI providers found", + "aiResourceProvidersRequired": "Select at least one AI provider", + "aiResourceProvidersUpdated": "Providers updated", + "aiResourceProvidersErrorUpdate": "Failed to update providers", + "aiResourceAliasRequired": "Alias is required for inference resources", "sidebarApiKeys": "API Keys", "sidebarProvisioning": "Provisioning", "sidebarSettings": "Settings", diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index b323ae411..868b2a121 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -232,9 +232,9 @@ export const resourceAiProviders = pgTable( .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelAccessMode: varchar("modelAccessMode") - .$type<"passthrough" | "catalog" | "allowlist">() + .$type<"catalog" | "allowlist">() .notNull() - .default("passthrough") + .default("catalog") }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -523,9 +523,9 @@ export const siteResourceAiProviders = pgTable( .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelAccessMode: varchar("modelAccessMode") - .$type<"passthrough" | "catalog" | "allowlist">() + .$type<"catalog" | "allowlist">() .notNull() - .default("passthrough") + .default("catalog") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index c8f25dfed..2ae941895 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -229,9 +229,9 @@ export const resourceAiProviders = sqliteTable( .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelAccessMode: text("modelAccessMode") - .$type<"passthrough" | "catalog" | "allowlist">() + .$type<"catalog" | "allowlist">() .notNull() - .default("passthrough") + .default("catalog") }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -508,9 +508,9 @@ export const siteResourceAiProviders = sqliteTable( .notNull() .references(() => aiProviders.providerId, { onDelete: "cascade" }), modelAccessMode: text("modelAccessMode") - .$type<"passthrough" | "catalog" | "allowlist">() + .$type<"catalog" | "allowlist">() .notNull() - .default("passthrough") + .default("catalog") }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); diff --git a/server/lib/aiInferenceResource.ts b/server/lib/aiInferenceResource.ts index 4175a11e8..6e0f490d8 100644 --- a/server/lib/aiInferenceResource.ts +++ b/server/lib/aiInferenceResource.ts @@ -13,11 +13,7 @@ import { z } from "zod"; type DbOrTrx = Transaction | typeof db; -export const modelAccessModeSchema = z.enum([ - "passthrough", - "catalog", - "allowlist" -]); +export const modelAccessModeSchema = z.enum(["catalog", "allowlist"]); export type ModelAccessMode = z.infer; @@ -50,10 +46,7 @@ function normalizeAttachments( ): ResourceAiProviderAttachment[] { const byProvider = new Map(); for (const input of inputs) { - byProvider.set( - input.providerId, - input.modelAccessMode ?? "passthrough" - ); + byProvider.set(input.providerId, input.modelAccessMode ?? "catalog"); } return [...byProvider.entries()].map(([providerId, modelAccessMode]) => ({ providerId, @@ -61,9 +54,61 @@ function normalizeAttachments( })); } +/** + * Ensure enabled catalog modelKeys are unique across attached providers. + * Catalog attachments contribute all enabled models on the provider. + * Allowlist attachments contribute nothing until models are allowlisted + * (those are checked when the allowlist is set). + */ +export async function assertNoOverlappingModelKeys( + attachments: ResourceAiProviderAttachment[], + trx: DbOrTrx = db +): Promise { + const catalogProviderIds = attachments + .filter((a) => a.modelAccessMode === "catalog") + .map((a) => a.providerId); + + if (catalogProviderIds.length < 2) { + return null; + } + + const models = await trx + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey + }) + .from(aiModels) + .where( + and( + inArray(aiModels.providerId, catalogProviderIds), + eq(aiModels.enabled, true) + ) + ); + + const keyToProviders = new Map(); + for (const model of models) { + const existing = keyToProviders.get(model.modelKey) ?? []; + if (!existing.includes(model.providerId)) { + existing.push(model.providerId); + } + keyToProviders.set(model.modelKey, existing); + } + + const overlaps = [...keyToProviders.entries()].filter( + ([, providerIds]) => providerIds.length > 1 + ); + if (overlaps.length === 0) { + return null; + } + + const keys = overlaps.map(([key]) => key).sort(); + return { + error: `Model keys must be unique across providers on a resource. Overlapping keys: ${keys.join(", ")}` + }; +} + /** * Validate provider attachments for an org. - * At most one passthrough provider is allowed per resource. */ export async function resolveProviderAttachments(input: { orgId: string; @@ -78,15 +123,6 @@ export async function resolveProviderAttachments(input: { }; } - const passthroughCount = attachments.filter( - (a) => a.modelAccessMode === "passthrough" - ).length; - if (passthroughCount > 1) { - return { - error: "A resource may have at most one AI provider in passthrough mode" - }; - } - if (attachments.length === 0) { return []; } @@ -119,6 +155,11 @@ export async function resolveProviderAttachments(input: { }; } + const overlapError = await assertNoOverlappingModelKeys(attachments); + if (overlapError) { + return overlapError; + } + return attachments; } diff --git a/server/routers/aiGateway/chatCompletions.ts b/server/routers/aiGateway/chatCompletions.ts index 51508778a..f1ac231ff 100644 --- a/server/routers/aiGateway/chatCompletions.ts +++ b/server/routers/aiGateway/chatCompletions.ts @@ -331,10 +331,6 @@ async function providerMatchesModel( requestedModel: string, allowedModelIds: number[] ): Promise { - if (attachment.modelAccessMode === "passthrough") { - return true; - } - const [matchedModel] = await db .select({ modelId: aiModels.modelId, @@ -366,26 +362,7 @@ async function selectProvider( allowedModelIds: number[], requestedModel: string | undefined ): Promise { - const passthroughAttachments = attachments.filter( - (a) => a.modelAccessMode === "passthrough" - ); - const hasRestricted = attachments.some( - (a) => - a.modelAccessMode === "catalog" || a.modelAccessMode === "allowlist" - ); - if (!requestedModel) { - if (hasRestricted) { - return { - ok: false, - status: HttpCode.FORBIDDEN, - message: - "This resource restricts access to specific models; a model must be specified" - }; - } - if (passthroughAttachments.length === 1) { - return { ok: true, provider: passthroughAttachments[0].provider }; - } return { ok: false, status: HttpCode.FORBIDDEN, @@ -395,10 +372,6 @@ async function selectProvider( const candidates: ProviderAttachment[] = []; for (const attachment of attachments) { - if (attachment.modelAccessMode === "passthrough") { - candidates.push(attachment); - continue; - } if ( await providerMatchesModel( attachment, @@ -422,11 +395,6 @@ async function selectProvider( }; } - // Zero candidates: fall back to a single passthrough attachment if present - if (passthroughAttachments.length === 1) { - return { ok: true, provider: passthroughAttachments[0].provider }; - } - return { ok: false, status: HttpCode.FORBIDDEN, diff --git a/server/routers/resource/createResource.ts b/server/routers/resource/createResource.ts index 4c880c3af..4d546468f 100644 --- a/server/routers/resource/createResource.ts +++ b/server/routers/resource/createResource.ts @@ -109,7 +109,7 @@ const createHttpResourceSchema = z .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed." + "For inference-mode resources: AI providers to attach. Each entry may set modelAccessMode (catalog or allowlist); defaults to catalog. Model keys must be unique across attached catalog providers." ) }) .refine( @@ -397,9 +397,7 @@ async function createHttpResource( requireAtLeastOne: true }); if (isInferenceFieldsError(resolved)) { - return next( - createHttpError(HttpCode.BAD_REQUEST, resolved.error) - ); + return next(createHttpError(HttpCode.BAD_REQUEST, resolved.error)); } providerAttachments = resolved; } else if (aiProviderInputs && aiProviderInputs.length > 0) { diff --git a/server/routers/resource/setResourceAiProviders.ts b/server/routers/resource/setResourceAiProviders.ts index a1de8f176..0cd0d5d84 100644 --- a/server/routers/resource/setResourceAiProviders.ts +++ b/server/routers/resource/setResourceAiProviders.ts @@ -27,7 +27,7 @@ registry.registerPath({ method: "post", path: "/resource/{resourceId}/ai-providers", description: - "Replace the AI providers attached to an inference resource. At least one provider is required. At most one may use passthrough mode.", + "Replace the AI providers attached to an inference resource. At least one provider is required. Model keys must be unique across attached catalog providers.", tags: [OpenAPITags.PublicResource], request: { params: setResourceAiProvidersParamsSchema, @@ -116,7 +116,9 @@ export async function setResourceAiProviders( requireAtLeastOne: true }); if (isInferenceFieldsError(attachments)) { - return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + return next( + createHttpError(HttpCode.BAD_REQUEST, attachments.error) + ); } await setPublicResourceAiProviders(resourceId, attachments); diff --git a/server/routers/siteResource/createSiteResource.ts b/server/routers/siteResource/createSiteResource.ts index 47dcfc59b..77d9d5f51 100644 --- a/server/routers/siteResource/createSiteResource.ts +++ b/server/routers/siteResource/createSiteResource.ts @@ -90,7 +90,7 @@ const createSiteResourceSchema = z .array(resourceAiProviderAttachmentSchema) .optional() .describe( - "For inference-mode site resources: AI providers to attach. Each entry may set modelAccessMode (passthrough, catalog, or allowlist); defaults to passthrough. At most one passthrough provider is allowed." + "For inference-mode site resources: AI providers to attach. Each entry may set modelAccessMode (catalog or allowlist); defaults to catalog. Model keys must be unique across attached catalog providers." ) }) .strict() diff --git a/server/routers/siteResource/setSiteResourceAiProviders.ts b/server/routers/siteResource/setSiteResourceAiProviders.ts index 93bf8a866..c6e708064 100644 --- a/server/routers/siteResource/setSiteResourceAiProviders.ts +++ b/server/routers/siteResource/setSiteResourceAiProviders.ts @@ -27,7 +27,7 @@ registry.registerPath({ method: "post", path: "/site-resource/{siteResourceId}/ai-providers", description: - "Replace the AI providers attached to an inference site resource. At least one provider is required. At most one may use passthrough mode.", + "Replace the AI providers attached to an inference site resource. At least one provider is required. Model keys must be unique across attached catalog providers.", tags: [OpenAPITags.PrivateResource], request: { params: setSiteResourceAiProvidersParamsSchema, @@ -118,7 +118,9 @@ export async function setSiteResourceAiProviders( requireAtLeastOne: true }); if (isInferenceFieldsError(attachments)) { - return next(createHttpError(HttpCode.BAD_REQUEST, attachments.error)); + return next( + createHttpError(HttpCode.BAD_REQUEST, attachments.error) + ); } await replaceAttachments(siteResourceId, attachments); diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/layout.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/layout.tsx index 5846ad86d..4647347f5 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/layout.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/layout.tsx @@ -69,6 +69,10 @@ export default async function AiProviderLayout({ children, params }: Props) { title: t("aiProviderNetworkSettings"), href: "/{orgId}/settings/ai-providers/{providerId}/network" }, + { + title: t("aiProviderModels"), + href: "/{orgId}/settings/ai-providers/{providerId}/models" + }, { title: t("aiProviderAuthSettings"), href: "/{orgId}/settings/ai-providers/{providerId}/authentication" diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx new file mode 100644 index 000000000..f90143492 --- /dev/null +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/models/page.tsx @@ -0,0 +1,150 @@ +"use client"; + +import { + SettingsContainer, + SettingsSection, + SettingsSectionBody, + SettingsSectionDescription, + SettingsSectionFooter, + SettingsSectionForm, + SettingsSectionHeader, + SettingsSectionTitle +} from "@app/components/Settings"; +import { TagInput, type Tag } from "@app/components/tags/tag-input"; +import { Button } from "@app/components/ui/button"; +import { useAiProviderContext } from "@app/hooks/useAiProviderContext"; +import { useEnvContext } from "@app/hooks/useEnvContext"; +import { toast } from "@app/hooks/useToast"; +import { createApiClient, formatAxiosError } from "@app/lib/api"; +import { aiProviderQueries } from "@app/lib/queries"; +import { useQuery, useQueryClient } from "@tanstack/react-query"; +import { useTranslations } from "next-intl"; +import { useEffect, useState } from "react"; + +export default function AiProviderModelsPage() { + const { provider } = useAiProviderContext(); + const { env } = useEnvContext(); + const api = createApiClient({ env }); + const queryClient = useQueryClient(); + const t = useTranslations(); + const [saveLoading, setSaveLoading] = useState(false); + const [tags, setTags] = useState([]); + const [activeTagIndex, setActiveTagIndex] = useState(null); + + const modelsQuery = useQuery( + aiProviderQueries.providerModels({ providerId: provider.providerId }) + ); + + useEffect(() => { + if (!modelsQuery.data) return; + setTags( + modelsQuery.data.map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })) + ); + }, [modelsQuery.data]); + + async function onSave() { + setSaveLoading(true); + try { + const existing = modelsQuery.data ?? []; + const existingByKey = new Map( + existing.map((model) => [model.modelKey, model]) + ); + const nextKeys = new Set( + tags.map((tag) => tag.text.trim()).filter(Boolean) + ); + + const toCreate = [...nextKeys].filter( + (key) => !existingByKey.has(key) + ); + const toDelete = existing.filter( + (model) => !nextKeys.has(model.modelKey) + ); + + await Promise.all([ + ...toCreate.map((modelKey) => + api.put(`/ai-provider/${provider.providerId}/model`, { + modelKey, + name: modelKey + }) + ), + ...toDelete.map((model) => + api.delete(`/ai-model/${model.modelId}`) + ) + ]); + + await queryClient.invalidateQueries( + aiProviderQueries.providerModels({ + providerId: provider.providerId + }) + ); + + toast({ + title: t("success"), + description: t("aiProviderModelsUpdated") + }); + } catch (e) { + toast({ + variant: "destructive", + title: t("aiProviderModelsErrorUpdate"), + description: formatAxiosError( + e, + t("aiProviderModelsErrorUpdate") + ) + }); + } finally { + setSaveLoading(false); + } + } + + return ( + + + + + {t("aiProviderModels")} + + + {t("aiProviderModelsDescription")} + + + + + + { + const next = + typeof newTags === "function" + ? newTags(tags) + : newTags; + setTags(next as Tag[]); + }} + allowDuplicates={false} + sortTags + delimiterList={[",", "Enter"]} + disabled={modelsQuery.isLoading || saveLoading} + /> + + + + + + + + + ); +} diff --git a/src/app/[orgId]/settings/resources/private/[niceId]/general/page.tsx b/src/app/[orgId]/settings/resources/private/[niceId]/general/page.tsx index 8dbdb3809..b0af75ec0 100644 --- a/src/app/[orgId]/settings/resources/private/[niceId]/general/page.tsx +++ b/src/app/[orgId]/settings/resources/private/[niceId]/general/page.tsx @@ -24,7 +24,9 @@ import { } from "@app/components/ui/form"; import { Input } from "@app/components/ui/input"; import { SwitchInput } from "@app/components/SwitchInput"; +import { PrivateResourceAliasField } from "@app/components/PrivateResourceDestinationFields"; import { createGeneralFormSchema } from "@app/lib/privateResourceForm"; +import { asAnyControl, asAnyWatch } from "@app/lib/formControlUtils"; import { zodResolver } from "@hookform/resolvers/zod"; import { useTranslations } from "next-intl"; import { useActionState, useMemo } from "react"; @@ -35,8 +37,12 @@ import { useSaveSiteResource } from "@app/hooks/useSaveSiteResource"; export default function PrivateResourceGeneralPage() { const t = useTranslations(); const { save, siteResource } = useSaveSiteResource(); + const isInference = siteResource.mode === "inference"; - const formSchema = useMemo(() => createGeneralFormSchema(t), [t]); + const formSchema = useMemo( + () => createGeneralFormSchema(t, { requireAlias: isInference }), + [t, isInference] + ); type FormValues = z.infer; const form = useForm({ @@ -44,7 +50,8 @@ export default function PrivateResourceGeneralPage() { defaultValues: { name: siteResource.name, niceId: siteResource.niceId, - enabled: siteResource.enabled + enabled: siteResource.enabled, + alias: siteResource.alias ?? null } }); @@ -56,7 +63,13 @@ export default function PrivateResourceGeneralPage() { await save({ name: data.name, niceId: data.niceId, - enabled: data.enabled + enabled: data.enabled, + ...(isInference + ? { + mode: "inference" as const, + alias: data.alias + } + : {}) }); }, null); @@ -152,6 +165,17 @@ export default function PrivateResourceGeneralPage() { )} /> + {isInference && ( + + + + )} diff --git a/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx b/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx index 24aab5067..34534fdd4 100644 --- a/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx +++ b/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx @@ -1,105 +1,15 @@ -"use client"; +import type { Metadata } from "next"; +import { redirect } from "next/navigation"; -import { - SettingsContainer, - SettingsFormCell, - SettingsFormGrid, - SettingsSection, - SettingsSectionBody, - SettingsSectionDescription, - SettingsSectionFooter, - SettingsSectionForm, - SettingsSectionHeader, - SettingsSectionTitle -} from "@app/components/Settings"; -import { Button } from "@app/components/ui/button"; -import { Form } from "@app/components/ui/form"; -import { createInferenceFormSchema } from "@app/lib/privateResourceForm"; -import { zodResolver } from "@hookform/resolvers/zod"; -import { useTranslations } from "next-intl"; -import { useActionState, useMemo, useState } from "react"; -import { useForm } from "react-hook-form"; -import { z } from "zod"; -import { PrivateResourceInferenceDestinationFields } from "@app/components/PrivateResourceDestinationFields"; -import { useSaveSiteResource } from "@app/hooks/useSaveSiteResource"; -import { - asAnyControl, - asAnySetValue, - asAnyWatch -} from "@app/lib/formControlUtils"; -import { buildSelectedSitesForResource } from "@app/lib/privateResourceUtils"; +export const metadata: Metadata = { + title: "Private Resource" +}; -export default function PrivateResourceInferencePage() { - const t = useTranslations(); - const { save, siteResource } = useSaveSiteResource(); - const [selectedSites, setSelectedSites] = useState(() => - buildSelectedSitesForResource(siteResource) - ); - - const formSchema = useMemo(() => createInferenceFormSchema(t), [t]); - type FormValues = z.infer; - - const form = useForm({ - resolver: zodResolver(formSchema), - defaultValues: { - mode: "inference", - alias: siteResource.alias ?? null - } - }); - - const [, formAction, saveLoading] = useActionState(async () => { - const isValid = await form.trigger(); - if (!isValid) return; - - const data = form.getValues(); - await save({ - mode: "inference", - alias: data.alias - }); - }, null); - - return ( - - - - - {t("hostSettings")} - - - {t("editInternalResourceDialogDestinationDescription")} - - - - - -
- - - - - - -
- -
-
- - - - -
-
+export default async function PrivateResourceInferencePage(props: { + params: Promise<{ niceId: string; orgId: string }>; +}) { + const params = await props.params; + redirect( + `/${params.orgId}/settings/resources/private/${params.niceId}/providers` ); } diff --git a/src/app/[orgId]/settings/resources/private/[niceId]/layout.tsx b/src/app/[orgId]/settings/resources/private/[niceId]/layout.tsx index f07b92673..57064e545 100644 --- a/src/app/[orgId]/settings/resources/private/[niceId]/layout.tsx +++ b/src/app/[orgId]/settings/resources/private/[niceId]/layout.tsx @@ -55,20 +55,37 @@ export default async function PrivateResourceLayout( | "sshSettings" | "inferenceSettings"; - const navItems = [ - { - title: t("general"), - href: `/{orgId}/settings/resources/private/{niceId}/general` - }, - { - title: t(modeSettingsKey), - href: `/{orgId}/settings/resources/private/{niceId}/${siteResource.mode}` - }, - { - title: t("authentication"), - href: `/{orgId}/settings/resources/private/{niceId}/access` - } - ]; + const isInference = siteResource.mode === "inference"; + + const navItems = isInference + ? [ + { + title: t("general"), + href: `/{orgId}/settings/resources/private/{niceId}/general` + }, + { + title: t("aiResourceProviders"), + href: `/{orgId}/settings/resources/private/{niceId}/providers` + }, + { + title: t("authentication"), + href: `/{orgId}/settings/resources/private/{niceId}/access` + } + ] + : [ + { + title: t("general"), + href: `/{orgId}/settings/resources/private/{niceId}/general` + }, + { + title: t(modeSettingsKey), + href: `/{orgId}/settings/resources/private/{niceId}/${siteResource.mode}` + }, + { + title: t("authentication"), + href: `/{orgId}/settings/resources/private/{niceId}/access` + } + ]; return ( <> diff --git a/src/app/[orgId]/settings/resources/private/[niceId]/providers/page.tsx b/src/app/[orgId]/settings/resources/private/[niceId]/providers/page.tsx new file mode 100644 index 000000000..ead7ab24b --- /dev/null +++ b/src/app/[orgId]/settings/resources/private/[niceId]/providers/page.tsx @@ -0,0 +1,232 @@ +"use client"; + +import { + SettingsContainer, + SettingsFormCell, + SettingsFormGrid, + SettingsSection, + SettingsSectionBody, + SettingsSectionDescription, + SettingsSectionFooter, + SettingsSectionForm, + SettingsSectionHeader, + SettingsSectionTitle +} from "@app/components/Settings"; +import { + AiProvidersSelector, + type SelectedAiProvider +} from "@app/components/AiProvidersSelector"; +import { Button } from "@app/components/ui/button"; +import { + Form, + FormControl, + FormDescription, + FormField, + FormItem, + FormLabel, + FormMessage +} from "@app/components/ui/form"; +import { useEnvContext } from "@app/hooks/useEnvContext"; +import { useSiteResourceContext } from "@app/hooks/useSiteResourceContext"; +import { toast } from "@app/hooks/useToast"; +import { createApiClient, formatAxiosError } from "@app/lib/api"; +import { resourceQueries } from "@app/lib/queries"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { useQuery, useQueryClient } from "@tanstack/react-query"; +import { useTranslations } from "next-intl"; +import { useRouter } from "next/navigation"; +import { useActionState, useEffect, useMemo, useState } from "react"; +import { useForm } from "react-hook-form"; +import { z } from "zod"; + +export default function PrivateResourceProvidersPage() { + const t = useTranslations(); + const router = useRouter(); + const { env } = useEnvContext(); + const api = createApiClient({ env }); + const queryClient = useQueryClient(); + const { siteResource } = useSiteResourceContext(); + + useEffect(() => { + if (siteResource.mode !== "inference") { + router.replace( + `/${siteResource.orgId}/settings/resources/private/${siteResource.niceId}/general` + ); + } + }, [router, siteResource.mode, siteResource.niceId, siteResource.orgId]); + + const formSchema = useMemo( + () => + z.object({ + providerIds: z + .array(z.number().int().positive()) + .min(1, t("aiResourceProvidersRequired")) + }), + [t] + ); + type FormValues = z.infer; + + const [selectedProviders, setSelectedProviders] = useState< + SelectedAiProvider[] + >([]); + + const attachedQuery = useQuery({ + ...resourceQueries.siteResourceAiProviders({ + siteResourceId: siteResource.id + }), + enabled: siteResource.mode === "inference" + }); + + const form = useForm({ + resolver: zodResolver(formSchema), + defaultValues: { + providerIds: [] + } + }); + + useEffect(() => { + if (!attachedQuery.data) return; + const providers = attachedQuery.data.map((provider) => ({ + id: String(provider.providerId), + text: provider.name + })); + setSelectedProviders(providers); + form.reset({ + providerIds: attachedQuery.data.map((p) => p.providerId) + }); + }, [attachedQuery.data, form]); + + const [, formAction, saveLoading] = useActionState(async () => { + const isValid = await form.trigger(); + if (!isValid) return; + + const data = form.getValues(); + try { + await api.post(`/site-resource/${siteResource.id}/ai-providers`, { + providers: data.providerIds.map((providerId) => ({ + providerId, + modelAccessMode: "catalog" + })) + }); + + await queryClient.invalidateQueries( + resourceQueries.siteResourceAiProviders({ + siteResourceId: siteResource.id + }) + ); + + toast({ + title: t("success"), + description: t("aiResourceProvidersUpdated") + }); + } catch (error) { + toast({ + variant: "destructive", + title: t("aiResourceProvidersErrorUpdate"), + description: formatAxiosError( + error, + t("aiResourceProvidersErrorUpdate") + ) + }); + } + }, null); + + if (siteResource.mode !== "inference") { + return null; + } + + return ( + + + + + {t("aiResourceProviders")} + + + {t("aiResourceProvidersDescription")} + + + + + +
+ + + + ( + + + {t( + "aiResourceProviders" + )} + + + { + setSelectedProviders( + providers + ); + form.setValue( + "providerIds", + providers.map( + (p) => + parseInt( + p.id, + 10 + ) + ), + { + shouldValidate: true + } + ); + }} + /> + + + {t( + "aiResourceProvidersHelp" + )} + + + + )} + /> + + +
+ +
+
+ + + + +
+
+ ); +} diff --git a/src/app/[orgId]/settings/resources/private/create/page.tsx b/src/app/[orgId]/settings/resources/private/create/page.tsx index 936409d40..9420da17b 100644 --- a/src/app/[orgId]/settings/resources/private/create/page.tsx +++ b/src/app/[orgId]/settings/resources/private/create/page.tsx @@ -64,6 +64,10 @@ import { asAnySetValue, asAnyWatch } from "@app/lib/formControlUtils"; +import { + AiProvidersSelector, + type SelectedAiProvider +} from "@app/components/AiProvidersSelector"; export default function CreatePrivateResourcePage() { const params = useParams(); @@ -88,6 +92,9 @@ export default function CreatePrivateResourcePage() { : null; const [selectedSites, setSelectedSites] = useState([]); + const [selectedProviders, setSelectedProviders] = useState< + SelectedAiProvider[] + >([]); const formSchema = useMemo(() => createCreateFormSchema(t), [t]); type FormValues = z.infer; @@ -112,7 +119,8 @@ export default function CreatePrivateResourcePage() { pamMode: "passthrough", tcpPortRangeString: "*", udpPortRangeString: "*", - disableIcmp: false + disableIcmp: false, + providerIds: [] } }); @@ -196,7 +204,9 @@ export default function CreatePrivateResourcePage() { } router.push( - `/${orgId}/settings/resources/private/${created.niceId}/${created.mode}` + created.mode === "inference" + ? `/${orgId}/settings/resources/private/${created.niceId}/general` + : `/${orgId}/settings/resources/private/${created.niceId}/${created.mode}` ); } catch (error) { toast({ @@ -336,6 +346,13 @@ export default function CreatePrivateResourcePage() { "destinationPort", null ); + form.setValue( + "providerIds", + [] + ); + setSelectedProviders( + [] + ); } else { form.setValue( "destinationPort", @@ -637,6 +654,76 @@ export default function CreatePrivateResourcePage() { )} + {mode === "inference" && ( + + + + {t("aiResourceProviders")} + + + {t("aiResourceProvidersDescription")} + + + + + + + ( + + + {t( + "aiResourceProviders" + )} + + + { + setSelectedProviders( + providers + ); + form.setValue( + "providerIds", + providers.map( + ( + p + ) => + parseInt( + p.id, + 10 + ) + ), + { + shouldValidate: true + } + ); + }} + /> + + + {t( + "aiResourceProvidersHelp" + )} + + + + )} + /> + + + + + + )} +