import { Request, Response } from "express"; import { and, eq } from "drizzle-orm"; import { AiProvider, db, exitNodes, sites, targets } from "@server/db"; import config from "@server/lib/config"; import { decrypt } from "@server/lib/crypto"; import { localCache } from "@server/lib/cache"; import { AiProviderAuthType, applyAiProviderAuthHeaders, applyAiProviderCustomHeaders, authTypeRequiresApiKey } from "@server/lib/aiProviderDefaults"; import logger from "@server/logger"; import HttpCode from "@server/types/HttpCode"; import { applyRequestUserHeaders, type RequestUser } from "@server/routers/aiGateway/pipeline"; // Short TTL: long enough to spare the DB on a burst of requests, short // enough that target/site changes (added, removed, exit node moved) show up // almost immediately without needing explicit cache invalidation. const PROVIDER_TARGETS_TTL_SEC = 7; // Header gerbil reads to know which scheme://host:port (reachable over the // WireGuard network) to rewrite an incoming /router/* request to. Must // match gerbil's `pangolinDestHeader` constant. const PANGOLIN_DEST_HEADER = "p-dest-header"; // Header gerbil reads for the Host header value to send to the destination, // when it should differ from PANGOLIN_DEST_HEADER (the target's configured // ip rather than the WireGuard routing address). Must match gerbil's // `pangolinHostHeader` constant. const PANGOLIN_HOST_HEADER = "p-dest-host-header"; const SKIP_HEADERS = new Set([ "p-host", "host", "connection", "keep-alive", "proxy-authenticate", "proxy-authorization", "te", "trailers", "transfer-encoding", "upgrade", "content-length", "accept-encoding" ]); type ResolvedProviderTarget = { targetId: number; // "://:", passed to // gerbil as the destination to proxy the request to over the WireGuard // tunnel. destination: string; // The target's configured ip, passed to gerbil as the Host header to // send to the destination (which may differ from the WireGuard routing // address above, e.g. for vhost-based targets). hostHeader: string; // The target's site's exit node HTTP API base URL (gerbil's /router/*). gerbilBaseUrl: string; }; async function fetchProviderTargets( providerId: number ): Promise { const rows = await db .select({ targetId: targets.targetId, ip: targets.ip, internalPort: targets.internalPort, port: targets.port, method: targets.method, exitNodeSubnet: sites.exitNodeSubnet, reachableAt: exitNodes.reachableAt }) .from(targets) .innerJoin(sites, eq(targets.siteId, sites.siteId)) .innerJoin(exitNodes, eq(sites.exitNodeId, exitNodes.exitNodeId)) .where( and(eq(targets.providerId, providerId), eq(targets.enabled, true)) ); const resolved: ResolvedProviderTarget[] = []; for (const row of rows) { // Sites not yet connected to an exit node (no subnet assigned) or // whose exit node has no known HTTP address can't be routed to. if (!row.exitNodeSubnet || !row.reachableAt) { continue; } const host = row.exitNodeSubnet.split("/")[0]; const port = row.internalPort ?? row.port; const scheme = row.method?.toLowerCase() ?? "https"; resolved.push({ targetId: row.targetId, destination: `${scheme}://${host}:${port}`, hostHeader: row.ip, gerbilBaseUrl: row.reachableAt }); } return resolved; } async function getProviderTargets( providerId: number ): Promise { const cacheKey = `aiGateway:providerTargets:${providerId}`; const cached = localCache.get(cacheKey); if (cached !== undefined) { return cached; } const resolved = await fetchProviderTargets(providerId); localCache.set(cacheKey, resolved, PROVIDER_TARGETS_TTL_SEC); return resolved; } // Round-robin cursor per provider. Process-local and unpersisted - fine // since it only needs to spread load across targets, not guarantee a // perfectly even distribution across restarts or multiple server instances. const roundRobinCursors = new Map(); function pickTarget( providerId: number, providerTargets: ResolvedProviderTarget[] ): ResolvedProviderTarget { const cursor = roundRobinCursors.get(providerId) ?? 0; roundRobinCursors.set(providerId, cursor + 1); return providerTargets[cursor % providerTargets.length]; } function pathFromRequest(req: Request): string { // Query string is preserved - some providers use it to select the // streaming response format (e.g. Gemini's `?alt=sse`), and gerbil's // /router/* forwards it through untouched. const raw = req.originalUrl || req.url || req.path; return raw.startsWith("/") ? raw : `/${raw}`; } /** * Proxies an AI gateway request to one of a "custom" / "target" routing-mode * provider's site targets, via that site's gerbil sidecar. Gerbil's * /router/* endpoint forwards the request (untouched body, same path minus * the /router prefix, and all headers besides PANGOLIN_DEST_HEADER and * PANGOLIN_HOST_HEADER) over the WireGuard tunnel to the destination named * in PANGOLIN_DEST_HEADER, sending PANGOLIN_HOST_HEADER as the Host header. * Always writes a response to `res`, including on failure. */ export async function proxyAiGatewayToSiteTarget( req: Request, res: Response, provider: AiProvider, requestUser: RequestUser | null ): 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" } }); return; } const target = pickTarget(provider.providerId, providerTargets); const gerbilUrl = `${target.gerbilBaseUrl.replace(/\/+$/, "")}/router${pathFromRequest(req)}`; const headers: Record = {}; for (const [key, value] of Object.entries(req.headers)) { if (SKIP_HEADERS.has(key.toLowerCase()) || value === undefined) { continue; } headers[key] = Array.isArray(value) ? value.join(", ") : value; } const authType = provider.authType as AiProviderAuthType; 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" } }); return; } const secret = config.getRawConfig().server.secret!; apiKey = decrypt(provider.apiKey, secret); } applyAiProviderCustomHeaders( headers, provider.headers, config.getRawConfig().server.secret! ); applyAiProviderAuthHeaders(headers, authType, apiKey); applyRequestUserHeaders(headers, requestUser); headers[PANGOLIN_DEST_HEADER] = target.destination; headers[PANGOLIN_HOST_HEADER] = target.hostHeader; const body = JSON.stringify(req.body); logger.debug("AI gateway target-routed request", { providerId: provider.providerId, targetId: target.targetId, destination: target.destination, hostHeader: target.hostHeader, url: gerbilUrl, headers, body: req.body }); // Cancel the request to gerbil (which cascades to gerbil cancelling its // proxied request to the actual site target, since gerbil's reverse // proxy derives the outbound request's context from the inbound one) if // the client goes away before we're done. const abortController = new AbortController(); const onClientClose = () => { if (!res.writableEnded) { abortController.abort(); } }; res.on("close", onClientClose); let upstreamRes: globalThis.Response; try { upstreamRes = await fetch(gerbilUrl, { method: "POST", headers, body, signal: abortController.signal }); } catch (fetchError) { res.off("close", onClientClose); if (abortController.signal.aborted) { // Client already disconnected; nothing left to respond to. return; } logger.error({ message: "AI gateway target proxy request failed", url: gerbilUrl, targetId: target.targetId, error: fetchError, cause: fetchError instanceof Error ? (fetchError as Error & { cause?: unknown }).cause : undefined }); res.status(HttpCode.BAD_GATEWAY).json({ error: { message: "Failed to reach AI provider target" } }); return; } 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"); res.status(upstreamRes.status); res.setHeader("Content-Type", contentType || "application/json"); if (isStream && upstreamRes.body) { res.flushHeaders(); const reader = upstreamRes.body.getReader(); try { while (!abortController.signal.aborted) { const { done, value } = await reader.read(); if (done) break; res.write(value); } } finally { await reader.cancel().catch(() => {}); res.off("close", onClientClose); } if (!res.writableEnded) { res.end(); } return; } res.off("close", onClientClose); const text = await upstreamRes.text(); res.send(text); }