diff --git a/server/lib/aiCapabilities.ts b/server/lib/aiCapabilities.ts index c9f78f10e..b43cf3c1f 100644 --- a/server/lib/aiCapabilities.ts +++ b/server/lib/aiCapabilities.ts @@ -8,8 +8,11 @@ export type AiCapabilityRoute = { path: string; }; +export type AiProtocolFamily = "openai" | "anthropic" | "google" | "bedrock"; + export type AiCapabilityDefinition = { id: AiCapability; + protocolFamily: AiProtocolFamily; routes: AiCapabilityRoute[]; extractModel: (req: Request) => string | undefined; resolveUpstreamUrl: ( @@ -18,41 +21,8 @@ export type AiCapabilityDefinition = { model: string ) => string; isStreaming: (req: Request, contentType: string) => boolean; - /** Protocol-shaped body returned when gateway auth fails for this capability. */ - authErrorBody: Record; }; -const AUTH_MESSAGE = "Invalid API key provided."; - -export const OPENAI_AUTH_ERROR_BODY = { - error: { - message: AUTH_MESSAGE, - type: "authentication_error", - param: null, - code: "invalid_api_key" - } -} as const; - -const ANTHROPIC_AUTH_ERROR_BODY = { - type: "error", - error: { - type: "authentication_error", - message: AUTH_MESSAGE - } -} as const; - -const GOOGLE_AUTH_ERROR_BODY = { - error: { - code: 401, - message: AUTH_MESSAGE, - status: "UNAUTHENTICATED" - } -} as const; - -const BEDROCK_AUTH_ERROR_BODY = { - message: "The security token included in the request is invalid." -} as const; - function bodyModel(req: Request): string | undefined { return typeof req.body?.model === "string" ? req.body.model : undefined; } @@ -137,6 +107,7 @@ export const AI_CAPABILITY_DEFS: Record = { openai_chat: { id: "openai_chat", + protocolFamily: "openai", routes: [ { method: "POST", path: "/v1/chat/completions" }, { method: "POST", path: "/chat/completions" } @@ -144,29 +115,29 @@ export const AI_CAPABILITY_DEFS: Record = extractModel: bodyModel, resolveUpstreamUrl: (base, req) => joinUpstreamUrl(base, pathFromRequest(req)), - isStreaming: isBodyOrSseStreaming, - authErrorBody: OPENAI_AUTH_ERROR_BODY + isStreaming: isBodyOrSseStreaming }, openai_responses: { id: "openai_responses", + protocolFamily: "openai", routes: [{ method: "POST", path: "/v1/responses" }], extractModel: bodyModel, resolveUpstreamUrl: (base, req) => joinUpstreamUrl(base, pathFromRequest(req)), - isStreaming: isBodyOrSseStreaming, - authErrorBody: OPENAI_AUTH_ERROR_BODY + isStreaming: isBodyOrSseStreaming }, anthropic_messages: { id: "anthropic_messages", + protocolFamily: "anthropic", routes: [{ method: "POST", path: "/v1/messages" }], extractModel: bodyModel, resolveUpstreamUrl: (base, req) => joinUpstreamUrl(base, pathFromRequest(req)), - isStreaming: isBodyOrSseStreaming, - authErrorBody: ANTHROPIC_AUTH_ERROR_BODY + isStreaming: isBodyOrSseStreaming }, gemini_generate_content: { id: "gemini_generate_content", + protocolFamily: "google", routes: [ { method: "POST", @@ -180,11 +151,11 @@ export const AI_CAPABILITY_DEFS: Record = extractModel: paramModel, resolveUpstreamUrl: (base, req) => joinUpstreamUrl(base, pathFromRequest(req)), - isStreaming: isGeminiStyleStreaming, - authErrorBody: GOOGLE_AUTH_ERROR_BODY + isStreaming: isGeminiStyleStreaming }, google_generate_content: { id: "google_generate_content", + protocolFamily: "google", routes: [ { method: "POST", @@ -199,11 +170,11 @@ export const AI_CAPABILITY_DEFS: Record = extractModel: paramModel, resolveUpstreamUrl: (base, req) => joinUpstreamUrl(base, pathFromRequest(req)), - isStreaming: isGeminiStyleStreaming, - authErrorBody: GOOGLE_AUTH_ERROR_BODY + isStreaming: isGeminiStyleStreaming }, google_raw_predict: { id: "google_raw_predict", + protocolFamily: "google", routes: [ { method: "POST", @@ -220,11 +191,11 @@ export const AI_CAPABILITY_DEFS: Record = isStreaming: (req, contentType) => pathIncludes(req, "streamRawPredict") || pathIncludes(req, "alt=sse") || - contentTypeIsSse(contentType), - authErrorBody: GOOGLE_AUTH_ERROR_BODY + contentTypeIsSse(contentType) }, bedrock_model_invoke: { id: "bedrock_model_invoke", + protocolFamily: "bedrock", routes: [ { method: "POST", path: "/model/:model/invoke" }, { @@ -238,11 +209,11 @@ export const AI_CAPABILITY_DEFS: Record = isStreaming: (req, contentType) => pathIncludes(req, "invoke-with-response-stream") || contentTypeIsAmazonEventStream(contentType) || - contentTypeIsSse(contentType), - authErrorBody: BEDROCK_AUTH_ERROR_BODY + contentTypeIsSse(contentType) }, bedrock_converse: { id: "bedrock_converse", + protocolFamily: "bedrock", routes: [ { method: "POST", path: "/model/:model/converse" }, { method: "POST", path: "/model/:model/converse-stream" } @@ -253,8 +224,7 @@ export const AI_CAPABILITY_DEFS: Record = isStreaming: (req, contentType) => pathIncludes(req, "converse-stream") || contentTypeIsAmazonEventStream(contentType) || - contentTypeIsSse(contentType), - authErrorBody: BEDROCK_AUTH_ERROR_BODY + contentTypeIsSse(contentType) } }; diff --git a/server/lib/aiGatewayAuthError.ts b/server/lib/aiGatewayAuthError.ts index 55e4d5e28..ea2fba4ff 100644 --- a/server/lib/aiGatewayAuthError.ts +++ b/server/lib/aiGatewayAuthError.ts @@ -1,7 +1,7 @@ import { AI_CAPABILITY_DEFS, - OPENAI_AUTH_ERROR_BODY, - type AiCapability + type AiCapability, + type AiProtocolFamily } from "@server/lib/aiCapabilities"; import HttpCode from "@server/types/HttpCode"; @@ -11,17 +11,129 @@ export type ClientErrorResponse = { body: string; }; +export type AiCapabilityErrorKind = + | "authentication" + | "invalid_request" + | "not_found" + | "permission" + | "rate_limit" + | "internal"; + +const AUTH_MESSAGE = "Invalid API key provided."; + +type KindFields = { + openaiType: string; + openaiCode: string | null; + anthropicType: string; + googleStatus: string; +}; + +const KIND_FIELDS: Record = { + authentication: { + openaiType: "authentication_error", + openaiCode: "invalid_api_key", + anthropicType: "authentication_error", + googleStatus: "UNAUTHENTICATED" + }, + invalid_request: { + openaiType: "invalid_request_error", + openaiCode: null, + anthropicType: "invalid_request_error", + googleStatus: "INVALID_ARGUMENT" + }, + not_found: { + openaiType: "invalid_request_error", + openaiCode: null, + anthropicType: "not_found_error", + googleStatus: "NOT_FOUND" + }, + permission: { + openaiType: "invalid_request_error", + openaiCode: null, + anthropicType: "permission_error", + googleStatus: "PERMISSION_DENIED" + }, + rate_limit: { + openaiType: "rate_limit_error", + openaiCode: "rate_limit_exceeded", + anthropicType: "rate_limit_error", + googleStatus: "RESOURCE_EXHAUSTED" + }, + internal: { + openaiType: "api_error", + openaiCode: null, + anthropicType: "api_error", + googleStatus: "INTERNAL" + } +}; + +function resolveProtocolFamily( + capability: AiCapability | null +): AiProtocolFamily { + if (capability == null) { + return "openai"; + } + return AI_CAPABILITY_DEFS[capability].protocolFamily; +} + +/** + * Build a protocol-native error body for the given capability. + * Message stays contextual; only the envelope/machine fields follow the + * capability's native API shape. + */ +export function buildAiCapabilityErrorBody( + capability: AiCapability | null, + kind: AiCapabilityErrorKind, + message: string, + httpStatus?: number +): Record { + const family = resolveProtocolFamily(capability); + const fields = KIND_FIELDS[kind]; + + switch (family) { + case "openai": + return { + error: { + message, + type: fields.openaiType, + param: null, + code: fields.openaiCode + } + }; + case "anthropic": + return { + type: "error", + error: { + type: fields.anthropicType, + message + } + }; + case "google": + return { + error: { + code: httpStatus ?? HttpCode.BAD_REQUEST, + message, + status: fields.googleStatus + } + }; + case "bedrock": + return { message }; + } +} + export function buildInferenceAuthClientError( capability: AiCapability | null ): ClientErrorResponse { - const body = - capability != null - ? AI_CAPABILITY_DEFS[capability].authErrorBody - : OPENAI_AUTH_ERROR_BODY; - return { statusCode: HttpCode.UNAUTHORIZED, contentType: "application/json", - body: JSON.stringify(body) + body: JSON.stringify( + buildAiCapabilityErrorBody( + capability, + "authentication", + AUTH_MESSAGE, + HttpCode.UNAUTHORIZED + ) + ) }; } diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index 06ca9cdae..ea6676c70 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -31,6 +31,10 @@ import { providerHasCapability, type AiCapability } from "@server/lib/aiCapabilities"; +import { + buildAiCapabilityErrorBody, + type AiCapabilityErrorKind +} from "@server/lib/aiGatewayAuthError"; import { proxyAiGatewayToSiteTarget } from "@server/routers/aiGateway/targetRouting"; import { SESSION_COOKIE_NAME, @@ -150,7 +154,12 @@ type ResolvedTarget = { type ProviderSelection = | { ok: true; provider: AiProvider } - | { ok: false; status: number; message: string }; + | { + ok: false; + status: number; + kind: AiCapabilityErrorKind; + message: string; + }; export type RequestUser = { userId: string; @@ -241,10 +250,6 @@ async function resolveRequestUser( const virtualApiKeyId = getRequestHeader(req, remoteHeaders.virtual_api_key_id) || null; const userId = getRequestHeader(req, remoteHeaders.user_id); - logger.debug("+++++++AI gateway request identity from trust header", { - virtualApiKeyId, - userId - }); if (userId) { const username = getRequestHeader(req, remoteHeaders.user) || userId; @@ -499,6 +504,7 @@ async function selectProvider( return { ok: false, status: HttpCode.FORBIDDEN, + kind: "invalid_request", message: "A model must be specified for this resource" }; } @@ -511,6 +517,7 @@ async function selectProvider( return { ok: false, status: HttpCode.FORBIDDEN, + kind: "permission", message: `Model "${requestedModel}" is not permitted on this resource` }; } @@ -571,6 +578,7 @@ async function selectProvider( return { ok: false, status: HttpCode.FORBIDDEN, + kind: "permission", message: `Model "${requestedModel}" is not permitted on this resource` }; } @@ -607,6 +615,7 @@ async function selectProvider( return { ok: false, status: HttpCode.FORBIDDEN, + kind: "permission", message: `Model "${requestedModel}" is ambiguous across multiple AI providers on this resource. Ask your administrator to configure a more specific allow pattern for this model.` }; } @@ -745,21 +754,29 @@ export async function handleAiGatewayProxy( if (!host) { return res .status(HttpCode.BAD_REQUEST) - .json({ error: { message: "Missing Host header" } }); + .json( + buildAiCapabilityErrorBody( + capability, + "invalid_request", + "Missing Host header", + HttpCode.BAD_REQUEST + ) + ); } - logger.info(`AI gateway ${capability} request for host: ${host}`); - logger.debug("AI gateway request headers", { - headers: req.headers - }); - + logger.debug(`AI gateway ${capability} request for host: ${host}`); const target = await resolveTarget(host); if (!target) { - return res.status(HttpCode.NOT_FOUND).json({ - error: { - message: "No inference resource found for this host" - } - }); + return res + .status(HttpCode.NOT_FOUND) + .json( + buildAiCapabilityErrorBody( + capability, + "not_found", + "No inference resource found for this host", + HttpCode.NOT_FOUND + ) + ); } const { @@ -777,20 +794,29 @@ export async function handleAiGatewayProxy( resourceId != null && !isAiGatewayTrustHeaderValid(req.headers as Record) ) { - return res.status(HttpCode.UNAUTHORIZED).json({ - error: { - message: - "Request must be authenticated via the inference resource" - } - }); + return res + .status(HttpCode.UNAUTHORIZED) + .json( + buildAiCapabilityErrorBody( + capability, + "authentication", + "Request must be authenticated via the inference resource", + HttpCode.UNAUTHORIZED + ) + ); } if (attachments.length === 0) { - return res.status(HttpCode.FORBIDDEN).json({ - error: { - message: "No AI providers configured for this resource" - } - }); + return res + .status(HttpCode.FORBIDDEN) + .json( + buildAiCapabilityErrorBody( + capability, + "permission", + "No AI providers configured for this resource", + HttpCode.FORBIDDEN + ) + ); } const capableAttachments = attachments.filter((a) => @@ -798,11 +824,16 @@ export async function handleAiGatewayProxy( ); if (capableAttachments.length === 0) { - return res.status(HttpCode.FORBIDDEN).json({ - error: { - message: `No AI provider on this resource supports ${capability}` - } - }); + return res + .status(HttpCode.FORBIDDEN) + .json( + buildAiCapabilityErrorBody( + capability, + "permission", + `No AI provider on this resource supports ${capability}`, + HttpCode.FORBIDDEN + ) + ); } const requestedModel = def.extractModel(req); @@ -824,9 +855,16 @@ export async function handleAiGatewayProxy( } if (!selection.ok) { - return res.status(selection.status).json({ - error: { message: selection.message } - }); + return res + .status(selection.status) + .json( + buildAiCapabilityErrorBody( + capability, + selection.kind, + selection.message, + selection.status + ) + ); } const { provider } = selection; @@ -855,11 +893,16 @@ export async function handleAiGatewayProxy( siteResourceId, userId: requestUser?.userId ?? null }); - return res.status(HttpCode.TOO_MANY_REQUESTS).json({ - error: { - message: "AI usage budget exceeded for this request" - } - }); + return res + .status(HttpCode.TOO_MANY_REQUESTS) + .json( + buildAiCapabilityErrorBody( + capability, + "rate_limit", + "AI usage budget exceeded for this request", + HttpCode.TOO_MANY_REQUESTS + ) + ); } } @@ -885,21 +928,31 @@ export async function handleAiGatewayProxy( const authType = provider.authType as AiProviderAuthType; if (!upstreamUrl) { - return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ - error: { - message: "AI provider has no upstream URL configured" - } - }); + return res + .status(HttpCode.INTERNAL_SERVER_ERROR) + .json( + buildAiCapabilityErrorBody( + capability, + "internal", + "AI provider has no upstream URL configured", + HttpCode.INTERNAL_SERVER_ERROR + ) + ); } let apiKey: string | null = null; if (authTypeRequiresApiKey(authType)) { if (!provider.apiKey) { - return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ - error: { - message: "AI provider has no API key configured" - } - }); + return res + .status(HttpCode.INTERNAL_SERVER_ERROR) + .json( + buildAiCapabilityErrorBody( + capability, + "internal", + "AI provider has no API key configured", + HttpCode.INTERNAL_SERVER_ERROR + ) + ); } const secret = config.getRawConfig().server.secret!; apiKey = decrypt(provider.apiKey, secret); @@ -1034,8 +1087,15 @@ export async function handleAiGatewayProxy( return; } catch (error) { logger.error(error); - return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ - error: { message: "Failed to proxy inference request" } - }); + return res + .status(HttpCode.INTERNAL_SERVER_ERROR) + .json( + buildAiCapabilityErrorBody( + capability, + "internal", + "Failed to proxy inference request", + HttpCode.INTERNAL_SERVER_ERROR + ) + ); } } diff --git a/server/routers/aiGateway/targetRouting.ts b/server/routers/aiGateway/targetRouting.ts index 8c4f244aa..957890a1f 100644 --- a/server/routers/aiGateway/targetRouting.ts +++ b/server/routers/aiGateway/targetRouting.ts @@ -21,6 +21,7 @@ import { AI_CAPABILITY_DEFS, type AiCapability } from "@server/lib/aiCapabilities"; +import { buildAiCapabilityErrorBody } from "@server/lib/aiGatewayAuthError"; import { needsStreamUsageInjection, withStreamUsageOption @@ -184,11 +185,14 @@ export async function proxyAiGatewayToSiteTarget( ): Promise { const providerTargets = await getProviderTargets(provider.providerId); if (providerTargets.length === 0) { - res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ - error: { - message: "AI provider has no reachable site targets configured" - } - }); + res.status(HttpCode.INTERNAL_SERVER_ERROR).json( + buildAiCapabilityErrorBody( + capability, + "internal", + "AI provider has no reachable site targets configured", + HttpCode.INTERNAL_SERVER_ERROR + ) + ); return; } @@ -207,11 +211,14 @@ export async function proxyAiGatewayToSiteTarget( let apiKey: string | null = null; if (authTypeRequiresApiKey(authType)) { if (!provider.apiKey) { - res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ - error: { - message: "AI provider has no API key configured" - } - }); + res.status(HttpCode.INTERNAL_SERVER_ERROR).json( + buildAiCapabilityErrorBody( + capability, + "internal", + "AI provider has no API key configured", + HttpCode.INTERNAL_SERVER_ERROR + ) + ); return; } const secret = config.getRawConfig().server.secret!; @@ -286,9 +293,14 @@ export async function proxyAiGatewayToSiteTarget( ? (fetchError as Error & { cause?: unknown }).cause : undefined }); - res.status(HttpCode.BAD_GATEWAY).json({ - error: { message: "Failed to reach AI provider target" } - }); + res.status(HttpCode.BAD_GATEWAY).json( + buildAiCapabilityErrorBody( + capability, + "internal", + "Failed to reach AI provider target", + HttpCode.BAD_GATEWAY + ) + ); return; }