diff --git a/server/lib/aiCapabilities.ts b/server/lib/aiCapabilities.ts index 2b3691493..89fbd3d75 100644 --- a/server/lib/aiCapabilities.ts +++ b/server/lib/aiCapabilities.ts @@ -27,6 +27,7 @@ export type AiCapabilityDefinition = { req: Request, model: string ) => string; + isStreaming: (req: Request, contentType: string) => boolean; }; function bodyModel(req: Request): string | undefined { @@ -38,9 +39,6 @@ function paramModel(req: Request): string | undefined { return typeof model === "string" && model.length > 0 ? model : undefined; } -/** - * Join a provider base URL with an inbound request path. - */ export function joinUpstreamUrl(baseUrl: string, path: string): string { const base = baseUrl.replace(/\/+$/, ""); let suffix = path.startsWith("/") ? path : `/${path}`; @@ -80,13 +78,38 @@ export function joinUpstreamUrl(baseUrl: string, path: string): string { } function pathFromRequest(req: Request): string { - // Prefer originalUrl (includes mounted path) over req.url when available. - // Query string is preserved - some providers use it to select the - // streaming response format (e.g. Gemini's `?alt=sse`). const raw = req.originalUrl || req.url || req.path; return raw.startsWith("/") ? raw : `/${raw}`; } +function bodyRequestsStream(req: Request): boolean { + return req.body?.stream === true; +} + +function contentTypeIsSse(contentType: string): boolean { + return contentType.includes("text/event-stream"); +} + +function contentTypeIsAmazonEventStream(contentType: string): boolean { + return contentType.includes("application/vnd.amazon.eventstream"); +} + +function pathIncludes(req: Request, fragment: string): boolean { + return pathFromRequest(req).includes(fragment); +} + +function isBodyOrSseStreaming(req: Request, contentType: string): boolean { + return bodyRequestsStream(req) || contentTypeIsSse(contentType); +} + +function isGeminiStyleStreaming(req: Request, contentType: string): boolean { + return ( + pathIncludes(req, "streamGenerateContent") || + pathIncludes(req, "alt=sse") || + contentTypeIsSse(contentType) + ); +} + export const AI_CAPABILITY_DEFS: Record = { openai_chat: { @@ -97,21 +120,24 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: bodyModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: isBodyOrSseStreaming }, openai_responses: { id: "openai_responses", routes: [{ method: "POST", path: "/v1/responses" }], extractModel: bodyModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: isBodyOrSseStreaming }, anthropic_messages: { id: "anthropic_messages", routes: [{ method: "POST", path: "/v1/messages" }], extractModel: bodyModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: isBodyOrSseStreaming }, gemini_generate_content: { id: "gemini_generate_content", @@ -127,7 +153,8 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: paramModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: isGeminiStyleStreaming }, google_generate_content: { id: "google_generate_content", @@ -144,7 +171,8 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: paramModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: isGeminiStyleStreaming }, google_raw_predict: { id: "google_raw_predict", @@ -160,7 +188,11 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: paramModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: (req, contentType) => + pathIncludes(req, "streamRawPredict") || + pathIncludes(req, "alt=sse") || + contentTypeIsSse(contentType) }, bedrock_model_invoke: { id: "bedrock_model_invoke", @@ -173,7 +205,11 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: paramModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: (req, contentType) => + pathIncludes(req, "invoke-with-response-stream") || + contentTypeIsAmazonEventStream(contentType) || + contentTypeIsSse(contentType) }, bedrock_converse: { id: "bedrock_converse", @@ -183,7 +219,11 @@ export const AI_CAPABILITY_DEFS: Record = ], extractModel: paramModel, resolveUpstreamUrl: (base, req) => - joinUpstreamUrl(base, pathFromRequest(req)) + joinUpstreamUrl(base, pathFromRequest(req)), + isStreaming: (req, contentType) => + pathIncludes(req, "converse-stream") || + contentTypeIsAmazonEventStream(contentType) || + contentTypeIsSse(contentType) } }; diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index 724102f07..e84165620 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -62,11 +62,6 @@ import { type AiUsage } from "@server/lib/aiUsageExtraction"; -// Short-lived local caches so a burst of requests from the same IP/user -// doesn't hit the database on every single request. None of this is -// security-critical to cache aggressively (identity is re-derived from the -// session cookie or from a client's exit-node-scoped subnet each time), so -// a small TTL is just an efficiency win, not a trust boundary. const EXIT_NODE_RANGES_CACHE_KEY = "aiGateway:exitNodeRanges"; const EXIT_NODE_RANGES_TTL_SEC = 6000; const CLIENT_BY_IP_TTL_SEC = 30; @@ -638,7 +633,8 @@ export async function handleAiGatewayProxy( req, res, provider, - requestUser + requestUser, + capability ); } @@ -724,10 +720,6 @@ export async function handleAiGatewayProxy( skipTlsVerification: provider.skipTlsVerification }); - // Cancel the upstream request (and, transitively, anything it fans - // out to) if the client goes away before we're done - otherwise a - // client-cancelled streaming chat completion keeps running upstream - // to completion, wasting the connection and any per-token billing. const abortController = new AbortController(); const onClientClose = () => { if (!res.writableEnded) { @@ -764,13 +756,7 @@ export async function handleAiGatewayProxy( } const contentType = upstreamRes.headers.get("content-type") || ""; - const isStream = - req.body?.stream === true || - 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"); + const isStream = def.isStreaming(req, contentType); res.status(upstreamRes.status); res.setHeader("Content-Type", contentType || "application/json"); diff --git a/server/routers/aiGateway/targetRouting.ts b/server/routers/aiGateway/targetRouting.ts index b091bf8f5..a6d08d925 100644 --- a/server/routers/aiGateway/targetRouting.ts +++ b/server/routers/aiGateway/targetRouting.ts @@ -10,6 +10,10 @@ import { applyAiProviderCustomHeaders, authTypeRequiresApiKey } from "@server/lib/aiProviderDefaults"; +import { + AI_CAPABILITY_DEFS, + type AiCapability +} from "@server/lib/aiCapabilities"; import logger from "@server/logger"; import HttpCode from "@server/types/HttpCode"; import { @@ -152,14 +156,14 @@ export async function proxyAiGatewayToSiteTarget( req: Request, res: Response, provider: AiProvider, - requestUser: RequestUser | null + requestUser: RequestUser | null, + capability: AiCapability ): 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" + message: "AI provider has no reachable site targets configured" } }); return; @@ -256,13 +260,10 @@ export async function proxyAiGatewayToSiteTarget( } const contentType = upstreamRes.headers.get("content-type") || ""; - const isStream = - req.body?.stream === true || - contentType.includes("text/event-stream") || - pathFromRequest(req).includes("streamGenerateContent") || - pathFromRequest(req).includes("streamRawPredict") || - pathFromRequest(req).includes("converse-stream") || - pathFromRequest(req).includes("invoke-with-response-stream"); + const isStream = AI_CAPABILITY_DEFS[capability].isStreaming( + req, + contentType + ); res.status(upstreamRes.status); res.setHeader("Content-Type", contentType || "application/json");