From 5aa3b1a1f5e877c70f22b2d9021631f5b125d709 Mon Sep 17 00:00:00 2001 From: QA Runner Date: Fri, 17 Jul 2026 09:50:17 -0700 Subject: [PATCH] feat: Vertex AI provider Route Gemini model IDs through Google Cloud Vertex AI instead of the AI Studio endpoint. setupVertexAI() re-registers the "gemini" provider id in the provider registry with a Vertex AI-backed adapter that authenticates via Application Default Credentials (no API key). Model IDs are unchanged, so routing is transparent to all call sites. Nothing changes unless setupVertexAI() is called at startup. Mechanical port of apps/api/src/lib/llm/providers/vertexAI.ts and its test from amal66/mike@main (commit b3166dd), verbatim; stacks on the provider-registry branch it registers into. Co-Authored-By: Claude Fable 5 --- .../src/lib/llm/__tests__/vertexAI.test.ts | 82 ++++++ backend/src/lib/llm/providers/vertexAI.ts | 264 ++++++++++++++++++ 2 files changed, 346 insertions(+) create mode 100644 backend/src/lib/llm/__tests__/vertexAI.test.ts create mode 100644 backend/src/lib/llm/providers/vertexAI.ts diff --git a/backend/src/lib/llm/__tests__/vertexAI.test.ts b/backend/src/lib/llm/__tests__/vertexAI.test.ts new file mode 100644 index 000000000..38a799534 --- /dev/null +++ b/backend/src/lib/llm/__tests__/vertexAI.test.ts @@ -0,0 +1,82 @@ +import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"; +import { _resetRegistryForTesting, getRegisteredProvider } from "../registry"; + +// Reset env and registry around each test. +beforeEach(() => { + _resetRegistryForTesting(); + process.env.VERTEX_AI_PROJECT = "test-project"; + process.env.VERTEX_AI_LOCATION = "us-central1"; +}); + +afterEach(() => { + _resetRegistryForTesting(); + delete process.env.VERTEX_AI_PROJECT; + delete process.env.VERTEX_AI_LOCATION; +}); + +describe("setupVertexAI", () => { + it("registers under the 'gemini' provider id", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + + const provider = getRegisteredProvider("gemini"); + expect(provider).toBeDefined(); + expect(provider!.id).toBe("gemini"); + }); + + it("matchesModel returns true for all built-in Gemini model IDs", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + + const provider = getRegisteredProvider("gemini")!; + expect(provider.matchesModel("gemini-3.1-pro-preview")).toBe(true); + expect(provider.matchesModel("gemini-3-flash-preview")).toBe(true); + expect(provider.matchesModel("gemini-3.1-flash-lite-preview")).toBe(true); + }); + + it("matchesModel returns true for any gemini- prefixed model (future models)", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + + const provider = getRegisteredProvider("gemini")!; + expect(provider.matchesModel("gemini-future-ultra")).toBe(true); + }); + + it("matchesModel returns false for non-Gemini models", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + + const provider = getRegisteredProvider("gemini")!; + expect(provider.matchesModel("claude-sonnet-4-6")).toBe(false); + expect(provider.matchesModel("gpt-5.5")).toBe(false); + }); + + it("registers extra models passed via options", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI({ extraModels: ["gemini-experimental-xyz"] }); + + const provider = getRegisteredProvider("gemini")!; + expect(provider.matchesModel("gemini-experimental-xyz")).toBe(true); + }); + + it("models lists include main/mid/low tiers", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + + const provider = getRegisteredProvider("gemini")!; + expect(provider.models.main.length).toBeGreaterThan(0); + expect(provider.models.mid.length).toBeGreaterThan(0); + expect(provider.models.low.length).toBeGreaterThan(0); + }); + + it("re-registering replaces the previous adapter", async () => { + const { setupVertexAI } = await import("../providers/vertexAI"); + setupVertexAI(); + const first = getRegisteredProvider("gemini"); + setupVertexAI({ extraModels: ["gemini-extra"] }); + const second = getRegisteredProvider("gemini"); + // Both are registered under "gemini"; second call replaces first. + expect(second).not.toBe(first); + expect(second!.matchesModel("gemini-extra")).toBe(true); + }); +}); diff --git a/backend/src/lib/llm/providers/vertexAI.ts b/backend/src/lib/llm/providers/vertexAI.ts new file mode 100644 index 000000000..e3daae9ef --- /dev/null +++ b/backend/src/lib/llm/providers/vertexAI.ts @@ -0,0 +1,264 @@ +/** + * Vertex AI provider — routes Gemini model IDs through Google Cloud Vertex AI + * instead of the default Gemini AI Studio endpoint. + * + * Why use this instead of the built-in Gemini provider? + * - Enterprise billing through a Google Cloud project (not AI Studio quota) + * - Data residency / VPC Service Controls compliance + * - Cloud IAM-gated access (no API key distributed to servers) + * - Workload Identity support on GKE / Cloud Run (zero secrets) + * + * Auth uses Application Default Credentials (ADC) — the standard GCP chain: + * 1. GOOGLE_APPLICATION_CREDENTIALS (path to service account JSON key) + * 2. Workload Identity (GKE, Cloud Run, Compute Engine, Cloud Functions) + * 3. gcloud CLI: `gcloud auth application-default login` (local dev) + * + * Required env vars: + * VERTEX_AI_PROJECT — GCP project ID (e.g. "my-project-123") + * VERTEX_AI_LOCATION — region (default: "us-central1") + * + * Call setupVertexAI() once at application startup to replace the default + * Gemini provider with a Vertex AI-backed one. All Gemini model IDs remain + * unchanged — the routing is transparent to the rest of the application. + * + * import { setupVertexAI } from "lib/llm/providers/vertexAI"; + * setupVertexAI(); + * + * The same Gemini model IDs are supported: + * gemini-3.1-pro-preview, gemini-3-flash-preview, gemini-3.1-flash-lite-preview + */ + +import type { StreamChatParams, StreamChatResult, CompleteTextParams } from "../types"; +import { toGeminiTools } from "../tools"; +import { registerProvider } from "../registry"; +import { + GEMINI_MAIN_MODELS, + GEMINI_MID_MODELS, + GEMINI_LOW_MODELS, +} from "../models"; + +// --------------------------------------------------------------------------- +// Internal types (mirrors gemini.ts — Vertex AI uses the same content format) +// --------------------------------------------------------------------------- + +type GoogleGenAIConstructor = typeof import("@google/genai").GoogleGenAI; +type GoogleGenAIClient = InstanceType; + +const importEsm = new Function("specifier", "return import(specifier)") as ( + specifier: string, +) => Promise<{ GoogleGenAI: GoogleGenAIConstructor }>; + +type GeminiPart = { + text?: string; + thought?: boolean; + functionCall?: { + id?: string; + name: string; + args?: Record; + }; + functionResponse?: { + id?: string; + name: string; + response: Record; + }; + thoughtSignature?: string; +}; + +type GeminiContent = { + role: "user" | "model"; + parts: GeminiPart[]; +}; + +// --------------------------------------------------------------------------- +// Vertex AI client (ADC — no API key) +// --------------------------------------------------------------------------- + +function requireConfig(): { project: string; location: string } { + const project = process.env.VERTEX_AI_PROJECT?.trim(); + if (!project) { + throw new Error( + "VERTEX_AI_PROJECT must be set to use the Vertex AI Gemini provider.", + ); + } + const location = process.env.VERTEX_AI_LOCATION?.trim() || "us-central1"; + return { project, location }; +} + +async function vertexClient(): Promise { + const { project, location } = requireConfig(); + const { GoogleGenAI } = await importEsm("@google/genai"); + // vertexai: true tells the SDK to use Vertex AI endpoints and ADC auth + // instead of the apiKey-based AI Studio endpoint. + return new GoogleGenAI({ vertexai: true, project, location } as never); +} + +function throwIfAborted(signal?: AbortSignal): void { + if (!signal?.aborted) return; + // Match the abort shape isAbortError() detects across providers + // (name "AbortError" / message "Stream aborted."). + const err = new Error("Stream aborted."); + err.name = "AbortError"; + throw err; +} + +// --------------------------------------------------------------------------- +// Stream +// --------------------------------------------------------------------------- + +async function streamVertexGemini(params: StreamChatParams): Promise { + const { + model, + systemPrompt, + tools = [], + callbacks = {}, + runTools, + enableThinking, + } = params; + const maxIter = params.maxIterations ?? 10; + const ai = await vertexClient(); + const functionDeclarations = toGeminiTools(tools); + + const contents: GeminiContent[] = params.messages.map((m) => ({ + role: m.role === "assistant" ? "model" : "user", + parts: [{ text: m.content }], + })); + let fullText = ""; + + for (let iter = 0; iter < maxIter; iter++) { + throwIfAborted(params.abortSignal); + const stream = await ai.models.generateContentStream({ + model, + contents: contents as never, + config: { + systemInstruction: systemPrompt, + tools: functionDeclarations.length + ? [{ functionDeclarations } as never] + : undefined, + thinkingConfig: enableThinking + ? { includeThoughts: true } + : { thinkingBudget: 0 }, + }, + }); + + const textParts: string[] = []; + const callParts: GeminiPart[] = []; + const toolCalls: import("../types").NormalizedToolCall[] = []; + let sawThinking = false; + + for await (const chunk of stream) { + throwIfAborted(params.abortSignal); + const parts = + (chunk as { candidates?: { content?: { parts?: GeminiPart[] } }[] }) + .candidates?.[0]?.content?.parts ?? []; + + for (const part of parts) { + if (part.text) { + if (part.thought) { + sawThinking = true; + callbacks.onReasoningDelta?.(part.text); + } else { + textParts.push(part.text); + callbacks.onContentDelta?.(part.text); + } + } + if (part.functionCall) { + callParts.push(part); + const call: import("../types").NormalizedToolCall = { + id: part.functionCall.id ?? `${part.functionCall.name}-${toolCalls.length}`, + name: part.functionCall.name, + input: part.functionCall.args ?? {}, + }; + callbacks.onToolCallStart?.(call); + toolCalls.push(call); + } + } + } + + if (sawThinking) callbacks.onReasoningBlockEnd?.(); + fullText += textParts.join(""); + + if (!toolCalls.length || !runTools) break; + + const results = await runTools(toolCalls); + + const modelParts: GeminiPart[] = []; + if (textParts.length) modelParts.push({ text: textParts.join("") }); + for (const cp of callParts) modelParts.push(cp); + contents.push({ role: "model", parts: modelParts }); + + contents.push({ + role: "user", + parts: results.map((r) => { + const match = toolCalls.find((c) => c.id === r.tool_use_id); + return { + functionResponse: { + ...(r.tool_use_id && !r.tool_use_id.startsWith(match?.name ?? "") + ? { id: r.tool_use_id } + : {}), + name: match?.name ?? "tool", + response: { output: r.content }, + }, + }; + }), + }); + } + + return { fullText }; +} + +// --------------------------------------------------------------------------- +// Complete (non-streaming) +// --------------------------------------------------------------------------- + +async function completeVertexGeminiText(params: CompleteTextParams): Promise { + const ai = await vertexClient(); + const resp = await ai.models.generateContent({ + model: params.model, + contents: [{ role: "user", parts: [{ text: params.user }] }], + ...(params.systemPrompt ? { config: { systemInstruction: params.systemPrompt } } : {}), + }); + return (resp as { text?: string }).text ?? ""; +} + +// --------------------------------------------------------------------------- +// Registration +// --------------------------------------------------------------------------- + +export interface VertexAISetupOptions { + /** + * Additional model IDs to register beyond the built-in Gemini models. + * Useful when Vertex AI grants access to preview models not in the list. + */ + extraModels?: string[]; +} + +/** + * Replaces the built-in Gemini provider with a Vertex AI-backed one. + * + * After calling this, all requests that would have gone to Gemini AI Studio + * are instead routed to your Google Cloud project via ADC. The model IDs + * (gemini-3.1-pro-preview, etc.) remain unchanged. + */ +export function setupVertexAI(options: VertexAISetupOptions = {}): void { + const allModels = [ + ...GEMINI_MAIN_MODELS, + ...GEMINI_MID_MODELS, + ...GEMINI_LOW_MODELS, + ...(options.extraModels ?? []), + ]; + const modelSet = new Set(allModels); + + // Re-registers under the same id "gemini" — replaces the built-in adapter. + // No API key provider registration needed: Vertex AI uses ADC, not a key. + registerProvider({ + id: "gemini", + matchesModel: (m) => modelSet.has(m) || m.startsWith("gemini"), + stream: streamVertexGemini, + complete: completeVertexGeminiText, + models: { + main: [...GEMINI_MAIN_MODELS], + mid: [...GEMINI_MID_MODELS], + low: [...GEMINI_LOW_MODELS], + }, + }); +}