diff --git a/packages/cli/src/provider/provider.ts b/packages/cli/src/provider/provider.ts index 866cb59..b4cc47a 100644 --- a/packages/cli/src/provider/provider.ts +++ b/packages/cli/src/provider/provider.ts @@ -38,17 +38,50 @@ export namespace Provider { return isGpt5OrLater(modelID) && !modelID.startsWith("gpt-5-mini") } - function googleVertexVars(options: Record) { + // A DNS label can have 63 characters; regional hosts append "-aiplatform". + const MAX_VERTEX_LOCATION_LENGTH = 63 - "-aiplatform".length + + function googleVertexLocation(options: Record) { + const raw = options["location"] ?? Env.get("GOOGLE_CLOUD_LOCATION") ?? Env.get("VERTEX_LOCATION") ?? "us-central1" + const location = typeof raw === "string" ? raw.trim().toLowerCase() : "" + if (location.length > MAX_VERTEX_LOCATION_LENGTH || !/^(?:global|us|eu|[a-z]+(?:-[a-z]+)+[0-9]+)$/.test(location)) + throw vertexError("Invalid Google Vertex location. Use global, us, eu, or a region such as us-central1.") + return location + } + + function vertexError(reason: string) { + const error = new VertexConfigError({ reason }) + error.message = reason + return error + } + + function googleVertexProject(options: Record) { const project = - options["project"] ?? Env.get("GOOGLE_CLOUD_PROJECT") ?? Env.get("GCP_PROJECT") ?? Env.get("GCLOUD_PROJECT") - const location = - options["location"] ?? Env.get("GOOGLE_CLOUD_LOCATION") ?? Env.get("VERTEX_LOCATION") ?? "us-central1" - const endpoint = location === "global" ? "aiplatform.googleapis.com" : `${location}-aiplatform.googleapis.com` + options["project"] ?? + Env.get("GOOGLE_CLOUD_PROJECT") ?? + Env.get("GCP_PROJECT") ?? + Env.get("GCLOUD_PROJECT") ?? + Env.get("GOOGLE_VERTEX_PROJECT") + if (project === undefined) return undefined + if (typeof project !== "string" || !/^(?:[a-z][a-z0-9-]{4,28}[a-z0-9]|[0-9]{4,30})$/.test(project)) + throw vertexError("Invalid Google Vertex project. Use a project ID or project number.") + return project + } + + function googleVertexEndpoint(location: string) { + if (location === "global") return "aiplatform.googleapis.com" + if (location === "eu" || location === "us") return `aiplatform.${location}.rep.googleapis.com` + return `${location}-aiplatform.googleapis.com` + } + + function googleVertexVars(options: Record) { + const project = googleVertexProject(options) + const location = googleVertexLocation(options) return { GOOGLE_VERTEX_PROJECT: project, GOOGLE_VERTEX_LOCATION: location, - GOOGLE_VERTEX_ENDPOINT: endpoint, + GOOGLE_VERTEX_ENDPOINT: googleVertexEndpoint(location), } } @@ -57,7 +90,7 @@ export namespace Provider { if (typeof raw !== "string") return raw const vars = model.providerID === "google-vertex" ? googleVertexVars(options) : undefined return raw.replace(/\$\{([^}]+)\}/g, (match, key) => { - const val = Env.get(String(key)) ?? vars?.[String(key) as keyof typeof vars] + const val = vars?.[String(key) as keyof typeof vars] ?? Env.get(String(key)) return val ?? match }) } @@ -359,27 +392,24 @@ export namespace Provider { } }, "google-vertex": async (provider) => { - const project = - provider.options?.project ?? - Env.get("GOOGLE_CLOUD_PROJECT") ?? - Env.get("GCP_PROJECT") ?? - Env.get("GCLOUD_PROJECT") - - const location = - provider.options?.location ?? Env.get("GOOGLE_CLOUD_LOCATION") ?? Env.get("VERTEX_LOCATION") ?? "us-central1" + const project = googleVertexProject(provider.options ?? {}) const autoload = Boolean(project) if (!autoload) return { autoload: false } + const location = googleVertexLocation(provider.options ?? {}) + let auth: Promise> | undefined return { autoload: true, options: { project, location, fetch: async (input: RequestInfo | URL, init?: RequestInit) => { - const { GoogleAuth } = await import("google-auth-library") - const auth = new GoogleAuth() - const client = await auth.getApplicationDefault() - const token = await client.credential.getAccessToken() + // Import and resolve ADC only when a Vertex request is made; share the client across requests. + auth ??= import("google-auth-library").then( + ({ GoogleAuth }) => new GoogleAuth({ scopes: ["https://www.googleapis.com/auth/cloud-platform"] }), + ) + const client = await (await auth).getClient() + const token = await client.getAccessToken() const headers = new Headers(init?.headers) headers.set("Authorization", `Bearer ${token.token}`) @@ -393,16 +423,20 @@ export namespace Provider { }, } }, - "google-vertex-anthropic": async () => { - const project = Env.get("GOOGLE_CLOUD_PROJECT") ?? Env.get("GCP_PROJECT") ?? Env.get("GCLOUD_PROJECT") - const location = Env.get("GOOGLE_CLOUD_LOCATION") ?? Env.get("VERTEX_LOCATION") ?? "global" + "google-vertex-anthropic": async (provider) => { + const project = googleVertexProject(provider.options ?? {}) const autoload = Boolean(project) if (!autoload) return { autoload: false } + const location = googleVertexLocation({ + location: + provider.options?.location ?? Env.get("GOOGLE_CLOUD_LOCATION") ?? Env.get("VERTEX_LOCATION") ?? "global", + }) return { autoload: true, options: { project, location, + baseURL: `https://${googleVertexEndpoint(location)}/v1/projects/${project}/locations/${location}/publishers/anthropic/models`, }, async getModel(sdk: any, modelID) { const id = String(modelID).trim() @@ -872,6 +906,7 @@ export namespace Provider { } const providers: { [providerID: string]: Info } = {} + const errors: Record = {} const languages = new Map() const modelLoaders: { [providerID: string]: CustomModelLoader @@ -1069,7 +1104,12 @@ export namespace Provider { log.error("Provider does not exist in model list " + providerID) continue } - const result = await fn(data) + const result = await fn(data).catch((error) => { + if (!VertexConfigError.isInstance(error)) throw error + errors[providerID] = error + log.error("invalid provider configuration", { providerID, error: error.message }) + return undefined + }) if (result && (result.autoload || providers[providerID])) { if (result.getModel) modelLoaders[providerID] = result.getModel const opts = result.options ?? {} @@ -1093,6 +1133,28 @@ export namespace Provider { continue } + // Config options are merged after custom loaders; normalize the final value + // for native SDKs as well as templated OpenAI-compatible endpoints. + if (providerID === "google-vertex" || providerID === "google-vertex-anthropic") { + try { + const project = googleVertexProject(provider.options) + const location = googleVertexLocation({ + location: provider.options.location ?? (providerID === "google-vertex-anthropic" ? "global" : undefined), + }) + if (project) provider.options.project = project + provider.options.location = location + if (providerID === "google-vertex-anthropic" && project) + provider.options.baseURL = `https://${googleVertexEndpoint(location)}/v1/projects/${project}/locations/${location}/publishers/anthropic/models` + delete errors[providerID] + } catch (error) { + if (!VertexConfigError.isInstance(error)) throw error + errors[providerID] = error + log.error("invalid provider configuration", { providerID, error: error.message }) + delete providers[providerID] + continue + } + } + const configProvider = config.provider?.[providerID] for (const [modelID, model] of Object.entries(provider.models)) { @@ -1131,6 +1193,7 @@ export namespace Provider { return { models: languages, providers, + errors, sdk, modelLoaders, } @@ -1146,6 +1209,7 @@ export namespace Provider { providerID: model.providerID, }) const s = await state() + if (s.errors[model.providerID]) throw s.errors[model.providerID] const provider = s.providers[model.providerID] const options = { ...provider.options } @@ -1247,11 +1311,15 @@ export namespace Provider { } export async function getProvider(providerID: string) { - return state().then((s) => s.providers[providerID]) + return state().then((s) => { + if (s.errors[providerID]) throw s.errors[providerID] + return s.providers[providerID] + }) } export async function getModel(providerID: string, modelID: string) { const s = await state() + if (s.errors[providerID]) throw s.errors[providerID] const provider = s.providers[providerID] if (!provider) { const availableProviders = Object.keys(s.providers) @@ -1272,6 +1340,7 @@ export namespace Provider { export async function getLanguage(model: Model): Promise { const s = await state() + if (s.errors[model.providerID]) throw s.errors[model.providerID] const key = `${model.providerID}/${model.id}` if (s.models.has(key)) return s.models.get(key)! @@ -1457,4 +1526,6 @@ export namespace Provider { providerID: z.string(), }), ) + + export const VertexConfigError = NamedError.create("VertexConfigError", z.object({ reason: z.string() })) } diff --git a/packages/cli/test/provider/google-vertex.test.ts b/packages/cli/test/provider/google-vertex.test.ts new file mode 100644 index 0000000..61f51fb --- /dev/null +++ b/packages/cli/test/provider/google-vertex.test.ts @@ -0,0 +1,300 @@ +import { expect, spyOn, test } from "bun:test" +import { GoogleAuth, OAuth2Client } from "google-auth-library" +import path from "path" + +import { tmpdir } from "../fixture/fixture" +import { Instance } from "../../src/project/instance" +import { Env } from "../../src/env" + +test.each([ + ["global", "aiplatform.googleapis.com"], + ["us", "aiplatform.us.rep.googleapis.com"], + ["eu", "aiplatform.eu.rep.googleapis.com"], + ["europe-west1", "europe-west1-aiplatform.googleapis.com"], + [" EU ", "aiplatform.eu.rep.googleapis.com"], + ["US", "aiplatform.us.rep.googleapis.com"], + [" Global\t", "aiplatform.googleapis.com"], + [" EUROPE-WEST1 ", "europe-west1-aiplatform.googleapis.com"], +])("Google Vertex resolves the %s endpoint", async (location, endpoint) => { + await using tmp = await tmpdir({ + init: async (dir) => { + await Bun.write( + path.join(dir, "aictrl.json"), + JSON.stringify({ + $schema: "https://aictrl.ai/config.json", + provider: { + "google-vertex": { + options: { + project: "test-project", + location, + }, + models: { + "test-model": { + name: "Test Model", + tool_call: true, + provider: { + npm: "@ai-sdk/openai-compatible", + api: "https://${GOOGLE_VERTEX_ENDPOINT}/v1/projects/${GOOGLE_VERTEX_PROJECT}/locations/${GOOGLE_VERTEX_LOCATION}", + }, + }, + }, + }, + }, + }), + ) + }, + }) + + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + const model = await Provider.getModel("google-vertex", "test-model") + expect((await Provider.getProvider("google-vertex")).options.location).toBe(location.trim().toLowerCase()) + const language = (await Provider.getLanguage(model)) as unknown as { + config: { url(input: { path: string }): string } + } + + expect(language.config.url({ path: "/chat/completions" })).toBe( + `https://${endpoint}/v1/projects/test-project/locations/${location.trim().toLowerCase()}/chat/completions`, + ) + }, + }) +}) + +test.each([ + "attacker.com/", + "us@attacker.com", + "us\\attacker.com", + "us?x=1", + "us#fragment", + "us%2fhost", + "", + 123, + "a".repeat(60) + "-west1", +])("Google Vertex rejects unsafe location %s before resolving a client or constructing an SDK", async (location) => { + await using tmp = await tmpdir({ + config: { + provider: { + "google-vertex": { options: { project: "test-project", location } }, + anthropic: { options: { apiKey: "synthetic" } }, + }, + }, + }) + const client = spyOn(GoogleAuth.prototype, "getClient") + try { + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + await expect(Provider.getProvider("google-vertex")).rejects.toThrow("Invalid Google Vertex location") + await expect(Provider.getProvider("google-vertex")).rejects.toBeInstanceOf(Provider.VertexConfigError) + expect(await Provider.list()).toHaveProperty("anthropic") + expect(await Provider.getProvider("anthropic")).toBeDefined() + expect(client).not.toHaveBeenCalled() + }, + }) + } finally { + client.mockRestore() + } +}) + +test.each([ + ["global", "aiplatform.googleapis.com"], + [" US ", "aiplatform.us.rep.googleapis.com"], + ["eu", "aiplatform.eu.rep.googleapis.com"], + ["us-central1", "us-central1-aiplatform.googleapis.com"], +])("Google Vertex Anthropic resolves the %s endpoint", async (location, endpoint) => { + await using tmp = await tmpdir({ + config: { provider: { "google-vertex-anthropic": { options: { project: "test-project", location } } } }, + }) + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + const provider = await Provider.getProvider("google-vertex-anthropic") + expect(provider.options.location).toBe(location.trim().toLowerCase()) + expect(provider.options.baseURL).toBe( + `https://${endpoint}/v1/projects/test-project/locations/${location.trim().toLowerCase()}/publishers/anthropic/models`, + ) + }, + }) +}) + +test.each([ + ["google-vertex", "location", "us-central1-a"], + ["google-vertex-anthropic", "location", "attacker.com/"], + ["google-vertex", "project", "test-project/other"], + ["google-vertex-anthropic", "project", "test-project?x=1"], +])("%s rejects invalid %s without breaking other providers", async (id, key, value) => { + await using tmp = await tmpdir({ + config: { + provider: { + [id]: { options: { project: "test-project", location: "us-central1", [key]: value } }, + anthropic: { options: { apiKey: "synthetic" } }, + }, + }, + }) + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + await expect(Provider.getProvider(id)).rejects.toThrow(`Invalid Google Vertex ${key}`) + expect(await Provider.list()).toHaveProperty("anthropic") + expect(await Provider.getProvider("anthropic")).toBeDefined() + }, + }) +}) + +test("Google Vertex template ignores an unsafe endpoint override", async () => { + await using tmp = await tmpdir({ + config: { + provider: { + "google-vertex": { + options: { project: "test-project", location: "global" }, + models: { + "test-model": { + name: "Test Model", + provider: { + npm: "@ai-sdk/openai-compatible", + api: "https://${GOOGLE_VERTEX_ENDPOINT}/v1/projects/${GOOGLE_VERTEX_PROJECT}/locations/${GOOGLE_VERTEX_LOCATION}", + }, + }, + }, + }, + }, + }, + }) + await Instance.provide({ + directory: tmp.path, + init: async () => Env.set("GOOGLE_VERTEX_ENDPOINT", "attacker.test"), + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + const model = await Provider.getModel("google-vertex", "test-model") + const language = (await Provider.getLanguage(model)) as unknown as { + config: { url(input: { path: string }): string } + } + expect(language.config.url({ path: "/chat/completions" })).toStartWith("https://aiplatform.googleapis.com/") + }, + }) +}) + +test.each([ + ["GOOGLE_CLOUD_LOCATION", " EU ", "eu"], + ["VERTEX_LOCATION", " US ", "us"], + ["GOOGLE_CLOUD_LOCATION", "attacker.com/", undefined], + ["VERTEX_LOCATION", "us@attacker.com", undefined], +] as const)("Google Vertex validates location from %s", async (key, value, expected) => { + await using tmp = await tmpdir({ + config: { provider: { "google-vertex": { options: { project: "test-project" } } } }, + }) + await Instance.provide({ + directory: tmp.path, + init: async () => { + Env.remove("GOOGLE_CLOUD_LOCATION") + Env.remove("VERTEX_LOCATION") + Env.set(key, value) + }, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + if (expected === undefined) { + await expect(Provider.getProvider("google-vertex")).rejects.toThrow("Invalid Google Vertex location") + return + } + expect((await Provider.getProvider("google-vertex")).options.location).toBe(expected) + }, + }) +}) + +test("Google Vertex reuses its auth instance and cached client across requests", async () => { + await using tmp = await tmpdir({ + config: { + provider: { "google-vertex": { options: { project: "test-project", location: "us-central1" } } }, + }, + }) + const credential = new OAuth2Client() + credential.setCredentials({ access_token: "synthetic-cached-token", expiry_date: Date.now() + 3600000 }) + const instances = new Set() + const original = GoogleAuth.prototype.getClient + const client = spyOn(GoogleAuth.prototype, "getClient").mockImplementation(function (this: GoogleAuth) { + instances.add(this) + this.cachedCredential ??= credential + return original.call(this) + }) + const tokens = spyOn(credential, "getAccessToken") + const server = Bun.serve({ + port: 0, + fetch(request) { + expect(request.headers.get("authorization")).toBe("Bearer synthetic-cached-token") + return new Response("ok") + }, + }) + try { + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + const provider = await Provider.getProvider("google-vertex") + expect(client).not.toHaveBeenCalled() + const requests = await Promise.all([provider.options.fetch(server.url), provider.options.fetch(server.url)]) + expect(await Promise.all(requests.map((response: Response) => response.text()))).toEqual(["ok", "ok"]) + expect(instances.size).toBe(1) + expect(client).toHaveBeenCalledTimes(2) + expect(tokens).toHaveBeenCalledTimes(2) + expect(await Promise.all(tokens.mock.results.map((result) => result.value))).toEqual([ + { token: "synthetic-cached-token" }, + { token: "synthetic-cached-token" }, + ]) + }, + }) + } finally { + server.stop(true) + tokens.mockRestore() + client.mockRestore() + } +}) + +test("Google Vertex requests the cloud-platform OAuth scope", async () => { + await using tmp = await tmpdir({ + init: async (dir) => { + await Bun.write( + path.join(dir, "aictrl.json"), + JSON.stringify({ + $schema: "https://aictrl.ai/config.json", + provider: { + "google-vertex": { + options: { + project: "test-project", + location: "us-central1", + }, + }, + }, + }), + ) + }, + }) + + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const { Provider } = await import("../../src/provider/provider") + const provider = await Provider.getProvider("google-vertex") + + const client = spyOn(GoogleAuth.prototype, "getClient").mockImplementation(async function (this: GoogleAuth) { + expect(Reflect.get(this, "scopes")).toEqual(["https://www.googleapis.com/auth/cloud-platform"]) + throw new Error("stop after resolving auth client") + }) + const defaults = spyOn(GoogleAuth.prototype, "getApplicationDefault").mockImplementation(() => { + throw new Error("unexpected application default lookup") + }) + try { + await expect(provider.options.fetch("https://example.test")).rejects.toThrow("stop after resolving auth client") + expect(client).toHaveBeenCalledTimes(1) + expect(defaults).not.toHaveBeenCalled() + } finally { + client.mockRestore() + defaults.mockRestore() + } + }, + }) +})