diff --git a/messages/en-US.json b/messages/en-US.json index 8cc780829..2d77bb365 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1729,7 +1729,7 @@ "aiProviderErrorRoutingModeTarget": "Site targets routing is only available for custom providers", "aiProviderErrorCapabilitiesRequired": "Select at least one API capability", "aiProviderCapabilities": "API Capabilities", - "aiProviderCapabilitiesDescription": "Which API formats this provider accepts. Built-in providers use fixed capabilities.", + "aiProviderCapabilitiesDescription": "Select which API formats this provider can handle. Known providers start with recommended defaults.", "aiProviderCapabilitiesCustomDescription": "Select which API formats this custom provider can handle", "aiProviderCapabilitiesSelect": "Select capabilities", "aiProviderCapabilitiesEmpty": "No capabilities found", diff --git a/server/lib/aiCapabilities.ts b/server/lib/aiCapabilities.ts index f7b370691..356122e7f 100644 --- a/server/lib/aiCapabilities.ts +++ b/server/lib/aiCapabilities.ts @@ -261,8 +261,11 @@ export function resolveCapabilitiesForCreate(input: { type: AiProviderType; capabilities?: AiCapability[] | null; }): AiCapability[] { + if (input.capabilities != null) { + return parseCapabilities(input.capabilities); + } if (input.type === "custom") { - return parseCapabilities(input.capabilities ?? []); + return []; } return [...AI_PROVIDER_CAPABILITY_DEFAULTS[input.type]]; } diff --git a/server/routers/aiProvider/createAiProvider.ts b/server/routers/aiProvider/createAiProvider.ts index 995e9eaa0..195c8040e 100644 --- a/server/routers/aiProvider/createAiProvider.ts +++ b/server/routers/aiProvider/createAiProvider.ts @@ -124,6 +124,15 @@ export async function createAiProvider( capabilities }); + if (resolvedCapabilities.length === 0) { + return next( + createHttpError( + HttpCode.BAD_REQUEST, + "At least one capability is required" + ) + ); + } + const [provider] = await db .insert(aiProviders) .values({ diff --git a/server/routers/aiProvider/updateAiProvider.ts b/server/routers/aiProvider/updateAiProvider.ts index 7cbf1145a..1f22dc6ef 100644 --- a/server/routers/aiProvider/updateAiProvider.ts +++ b/server/routers/aiProvider/updateAiProvider.ts @@ -131,20 +131,9 @@ export async function updateAiProvider( ? body.authType : (existing.authType as AiProviderAuthType); - if (body.capabilities !== undefined && providerType !== "custom") { - return next( - createHttpError( - HttpCode.BAD_REQUEST, - "Capabilities can only be updated for custom providers" - ) - ); - } - const nextCapabilities = - providerType === "custom" - ? body.capabilities !== undefined - ? body.capabilities - : parseCapabilities(existing.capabilities) + body.capabilities !== undefined + ? body.capabilities : parseCapabilities(existing.capabilities); const validation = z @@ -195,7 +184,7 @@ export async function updateAiProvider( if (body.authType !== undefined) { updateData.authType = body.authType; } - if (providerType === "custom" && body.capabilities !== undefined) { + if (body.capabilities !== undefined) { updateData.capabilities = serializeCapabilities(body.capabilities); } diff --git a/server/routers/aiProvider/validation.ts b/server/routers/aiProvider/validation.ts index a56ccad1c..cc1346325 100644 --- a/server/routers/aiProvider/validation.ts +++ b/server/routers/aiProvider/validation.ts @@ -112,5 +112,15 @@ export function refineProviderUpstreamFields( path: ["capabilities"] }); } + } else if ( + data.capabilities !== undefined && + data.capabilities !== null && + data.capabilities.length === 0 + ) { + ctx.addIssue({ + code: "custom", + message: "At least one capability is required", + path: ["capabilities"] + }); } } diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx index 62ffa6eb6..6dabdbaad 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx @@ -12,10 +12,7 @@ import { SettingsSectionHeader, SettingsSectionTitle } from "@app/components/Settings"; -import { - AiProviderCapabilitiesSelect, - capabilityLabelKey -} from "@app/components/AiProviderCapabilitiesSelect"; +import { AiProviderCapabilitiesSelect } from "@app/components/AiProviderCapabilitiesSelect"; import { SwitchInput } from "@app/components/SwitchInput"; import { Button } from "@app/components/ui/button"; import { @@ -49,7 +46,6 @@ export default function AiProviderGeneralPage() { const router = useRouter(); const t = useTranslations(); const [saveLoading, setSaveLoading] = useState(false); - const isCustom = provider.type === "custom"; const generalSchema = useMemo( () => @@ -63,10 +59,7 @@ export default function AiProviderGeneralPage() { capabilities: z.array(z.enum(AI_CAPABILITIES)).optional() }) .superRefine((data, ctx) => { - if ( - isCustom && - (!data.capabilities || data.capabilities.length === 0) - ) { + if (!data.capabilities || data.capabilities.length === 0) { ctx.addIssue({ code: "custom", message: t("aiProviderErrorCapabilitiesRequired"), @@ -74,7 +67,7 @@ export default function AiProviderGeneralPage() { }); } }), - [t, isCustom] + [t] ); type GeneralFormValues = z.infer; @@ -97,11 +90,9 @@ export default function AiProviderGeneralPage() { capabilities?: AiCapability[]; } = { name: values.name.trim(), - enabled: values.enabled + enabled: values.enabled, + capabilities: values.capabilities ?? [] }; - if (isCustom) { - body.capabilities = values.capabilities ?? []; - } const res = await api.post< AxiosResponse @@ -211,46 +202,20 @@ export default function AiProviderGeneralPage() { )} - {isCustom ? ( - - ) : ( -
- {( - provider.capabilities ?? - [] - ).map((cap) => ( - - {t( - capabilityLabelKey( - cap - ) - )} - - ))} -
- )} +
- {isCustom - ? t( - "aiProviderCapabilitiesCustomDescription" - ) - : t( - "aiProviderCapabilitiesDescription" - )} + {t( + "aiProviderCapabilitiesDescription" + )} diff --git a/src/app/[orgId]/settings/ai-providers/create/page.tsx b/src/app/[orgId]/settings/ai-providers/create/page.tsx index 0d9541e54..999782bc3 100644 --- a/src/app/[orgId]/settings/ai-providers/create/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/create/page.tsx @@ -20,10 +20,7 @@ import { } from "@app/components/Settings"; import HeaderTitle from "@app/components/SettingsSectionTitle"; import { AiProviderAuthTypeSelect } from "@app/components/AiProviderAuthTypeSelect"; -import { - AiProviderCapabilitiesSelect, - capabilityLabelKey -} from "@app/components/AiProviderCapabilitiesSelect"; +import { AiProviderCapabilitiesSelect } from "@app/components/AiProviderCapabilitiesSelect"; import { AiProviderTypeSelect } from "@app/components/AiProviderTypeSelect"; import { HeadersInput } from "@app/components/HeadersInput"; import { StrategySelect } from "@app/components/StrategySelect"; @@ -93,14 +90,12 @@ export default function CreateAiProviderPage() { const providerType = form.watch("type"); const routingMode = form.watch("routingMode"); const authType = form.watch("authType"); - const capabilities = form.watch("capabilities"); const showUpstream = showsUpstreamUrlField(providerType, routingMode); const requireUpstream = upstreamUrlRequired(providerType, routingMode); const showRoutingMode = providerType === "custom"; const showTargets = providerType === "custom" && routingMode === "target"; const showApiKey = authTypeRequiresApiKey(authType ?? "bearer"); - const showCapabilitiesSelect = providerType === "custom"; async function createTargets( providerId: number, @@ -326,48 +321,20 @@ export default function CreateAiProviderPage() { )} - {showCapabilitiesSelect ? ( - - ) : ( -
- {( - capabilities ?? - defaultCapabilitiesForProvider( - providerType - ) - ).map((cap) => ( - - {t( - capabilityLabelKey( - cap - ) - )} - - ))} -
- )} +
- {showCapabilitiesSelect - ? t( - "aiProviderCapabilitiesCustomDescription" - ) - : t( - "aiProviderCapabilitiesDescription" - )} + {t( + "aiProviderCapabilitiesDescription" + )} diff --git a/src/lib/aiProviderFormSchema.ts b/src/lib/aiProviderFormSchema.ts index 794d8ad62..04e38a71a 100644 --- a/src/lib/aiProviderFormSchema.ts +++ b/src/lib/aiProviderFormSchema.ts @@ -97,10 +97,7 @@ export function createAiProviderFormSchema(t: TranslateFn) { }); } - if ( - data.type === "custom" && - (!data.capabilities || data.capabilities.length === 0) - ) { + if (!data.capabilities || data.capabilities.length === 0) { ctx.addIssue({ code: "custom", message: t("aiProviderErrorCapabilitiesRequired"), @@ -187,8 +184,7 @@ export function toAiProviderCreatePayload(values: AiProviderFormValues) { upstreamUrl, apiKey: values.apiKey?.trim() ? values.apiKey.trim() : undefined, authType: values.authType ?? "bearer", - capabilities: - values.type === "custom" ? (values.capabilities ?? []) : undefined, + capabilities: values.capabilities ?? [], headers: values.headers && values.headers.length > 0 ? values.headers : null, skipTlsVerification: values.skipTlsVerification, @@ -216,7 +212,7 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) { enabled: values.enabled ?? true }; - if (values.type === "custom" && values.capabilities) { + if (values.capabilities) { payload.capabilities = values.capabilities; }