Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 95 additions & 24 deletions packages/cli/src/provider/provider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -38,17 +38,50 @@ export namespace Provider {
return isGpt5OrLater(modelID) && !modelID.startsWith("gpt-5-mini")
}

function googleVertexVars(options: Record<string, any>) {
// A DNS label can have 63 characters; regional hosts append "-aiplatform".
const MAX_VERTEX_LOCATION_LENGTH = 63 - "-aiplatform".length

function googleVertexLocation(options: Record<string, any>) {
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<string, any>) {
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) {
Comment thread
byapparov marked this conversation as resolved.
Comment thread
byapparov marked this conversation as resolved.
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<string, any>) {
const project = googleVertexProject(options)
const location = googleVertexLocation(options)

return {
GOOGLE_VERTEX_PROJECT: project,
Comment thread
byapparov marked this conversation as resolved.
GOOGLE_VERTEX_LOCATION: location,
GOOGLE_VERTEX_ENDPOINT: endpoint,
GOOGLE_VERTEX_ENDPOINT: googleVertexEndpoint(location),
}
}

Expand All @@ -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
})
}
Expand Down Expand Up @@ -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<InstanceType<typeof import("google-auth-library").GoogleAuth>> | 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}`)
Expand All @@ -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()
Expand Down Expand Up @@ -872,6 +906,7 @@ export namespace Provider {
}

const providers: { [providerID: string]: Info } = {}
const errors: Record<string, Error> = {}
const languages = new Map<string, LanguageModelV2>()
const modelLoaders: {
[providerID: string]: CustomModelLoader
Expand Down Expand Up @@ -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 ?? {}
Expand All @@ -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)) {
Expand Down Expand Up @@ -1131,6 +1193,7 @@ export namespace Provider {
return {
models: languages,
providers,
errors,
sdk,
modelLoaders,
}
Expand All @@ -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 }

Expand Down Expand Up @@ -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)
Expand All @@ -1272,6 +1340,7 @@ export namespace Provider {

export async function getLanguage(model: Model): Promise<LanguageModelV2> {
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)!

Expand Down Expand Up @@ -1457,4 +1526,6 @@ export namespace Provider {
providerID: z.string(),
}),
)

export const VertexConfigError = NamedError.create("VertexConfigError", z.object({ reason: z.string() }))
}
Loading
Loading