diff --git a/messages/en-US.json b/messages/en-US.json index 2d77bb365..f18f99133 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1772,11 +1772,28 @@ "aiProviderModelsErrorUpdate": "Failed to update models", "aiResourceProviders": "Providers", "aiResourceProvidersDescription": "Choose which AI providers this inference resource can use", - "aiResourceProvidersHelp": "Each attached provider uses its own allow and block lists. Allow patterns that conflict across attached providers are not allowed.", + "aiResourceProvidersHelp": "Attach providers and choose inherit (use each provider's lists) or select (pick an allow list for this resource). Allow patterns that conflict across attached providers are not allowed.", "aiResourceProvidersSelect": "Select providers", "aiResourceProvidersEmpty": "No AI providers found", + "aiResourceProvidersNoneAttached": "No providers attached yet.", + "aiResourceProvidersAdd": "Add provider", + "aiResourceProvidersRemove": "Remove provider", + "aiResourceProviderToggleEnabled": "Enable or disable this provider on the resource", + "aiResourceProviderDisabled": "Disabled", "aiResourceProvidersUpdated": "Providers updated", "aiResourceProvidersErrorUpdate": "Failed to update providers", + "aiResourceProviderEditDescription": "Choose how this provider's models are exposed on this resource.", + "aiResourceProviderMode": "Access mode", + "aiResourceProviderModeInherit": "Inherit", + "aiResourceProviderModeSelect": "Select", + "aiResourceProviderModeSelectSummary": "Select ยท {count} models", + "aiResourceProviderModeInheritHelp": "Use this provider's allow and block lists as configured on the provider.", + "aiResourceProviderModeSelectHelp": "Choose a subset of this provider's allow-list models for this resource.", + "aiResourceProviderAllowModels": "Allow list", + "aiResourceProviderAllowModelsSelect": "Select models", + "aiResourceProviderAllowModelsSearch": "Search models...", + "aiResourceProviderAllowModelsEmpty": "No models found", + "aiResourceProviderAllowModelsHelp": "Only models from this provider's allow list can be selected.", "aiResourceAliasRequired": "Alias is required for inference resources", "aiResourceDomainConfiguration": "Domain configuration", "aiResourceDomainConfigurationDescription": "Choose the domain clients will use to reach this inference resource.", diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 90dd5abd9..1c3d21b10 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -234,7 +234,8 @@ export const resourceAiProviders = pgTable( accessMode: varchar("accessMode") .$type<"inherit" | "select">() .notNull() - .default("inherit") + .default("inherit"), + enabled: boolean("enabled").notNull().default(true) }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -529,7 +530,8 @@ export const siteResourceAiProviders = pgTable( accessMode: varchar("accessMode") .$type<"inherit" | "select">() .notNull() - .default("inherit") + .default("inherit"), + enabled: boolean("enabled").notNull().default(true) }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index dab445282..71fd197e7 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -231,7 +231,8 @@ export const resourceAiProviders = sqliteTable( accessMode: text("accessMode") .$type<"inherit" | "select">() .notNull() - .default("inherit") + .default("inherit"), + enabled: integer("enabled", { mode: "boolean" }).notNull().default(true) }, (t) => [primaryKey({ columns: [t.resourceId, t.providerId] })] ); @@ -514,7 +515,8 @@ export const siteResourceAiProviders = sqliteTable( accessMode: text("accessMode") .$type<"inherit" | "select">() .notNull() - .default("inherit") + .default("inherit"), + enabled: integer("enabled", { mode: "boolean" }).notNull().default(true) }, (t) => [primaryKey({ columns: [t.siteResourceId, t.providerId] })] ); diff --git a/server/lib/aiInferenceResource.ts b/server/lib/aiInferenceResource.ts index 51333473a..92a438ac2 100644 --- a/server/lib/aiInferenceResource.ts +++ b/server/lib/aiInferenceResource.ts @@ -24,7 +24,8 @@ export type AccessMode = z.infer; export const resourceAiProviderAttachmentSchema = z.strictObject({ providerId: z.number().int().positive(), - accessMode: accessModeSchema.optional().default("inherit") + accessMode: accessModeSchema.optional().default("inherit"), + enabled: z.boolean().optional().default(true) }); export type ResourceAiProviderInput = z.infer< @@ -34,6 +35,7 @@ export type ResourceAiProviderInput = z.infer< export type ResourceAiProviderAttachment = { providerId: number; accessMode: AccessMode; + enabled: boolean; }; export const resourceAiModelEntrySchema = z.strictObject({ @@ -79,14 +81,23 @@ export function resolveEffectiveLists(input: { function normalizeAttachments( inputs: ResourceAiProviderInput[] ): ResourceAiProviderAttachment[] { - const byProviderId = new Map(); + const byProviderId = new Map< + number, + { accessMode: AccessMode; enabled: boolean } + >(); for (const input of inputs) { - byProviderId.set(input.providerId, input.accessMode ?? "inherit"); + byProviderId.set(input.providerId, { + accessMode: input.accessMode ?? "inherit", + enabled: input.enabled ?? true + }); } - return [...byProviderId.entries()].map(([providerId, accessMode]) => ({ - providerId, - accessMode - })); + return [...byProviderId.entries()].map( + ([providerId, { accessMode, enabled }]) => ({ + providerId, + accessMode, + enabled + }) + ); } type EffectiveAllowRow = { @@ -110,14 +121,16 @@ export async function assertNoOverlappingModelKeys( ): Promise { const trx = options.trx ?? db; - if (attachments.length < 2) { + const activeAttachments = attachments.filter((a) => a.enabled); + + if (activeAttachments.length < 2) { return null; } - const inheritProviderIds = attachments + const inheritProviderIds = activeAttachments .filter((a) => a.accessMode === "inherit") .map((a) => a.providerId); - const selectProviderIds = attachments + const selectProviderIds = activeAttachments .filter((a) => a.accessMode === "select") .map((a) => a.providerId); @@ -318,7 +331,8 @@ export async function setPublicResourceAiProviders( attachments.map((a) => ({ resourceId, providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })) ); } @@ -344,7 +358,8 @@ export async function setSiteResourceAiProviders( attachments.map((a) => ({ siteResourceId, providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })) ); } @@ -473,7 +488,8 @@ export async function listPublicResourceAiProviders(resourceId: number) { providerId: resourceAiProviders.providerId, name: aiProviders.name, type: aiProviders.type, - enabled: aiProviders.enabled, + enabled: resourceAiProviders.enabled, + providerEnabled: aiProviders.enabled, accessMode: resourceAiProviders.accessMode }) .from(resourceAiProviders) @@ -490,7 +506,8 @@ export async function listSiteResourceAiProviders(siteResourceId: number) { providerId: siteResourceAiProviders.providerId, name: aiProviders.name, type: aiProviders.type, - enabled: aiProviders.enabled, + enabled: siteResourceAiProviders.enabled, + providerEnabled: aiProviders.enabled, accessMode: siteResourceAiProviders.accessMode }) .from(siteResourceAiProviders) @@ -575,7 +592,8 @@ export async function assertPublicResourceModelEntriesValid(input: { const attachments = await db .select({ providerId: resourceAiProviders.providerId, - accessMode: resourceAiProviders.accessMode + accessMode: resourceAiProviders.accessMode, + enabled: resourceAiProviders.enabled }) .from(resourceAiProviders) .innerJoin( @@ -610,7 +628,8 @@ export async function assertSiteResourceModelEntriesValid(input: { const attachments = await db .select({ providerId: siteResourceAiProviders.providerId, - accessMode: siteResourceAiProviders.accessMode + accessMode: siteResourceAiProviders.accessMode, + enabled: siteResourceAiProviders.enabled }) .from(siteResourceAiProviders) .innerJoin( diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index e84165620..bcff1a501 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -287,7 +287,8 @@ async function resolveTarget(host: string): Promise { resourceAiProviders.resourceId, resourceRow.resourceId ), - eq(aiProviders.enabled, true) + eq(aiProviders.enabled, true), + eq(resourceAiProviders.enabled, true) ) ), db @@ -342,7 +343,8 @@ async function resolveTarget(host: string): Promise { siteResourceAiProviders.siteResourceId, siteResourceRow.siteResourceId ), - eq(aiProviders.enabled, true) + eq(aiProviders.enabled, true), + eq(siteResourceAiProviders.enabled, true) ) ), db @@ -525,8 +527,15 @@ function logAiUsageAndCost(args: { isStream: boolean; headers: Headers; }): void { - const { capability, provider, requestedModel, requestBody, responseText, isStream, headers } = - args; + const { + capability, + provider, + requestedModel, + requestBody, + responseText, + isStream, + headers + } = args; let usage: AiUsage | null = extractUsage( capability, diff --git a/server/routers/resource/addAiProviderToResource.ts b/server/routers/resource/addAiProviderToResource.ts index bcc175c18..f2477d89d 100644 --- a/server/routers/resource/addAiProviderToResource.ts +++ b/server/routers/resource/addAiProviderToResource.ts @@ -118,9 +118,14 @@ export async function addAiProviderToResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })), - { providerId, accessMode: "inherit" as const } + { + providerId, + accessMode: "inherit" as const, + enabled: true as const + } ]; const attachments = await resolveProviderAttachments({ diff --git a/server/routers/resource/createResource.ts b/server/routers/resource/createResource.ts index e225458d1..45e66c2c3 100644 --- a/server/routers/resource/createResource.ts +++ b/server/routers/resource/createResource.ts @@ -397,7 +397,8 @@ async function createHttpResource( orgId, attachments: (aiProviderInputs ?? []).map((p) => ({ providerId: p.providerId, - accessMode: "inherit" as const + accessMode: "inherit" as const, + enabled: true as const })), requireAtLeastOne: false }); diff --git a/server/routers/resource/removeAiProviderFromResource.ts b/server/routers/resource/removeAiProviderFromResource.ts index a7f9298ee..5b99ff089 100644 --- a/server/routers/resource/removeAiProviderFromResource.ts +++ b/server/routers/resource/removeAiProviderFromResource.ts @@ -127,7 +127,8 @@ export async function removeAiProviderFromResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })); const attachments = await resolveProviderAttachments({ diff --git a/server/routers/siteResource/addAiProviderToSiteResource.ts b/server/routers/siteResource/addAiProviderToSiteResource.ts index 79a0ae093..18e8cd74d 100644 --- a/server/routers/siteResource/addAiProviderToSiteResource.ts +++ b/server/routers/siteResource/addAiProviderToSiteResource.ts @@ -118,9 +118,14 @@ export async function addAiProviderToSiteResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })), - { providerId, accessMode: "inherit" as const } + { + providerId, + accessMode: "inherit" as const, + enabled: true as const + } ]; const attachments = await resolveProviderAttachments({ diff --git a/server/routers/siteResource/createSiteResource.ts b/server/routers/siteResource/createSiteResource.ts index 15b82cb56..2038f9ce7 100644 --- a/server/routers/siteResource/createSiteResource.ts +++ b/server/routers/siteResource/createSiteResource.ts @@ -356,7 +356,8 @@ export async function createSiteResource( orgId, attachments: (aiProviderInputs ?? []).map((p) => ({ providerId: p.providerId, - accessMode: "inherit" as const + accessMode: "inherit" as const, + enabled: true as const })), requireAtLeastOne: false }); diff --git a/server/routers/siteResource/removeAiProviderFromSiteResource.ts b/server/routers/siteResource/removeAiProviderFromSiteResource.ts index ffb4f1ff8..29facaece 100644 --- a/server/routers/siteResource/removeAiProviderFromSiteResource.ts +++ b/server/routers/siteResource/removeAiProviderFromSiteResource.ts @@ -126,7 +126,8 @@ export async function removeAiProviderFromSiteResource( .filter((a) => a.providerId !== providerId) .map((a) => ({ providerId: a.providerId, - accessMode: a.accessMode + accessMode: a.accessMode, + enabled: a.enabled })); const attachments = await resolveProviderAttachments({ 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 f9ad664a0..e824cc105 100644 --- a/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx +++ b/src/app/[orgId]/settings/resources/private/[niceId]/inference/page.tsx @@ -16,9 +16,9 @@ import { SettingsSubsectionTitle } from "@app/components/Settings"; import { - AiProvidersSelector, - type SelectedAiProvider -} from "@app/components/AiProvidersSelector"; + AiProviderAttachments, + type AiProviderAttachmentValue +} from "@app/components/AiProviderAttachments"; import DomainPicker from "@app/components/DomainPicker"; import { SwitchInput } from "@app/components/SwitchInput"; import { Button } from "@app/components/ui/button"; @@ -41,7 +41,7 @@ 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 { useActionState, useEffect, useMemo } from "react"; import { useForm } from "react-hook-form"; import { z } from "zod"; @@ -65,7 +65,15 @@ export default function PrivateResourceInferencePage() { const formSchema = useMemo( () => z.object({ - providerIds: z.array(z.number().int().positive()), + providers: z.array( + z.object({ + providerId: z.number().int().positive(), + name: z.string(), + accessMode: z.enum(["inherit", "select"]), + enabled: z.boolean(), + selectedModelIds: z.array(z.number().int().positive()) + }) + ), httpConfigSubdomain: z.string().nullish(), httpConfigDomainId: z.string().nullish(), httpConfigFullDomain: z.string().nullish(), @@ -75,10 +83,6 @@ export default function PrivateResourceInferencePage() { ); type FormValues = z.infer; - const [selectedProviders, setSelectedProviders] = useState< - SelectedAiProvider[] - >([]); - const attachedQuery = useQuery({ ...resourceQueries.siteResourceAiProviders({ siteResourceId: siteResource.id @@ -86,10 +90,17 @@ export default function PrivateResourceInferencePage() { enabled: siteResource.mode === "inference" }); + const modelsQuery = useQuery({ + ...resourceQueries.siteResourceAiModels({ + siteResourceId: siteResource.id + }), + enabled: siteResource.mode === "inference" + }); + const form = useForm({ resolver: zodResolver(formSchema), defaultValues: { - providerIds: [], + providers: [], httpConfigSubdomain: siteResource.subdomain ?? null, httpConfigDomainId: siteResource.domainId ?? null, httpConfigFullDomain: siteResource.fullDomain ?? null, @@ -103,16 +114,33 @@ export default function PrivateResourceInferencePage() { useEffect(() => { if (!attachedQuery.data) return; - const providers = attachedQuery.data.map((provider) => ({ - id: String(provider.providerId), - text: provider.name - })); - setSelectedProviders(providers); - form.setValue( - "providerIds", - attachedQuery.data.map((p) => p.providerId) + const hasSelect = attachedQuery.data.some( + (provider) => provider.accessMode === "select" ); - }, [attachedQuery.data, form]); + if (hasSelect && modelsQuery.isLoading) return; + + const modelsByProvider = new Map(); + for (const model of modelsQuery.data ?? []) { + if (model.listType !== "allow") continue; + const existing = modelsByProvider.get(model.providerId) ?? []; + existing.push(model.modelId); + modelsByProvider.set(model.providerId, existing); + } + + form.setValue( + "providers", + attachedQuery.data.map((provider) => ({ + providerId: provider.providerId, + name: provider.name, + accessMode: provider.accessMode, + enabled: provider.enabled, + selectedModelIds: + provider.accessMode === "select" + ? (modelsByProvider.get(provider.providerId) ?? []) + : [] + })) + ); + }, [attachedQuery.data, modelsQuery.data, modelsQuery.isLoading, form]); const [, formAction, saveLoading] = useActionState(async () => { const isValid = await form.trigger(); @@ -129,16 +157,37 @@ export default function PrivateResourceInferencePage() { }); await api.post(`/site-resource/${siteResource.id}/ai-providers`, { - providers: data.providerIds.map((providerId) => ({ - providerId + providers: data.providers.map((provider) => ({ + providerId: provider.providerId, + accessMode: provider.accessMode, + enabled: provider.enabled })) }); + const selectProviders = data.providers.filter( + (provider) => provider.accessMode === "select" + ); + if (selectProviders.length > 0) { + await api.post(`/site-resource/${siteResource.id}/ai-models`, { + models: selectProviders.flatMap((provider) => + provider.selectedModelIds.map((modelId) => ({ + modelId, + listType: "allow" as const + })) + ) + }); + } + await queryClient.invalidateQueries( resourceQueries.siteResourceAiProviders({ siteResourceId: siteResource.id }) ); + await queryClient.invalidateQueries( + resourceQueries.siteResourceAiModels({ + siteResourceId: siteResource.id + }) + ); toast({ title: t("success"), @@ -160,6 +209,11 @@ export default function PrivateResourceInferencePage() { return null; } + const providersLoading = + attachedQuery.isLoading || + (attachedQuery.data?.some((p) => p.accessMode === "select") && + modelsQuery.isLoading); + return ( @@ -183,8 +237,8 @@ export default function PrivateResourceInferencePage() { ( + name="providers" + render={({ field }) => ( {t( @@ -192,44 +246,21 @@ export default function PrivateResourceInferencePage() { )} - { - setSelectedProviders( - providers - ); - form.setValue( - "providerIds", - providers.map( - (p) => - parseInt( - p.id, - 10 - ) - ), - { - shouldValidate: true - } - ); - }} /> - - {t( - "aiResourceProvidersHelp" - )} - )} @@ -335,7 +366,7 @@ export default function PrivateResourceInferencePage() { type="submit" form="private-resource-providers-form" loading={saveLoading} - disabled={attachedQuery.isLoading} + disabled={providersLoading || saveLoading} > {t("saveSettings")} diff --git a/src/app/[orgId]/settings/resources/private/create/page.tsx b/src/app/[orgId]/settings/resources/private/create/page.tsx index da3c968b6..a1b04f5e4 100644 --- a/src/app/[orgId]/settings/resources/private/create/page.tsx +++ b/src/app/[orgId]/settings/resources/private/create/page.tsx @@ -710,11 +710,6 @@ export default function CreatePrivateResourcePage() { }} /> - - {t( - "aiResourceProvidersHelp" - )} - )} diff --git a/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx b/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx index ad7230591..7de441995 100644 --- a/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx +++ b/src/app/[orgId]/settings/resources/public/[niceId]/inference/page.tsx @@ -13,9 +13,9 @@ import { SettingsSectionTitle } from "@app/components/Settings"; import { - AiProvidersSelector, - type SelectedAiProvider -} from "@app/components/AiProvidersSelector"; + AiProviderAttachments, + type AiProviderAttachmentValue +} from "@app/components/AiProviderAttachments"; import { Button } from "@app/components/ui/button"; import { Form, @@ -35,7 +35,7 @@ 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 { useActionState, useEffect, useMemo } from "react"; import { useForm } from "react-hook-form"; import { z } from "zod"; @@ -58,16 +58,20 @@ export default function PublicResourceInferencePage() { const formSchema = useMemo( () => z.object({ - providerIds: z.array(z.number().int().positive()) + providers: z.array( + z.object({ + providerId: z.number().int().positive(), + name: z.string(), + accessMode: z.enum(["inherit", "select"]), + enabled: z.boolean(), + selectedModelIds: z.array(z.number().int().positive()) + }) + ) }), [] ); type FormValues = z.infer; - const [selectedProviders, setSelectedProviders] = useState< - SelectedAiProvider[] - >([]); - const attachedQuery = useQuery({ ...resourceQueries.resourceAiProviders({ resourceId: resource.resourceId @@ -75,24 +79,53 @@ export default function PublicResourceInferencePage() { enabled: resource.mode === "inference" }); + const modelsQuery = useQuery({ + ...resourceQueries.resourceAiModels({ + resourceId: resource.resourceId + }), + enabled: resource.mode === "inference" + }); + const form = useForm({ resolver: zodResolver(formSchema), defaultValues: { - providerIds: [] + providers: [] } }); useEffect(() => { if (!attachedQuery.data) return; - const providers = attachedQuery.data.map((provider) => ({ - id: String(provider.providerId), - text: provider.name - })); - setSelectedProviders(providers); + const hasSelect = attachedQuery.data.some( + (provider) => provider.accessMode === "select" + ); + if (hasSelect && modelsQuery.isLoading) return; + + const modelsByProvider = new Map(); + for (const model of modelsQuery.data ?? []) { + if (model.listType !== "allow") continue; + const existing = modelsByProvider.get(model.providerId) ?? []; + existing.push(model.modelId); + modelsByProvider.set(model.providerId, existing); + } + form.reset({ - providerIds: attachedQuery.data.map((p) => p.providerId) + providers: attachedQuery.data.map((provider) => ({ + providerId: provider.providerId, + name: provider.name, + accessMode: provider.accessMode, + enabled: provider.enabled, + selectedModelIds: + provider.accessMode === "select" + ? (modelsByProvider.get(provider.providerId) ?? []) + : [] + })) }); - }, [attachedQuery.data, form]); + }, [ + attachedQuery.data, + modelsQuery.data, + modelsQuery.isLoading, + form + ]); const [, formAction, saveLoading] = useActionState(async () => { const isValid = await form.trigger(); @@ -101,16 +134,37 @@ export default function PublicResourceInferencePage() { const data = form.getValues(); try { await api.post(`/resource/${resource.resourceId}/ai-providers`, { - providers: data.providerIds.map((providerId) => ({ - providerId + providers: data.providers.map((provider) => ({ + providerId: provider.providerId, + accessMode: provider.accessMode, + enabled: provider.enabled })) }); + const selectProviders = data.providers.filter( + (provider) => provider.accessMode === "select" + ); + if (selectProviders.length > 0) { + await api.post(`/resource/${resource.resourceId}/ai-models`, { + models: selectProviders.flatMap((provider) => + provider.selectedModelIds.map((modelId) => ({ + modelId, + listType: "allow" as const + })) + ) + }); + } + await queryClient.invalidateQueries( resourceQueries.resourceAiProviders({ resourceId: resource.resourceId }) ); + await queryClient.invalidateQueries( + resourceQueries.resourceAiModels({ + resourceId: resource.resourceId + }) + ); toast({ title: t("success"), @@ -132,6 +186,11 @@ export default function PublicResourceInferencePage() { return null; } + const providersLoading = + attachedQuery.isLoading || + (attachedQuery.data?.some((p) => p.accessMode === "select") && + modelsQuery.isLoading); + return ( @@ -155,8 +214,8 @@ export default function PublicResourceInferencePage() { ( + name="providers" + render={({ field }) => ( {t( @@ -164,44 +223,21 @@ export default function PublicResourceInferencePage() { )} - { - setSelectedProviders( - providers - ); - form.setValue( - "providerIds", - providers.map( - (p) => - parseInt( - p.id, - 10 - ) - ), - { - shouldValidate: true - } - ); - }} /> - - {t( - "aiResourceProvidersHelp" - )} - )} @@ -218,7 +254,7 @@ export default function PublicResourceInferencePage() { type="submit" form="public-resource-providers-form" loading={saveLoading} - disabled={attachedQuery.isLoading} + disabled={providersLoading || saveLoading} > {t("saveSettings")} diff --git a/src/app/[orgId]/settings/resources/public/create/page.tsx b/src/app/[orgId]/settings/resources/public/create/page.tsx index 4d1c24efb..65529f74a 100644 --- a/src/app/[orgId]/settings/resources/public/create/page.tsx +++ b/src/app/[orgId]/settings/resources/public/create/page.tsx @@ -1437,11 +1437,6 @@ export default function Page() { ); }} /> -

- {t( - "aiResourceProvidersHelp" - )} -

diff --git a/src/components/AiProviderAttachments.tsx b/src/components/AiProviderAttachments.tsx new file mode 100644 index 000000000..09a4243e1 --- /dev/null +++ b/src/components/AiProviderAttachments.tsx @@ -0,0 +1,581 @@ +"use client"; + +import { + Credenza, + CredenzaBody, + CredenzaClose, + CredenzaContent, + CredenzaDescription, + CredenzaFooter, + CredenzaHeader, + CredenzaTitle +} from "@app/components/Credenza"; +import { type TagValue } from "@app/components/multi-select/multi-select-content"; +import { MultiSelectTagInput } from "@app/components/multi-select/multi-select-tag-input"; +import { Button } from "@app/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger +} from "@app/components/ui/dropdown-menu"; +import { Switch } from "@app/components/ui/switch"; +import { + Form, + FormControl, + FormDescription, + FormField, + FormItem, + FormLabel, + FormMessage +} from "@app/components/ui/form"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue +} from "@app/components/ui/select"; +import { cn } from "@app/lib/cn"; +import { aiProviderQueries } from "@app/lib/queries"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { useQuery } from "@tanstack/react-query"; +import { Plus, XIcon } from "lucide-react"; +import { useTranslations } from "next-intl"; +import { useEffect, useMemo, useRef, useState } from "react"; +import { useForm } from "react-hook-form"; +import { z } from "zod"; + +export type AiProviderAttachmentValue = { + providerId: number; + name: string; + accessMode: "inherit" | "select"; + enabled: boolean; + selectedModelIds: number[]; +}; + +export type AiProviderAttachmentsProps = { + orgId: string; + value: AiProviderAttachmentValue[]; + onChange: (value: AiProviderAttachmentValue[]) => void; + disabled?: boolean; +}; + +export function AiProviderAttachments({ + orgId, + value, + onChange, + disabled +}: AiProviderAttachmentsProps) { + const t = useTranslations(); + const [editingProviderId, setEditingProviderId] = useState( + null + ); + + const { data: providers = [] } = useQuery( + aiProviderQueries.orgProviders({ orgId }) + ); + + const attachedIds = useMemo( + () => new Set(value.map((v) => v.providerId)), + [value] + ); + + const availableProviders = providers + .filter((provider) => provider.enabled) + .filter((provider) => !attachedIds.has(provider.providerId)); + + const editing = value.find((v) => v.providerId === editingProviderId); + + function addProvider(providerId: number, name: string) { + if (value.some((v) => v.providerId === providerId)) { + return; + } + onChange([ + ...value, + { + providerId, + name, + accessMode: "inherit", + enabled: true, + selectedModelIds: [] + } + ]); + } + + function removeProvider(providerId: number) { + onChange(value.filter((v) => v.providerId !== providerId)); + } + + function updateProvider(updated: AiProviderAttachmentValue) { + onChange( + value.map((v) => + v.providerId === updated.providerId ? updated : v + ) + ); + setEditingProviderId(null); + } + + return ( +
+ {value.length === 0 ? ( +

+ {t("aiResourceProvidersNoneAttached")} +

+ ) : ( +
+ {value.map((attachment) => ( + + setEditingProviderId(attachment.providerId) + } + onRemove={() => + removeProvider(attachment.providerId) + } + onToggleEnabled={(enabled) => { + onChange( + value.map((v) => + v.providerId === attachment.providerId + ? { ...v, enabled } + : v + ) + ); + }} + /> + ))} +
+ )} + + + + + + + {availableProviders.map((provider) => ( + + addProvider(provider.providerId, provider.name) + } + > + {provider.name} + + ))} + + + + {editing && ( + { + if (!open) setEditingProviderId(null); + }} + onSave={updateProvider} + /> + )} +
+ ); +} + +function AttachmentRow({ + attachment, + disabled, + onEdit, + onRemove, + onToggleEnabled +}: { + attachment: AiProviderAttachmentValue; + disabled?: boolean; + onEdit: () => void; + onRemove: () => void; + onToggleEnabled: (enabled: boolean) => void; +}) { + const t = useTranslations(); + const summary = + attachment.accessMode === "inherit" + ? t("aiResourceProviderModeInherit") + : t("aiResourceProviderModeSelectSummary", { + count: attachment.selectedModelIds.length + }); + + return ( +
{ + if (e.key === "Enter" || e.key === " ") { + e.preventDefault(); + onEdit(); + } + } + } + role={disabled ? undefined : "button"} + tabIndex={disabled ? undefined : 0} + > +
+ + {attachment.name} + +

+ {attachment.enabled + ? summary + : t("aiResourceProviderDisabled")} +

+
+
e.stopPropagation()} + onKeyDown={(e) => e.stopPropagation()} + > + + + +
+
+ ); +} + +type EditFormValues = { + accessMode: "inherit" | "select"; + selectedModels: TagValue[]; +}; + +function EditAttachmentCredenza({ + attachment, + open, + onOpenChange, + onSave +}: { + attachment: AiProviderAttachmentValue; + open: boolean; + onOpenChange: (open: boolean) => void; + onSave: (value: AiProviderAttachmentValue) => void; +}) { + const t = useTranslations(); + const [modelSearch, setModelSearch] = useState(""); + + const editSchema = useMemo( + () => + z.object({ + accessMode: z.enum(["inherit", "select"]), + selectedModels: z.array( + z.object({ + id: z.string(), + text: z.string() + }) + ) + }), + [] + ); + + const modelsQuery = useQuery({ + ...aiProviderQueries.providerModels({ + providerId: attachment.providerId + }), + enabled: open + }); + + const allowCatalog = useMemo(() => { + const models = modelsQuery.data ?? []; + return models.filter( + (model) => model.enabled && (model.listType ?? "allow") === "allow" + ); + }, [modelsQuery.data]); + + const allowOptions: TagValue[] = useMemo(() => { + const query = modelSearch.trim().toLowerCase(); + return allowCatalog + .filter((model) => { + if (!query) return true; + return ( + model.modelKey.toLowerCase().includes(query) || + model.name.toLowerCase().includes(query) + ); + }) + .map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })); + }, [allowCatalog, modelSearch]); + + const form = useForm({ + resolver: zodResolver(editSchema), + defaultValues: { + accessMode: attachment.accessMode, + selectedModels: [] + } + }); + + const accessMode = form.watch("accessMode"); + const pendingSeedRef = useRef(false); + + useEffect(() => { + if (!open) return; + form.reset({ + accessMode: attachment.accessMode, + selectedModels: attachment.selectedModelIds.map((modelId) => { + const catalog = (modelsQuery.data ?? []).find( + (model) => model.modelId === modelId + ); + return { + id: String(modelId), + text: catalog?.modelKey ?? String(modelId) + }; + }) + }); + setModelSearch(""); + pendingSeedRef.current = false; + // Only re-init when opening or switching which attachment is edited. + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [open, attachment.providerId]); + + useEffect(() => { + if (!open || allowCatalog.length === 0) return; + + const current = form.getValues("selectedModels"); + const upgraded = current.map((model) => { + const catalog = allowCatalog.find( + (entry) => String(entry.modelId) === model.id + ); + return catalog ? { id: model.id, text: catalog.modelKey } : model; + }); + const changed = upgraded.some( + (model, index) => model.text !== current[index]?.text + ); + if (changed) { + form.setValue("selectedModels", upgraded); + } + + if (pendingSeedRef.current) { + form.setValue( + "selectedModels", + allowCatalog.map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })) + ); + pendingSeedRef.current = false; + } + }, [open, allowCatalog, form]); + + function handleAccessModeChange(next: "inherit" | "select") { + form.setValue("accessMode", next); + if (next === "inherit") { + form.setValue("selectedModels", []); + pendingSeedRef.current = false; + return; + } + if (attachment.accessMode === "select") { + form.setValue( + "selectedModels", + attachment.selectedModelIds.map((modelId) => { + const catalog = allowCatalog.find( + (model) => model.modelId === modelId + ); + return { + id: String(modelId), + text: catalog?.modelKey ?? String(modelId) + }; + }) + ); + pendingSeedRef.current = false; + return; + } + if (allowCatalog.length > 0) { + form.setValue( + "selectedModels", + allowCatalog.map((model) => ({ + id: String(model.modelId), + text: model.modelKey + })) + ); + pendingSeedRef.current = false; + return; + } + form.setValue("selectedModels", []); + pendingSeedRef.current = true; + } + + function onSubmit(values: EditFormValues) { + onSave({ + providerId: attachment.providerId, + name: attachment.name, + accessMode: values.accessMode, + enabled: attachment.enabled, + selectedModelIds: + values.accessMode === "select" + ? values.selectedModels.map((model) => + parseInt(model.id, 10) + ) + : [] + }); + } + + return ( + + + + {attachment.name} + + {t("aiResourceProviderEditDescription")} + + + +
+ + ( + + + {t("aiResourceProviderMode")} + + + + {field.value === "inherit" + ? t( + "aiResourceProviderModeInheritHelp" + ) + : t( + "aiResourceProviderModeSelectHelp" + )} + + + + )} + /> + + {accessMode === "select" && ( + ( + + + {t( + "aiResourceProviderAllowModels" + )} + + + + + + {t( + "aiResourceProviderAllowModelsHelp" + )} + + + + )} + /> + )} + + +
+ + + + + + +
+
+ ); +} diff --git a/src/lib/queries.ts b/src/lib/queries.ts index 6b7070f5b..e1e06f928 100644 --- a/src/lib/queries.ts +++ b/src/lib/queries.ts @@ -38,6 +38,7 @@ import type { GetResourcePolicyResponse } from "@server/routers/policy"; import type { GetResourcePoliciesResponse, GetResourceWhitelistResponse, + ListResourceAiModelsResponse, ListResourceNamesResponse, ListResourceRolesResponse, ListResourceRulesResponse, @@ -51,6 +52,7 @@ import type { ListRolesResponse } from "@server/routers/role"; import type { ListSitesResponse } from "@server/routers/site"; import type { ListAllSiteResourcesByOrgResponse, + ListSiteResourceAiModelsResponse, ListSiteResourceClientsResponse, ListSiteResourceRolesResponse, ListSiteResourceUsersResponse @@ -1291,6 +1293,7 @@ export const resourceQueries = { name: string; type: string; enabled: boolean; + providerEnabled: boolean; accessMode: "inherit" | "select"; }>; }> @@ -1311,6 +1314,7 @@ export const resourceQueries = { name: string; type: string; enabled: boolean; + providerEnabled: boolean; accessMode: "inherit" | "select"; }>; }> @@ -1320,6 +1324,26 @@ export const resourceQueries = { return res.data.data.providers; } }), + resourceAiModels: ({ resourceId }: { resourceId: number }) => + queryOptions({ + queryKey: ["RESOURCES", resourceId, "AI_MODELS"] as const, + queryFn: async ({ signal, meta }) => { + const res = await meta!.api.get< + AxiosResponse + >(`/resource/${resourceId}/ai-models`, { signal }); + return res.data.data.models; + } + }), + siteResourceAiModels: ({ siteResourceId }: { siteResourceId: number }) => + queryOptions({ + queryKey: ["SITE_RESOURCES", siteResourceId, "AI_MODELS"] as const, + queryFn: async ({ signal, meta }) => { + const res = await meta!.api.get< + AxiosResponse + >(`/site-resource/${siteResourceId}/ai-models`, { signal }); + return res.data.data.models; + } + }), resourceTargets: ({ resourceId }: { resourceId: number }) => queryOptions({ queryKey: ["RESOURCES", resourceId, "TARGETS"] as const,