From 39e06f2b6dddde344e33510d503800bd12118ec2 Mon Sep 17 00:00:00 2001 From: miloschwartz Date: Wed, 5 Aug 2026 16:55:48 -0400 Subject: [PATCH] add api capabilities --- messages/en-US.json | 23 ++ server/aiGatewayServer.ts | 4 +- server/db/pg/schema/schema.ts | 1 + server/db/sqlite/schema/schema.ts | 1 + server/lib/aiCapabilities.ts | 257 ++++++++++++++++++ .../aiGateway/createAiGatewayRouter.ts | 21 ++ server/routers/aiGateway/index.ts | 3 +- .../{chatCompletions.ts => pipeline.ts} | 73 ++--- server/routers/aiProvider/createAiProvider.ts | 12 + server/routers/aiProvider/types.ts | 14 +- server/routers/aiProvider/updateAiProvider.ts | 31 ++- server/routers/aiProvider/validation.ts | 18 ++ .../[providerId]/general/page.tsx | 119 +++++++- .../settings/ai-providers/create/page.tsx | 75 +++++ .../AiProviderCapabilitiesSelect.tsx | 89 ++++++ src/lib/aiProviderFormSchema.ts | 31 +++ 16 files changed, 715 insertions(+), 57 deletions(-) create mode 100644 server/lib/aiCapabilities.ts create mode 100644 server/routers/aiGateway/createAiGatewayRouter.ts rename server/routers/aiGateway/{chatCompletions.ts => pipeline.ts} (88%) create mode 100644 src/components/AiProviderCapabilitiesSelect.tsx diff --git a/messages/en-US.json b/messages/en-US.json index 39a3b54e6..ac6131abc 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1726,6 +1726,29 @@ "aiProviderErrorAuthTypeRequired": "Auth type is required", "aiProviderErrorApiKeyRequired": "API key is required", "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.", + "aiProviderCapabilitiesCustomDescription": "Select which API formats this custom provider can handle", + "aiProviderCapabilitiesSelect": "Select capabilities", + "aiProviderCapabilitiesEmpty": "No capabilities found", + "aiProviderCapabilitiesSearch": "Search capabilities...", + "aiCapabilityOpenaiChat": "OpenAI Chat Completions", + "aiCapabilityOpenaiChatDescription": "Supports /v1/chat/completions", + "aiCapabilityOpenaiResponses": "OpenAI Responses", + "aiCapabilityOpenaiResponsesDescription": "Supports /v1/responses", + "aiCapabilityAnthropicMessages": "Anthropic Messages", + "aiCapabilityAnthropicMessagesDescription": "Supports /v1/messages", + "aiCapabilityGeminiGenerateContent": "Gemini Generate Content", + "aiCapabilityGeminiGenerateContentDescription": "Supports the direct Gemini API", + "aiCapabilityBedrockModelInvoke": "Bedrock Model Invoke", + "aiCapabilityBedrockModelInvokeDescription": "Supports Amazon Bedrock InvokeModel", + "aiCapabilityGoogleGenerateContent": "Vertex Generate Content", + "aiCapabilityGoogleGenerateContentDescription": "Supports Vertex AI Gemini format", + "aiCapabilityGoogleRawPredict": "Vertex Raw Predict", + "aiCapabilityGoogleRawPredictDescription": "Supports Vertex AI rawPredict for Anthropic models", + "aiCapabilityBedrockConverse": "Bedrock Converse", + "aiCapabilityBedrockConverseDescription": "Supports Amazon Bedrock Converse API", "aiProviderCreated": "AI provider created", "aiProviderUpdated": "AI provider updated", "aiProviderDeleted": "AI provider deleted", diff --git a/server/aiGatewayServer.ts b/server/aiGatewayServer.ts index db8b12324..87d57b86b 100644 --- a/server/aiGatewayServer.ts +++ b/server/aiGatewayServer.ts @@ -7,7 +7,7 @@ import { errorHandlerMiddleware, notFoundMiddleware } from "@server/middlewares"; -import * as aiGateway from "@server/routers/aiGateway"; +import { createAiGatewayRouter } from "@server/routers/aiGateway"; const aiGatewayPort = config.getRawConfig().server.ai_gateway_port; @@ -23,7 +23,7 @@ export function createAiGatewayServer() { aiGatewayServer.use(cors()); aiGatewayServer.use(express.json()); - aiGatewayServer.post("/chat/completions", aiGateway.chatCompletions); + aiGatewayServer.use(createAiGatewayRouter()); aiGatewayServer.use(notFoundMiddleware); aiGatewayServer.use(errorHandlerMiddleware); diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 9bf11a21d..3091ed618 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -1659,6 +1659,7 @@ export const aiProviders = pgTable("aiProviders", { .$type<"url" | "target">() .notNull() .default("url"), + capabilities: text("capabilities").notNull().default("[]"), skipTlsVerification: boolean("skipTlsVerification") .notNull() .default(false), diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 35028e815..5c115cb2e 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -1641,6 +1641,7 @@ export const aiProviders = sqliteTable("aiProviders", { .$type<"url" | "target">() .notNull() .default("url"), + capabilities: text("capabilities").notNull().default("[]"), skipTlsVerification: integer("skipTlsVerification", { mode: "boolean" }) .notNull() .default(false), diff --git a/server/lib/aiCapabilities.ts b/server/lib/aiCapabilities.ts new file mode 100644 index 000000000..b6d2f7984 --- /dev/null +++ b/server/lib/aiCapabilities.ts @@ -0,0 +1,257 @@ +import type { Request } from "express"; +import type { AiProviderType } from "@server/lib/aiProviderDefaults"; + +export const AI_CAPABILITIES = [ + "openai_chat", + "openai_responses", + "anthropic_messages", + "gemini_generate_content", + "bedrock_model_invoke", + "google_generate_content", + "google_raw_predict", + "bedrock_converse" +] as const; + +export type AiCapability = (typeof AI_CAPABILITIES)[number]; + +export type AiCapabilityRoute = { + method: "POST"; + path: string; +}; + +export type AiCapabilityDefinition = { + id: AiCapability; + routes: AiCapabilityRoute[]; + extractModel: (req: Request) => string | undefined; + resolveUpstreamUrl: ( + baseUrl: string, + req: Request, + model: string + ) => string; +}; + +function bodyModel(req: Request): string | undefined { + return typeof req.body?.model === "string" ? req.body.model : undefined; +} + +function paramModel(req: Request): string | undefined { + const model = req.params?.model; + return typeof model === "string" && model.length > 0 ? model : undefined; +} + +/** + * Join base URL with a path, avoiding double slashes and a duplicated trailing + * /v1 when the inbound path already starts with /v1 and the base ends with /v1. + */ +export function joinUpstreamUrl(baseUrl: string, path: string): string { + const base = baseUrl.replace(/\/+$/, ""); + let suffix = path.startsWith("/") ? path : `/${path}`; + + if ( + base.endsWith("/v1") && + (suffix === "/v1" || suffix.startsWith("/v1/")) + ) { + suffix = suffix.slice("/v1".length) || "/"; + } + + if (suffix === "/") { + return base; + } + + return `${base}${suffix}`; +} + +function pathFromRequest(req: Request): string { + // Prefer originalUrl path (includes mounted path) over req.path when available. + const raw = + req.originalUrl?.split("?")[0] || req.url?.split("?")[0] || req.path; + return raw.startsWith("/") ? raw : `/${raw}`; +} + +export const AI_CAPABILITY_DEFS: Record = + { + openai_chat: { + id: "openai_chat", + routes: [ + { method: "POST", path: "/v1/chat/completions" }, + { method: "POST", path: "/chat/completions" } + ], + extractModel: bodyModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + openai_responses: { + id: "openai_responses", + routes: [{ method: "POST", path: "/v1/responses" }], + extractModel: bodyModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + anthropic_messages: { + id: "anthropic_messages", + routes: [{ method: "POST", path: "/v1/messages" }], + extractModel: bodyModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + gemini_generate_content: { + id: "gemini_generate_content", + routes: [ + { + method: "POST", + path: "/v1beta/models/:model\\:generateContent" + }, + { + method: "POST", + path: "/v1beta/models/:model\\:streamGenerateContent" + } + ], + extractModel: paramModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + google_generate_content: { + id: "google_generate_content", + routes: [ + { + method: "POST", + // Vertex publisher model generateContent + path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:generateContent" + }, + { + method: "POST", + path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:streamGenerateContent" + } + ], + extractModel: paramModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + google_raw_predict: { + id: "google_raw_predict", + routes: [ + { + method: "POST", + path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:rawPredict" + }, + { + method: "POST", + path: "/v1/projects/:project/locations/:location/publishers/:publisher/models/:model\\:streamRawPredict" + } + ], + extractModel: paramModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + bedrock_model_invoke: { + id: "bedrock_model_invoke", + routes: [ + { method: "POST", path: "/model/:model/invoke" }, + { + method: "POST", + path: "/model/:model/invoke-with-response-stream" + } + ], + extractModel: paramModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + }, + bedrock_converse: { + id: "bedrock_converse", + routes: [ + { method: "POST", path: "/model/:model/converse" }, + { method: "POST", path: "/model/:model/converse-stream" } + ], + extractModel: paramModel, + resolveUpstreamUrl: (base, req) => + joinUpstreamUrl(base, pathFromRequest(req)) + } + }; + +export const AI_PROVIDER_CAPABILITY_DEFAULTS: Record< + Exclude, + readonly AiCapability[] +> = { + openai: ["openai_chat"], + anthropic: ["anthropic_messages"], + googleGemini: ["openai_chat"], + vertexAi: ["google_generate_content"], + bedrock: ["bedrock_converse"], + microsoftFoundry: ["openai_chat"], + openRouter: ["openai_chat"], + vercelAiGateway: ["openai_chat"] +}; + +export function isAiCapability(value: unknown): value is AiCapability { + return ( + typeof value === "string" && + (AI_CAPABILITIES as readonly string[]).includes(value) + ); +} + +export function parseCapabilities(raw: unknown): AiCapability[] { + if (raw == null) { + return []; + } + + let parsed: unknown = raw; + if (typeof raw === "string") { + const trimmed = raw.trim(); + if (!trimmed) { + return []; + } + try { + parsed = JSON.parse(trimmed); + } catch { + return []; + } + } + + if (!Array.isArray(parsed)) { + return []; + } + + const out: AiCapability[] = []; + const seen = new Set(); + for (const item of parsed) { + if (isAiCapability(item) && !seen.has(item)) { + seen.add(item); + out.push(item); + } + } + return out; +} + +export function serializeCapabilities(capabilities: AiCapability[]): string { + return JSON.stringify(capabilities); +} + +export function providerHasCapability( + capabilities: AiCapability[] | string | null | undefined, + capability: AiCapability +): boolean { + const list = + typeof capabilities === "string" || capabilities == null + ? parseCapabilities(capabilities) + : capabilities; + return list.includes(capability); +} + +export function resolveCapabilitiesForCreate(input: { + type: AiProviderType; + capabilities?: AiCapability[] | null; +}): AiCapability[] { + if (input.type === "custom") { + return parseCapabilities(input.capabilities ?? []); + } + return [...AI_PROVIDER_CAPABILITY_DEFAULTS[input.type]]; +} + +export function defaultsForProviderType( + type: AiProviderType +): readonly AiCapability[] { + if (type === "custom") { + return []; + } + return AI_PROVIDER_CAPABILITY_DEFAULTS[type]; +} diff --git a/server/routers/aiGateway/createAiGatewayRouter.ts b/server/routers/aiGateway/createAiGatewayRouter.ts new file mode 100644 index 000000000..cc62ade28 --- /dev/null +++ b/server/routers/aiGateway/createAiGatewayRouter.ts @@ -0,0 +1,21 @@ +import { Router } from "express"; +import { + AI_CAPABILITY_DEFS, + type AiCapability +} from "@server/lib/aiCapabilities"; +import { handleAiGatewayProxy } from "@server/routers/aiGateway/pipeline"; + +export function createAiGatewayRouter() { + const router = Router(); + + for (const def of Object.values(AI_CAPABILITY_DEFS)) { + const capability = def.id as AiCapability; + for (const route of def.routes) { + router.post(route.path, (req, res) => + handleAiGatewayProxy(req, res, capability) + ); + } + } + + return router; +} diff --git a/server/routers/aiGateway/index.ts b/server/routers/aiGateway/index.ts index 691ebebe8..6eea36d60 100644 --- a/server/routers/aiGateway/index.ts +++ b/server/routers/aiGateway/index.ts @@ -1 +1,2 @@ -export * from "./chatCompletions"; +export { handleAiGatewayProxy } from "./pipeline"; +export { createAiGatewayRouter } from "./createAiGatewayRouter"; diff --git a/server/routers/aiGateway/chatCompletions.ts b/server/routers/aiGateway/pipeline.ts similarity index 88% rename from server/routers/aiGateway/chatCompletions.ts rename to server/routers/aiGateway/pipeline.ts index d2671aa63..cf031c381 100644 --- a/server/routers/aiGateway/chatCompletions.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -22,6 +22,11 @@ import { applyAiProviderAuthHeaders, authTypeRequiresApiKey } from "@server/lib/aiProviderDefaults"; +import { + AI_CAPABILITY_DEFS, + providerHasCapability, + type AiCapability +} from "@server/lib/aiCapabilities"; import { SESSION_COOKIE_NAME, validateSessionToken @@ -45,10 +50,6 @@ const REQUEST_USER_TTL_SEC = 30; type CachedClient = { clientId: number; userId: string | null } | null; -// The set of CIDRs an exit node manages; client exitNodeSubnets are always -// /32s carved out of one of these ranges. Checking against this small, -// cacheable list lets us skip the (much more frequent) per-IP client lookup -// entirely for traffic that could never match a client anyway. async function getExitNodeRanges(): Promise { const cached = localCache.get(EXIT_NODE_RANGES_CACHE_KEY); if (cached) { @@ -96,8 +97,6 @@ type ResolvedTarget = { siteResourceId: number | null; orgId: string | null; attachments: ProviderAttachment[]; - // Model IDs on the resource allowlist that belong to allowlist-mode - // providers. Empty when no attached provider uses allowlist mode. allowlistedModelIds: Set; }; @@ -150,13 +149,9 @@ async function buildRequestUser( async function resolveRequestUser( req: Request, - resourceId: number | null, + _resourceId: number | null, orgId: string | null ): Promise { - // Public resources behind badger: badger passes the resource session - // cookie through to the backend (same mechanism the browser gateway, - // e.g. the SSH page, relies on), so we can validate it exactly like - // verifySessionUserMiddleware does for the dashboard. const sessionToken = req.cookies?.[SESSION_COOKIE_NAME]; if (sessionToken) { const { session, user } = await validateSessionToken(sessionToken); @@ -167,10 +162,6 @@ async function resolveRequestUser( // TODO: MAKE SURE THIS CAN NOT BE SPOOFED AND CAN BE TRUSTED AS AN INTERNAL ADDRESS FROM A NODE - // No session cookie - fall back to identifying the caller by source IP. - // A client's exitNodeSubnet is a /32 handed out from one of our exit - // node's address ranges, so an IP that isn't inside any of those ranges - // can never belong to a client and we can skip the DB entirely. const ip = req.ip; if (!ip) { return null; @@ -193,9 +184,6 @@ async function resolveRequestUser( } async function resolveTarget(host: string): Promise { - // TODO: eventually we need to know if it's a private or public resource - // and not just simply check the fullDomain in case there is a private resource with the same fullDomain - const [resourceRow] = await db .select({ resourceId: resources.resourceId, @@ -378,7 +366,6 @@ async function selectProvider( }; } - // One lookup for the requested model key across all attached providers. const matchingModels = await db .select({ modelId: aiModels.modelId, @@ -407,7 +394,6 @@ async function selectProvider( continue; } - // allowlist: only models explicitly attached to the resource if (allowlistedModelIds.has(model.modelId)) { candidates.push(attachment.provider); } @@ -432,11 +418,14 @@ async function selectProvider( }; } -export async function chatCompletions( +export async function handleAiGatewayProxy( req: Request, - res: Response + res: Response, + capability: AiCapability ): Promise { try { + const def = AI_CAPABILITY_DEFS[capability]; + const host = ( (req.headers["p-host"] as string | undefined) || req.headers.host || @@ -448,7 +437,7 @@ export async function chatCompletions( .json({ error: { message: "Missing Host header" } }); } - logger.info(`AI gateway request for host: ${host}`); + logger.info(`AI gateway ${capability} request for host: ${host}`); const target = await resolveTarget(host); if (!target) { @@ -461,11 +450,6 @@ export async function chatCompletions( const { attachments, allowlistedModelIds, resourceId, orgId } = target; - logger.debug("+++++ gateway target: ", target); - - // Best-effort identity resolution - not yet enforced, but lets us - // start making per-user access decisions (e.g. model/role-based - // restrictions) without another round of plumbing later. const requestUser = await resolveRequestUser(req, resourceId, orgId); if (requestUser) { logger.debug( @@ -473,11 +457,22 @@ export async function chatCompletions( ); } - const requestedModel = - typeof req.body?.model === "string" ? req.body.model : undefined; + const capableAttachments = attachments.filter((a) => + providerHasCapability(a.provider.capabilities, capability) + ); + + if (capableAttachments.length === 0) { + return res.status(HttpCode.FORBIDDEN).json({ + error: { + message: `No AI provider on this resource supports ${capability}` + } + }); + } + + const requestedModel = def.extractModel(req); const selection = await selectProvider( - attachments, + capableAttachments, allowlistedModelIds, requestedModel ); @@ -513,11 +508,12 @@ export async function chatCompletions( apiKey = decrypt(provider.apiKey, secret); } - const targetUrl = `${upstreamUrl.replace(/\/$/, "")}`; + const targetUrl = def.resolveUpstreamUrl( + upstreamUrl, + req, + requestedModel! + ); - // Drop hop-by-hop / proxy-only headers. Forwarding Host especially - // breaks Node fetch (TLS/SNI targets the upstream URL while Host - // still says localhost). const skipHeaders = new Set([ "p-host", "host", @@ -554,6 +550,7 @@ export async function chatCompletions( const body = JSON.stringify(req.body); logger.debug("AI gateway upstream request", { + capability, url: targetUrl, method: "POST", headers, @@ -591,7 +588,11 @@ export async function chatCompletions( const contentType = upstreamRes.headers.get("content-type") || ""; const isStream = req.body?.stream === true || - contentType.includes("text/event-stream"); + contentType.includes("text/event-stream") || + req.path.includes("streamGenerateContent") || + req.path.includes("streamRawPredict") || + req.path.includes("converse-stream") || + req.path.includes("invoke-with-response-stream"); res.status(upstreamRes.status); res.setHeader("Content-Type", contentType || "application/json"); diff --git a/server/routers/aiProvider/createAiProvider.ts b/server/routers/aiProvider/createAiProvider.ts index 6188601a3..2fb61bf88 100644 --- a/server/routers/aiProvider/createAiProvider.ts +++ b/server/routers/aiProvider/createAiProvider.ts @@ -14,10 +14,15 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, + aiCapabilitiesSchema, aiProviderTypeSchema, aiRoutingModeSchema, refineProviderUpstreamFields } from "@server/routers/aiProvider/validation"; +import { + resolveCapabilitiesForCreate, + serializeCapabilities +} from "@server/lib/aiCapabilities"; const paramsSchema = z.strictObject({ orgId: z.string().nonempty() @@ -31,6 +36,7 @@ const bodySchema = z apiKey: z.string().optional(), authType: aiAuthTypeSchema.optional(), routingMode: aiRoutingModeSchema.optional(), + capabilities: aiCapabilitiesSchema.optional(), skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional() }) @@ -94,6 +100,7 @@ export async function createAiProvider( apiKey, authType, routingMode, + capabilities, skipTlsVerification, enabled } = parsedBody.data; @@ -108,6 +115,10 @@ export async function createAiProvider( authType, routingMode }); + const resolvedCapabilities = resolveCapabilitiesForCreate({ + type, + capabilities + }); const [provider] = await db .insert(aiProviders) @@ -120,6 +131,7 @@ export async function createAiProvider( apiKeyLastChars, authType: resolved.authType, routingMode: resolved.routingMode, + capabilities: serializeCapabilities(resolvedCapabilities), skipTlsVerification: skipTlsVerification ?? false, enabled: enabled ?? true, createdAt: now, diff --git a/server/routers/aiProvider/types.ts b/server/routers/aiProvider/types.ts index 7ac492056..5a1a41249 100644 --- a/server/routers/aiProvider/types.ts +++ b/server/routers/aiProvider/types.ts @@ -1,12 +1,17 @@ import type { AiModel, AiProvider } from "@server/db"; import type { PaginatedResponse } from "@server/types/Pagination"; import type { AiProviderAuthType } from "@server/lib/aiProviderDefaults"; +import { + parseCapabilities, + type AiCapability +} from "@server/lib/aiCapabilities"; import { decrypt } from "@server/lib/crypto"; import config from "@server/lib/config"; -export type AiProviderPublic = Omit & { +export type AiProviderPublic = Omit & { /** Decrypted API key. Only included on get/create/update of a single provider. */ apiKey?: string | null; + capabilities: AiCapability[]; effectiveUpstreamUrl: string | null; effectiveAuthType: AiProviderAuthType; }; @@ -39,7 +44,11 @@ export function toPublicAiProvider( provider: AiProvider, options?: { includeApiKey?: boolean } ): AiProviderPublic { - const { apiKey: encryptedApiKey, ...rest } = provider; + const { + apiKey: encryptedApiKey, + capabilities: rawCapabilities, + ...rest + } = provider; let apiKey: string | null | undefined; if (options?.includeApiKey) { @@ -56,6 +65,7 @@ export function toPublicAiProvider( return { ...rest, ...(options?.includeApiKey ? { apiKey } : {}), + capabilities: parseCapabilities(rawCapabilities), effectiveUpstreamUrl: provider.upstreamUrl, effectiveAuthType: provider.authType as AiProviderAuthType }; diff --git a/server/routers/aiProvider/updateAiProvider.ts b/server/routers/aiProvider/updateAiProvider.ts index 82224c28b..9dc28be60 100644 --- a/server/routers/aiProvider/updateAiProvider.ts +++ b/server/routers/aiProvider/updateAiProvider.ts @@ -14,6 +14,7 @@ import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, + aiCapabilitiesSchema, aiProviderTypeSchema, aiRoutingModeSchema, refineProviderUpstreamFields @@ -23,6 +24,10 @@ import type { AiProviderRoutingMode, AiProviderType } from "@server/lib/aiProviderDefaults"; +import { + parseCapabilities, + serializeCapabilities +} from "@server/lib/aiCapabilities"; const paramsSchema = z.strictObject({ providerId: z.coerce.number().int().positive() @@ -34,6 +39,7 @@ const bodySchema = z.strictObject({ apiKey: z.string().optional(), authType: aiAuthTypeSchema.optional(), routingMode: aiRoutingModeSchema.optional(), + capabilities: aiCapabilitiesSchema.optional(), skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional() }); @@ -122,19 +128,37 @@ 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) + : parseCapabilities(existing.capabilities); + const validation = z .object({ type: aiProviderTypeSchema, upstreamUrl: z.string().nullable().optional(), authType: aiAuthTypeSchema, - routingMode: aiRoutingModeSchema.optional() + routingMode: aiRoutingModeSchema.optional(), + capabilities: aiCapabilitiesSchema.optional() }) .superRefine((data, ctx) => refineProviderUpstreamFields(data, ctx)) .safeParse({ type: providerType, upstreamUrl: nextUpstreamUrl, authType: nextAuthType, - routingMode: nextRoutingMode + routingMode: nextRoutingMode, + capabilities: nextCapabilities }); if (!validation.success) { @@ -168,6 +192,9 @@ export async function updateAiProvider( if (body.authType !== undefined) { updateData.authType = body.authType; } + if (providerType === "custom" && body.capabilities !== undefined) { + updateData.capabilities = serializeCapabilities(body.capabilities); + } if (body.apiKey !== undefined) { const key = config.getRawConfig().server.secret!; diff --git a/server/routers/aiProvider/validation.ts b/server/routers/aiProvider/validation.ts index b4ab104aa..957a3914f 100644 --- a/server/routers/aiProvider/validation.ts +++ b/server/routers/aiProvider/validation.ts @@ -6,6 +6,7 @@ import { type AiProviderRoutingMode, type AiProviderType } from "@server/lib/aiProviderDefaults"; +import { AI_CAPABILITIES } from "@server/lib/aiCapabilities"; export const aiProviderTypeSchema = z.enum([ "openai", @@ -23,12 +24,17 @@ export const aiAuthTypeSchema = z.enum(AI_PROVIDER_AUTH_TYPES); export const aiRoutingModeSchema = z.enum(["url", "target"]); +export const aiCapabilitySchema = z.enum(AI_CAPABILITIES); + +export const aiCapabilitiesSchema = z.array(aiCapabilitySchema); + export function refineProviderUpstreamFields( data: { type: AiProviderType; upstreamUrl?: string | null; authType?: AiProviderAuthType | null; routingMode?: AiProviderRoutingMode | null; + capabilities?: z.infer | null; }, ctx: z.RefinementCtx ) { @@ -52,4 +58,16 @@ export function refineProviderUpstreamFields( path: ["upstreamUrl"] }); } + + if (data.type === "custom") { + const caps = data.capabilities; + if (!caps || caps.length === 0) { + ctx.addIssue({ + code: "custom", + message: + "At least one capability is required for custom providers", + 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 1018f8e63..62ffa6eb6 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/general/page.tsx @@ -12,11 +12,16 @@ import { SettingsSectionHeader, SettingsSectionTitle } from "@app/components/Settings"; +import { + AiProviderCapabilitiesSelect, + capabilityLabelKey +} from "@app/components/AiProviderCapabilitiesSelect"; import { SwitchInput } from "@app/components/SwitchInput"; import { Button } from "@app/components/ui/button"; import { Form, FormControl, + FormDescription, FormField, FormItem, FormLabel, @@ -28,6 +33,7 @@ import { useEnvContext } from "@app/hooks/useEnvContext"; import { toast } from "@app/hooks/useToast"; import { createApiClient, formatAxiosError } from "@app/lib/api"; import { zodResolver } from "@hookform/resolvers/zod"; +import { AI_CAPABILITIES, type AiCapability } from "@server/lib/aiCapabilities"; import type { CreateOrEditAiProviderResponse } from "@server/routers/aiProvider/types"; import type { AxiosResponse } from "axios"; import { useTranslations } from "next-intl"; @@ -43,17 +49,32 @@ export default function AiProviderGeneralPage() { const router = useRouter(); const t = useTranslations(); const [saveLoading, setSaveLoading] = useState(false); + const isCustom = provider.type === "custom"; const generalSchema = useMemo( () => - z.object({ - name: z - .string() - .trim() - .min(1, { message: t("nameRequired") }), - enabled: z.boolean() - }), - [t] + z + .object({ + name: z + .string() + .trim() + .min(1, { message: t("nameRequired") }), + enabled: z.boolean(), + capabilities: z.array(z.enum(AI_CAPABILITIES)).optional() + }) + .superRefine((data, ctx) => { + if ( + isCustom && + (!data.capabilities || data.capabilities.length === 0) + ) { + ctx.addIssue({ + code: "custom", + message: t("aiProviderErrorCapabilitiesRequired"), + path: ["capabilities"] + }); + } + }), + [t, isCustom] ); type GeneralFormValues = z.infer; @@ -62,24 +83,35 @@ export default function AiProviderGeneralPage() { resolver: zodResolver(generalSchema), defaultValues: { name: provider.name, - enabled: provider.enabled + enabled: provider.enabled, + capabilities: provider.capabilities ?? [] } }); async function onSubmit(values: GeneralFormValues) { setSaveLoading(true); try { - const res = await api.post< - AxiosResponse - >(`/ai-provider/${provider.providerId}`, { + const body: { + name: string; + enabled: boolean; + capabilities?: AiCapability[]; + } = { name: values.name.trim(), enabled: values.enabled - }); + }; + if (isCustom) { + body.capabilities = values.capabilities ?? []; + } + + const res = await api.post< + AxiosResponse + >(`/ai-provider/${provider.providerId}`, body); const updated = res.data.data.provider; updateProvider(updated); form.reset({ name: updated.name, - enabled: updated.enabled + enabled: updated.enabled, + capabilities: updated.capabilities ?? [] }); toast({ title: t("success"), @@ -166,6 +198,65 @@ export default function AiProviderGeneralPage() { )} /> + + + ( + + + {t( + "aiProviderCapabilities" + )} + + + {isCustom ? ( + + ) : ( +
+ {( + provider.capabilities ?? + [] + ).map((cap) => ( + + {t( + capabilityLabelKey( + cap + ) + )} + + ))} +
+ )} +
+ + {isCustom + ? t( + "aiProviderCapabilitiesCustomDescription" + ) + : 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 1947c3f85..34fac7b55 100644 --- a/src/app/[orgId]/settings/ai-providers/create/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/create/page.tsx @@ -20,6 +20,10 @@ 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 { AiProviderTypeSelect } from "@app/components/AiProviderTypeSelect"; import { StrategySelect } from "@app/components/StrategySelect"; import { SwitchInput } from "@app/components/SwitchInput"; @@ -40,6 +44,7 @@ import { createApiClient, formatAxiosError } from "@app/lib/api"; import { createAiProviderCreateFormSchema, defaultAuthTypeForProvider, + defaultCapabilitiesForProvider, emptyUpstreamForType, showsUpstreamUrlField, toAiProviderCreatePayload, @@ -76,6 +81,7 @@ export default function CreateAiProviderPage() { apiKey: "", authType: defaultAuthTypeForProvider("openai"), routingMode: "url", + capabilities: defaultCapabilitiesForProvider("openai"), skipTlsVerification: false, enabled: true } @@ -84,12 +90,14 @@ 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, @@ -277,6 +285,12 @@ export default function CreateAiProviderPage() { value ) ); + form.setValue( + "capabilities", + defaultCapabilitiesForProvider( + value + ) + ); if ( value !== "custom" @@ -296,6 +310,67 @@ export default function CreateAiProviderPage() { )} /> + + + ( + + + {t( + "aiProviderCapabilities" + )} + + + {showCapabilitiesSelect ? ( + + ) : ( +
+ {( + capabilities ?? + defaultCapabilitiesForProvider( + providerType + ) + ).map((cap) => ( + + {t( + capabilityLabelKey( + cap + ) + )} + + ))} +
+ )} +
+ + {showCapabilitiesSelect + ? t( + "aiProviderCapabilitiesCustomDescription" + ) + : t( + "aiProviderCapabilitiesDescription" + )} + + +
+ )} + /> +
diff --git a/src/components/AiProviderCapabilitiesSelect.tsx b/src/components/AiProviderCapabilitiesSelect.tsx new file mode 100644 index 000000000..086b6f03e --- /dev/null +++ b/src/components/AiProviderCapabilitiesSelect.tsx @@ -0,0 +1,89 @@ +"use client"; + +import { MultiSelectTagInput } from "@app/components/multi-select/multi-select-tag-input"; +import { AI_CAPABILITIES, type AiCapability } from "@server/lib/aiCapabilities"; +import { useTranslations } from "next-intl"; +import { useMemo, useState } from "react"; + +export type CapabilityOption = { + id: string; + text: string; +}; + +export type AiProviderCapabilitiesSelectProps = { + value: AiCapability[]; + onChange: (capabilities: AiCapability[]) => void; + disabled?: boolean; +}; + +const CAPABILITY_LABEL_KEYS: Record = { + openai_chat: "aiCapabilityOpenaiChat", + openai_responses: "aiCapabilityOpenaiResponses", + anthropic_messages: "aiCapabilityAnthropicMessages", + gemini_generate_content: "aiCapabilityGeminiGenerateContent", + bedrock_model_invoke: "aiCapabilityBedrockModelInvoke", + google_generate_content: "aiCapabilityGoogleGenerateContent", + google_raw_predict: "aiCapabilityGoogleRawPredict", + bedrock_converse: "aiCapabilityBedrockConverse" +}; + +export function capabilityLabelKey(capability: AiCapability): string { + return CAPABILITY_LABEL_KEYS[capability]; +} + +export function AiProviderCapabilitiesSelect({ + value, + onChange, + disabled +}: AiProviderCapabilitiesSelectProps) { + const t = useTranslations(); + const [searchQuery, setSearchQuery] = useState(""); + + const options: CapabilityOption[] = useMemo( + () => + AI_CAPABILITIES.map((id) => ({ + id, + text: t(CAPABILITY_LABEL_KEYS[id]) + })), + [t] + ); + + const filtered = useMemo(() => { + const q = searchQuery.trim().toLowerCase(); + if (!q) { + return options; + } + return options.filter( + (o) => + o.text.toLowerCase().includes(q) || + o.id.toLowerCase().includes(q) + ); + }, [options, searchQuery]); + + const selected: CapabilityOption[] = value.map((id) => ({ + id, + text: t(CAPABILITY_LABEL_KEYS[id]) + })); + + return ( + + onChange( + next + .map((item) => item.id) + .filter((id): id is AiCapability => + (AI_CAPABILITIES as readonly string[]).includes(id) + ) + ) + } + onSearch={setSearchQuery} + disabled={disabled} + /> + ); +} diff --git a/src/lib/aiProviderFormSchema.ts b/src/lib/aiProviderFormSchema.ts index f484d1f05..b01728478 100644 --- a/src/lib/aiProviderFormSchema.ts +++ b/src/lib/aiProviderFormSchema.ts @@ -7,6 +7,11 @@ import { type AiProviderAuthType, type AiProviderType } from "@server/lib/aiProviderDefaults"; +import { + AI_CAPABILITIES, + defaultsForProviderType, + type AiCapability +} from "@server/lib/aiCapabilities"; type TranslateFn = (key: string) => string; @@ -22,6 +27,8 @@ export const aiProviderTypeValues = [ "custom" ] as const satisfies readonly AiProviderType[]; +export const aiCapabilityValues = AI_CAPABILITIES; + export function createAiProviderFormSchema(t: TranslateFn) { return z .object({ @@ -34,6 +41,7 @@ export function createAiProviderFormSchema(t: TranslateFn) { apiKey: z.string().optional(), authType: z.enum(AI_PROVIDER_AUTH_TYPES).optional().nullable(), routingMode: z.enum(["url", "target"]).optional(), + capabilities: z.array(z.enum(AI_CAPABILITIES)).optional(), skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional() }) @@ -84,6 +92,17 @@ export function createAiProviderFormSchema(t: TranslateFn) { path: ["authType"] }); } + + if ( + data.type === "custom" && + (!data.capabilities || data.capabilities.length === 0) + ) { + ctx.addIssue({ + code: "custom", + message: t("aiProviderErrorCapabilitiesRequired"), + path: ["capabilities"] + }); + } }); } @@ -114,6 +133,12 @@ export function defaultAuthTypeForProvider( return AI_PROVIDER_DEFAULTS[type].authType; } +export function defaultCapabilitiesForProvider( + type: AiProviderType +): AiCapability[] { + return [...defaultsForProviderType(type)]; +} + export function emptyUpstreamForType(type: AiProviderType): string { if (type === "custom") { return ""; @@ -158,6 +183,8 @@ 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, skipTlsVerification: values.skipTlsVerification, enabled: values.enabled ?? true }; @@ -183,6 +210,10 @@ export function toAiProviderUpdatePayload(values: AiProviderFormValues) { enabled: values.enabled ?? true }; + if (values.type === "custom" && values.capabilities) { + payload.capabilities = values.capabilities; + } + if (values.apiKey?.trim()) { payload.apiKey = values.apiKey.trim(); }