diff --git a/.env.example b/.env.example index 03eef16..ecf6d8f 100644 --- a/.env.example +++ b/.env.example @@ -61,6 +61,16 @@ GOOGLE_GENERATIVE_AI_API_KEY=AIza... # AXIOM_TOKEN= # AXIOM_DATASET= +# ============================================================================= +# Optional: Custom Providers +# ============================================================================= + +# Set these when using a custom provider registered via configure(). +# Reference them in your createProvider() factory function. +# Example for a corporate OpenAI-compatible LLM proxy: +# LLM_PROXY_API_KEY=your-proxy-key +# LLM_PROXY_API_BASE_URL=https://llm-proxy.internal.company.com + # ============================================================================= # Optional: Logging # ============================================================================= diff --git a/README.md b/README.md index 8a19892..37878d9 100644 --- a/README.md +++ b/README.md @@ -293,6 +293,73 @@ configure({ }); ``` +### Custom Providers + +If you use a self-hosted LLM proxy, corporate gateway, or any OpenAI-compatible endpoint, you can register custom providers via `configure()`. This lets you use any [Vercel AI SDK](https://ai-sdk.dev/providers)-compatible provider package. + +Install the SDK package you need (e.g. `@ai-sdk/openai-compatible` for generic OpenAI-compatible APIs): + +```bash +npm install @ai-sdk/openai-compatible +``` + +Then register the provider and use it in model IDs: + +```typescript +import { configure } from "passmark"; +import { createOpenAICompatible } from "@ai-sdk/openai-compatible"; + +configure({ + ai: { + providers: { + "llm-proxy": { + createProvider: () => createOpenAICompatible({ + name: "llm-proxy", + apiKey: process.env.LLM_PROXY_API_KEY!, + baseURL: `${process.env.LLM_PROXY_API_BASE_URL}/api/v1/proxy/openai/auto`, + }), + }, + }, + models: { + stepExecution: "llm-proxy/google/gemini-3.5-flash", + assertionPrimary: "llm-proxy/anthropic/claude-haiku-4.5", + assertionSecondary: "llm-proxy/google/gemini-3-flash", + utility: "llm-proxy/google/gemini-2.5-flash", + }, + }, +}); +``` + +You can also route **all** models through a custom provider by setting `gateway` to the provider name: + +```typescript +configure({ + ai: { + gateway: "llm-proxy", // all models route through this provider + providers: { + "llm-proxy": { + createProvider: () => createOpenAICompatible({ + name: "llm-proxy", + apiKey: process.env.LLM_PROXY_API_KEY!, + baseURL: process.env.LLM_PROXY_API_BASE_URL!, + }), + // Optional: remap passmark's default model IDs to your proxy's model names + models: { + "google/gemini-3-flash": "gemini-3-flash", + "anthropic/claude-haiku-4.5": "claude-haiku-4-5", + }, + }, + }, + }, +}); +``` + +Notes: +- Provider instances are created lazily and cached — `createProvider()` is called only once per provider name. +- Custom providers work with per-step and per-call `ai` overrides, the same as built-in providers. +- Any `@ai-sdk/*` package can be used (e.g. `@ai-sdk/openai`, `@ai-sdk/anthropic`, `@ai-sdk/google`), not just `@ai-sdk/openai-compatible`. +- Video assertions still require `GOOGLE_GENERATIVE_AI_API_KEY` regardless of custom providers (Gemini Files API is used directly). + ## Environment Variables | Variable | Required | Default | Description | diff --git a/src/__tests__/config.test.ts b/src/__tests__/config.test.ts index 8525994..13a7add 100644 --- a/src/__tests__/config.test.ts +++ b/src/__tests__/config.test.ts @@ -1,5 +1,13 @@ import { describe, it, expect, beforeEach } from "vitest"; -import { configure, getConfig, getModelId, resetConfig, DEFAULT_MODELS } from "../config"; +import { + configure, + getConfig, + getModelId, + resetConfig, + DEFAULT_MODELS, + resolveAI, + type CustomProviderConfig, +} from "../config"; describe("config", () => { beforeEach(() => { @@ -96,4 +104,71 @@ describe("config", () => { resetConfig(); expect(getConfig()).toEqual({}); }); + + describe("custom providers", () => { + const mockProvider: CustomProviderConfig = { + createProvider: () => ({ + languageModel: (id: string) => id as never, + }) as never, + }; + + it("configure stores custom providers", () => { + configure({ + ai: { + providers: { "my-proxy": mockProvider }, + }, + }); + expect(getConfig().ai?.providers).toBeDefined(); + expect(getConfig().ai?.providers?.["my-proxy"]).toBe(mockProvider); + }); + + it("configure stores custom providers with model aliases", () => { + const providerWithAliases: CustomProviderConfig = { + ...mockProvider, + models: { "gemini-flash": "google/gemini-3.5-flash" }, + }; + configure({ + ai: { + providers: { "corp-proxy": providerWithAliases }, + }, + }); + expect(getConfig().ai?.providers?.["corp-proxy"]?.models).toEqual({ + "gemini-flash": "google/gemini-3.5-flash", + }); + }); + + it("resolveAI merges providers from global config and overrides", () => { + const providerA: CustomProviderConfig = { + createProvider: () => ({ languageModel: () => "a" }) as never, + }; + const providerB: CustomProviderConfig = { + createProvider: () => ({ languageModel: () => "b" }) as never, + }; + configure({ + ai: { providers: { alpha: providerA } }, + }); + const resolved = resolveAI({ providers: { beta: providerB } }); + expect(resolved.providers?.["alpha"]).toBe(providerA); + expect(resolved.providers?.["beta"]).toBe(providerB); + }); + + it("resolveAI later overrides win on same provider key", () => { + const providerV1: CustomProviderConfig = { + createProvider: () => ({ languageModel: () => "v1" }) as never, + }; + const providerV2: CustomProviderConfig = { + createProvider: () => ({ languageModel: () => "v2" }) as never, + }; + configure({ + ai: { providers: { "my-proxy": providerV1 } }, + }); + const resolved = resolveAI({ providers: { "my-proxy": providerV2 } }); + expect(resolved.providers?.["my-proxy"]).toBe(providerV2); + }); + + it("resolveAI returns undefined providers when none configured", () => { + const resolved = resolveAI(); + expect(resolved.providers).toBeUndefined(); + }); + }); }); diff --git a/src/__tests__/models.test.ts b/src/__tests__/models.test.ts new file mode 100644 index 0000000..b25f826 --- /dev/null +++ b/src/__tests__/models.test.ts @@ -0,0 +1,150 @@ +import { describe, it, expect, beforeEach, vi } from "vitest"; +import { resetConfig, configure, type CustomProviderConfig } from "../config"; + +// Mock dependencies that resolveModel uses internally +vi.mock("@ai-sdk/anthropic", () => ({ + createAnthropic: vi.fn(), +})); +vi.mock("@ai-sdk/google", () => ({ + createGoogleGenerativeAI: vi.fn(), +})); +vi.mock("@ai-sdk/openai", () => ({ + createOpenAI: vi.fn(), +})); +vi.mock("@openrouter/ai-sdk-provider", () => ({ + createOpenRouter: vi.fn(), +})); +vi.mock("ai", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + gateway: vi.fn(), + }; +}); +vi.mock("axiom/ai", () => ({ + wrapAISDKModel: vi.fn((model) => model), +})); +vi.mock("../instrumentation", () => ({ + isAxiomEnabled: vi.fn(() => false), +})); + +describe("resolveModel — custom providers", () => { + // We need to import resolveModel after mocks are set up + let resolveModel: typeof import("../models").resolveModel; + + beforeEach(async () => { + resetConfig(); + // Re-import to get fresh module state + const models = await import("../models"); + resolveModel = models.resolveModel; + }); + + function createMockProvider(returnValue?: string): CustomProviderConfig { + const languageModel = vi.fn((id: string) => returnValue ?? `resolved:${id}`); + const createProvider = vi.fn(() => ({ + languageModel, + textEmbeddingModel: vi.fn(), + })); + return { createProvider } as unknown as CustomProviderConfig; + } + + function createMockProviderWithAliases( + aliases: Record, + ): CustomProviderConfig { + const languageModel = vi.fn((id: string) => `resolved:${id}`); + const createProvider = vi.fn(() => ({ + languageModel, + textEmbeddingModel: vi.fn(), + })); + return { createProvider, models: aliases } as unknown as CustomProviderConfig; + } + + it("resolves model via custom provider prefix", () => { + const cp = createMockProvider(); + const result = resolveModel("my-proxy/gpt-4", "none", { "my-proxy": cp }); + expect(result).toBe("resolved:gpt-4"); + expect(cp.createProvider).toHaveBeenCalledOnce(); + }); + + it("resolves model via custom provider in gateway mode", () => { + const cp = createMockProvider(); + const result = resolveModel( + "google/gemini-3-flash", + "llm-proxy", + { "llm-proxy": cp }, + ); + expect(result).toBe("resolved:google/gemini-3-flash"); + expect(cp.createProvider).toHaveBeenCalledOnce(); + }); + + it("applies model aliases in provider-prefix mode", () => { + const cp = createMockProviderWithAliases({ + "gemini-flash": "google/gemini-3.5-flash-internal", + }); + const result = resolveModel("my-proxy/gemini-flash", "none", { + "my-proxy": cp, + }); + expect(result).toBe("resolved:google/gemini-3.5-flash-internal"); + }); + + it("applies model aliases in gateway mode", () => { + const cp = createMockProviderWithAliases({ + "google/gemini-3-flash": "gemini-3-flash-corp", + }); + const result = resolveModel("google/gemini-3-flash", "my-proxy", { + "my-proxy": cp, + }); + expect(result).toBe("resolved:gemini-3-flash-corp"); + }); + + it("passes through model name when no alias matches", () => { + const cp = createMockProviderWithAliases({ + "some-model": "other-model", + }); + const result = resolveModel("my-proxy/gpt-4", "none", { "my-proxy": cp }); + // No alias for "gpt-4", so it passes through as-is + expect(result).toBe("resolved:gpt-4"); + }); + + it("caches provider instances (createProvider called once)", () => { + const cp = createMockProvider(); + resolveModel("my-proxy/model-a", "none", { "my-proxy": cp }); + resolveModel("my-proxy/model-b", "none", { "my-proxy": cp }); + expect(cp.createProvider).toHaveBeenCalledOnce(); + }); + + it("throws for unknown provider when no custom provider matches", () => { + expect(() => + resolveModel("unknown-provider/model", "none"), + ).toThrow("Unknown AI provider: unknown-provider"); + }); + + it("falls back to built-in providers when custom provider doesn't match prefix", () => { + const cp = createMockProvider(); + // "google/..." should NOT match custom provider "my-proxy" + expect(() => + resolveModel("google/gemini-3-flash", "none", { "my-proxy": cp }), + ).toThrow(); // Throws because GOOGLE_GENERATIVE_AI_API_KEY is not set in test + }); + + it("reads custom providers from global config when not passed explicitly", () => { + const cp = createMockProvider(); + configure({ + ai: { + providers: { "global-proxy": cp }, + }, + }); + const result = resolveModel("global-proxy/some-model"); + expect(result).toBe("resolved:some-model"); + expect(cp.createProvider).toHaveBeenCalledOnce(); + }); + + it("handles nested model names with slashes in provider-prefix mode", () => { + const cp = createMockProvider(); + const result = resolveModel("my-proxy/org/model-name", "none", { + "my-proxy": cp, + }); + // Model name should be "org/model-name" + expect(result).toBe("resolved:org/model-name"); + }); +}); diff --git a/src/config.ts b/src/config.ts index 072b1e4..c631b59 100644 --- a/src/config.ts +++ b/src/config.ts @@ -1,3 +1,4 @@ +import type { Provider } from "ai"; import { initTelemetry } from "./instrumentation"; export type EmailProvider = { @@ -11,7 +12,7 @@ export type EmailProvider = { extractContent: (params: { email: string; prompt: string }) => Promise; }; -export type AIGateway = "vercel" | "openrouter" | "opencodezen" | "cloudflare" | "none"; +export type AIGateway = "vercel" | "openrouter" | "opencodezen" | "cloudflare" | "none" | (string & {}); /** * Execution mode for browser automation. @@ -21,6 +22,49 @@ export type AIGateway = "vercel" | "openrouter" | "opencodezen" | "cloudflare" | */ export type AIMode = "snapshot" | "cua"; +/** + * Configuration for a custom AI provider. Register custom providers via + * `configure({ ai: { providers: { "my-proxy": { ... } } } })` and reference + * them in model IDs as `"my-proxy/model-name"`. + * + * @example + * ```typescript + * import { createOpenAICompatible } from "@ai-sdk/openai-compatible"; + * + * configure({ + * ai: { + * providers: { + * "llm-proxy": { + * createProvider: () => createOpenAICompatible({ + * name: "llm-proxy", + * apiKey: process.env.LLM_PROXY_API_KEY, + * baseURL: process.env.LLM_PROXY_BASE_URL, + * }), + * }, + * }, + * models: { + * stepExecution: "llm-proxy/gemini-3.5-flash", + * }, + * }, + * }); + * ``` + */ +export type CustomProviderConfig = { + /** + * A function that creates a Vercel AI SDK provider instance. + * Called lazily on first use and cached thereafter. + * Any `@ai-sdk/*` package (e.g. `@ai-sdk/openai-compatible`, + * `@ai-sdk/openai`, `@ai-sdk/anthropic`) can be used. + */ + createProvider: () => Provider; + /** + * Optional model alias map. Keys are the model names used in passmark + * config (e.g. "gemini-3.5-flash"), values are the actual model IDs + * to pass to the provider SDK. When absent, model names pass through as-is. + */ + models?: Record; +}; + export type ModelConfig = { /** Model for executing individual steps. Default: google/gemini-3-flash */ stepExecution?: string; @@ -65,6 +109,12 @@ export type AIOverride = { gateway?: AIGateway; mode?: AIMode; models?: ModelConfig; + /** + * Register custom AI providers by name. The name becomes the provider + * prefix in model IDs (e.g. `"my-proxy/gpt-4"`). Can also be used as a + * gateway value to route all models through one custom provider. + */ + providers?: Record; }; export type RedisConfig = { @@ -201,6 +251,8 @@ export type ResolvedAI = { mode: AIMode; gateway: AIGateway; getModelId: (key: keyof ModelConfig) => string; + /** Merged custom providers from all config layers (global → call → step). */ + providers?: Record; }; const CUA_LOCK_MESSAGE = @@ -241,10 +293,27 @@ export function resolveAI(...overrides: (AIOverride | undefined)[]): ResolvedAI } return DEFAULT_MODELS[key]; }; - return { mode, gateway, getModelId: getModelIdForKey }; + // Merge providers: later layers win on a per-key basis + let mergedProviders: Record | undefined; + for (const layer of layers) { + if (layer?.providers) { + mergedProviders = { ...(mergedProviders ?? {}), ...layer.providers }; + } + } + return { mode, gateway, getModelId: getModelIdForKey, providers: mergedProviders }; } /** @internal Reset config to empty state. Used for testing only. */ export function resetConfig() { globalConfig = {}; + // Also reset the provider instance cache so tests start fresh + resetProviderCache(); +} + +// Lazy import to avoid circular dependency — set by models.ts at module load +let resetProviderCache: () => void = () => {}; + +/** @internal Called by models.ts to register its cache-reset function. */ +export function _registerProviderCacheReset(fn: () => void) { + resetProviderCache = fn; } diff --git a/src/index.ts b/src/index.ts index 75ee9d4..35de42b 100644 --- a/src/index.ts +++ b/src/index.ts @@ -473,7 +473,7 @@ export const runSteps = async ({ } const stepModelId = effectiveAi.getModelId("stepExecution"); - const model = resolveModel(stepModelId, effectiveAi.gateway); + const model = resolveModel(stepModelId, effectiveAi.gateway, effectiveAi.providers); logger.debug( `Using model: ${stepModelId} for step execution / gateway: ${effectiveAi.gateway}`, ); @@ -724,7 +724,7 @@ export const runUserFlow = async ({ if (assertion) { const { output } = await generateText({ - model: resolveModel(effectiveAi.getModelId("utility"), effectiveAi.gateway), + model: resolveModel(effectiveAi.getModelId("utility"), effectiveAi.gateway, effectiveAi.providers), prompt: `Convert the following text output into a valid JSON object with the specified properties:\n\n${text}`, output: Output.object({ schema: z.object({ @@ -750,8 +750,8 @@ export const runUserFlow = async ({ const model = effort === "low" - ? resolveModel(effectiveAi.getModelId("userFlowLow"), effectiveAi.gateway) - : resolveModel(effectiveAi.getModelId("userFlowHigh"), effectiveAi.gateway); + ? resolveModel(effectiveAi.getModelId("userFlowLow"), effectiveAi.gateway, effectiveAi.providers) + : resolveModel(effectiveAi.getModelId("userFlowHigh"), effectiveAi.gateway, effectiveAi.providers); const { tools } = getAItools(page, { abortController, @@ -803,7 +803,7 @@ export const runUserFlow = async ({ if (assertion) { const { output } = await generateText({ - model: resolveModel(effectiveAi.getModelId("utility"), effectiveAi.gateway), + model: resolveModel(effectiveAi.getModelId("utility"), effectiveAi.gateway, effectiveAi.providers), prompt: `Convert the following text output into a valid JSON object with the specified properties:\n\n${text}`, output: Output.object({ schema: z.object({ @@ -868,7 +868,7 @@ export const executeWithAutoHealing = async (config: { }; export { configure } from "./config"; -export type { EmailProvider } from "./config"; +export type { EmailProvider, CustomProviderConfig } from "./config"; export { emailsinkProvider } from "./providers/emailsink"; export { extractEmailContent, generateEmail } from "./email"; diff --git a/src/models.ts b/src/models.ts index 47d6b88..451525e 100644 --- a/src/models.ts +++ b/src/models.ts @@ -3,9 +3,9 @@ import { createAnthropic } from "@ai-sdk/anthropic"; import { createGoogleGenerativeAI } from "@ai-sdk/google"; import { createOpenAI } from "@ai-sdk/openai"; import { createOpenRouter } from "@openrouter/ai-sdk-provider"; -import { gateway, type LanguageModel } from "ai"; +import { gateway, type LanguageModel, type Provider } from "ai"; import { wrapAISDKModel } from "axiom/ai"; -import { type AIGateway, getConfig } from "./config"; +import { type AIGateway, type CustomProviderConfig, getConfig, _registerProviderCacheReset } from "./config"; import { isAxiomEnabled } from "./instrumentation"; function wrapModel(model: LanguageModel): LanguageModel { @@ -20,6 +20,26 @@ let _opencodezen: ReturnType | null = null; let _cloudflareGoogle: ReturnType | null = null; let _cloudflareAnthropic: ReturnType | null = null; +/** + * Cache for custom provider instances. Ensures `createProvider()` is called + * only once per provider name per process lifetime. + */ +const _customProviderCache = new Map(); + +function getCachedProvider(name: string, config: CustomProviderConfig): Provider { + let provider = _customProviderCache.get(name); + if (!provider) { + provider = config.createProvider(); + _customProviderCache.set(name, provider); + } + return provider; +} + +// Register the cache-reset function so resetConfig() can clear it during tests +_registerProviderCacheReset(() => { + _customProviderCache.clear(); +}); + function getGoogleProvider() { if (!_google) { if (!process.env.GOOGLE_GENERATIVE_AI_API_KEY) { @@ -217,16 +237,57 @@ function resolveOpenCodeZenModelId(modelId: string): string { * provider-native paths (google-ai-studio, anthropic) so provider-specific fields * like Gemini's thought_signature pass through unchanged. * When gateway is "none" (default), creates a direct provider instance with alias resolution. + * When gateway matches a custom provider name, all models route through that provider. * All paths wrap the model with wrapAISDKModel for tracing when Axiom is enabled. * * @param modelId - Canonical model id, e.g. "google/gemini-3-flash". * @param gatewayOverride - Optional resolved gateway for this call. When omitted, * falls back to the global `configure()` value. Pass this when a per-step or * per-call `ai` override changes the gateway for a single resolution. + * @param customProviders - Optional map of custom provider configurations. + * When omitted, falls back to the global `configure()` providers. */ -export function resolveModel(modelId: string, gatewayOverride?: AIGateway): LanguageModel { +export function resolveModel( + modelId: string, + gatewayOverride?: AIGateway, + customProviders?: Record, +): LanguageModel { const gatewayConfig = gatewayOverride ?? getConfig().ai?.gateway ?? "none"; + const providers = customProviders ?? getConfig().ai?.providers; + + // --- Custom provider resolution (checked before built-in providers) --- + if (providers) { + // Gateway mode: the gateway value itself names a custom provider, + // so all models route through it (e.g. gateway: "llm-proxy") + if ( + gatewayConfig !== "none" && + gatewayConfig !== "vercel" && + gatewayConfig !== "openrouter" && + gatewayConfig !== "opencodezen" && + gatewayConfig !== "cloudflare" && + providers[gatewayConfig] + ) { + const cp = providers[gatewayConfig]; + const provider = getCachedProvider(gatewayConfig, cp); + const resolvedModelName = cp.models?.[modelId] ?? modelId; + return wrapModel(provider.languageModel(resolvedModelName)); + } + + // Provider-prefix mode: "my-proxy/gpt-4" → providerKey="my-proxy", model="gpt-4" + const slashIdx = modelId.indexOf("/"); + if (slashIdx !== -1) { + const providerKey = modelId.slice(0, slashIdx); + if (providers[providerKey]) { + const cp = providers[providerKey]; + const modelName = modelId.slice(slashIdx + 1); + const resolvedModelName = cp.models?.[modelName] ?? modelName; + const provider = getCachedProvider(providerKey, cp); + return wrapModel(provider.languageModel(resolvedModelName)); + } + } + } + // --- Built-in provider resolution --- if (gatewayConfig === "vercel") { if (!process.env.AI_GATEWAY_API_KEY) { throw new ConfigurationError(