Explorar o código

refactor(core): generalize OpenAI-compatible routing

Aiden Cline hai 1 semana
pai
achega
16f144465b

+ 0 - 26
packages/ai/src/providers/cloudflare.ts

@@ -2,7 +2,6 @@ import type { Config, Redacted } from "effect"
 import * as OpenAICompatibleChat from "../protocols/openai-compatible-chat"
 import { Auth } from "../route/auth"
 import { AuthOptions, type AtLeastOne, type ProviderAuthOption } from "../route/auth-options"
-import type { ProviderPackage } from "../provider-package"
 import type { RouteDefaultsInput } from "../route/client"
 import { ProviderID, type ModelID } from "../schema"
 import type { OpenAIProviderOptionsInput } from "./openai-options"
@@ -40,13 +39,6 @@ export type WorkersAIOptions = WorkersAIURL &
     readonly providerOptions?: OpenAIProviderOptionsInput
   }
 
-export interface WorkersAISettings extends ProviderPackage.Settings {
-  readonly accountId?: string
-  readonly apiKey?: string
-  readonly baseURL?: string
-  readonly providerOptions?: OpenAIProviderOptionsInput
-}
-
 export const aiGatewayBaseURL = (input: GatewayURL) => {
   if (input.baseURL) return input.baseURL
   if (!input.accountId) throw new Error("CloudflareAIGateway.configure requires accountId unless baseURL is supplied")
@@ -139,21 +131,3 @@ export const CloudflareWorkersAI = {
   id: workersAIID,
   configure: configureWorkersAI,
 }
-
-export const workersAIModel: ProviderPackage.Definition<WorkersAISettings, OpenAIProviderOptionsInput>["model"] = (
-  modelID,
-  settings,
-) => {
-  const body = settings.body === undefined ? undefined : { ...settings.body }
-  if (body) delete body.accountId
-  const defaults = {
-    apiKey: settings.apiKey,
-    headers: settings.headers === undefined ? undefined : { ...settings.headers },
-    http: body === undefined ? undefined : { body },
-    limits: settings.limits,
-    providerOptions: settings.providerOptions,
-  }
-  if (settings.baseURL) return configureWorkersAI({ ...defaults, baseURL: settings.baseURL }).model(modelID)
-  if (settings.accountId) return configureWorkersAI({ ...defaults, accountId: settings.accountId }).model(modelID)
-  throw new Error("Cloudflare Workers AI requires accountId or baseURL")
-}

+ 0 - 2
packages/ai/src/providers/cloudflare/workers-ai.ts

@@ -1,2 +0,0 @@
-export { workersAIModel as model } from "../cloudflare"
-export type { WorkersAISettings as Settings } from "../cloudflare"

+ 7 - 2
packages/ai/src/providers/openai-compatible.ts

@@ -12,6 +12,7 @@ type GenericModelOptions = Omit<RouteDefaultsInput, "providerOptions"> &
   ProviderAuthOption<"optional"> & {
     readonly provider?: string
     readonly baseURL: string
+    readonly queryParams?: Readonly<Record<string, string>>
     readonly providerOptions?: OpenAIProviderOptionsInput
   }
 
@@ -19,6 +20,8 @@ export interface Settings extends ProviderPackage.Settings {
   readonly apiKey?: string
   readonly baseURL: string
   readonly provider?: string
+  readonly providerOptions?: OpenAIProviderOptionsInput
+  readonly queryParams?: Readonly<Record<string, string>>
 }
 
 export type FamilyModelOptions = Omit<RouteDefaultsInput, "providerOptions"> &
@@ -31,11 +34,11 @@ export const routes = [OpenAICompatibleChat.route]
 
 export const configure = (input: GenericModelOptions) => {
   const provider = input.provider ?? "openai-compatible"
-  const { provider: _, baseURL, apiKey: _apiKey, auth: _auth, ...rest } = input
+  const { provider: _, baseURL, apiKey: _apiKey, auth: _auth, queryParams, ...rest } = input
   const route = OpenAICompatibleChat.route.with({
     ...rest,
     provider,
-    endpoint: { baseURL },
+    endpoint: { baseURL, query: queryParams },
     auth: AuthOptions.bearer(input, []),
   })
   return {
@@ -75,6 +78,8 @@ export const model: ProviderPackage.Definition<Settings, OpenAIProviderOptionsIn
     http: settings.body === undefined ? undefined : { body: { ...settings.body } },
     limits: settings.limits,
     provider: settings.provider,
+    providerOptions: settings.providerOptions,
+    queryParams: settings.queryParams === undefined ? undefined : { ...settings.queryParams },
   }).model(modelID)
 
 export const baseten = define(profiles.baseten)

+ 18 - 20
packages/ai/test/provider-package.test.ts

@@ -26,7 +26,6 @@ describe("provider package entrypoints", () => {
       import("@opencode-ai/ai/providers/amazon-bedrock/mantle"),
       import("@opencode-ai/ai/providers/amazon-bedrock/mantle/chat"),
       import("@opencode-ai/ai/providers/amazon-bedrock/mantle/responses"),
-      import("@opencode-ai/ai/providers/cloudflare/workers-ai"),
     ])
 
     for (const module of modules) expect(module.model).toBeFunction()
@@ -36,25 +35,6 @@ describe("provider package entrypoints", () => {
     expect(modules[19].model).toBe(modules[20].model)
   })
 
-  test("maps Cloudflare Workers AI settings onto its native route", async () => {
-    const WorkersAI = await import("@opencode-ai/ai/providers/cloudflare/workers-ai")
-    const model = WorkersAI.model("@cf/meta/llama-3.1-8b-instruct", {
-      accountId: "account/id",
-      apiKey: "secret",
-      body: { custom: true, accountId: "account/id" },
-      limits: { context: 128_000, output: 8_192 },
-    })
-
-    expect(model.route).toMatchObject({
-      id: "cloudflare-workers-ai",
-      endpoint: { baseURL: "https://api.cloudflare.com/client/v4/accounts/account%2Fid/ai/v1" },
-      defaults: {
-        http: { body: { custom: true } },
-        limits: { context: 128_000, output: 8_192 },
-      },
-    })
-  })
-
   test("maps OpenRouter and xAI package settings onto executable models", async () => {
     const OpenRouter = await import("@opencode-ai/ai/providers/openrouter")
     const XAI = await import("@opencode-ai/ai/providers/xai")
@@ -84,6 +64,24 @@ describe("provider package entrypoints", () => {
     expect(xai.route.defaults.providerOptions).toMatchObject({ xai: { reasoningEffort: "high", store: false } })
   })
 
+  test("maps OpenAI-compatible package settings onto the executable model", async () => {
+    const OpenAICompatible = await import("@opencode-ai/ai/providers/openai-compatible")
+    const selected = OpenAICompatible.model("custom-model", {
+      apiKey: "fixture",
+      baseURL: "https://provider.example.test/v1",
+      provider: "example",
+      queryParams: { version: "preview" },
+      providerOptions: { openai: { reasoningEffort: "high" } },
+    })
+
+    expect(String(selected.provider)).toBe("example")
+    expect(selected.route.endpoint).toMatchObject({
+      baseURL: "https://provider.example.test/v1",
+      query: { version: "preview" },
+    })
+    expect(selected.route.defaults.providerOptions).toEqual({ openai: { reasoningEffort: "high" } })
+  })
+
   test("maps package settings onto the executable model", () => {
     const selected = model("gpt-5", {
       apiKey: "fixture",

+ 15 - 31
packages/core/src/aisdk-native.ts

@@ -53,9 +53,7 @@ export function map(input: MapInput): Mapping | undefined {
         },
       }
     case "@ai-sdk/openai-compatible":
-      return input.providerID === "cloudflare-workers-ai"
-        ? mapCloudflareWorkers(input, baseSettings)
-        : mapOpenAICompatible(input, baseSettings)
+      return mapOpenAICompatible(input, baseSettings)
     case "@openrouter/ai-sdk-provider":
       return mapOpenRouter(input.settings, baseSettings)
     case "@ai-sdk/xai":
@@ -68,49 +66,35 @@ export function map(input: MapInput): Mapping | undefined {
         },
       }
   }
-}
-
-function mapCloudflareWorkers(input: MapInput, baseSettings: Readonly<Record<string, unknown>>): Mapping {
-  const accountId = typeof input.settings.accountId === "string" ? input.settings.accountId : undefined
-  const configured = typeof baseSettings.baseURL === "string" ? baseSettings.baseURL : undefined
-  const baseURL =
-    configured && accountId
-      ? configured.replaceAll("${CLOUDFLARE_ACCOUNT_ID}", encodeURIComponent(accountId))
-      : configured
-  return {
-    package: "@opencode-ai/ai/providers/cloudflare/workers-ai",
-    settings: {
-      ...(baseURL?.includes("${CLOUDFLARE_ACCOUNT_ID}") ? {} : baseURL ? { baseURL } : {}),
-      ...mapAPIKey(input.settings),
-      ...(accountId ? { accountId } : {}),
-      ...mapOpenAICompatibleOptions(input.settings, ["accountId"]),
-    },
-  }
+  return undefined
 }
 
 function mapOpenAICompatible(
   input: MapInput,
   baseSettings: Readonly<Record<string, unknown>>,
 ): Mapping | undefined {
-  if (typeof baseSettings.baseURL !== "string") return
+  const accountId =
+    input.providerID === "cloudflare-workers-ai" && typeof input.settings.accountId === "string"
+      ? input.settings.accountId
+      : undefined
+  const baseURL =
+    typeof baseSettings.baseURL === "string" && accountId
+      ? baseSettings.baseURL.replaceAll("${CLOUDFLARE_ACCOUNT_ID}", encodeURIComponent(accountId))
+      : baseSettings.baseURL
+  if (typeof baseURL !== "string") return undefined
   return {
     package: "@opencode-ai/ai/providers/openai-compatible",
     settings: {
-      ...baseSettings,
+      baseURL,
       ...mapAPIKey(input.settings),
       provider: input.providerID,
-      ...mapOpenAICompatibleOptions(input.settings),
+      ...(isStringRecord(input.settings.queryParams) ? { queryParams: input.settings.queryParams } : {}),
+      ...mapOpenAIOptions(input.settings),
     },
+    ...(isStringRecord(input.settings.headers) ? { headers: input.settings.headers } : {}),
   }
 }
 
-function mapOpenAICompatibleOptions(settings: Readonly<Record<string, unknown>>, exclude: readonly string[] = []) {
-  const options = Object.fromEntries(
-    Object.entries(settings).filter(([key]) => !["apiKey", "baseURL", ...exclude].includes(key)),
-  )
-  return Object.keys(options).length === 0 ? {} : { providerOptions: { openai: options } }
-}
-
 function mapBedrockMantle(input: MapInput, baseSettings: Readonly<Record<string, unknown>>): Mapping | undefined {
   const settings = input.settings
   const chat = input.modelID === "openai.gpt-oss-safeguard-20b" || input.modelID === "openai.gpt-oss-safeguard-120b"

+ 1 - 0
packages/core/src/model-resolver.ts

@@ -144,6 +144,7 @@ export const fromCatalogModel = (
     if (draft.settings?.apiKey === "") delete draft.settings.apiKey
     if (credential?.type === "key" && credential.metadata !== undefined)
       draft.body = Provider.mergeOverlay(draft.body, credential.metadata)
+    if (draft.providerID === "cloudflare-workers-ai" && draft.body) delete draft.body.accountId
   })
   const packageName = Provider.packageName(resolved.package)
   const key = apiKey(resolved, credential)

+ 0 - 4
packages/core/src/provider.ts

@@ -47,10 +47,6 @@ const builtins = new Map<string, () => Promise<unknown>>([
   ["@opencode-ai/ai/providers/azure", () => import("@opencode-ai/ai/providers/azure")],
   ["@opencode-ai/ai/providers/azure/chat", () => import("@opencode-ai/ai/providers/azure/chat")],
   ["@opencode-ai/ai/providers/azure/responses", () => import("@opencode-ai/ai/providers/azure/responses")],
-  [
-    "@opencode-ai/ai/providers/cloudflare/workers-ai",
-    () => import("@opencode-ai/ai/providers/cloudflare/workers-ai"),
-  ],
   ["@opencode-ai/ai/providers/google", () => import("@opencode-ai/ai/providers/google")],
   ["@opencode-ai/ai/providers/openai", () => import("@opencode-ai/ai/providers/openai")],
   ["@opencode-ai/ai/providers/openai/chat", () => import("@opencode-ai/ai/providers/openai/chat")],

+ 7 - 3
packages/core/test/aisdk-native.test.ts

@@ -45,7 +45,7 @@ describe("AISDKNative", () => {
     )
   })
 
-  test("maps Cloudflare Workers AI to its native provider", () => {
+  test("maps Cloudflare Workers AI to the generic OpenAI-compatible provider", () => {
     expect(
       map(
         "@ai-sdk/openai-compatible",
@@ -53,19 +53,23 @@ describe("AISDKNative", () => {
           accountId: "account/id",
           apiKey: "secret",
           baseURL: "https://api.cloudflare.com/client/v4/accounts/${CLOUDFLARE_ACCOUNT_ID}/ai/v1",
+          headers: { "x-custom": "value" },
+          queryParams: { version: "preview" },
           reasoningEffort: "high",
         },
         "@cf/model",
         "cloudflare-workers-ai",
       ),
     ).toEqual({
-      package: "@opencode-ai/ai/providers/cloudflare/workers-ai",
+      package: "@opencode-ai/ai/providers/openai-compatible",
       settings: {
-        accountId: "account/id",
         apiKey: "secret",
         baseURL: "https://api.cloudflare.com/client/v4/accounts/account%2Fid/ai/v1",
+        provider: "cloudflare-workers-ai",
+        queryParams: { version: "preview" },
         providerOptions: { openai: { reasoningEffort: "high" } },
       },
+      headers: { "x-custom": "value" },
     })
   })
 

+ 13 - 2
packages/core/test/model-resolver.test.ts

@@ -131,7 +131,7 @@ describe("ModelResolver", () => {
     }),
   )
 
-  it.effect("routes Cloudflare Workers AI through its native provider", () =>
+  it.effect("routes Cloudflare Workers AI through the generic OpenAI-compatible provider", () =>
     Effect.gen(function* () {
       const resolved = yield* ModelResolver.fromCatalogModel(
         model(Provider.aisdk("@ai-sdk/openai-compatible"), {
@@ -139,6 +139,8 @@ describe("ModelResolver", () => {
           modelID: "@cf/meta/llama-3.1-8b-instruct",
           settings: {
             baseURL: "https://api.cloudflare.com/client/v4/accounts/${CLOUDFLARE_ACCOUNT_ID}/ai/v1",
+            queryParams: { version: "preview" },
+            reasoningEffort: "high",
           },
         }),
         Credential.Key.make({ type: "key", key: "secret", metadata: { accountId: "account/id" } }),
@@ -152,9 +154,18 @@ describe("ModelResolver", () => {
         headers: Headers.empty,
       })
 
-      expect(resolved.route.id).toBe("cloudflare-workers-ai")
+      expect(resolved.route.id).toBe("openai-compatible-chat")
+      expect(String(resolved.provider)).toBe("cloudflare-workers-ai")
       expect(resolved.route.endpoint.baseURL).toBe("https://api.cloudflare.com/client/v4/accounts/account%2Fid/ai/v1")
+      expect(resolved.route.endpoint.query).toEqual({ version: "preview" })
+      expect(resolved.route.defaults.providerOptions).toEqual({ openai: { reasoningEffort: "high" } })
       expect(resolved.route.defaults.http?.body).toEqual({ custom_extension: { enabled: true } })
+      const prepared = yield* compileRequest(LLM.request({ model: resolved, prompt: "Hello" }))
+      expect(prepared.body).toMatchObject({
+        reasoning_effort: "high",
+        stream_options: { include_usage: true },
+      })
+      expect(prepared.body).not.toHaveProperty("accountId")
       expect(headers.authorization).toBe("Bearer secret")
     }),
   )

+ 19 - 0
packages/core/test/plugin/provider-cloudflare-workers-ai.test.ts

@@ -59,6 +59,25 @@ describe("CloudflareWorkersAIPlugin", () => {
     ),
   )
 
+  it.effect("resolves an account ID from provider settings", () =>
+    withEnv(undefined, () =>
+      Effect.gen(function* () {
+        const catalog = yield* Catalog.Service
+        yield* catalog.transform((draft) =>
+          draft.provider.update(providerID, (provider) => {
+            provider.package = Provider.aisdk("@ai-sdk/openai-compatible")
+            provider.settings = { accountId: "configured/account" }
+          }),
+        )
+        yield* addPlugin()
+
+        expect(required(yield* catalog.provider.get(providerID)).settings?.baseURL).toBe(
+          "https://api.cloudflare.com/client/v4/accounts/configured%2Faccount/ai/v1",
+        )
+      }),
+    ),
+  )
+
   it.effect("expands account placeholders and preserves configured endpoints", () =>
     withEnv("env-account", () =>
       Effect.gen(function* () {