diff --git a/cli/commands/rotateServerSecret.ts b/cli/commands/rotateServerSecret.ts index afac262b2..2edb7592d 100644 --- a/cli/commands/rotateServerSecret.ts +++ b/cli/commands/rotateServerSecret.ts @@ -1,5 +1,5 @@ import { CommandModule } from "yargs"; -import { db, idpOidcConfig, licenseKey, certificates, eventStreamingDestinations, alertWebhookActions } from "@server/db"; +import { db, idpOidcConfig, licenseKey, certificates, eventStreamingDestinations, alertWebhookActions, aiProviders } from "@server/db"; import { encrypt, decrypt } from "@server/lib/crypto"; import { configFilePath1, configFilePath2 } from "@server/lib/consts"; import { eq } from "drizzle-orm"; @@ -132,12 +132,14 @@ export const rotateServerSecret: CommandModule< const certs = await db.select().from(certificates); const streamingDestinations = await db.select().from(eventStreamingDestinations); const webhookActions = await db.select().from(alertWebhookActions); + const providers = await db.select().from(aiProviders); console.log(`Found ${idpConfigs.length} OIDC IdP configuration(s)`); console.log(`Found ${licenseKeys.length} license key(s)`); console.log(`Found ${certs.length} certificate(s)`); console.log(`Found ${streamingDestinations.length} event streaming destination(s)`); console.log(`Found ${webhookActions.length} alert webhook action(s)`); + console.log(`Found ${providers.length} AI provider(s)`); // Prepare all decrypted and re-encrypted values console.log("\nDecrypting and re-encrypting values..."); @@ -171,11 +173,18 @@ export const rotateServerSecret: CommandModule< encryptedConfig: string; }; + type AiProviderUpdate = { + providerId: number; + encryptedApiKey: string | null; + encryptedHeaders: string | null; + }; + const idpUpdates: IdpUpdate[] = []; const licenseKeyUpdates: LicenseKeyUpdate[] = []; const certUpdates: CertUpdate[] = []; const streamingDestinationUpdates: StreamingDestinationUpdate[] = []; const webhookActionUpdates: WebhookActionUpdate[] = []; + const aiProviderUpdates: AiProviderUpdate[] = []; // Process idpOidcConfig entries for (const idpConfig of idpConfigs) { @@ -306,6 +315,37 @@ export const rotateServerSecret: CommandModule< } } + // Process aiProviders entries (apiKey + headers) + for (const provider of providers) { + try { + if (!provider.apiKey && !provider.headers) { + continue; + } + + const encryptedApiKey = provider.apiKey + ? encrypt(decrypt(provider.apiKey, oldSecret), newSecret) + : null; + const encryptedHeaders = provider.headers + ? encrypt( + decrypt(provider.headers, oldSecret), + newSecret + ) + : null; + + aiProviderUpdates.push({ + providerId: provider.providerId, + encryptedApiKey, + encryptedHeaders + }); + } catch (error) { + console.error( + `Error processing AI provider ${provider.providerId}:`, + error + ); + throw error; + } + } + // Perform all database updates in a single transaction console.log("\nUpdating database in transaction..."); await db.transaction(async (trx) => { @@ -376,6 +416,17 @@ export const rotateServerSecret: CommandModule< ) ); } + + // Update AI provider entries + for (const update of aiProviderUpdates) { + await trx + .update(aiProviders) + .set({ + apiKey: update.encryptedApiKey, + headers: update.encryptedHeaders + }) + .where(eq(aiProviders.providerId, update.providerId)); + } }); console.log(`Rotated ${idpUpdates.length} OIDC IdP configuration(s)`); @@ -383,6 +434,7 @@ export const rotateServerSecret: CommandModule< console.log(`Rotated ${certUpdates.length} certificate(s)`); console.log(`Rotated ${streamingDestinationUpdates.length} event streaming destination(s)`); console.log(`Rotated ${webhookActionUpdates.length} alert webhook action(s)`); + console.log(`Rotated ${aiProviderUpdates.length} AI provider(s)`); // Update config file with new secret console.log("\nUpdating config file..."); @@ -402,6 +454,7 @@ export const rotateServerSecret: CommandModule< console.log(` - Certificates: ${certUpdates.length}`); console.log(` - Event streaming destinations: ${streamingDestinationUpdates.length}`); console.log(` - Alert webhook actions: ${webhookActionUpdates.length}`); + console.log(` - AI providers: ${aiProviderUpdates.length}`); console.log( `\n IMPORTANT: Restart the server for the new secret to take effect.` ); diff --git a/messages/en-US.json b/messages/en-US.json index ac6131abc..a10a45421 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -1680,6 +1680,7 @@ "aiProviderEffectiveUpstreamUrl": "Effective Upstream URL", "aiProviderApiKey": "API Key", "aiProviderApiKeyDescription": "API key used to authenticate requests to this provider", + "aiProviderCustomHeadersDescription": "Headers sent on every request to this provider. Newline separated: Header-Name: value", "aiProviderApiKeyLastChars": "API Key", "aiProviderAuthType": "Auth Type", "aiProviderAuthTypeSearch": "Search auth types...", diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 3091ed618..6bc038660 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -1660,6 +1660,7 @@ export const aiProviders = pgTable("aiProviders", { .notNull() .default("url"), capabilities: text("capabilities").notNull().default("[]"), + headers: text("headers"), // JSON array of { name, value } skipTlsVerification: boolean("skipTlsVerification") .notNull() .default(false), diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 5c115cb2e..8130af9ea 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -1642,6 +1642,7 @@ export const aiProviders = sqliteTable("aiProviders", { .notNull() .default("url"), capabilities: text("capabilities").notNull().default("[]"), + headers: text("headers"), // JSON array of { name, value } skipTlsVerification: integer("skipTlsVerification", { mode: "boolean" }) .notNull() .default(false), diff --git a/server/lib/aiProviderDefaults.ts b/server/lib/aiProviderDefaults.ts index 70b666954..e48c6f877 100644 --- a/server/lib/aiProviderDefaults.ts +++ b/server/lib/aiProviderDefaults.ts @@ -1,3 +1,5 @@ +import { decrypt, encrypt } from "@server/lib/crypto"; + export type AiProviderType = | "openai" | "anthropic" @@ -127,6 +129,53 @@ export function resolveAiProviderCreateFields(input: { }; } +export type AiProviderHeader = { name: string; value: string }; + +export function serializeAiProviderHeaders( + headers: AiProviderHeader[] | null | undefined, + secret: string +): string | null { + if (!headers || headers.length === 0) { + return null; + } + return encrypt(JSON.stringify(headers), secret); +} + +export function parseAiProviderHeaders( + raw: string | null | undefined, + secret: string +): AiProviderHeader[] { + if (!raw) { + return []; + } + try { + const decrypted = decrypt(raw, secret); + const parsed = JSON.parse(decrypted); + if (!Array.isArray(parsed)) { + return []; + } + return parsed.filter( + (h): h is AiProviderHeader => + h != null && + typeof h === "object" && + typeof h.name === "string" && + typeof h.value === "string" + ); + } catch { + return []; + } +} + +export function applyAiProviderCustomHeaders( + headers: Record, + raw: string | null | undefined, + secret: string +): void { + for (const { name, value } of parseAiProviderHeaders(raw, secret)) { + headers[name] = value; + } +} + /** * Apply provider auth to upstream headers. * - Injected modes: strip client auth headers, then set the provider key. diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index cf031c381..44e2aef66 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -20,6 +20,7 @@ import { decrypt } from "@server/lib/crypto"; import { AiProviderAuthType, applyAiProviderAuthHeaders, + applyAiProviderCustomHeaders, authTypeRequiresApiKey } from "@server/lib/aiProviderDefaults"; import { @@ -536,6 +537,11 @@ export async function handleAiGatewayProxy( } headers[key] = Array.isArray(value) ? value.join(", ") : value; } + applyAiProviderCustomHeaders( + headers, + provider.headers, + config.getRawConfig().server.secret! + ); applyAiProviderAuthHeaders(headers, authType, apiKey); // No dedicated per-request TLS agent is wired up (no extra deps for diff --git a/server/routers/aiProvider/createAiProvider.ts b/server/routers/aiProvider/createAiProvider.ts index 2fb61bf88..995e9eaa0 100644 --- a/server/routers/aiProvider/createAiProvider.ts +++ b/server/routers/aiProvider/createAiProvider.ts @@ -15,10 +15,12 @@ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, aiCapabilitiesSchema, + aiProviderHeadersSchema, aiProviderTypeSchema, aiRoutingModeSchema, refineProviderUpstreamFields } from "@server/routers/aiProvider/validation"; +import { serializeAiProviderHeaders } from "@server/lib/aiProviderDefaults"; import { resolveCapabilitiesForCreate, serializeCapabilities @@ -37,6 +39,7 @@ const bodySchema = z authType: aiAuthTypeSchema.optional(), routingMode: aiRoutingModeSchema.optional(), capabilities: aiCapabilitiesSchema.optional(), + headers: aiProviderHeadersSchema, skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional() }) @@ -101,6 +104,7 @@ export async function createAiProvider( authType, routingMode, capabilities, + headers, skipTlsVerification, enabled } = parsedBody.data; @@ -132,6 +136,7 @@ export async function createAiProvider( authType: resolved.authType, routingMode: resolved.routingMode, capabilities: serializeCapabilities(resolvedCapabilities), + headers: serializeAiProviderHeaders(headers, key), skipTlsVerification: skipTlsVerification ?? false, enabled: enabled ?? true, createdAt: now, diff --git a/server/routers/aiProvider/types.ts b/server/routers/aiProvider/types.ts index 5a1a41249..d2f314b30 100644 --- a/server/routers/aiProvider/types.ts +++ b/server/routers/aiProvider/types.ts @@ -1,6 +1,10 @@ import type { AiModel, AiProvider } from "@server/db"; import type { PaginatedResponse } from "@server/types/Pagination"; -import type { AiProviderAuthType } from "@server/lib/aiProviderDefaults"; +import { + parseAiProviderHeaders, + type AiProviderAuthType, + type AiProviderHeader +} from "@server/lib/aiProviderDefaults"; import { parseCapabilities, type AiCapability @@ -8,10 +12,13 @@ import { import { decrypt } from "@server/lib/crypto"; import config from "@server/lib/config"; -export type AiProviderPublic = Omit & { - /** Decrypted API key. Only included on get/create/update of a single provider. */ +export type AiProviderPublic = Omit< + AiProvider, + "apiKey" | "capabilities" | "headers" +> & { apiKey?: string | null; capabilities: AiCapability[]; + headers: AiProviderHeader[] | null; effectiveUpstreamUrl: string | null; effectiveAuthType: AiProviderAuthType; }; @@ -47,6 +54,7 @@ export function toPublicAiProvider( const { apiKey: encryptedApiKey, capabilities: rawCapabilities, + headers: rawHeaders, ...rest } = provider; @@ -62,10 +70,16 @@ export function toPublicAiProvider( } } + const parsedHeaders = parseAiProviderHeaders( + rawHeaders, + config.getRawConfig().server.secret! + ); + return { ...rest, ...(options?.includeApiKey ? { apiKey } : {}), capabilities: parseCapabilities(rawCapabilities), + headers: parsedHeaders.length > 0 ? parsedHeaders : null, effectiveUpstreamUrl: provider.upstreamUrl, effectiveAuthType: provider.authType as AiProviderAuthType }; diff --git a/server/routers/aiProvider/updateAiProvider.ts b/server/routers/aiProvider/updateAiProvider.ts index 9dc28be60..7cbf1145a 100644 --- a/server/routers/aiProvider/updateAiProvider.ts +++ b/server/routers/aiProvider/updateAiProvider.ts @@ -15,14 +15,16 @@ import { toPublicAiProvider } from "@server/routers/aiProvider/types"; import { aiAuthTypeSchema, aiCapabilitiesSchema, + aiProviderHeadersSchema, aiProviderTypeSchema, aiRoutingModeSchema, refineProviderUpstreamFields } from "@server/routers/aiProvider/validation"; -import type { - AiProviderAuthType, - AiProviderRoutingMode, - AiProviderType +import { + serializeAiProviderHeaders, + type AiProviderAuthType, + type AiProviderRoutingMode, + type AiProviderType } from "@server/lib/aiProviderDefaults"; import { parseCapabilities, @@ -40,6 +42,7 @@ const bodySchema = z.strictObject({ authType: aiAuthTypeSchema.optional(), routingMode: aiRoutingModeSchema.optional(), capabilities: aiCapabilitiesSchema.optional(), + headers: aiProviderHeadersSchema, skipTlsVerification: z.boolean().optional(), enabled: z.boolean().optional() }); @@ -202,6 +205,11 @@ export async function updateAiProvider( updateData.apiKeyLastChars = body.apiKey.slice(-4); } + if (body.headers !== undefined) { + const key = config.getRawConfig().server.secret!; + updateData.headers = serializeAiProviderHeaders(body.headers, key); + } + const [provider] = await db .update(aiProviders) .set(updateData) diff --git a/server/routers/aiProvider/validation.ts b/server/routers/aiProvider/validation.ts index 957a3914f..a56ccad1c 100644 --- a/server/routers/aiProvider/validation.ts +++ b/server/routers/aiProvider/validation.ts @@ -28,6 +28,49 @@ export const aiCapabilitySchema = z.enum(AI_CAPABILITIES); export const aiCapabilitiesSchema = z.array(aiCapabilitySchema); +const validHeaderName = /^[a-zA-Z0-9!#$%&'*+\-.^_`|~]+$/; +const validHeaderValue = /^[\t\x20-\x7E]*$/; +const templatePattern = /\{\{[^}]+\}\}/; + +export const aiProviderHeadersSchema = z + .array(z.strictObject({ name: z.string(), value: z.string() })) + .nullable() + .optional() + .superRefine((headers, ctx) => { + if (!headers) { + return; + } + for (const [index, header] of headers.entries()) { + if (!validHeaderName.test(header.name)) { + ctx.addIssue({ + code: "custom", + message: + "Header names may only contain valid HTTP token characters (letters, digits, and !#$%&'*+-.^_`|~).", + path: [index, "name"] + }); + } + if (!validHeaderValue.test(header.value)) { + ctx.addIssue({ + code: "custom", + message: + "Header values may only contain printable ASCII characters and horizontal whitespace.", + path: [index, "value"] + }); + } + if ( + templatePattern.test(header.name) || + templatePattern.test(header.value) + ) { + ctx.addIssue({ + code: "custom", + message: + "Header names and values must not contain template expressions such as {{value}}.", + path: [index] + }); + } + } + }); + export function refineProviderUpstreamFields( data: { type: AiProviderType; diff --git a/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx b/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx index 49f251cfa..8b8814303 100644 --- a/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx +++ b/src/app/[orgId]/settings/ai-providers/[providerId]/network/page.tsx @@ -21,6 +21,7 @@ import { } from "@app/components/Settings"; import { StrategySelect } from "@app/components/StrategySelect"; import { SwitchInput } from "@app/components/SwitchInput"; +import { HeadersInput } from "@app/components/HeadersInput"; import { Button } from "@app/components/ui/button"; import { Form, @@ -66,6 +67,7 @@ export default function AiProviderNetworkPage() { const router = useRouter(); const t = useTranslations(); const [saveLoading, setSaveLoading] = useState(false); + const [headersValid, setHeadersValid] = useState(true); const targetsFormRef = useRef(null); const formSchema = useMemo(() => createAiProviderFormSchema(t), [t]); @@ -79,6 +81,7 @@ export default function AiProviderNetworkPage() { apiKey: "", authType: (provider.authType as AiProviderAuthType) ?? "bearer", routingMode: (provider.routingMode as "url" | "target") ?? "url", + headers: provider.headers ?? [], skipTlsVerification: provider.skipTlsVerification, enabled: provider.enabled } @@ -122,6 +125,7 @@ export default function AiProviderNetworkPage() { apiKey: "", authType: (updated.authType as AiProviderAuthType) ?? "bearer", routingMode: (updated.routingMode as "url" | "target") ?? "url", + headers: updated.headers ?? [], skipTlsVerification: updated.skipTlsVerification, enabled: updated.enabled }); @@ -310,6 +314,38 @@ export default function AiProviderNetworkPage() { /> )} + + + ( + + + {t("customHeaders")} + + + + + + {t( + "aiProviderCustomHeadersDescription" + )} + + + + )} + /> + @@ -346,7 +382,7 @@ export default function AiProviderNetworkPage() {