provider-cloudflare-workers-ai.test.ts 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291
  1. import { AISDK } from "@opencode-ai/core/aisdk"
  2. import { describe, expect } from "bun:test"
  3. import { Effect } from "effect"
  4. import { Catalog } from "@opencode-ai/core/catalog"
  5. import { Model } from "@opencode-ai/core/model"
  6. import { Plugin } from "@opencode-ai/core/plugin"
  7. import { PluginHost } from "@opencode-ai/core/plugin/host"
  8. import { CloudflareWorkersAIPlugin } from "@opencode-ai/core/plugin/provider/cloudflare-workers-ai"
  9. import { Provider } from "@opencode-ai/core/provider"
  10. import { Integration } from "@opencode-ai/core/integration"
  11. import type { LanguageModelV3 } from "@ai-sdk/provider"
  12. import { testEffect } from "../lib/effect"
  13. import { PluginTestLayer } from "./fixture"
  14. const it = testEffect(PluginTestLayer)
  15. const addPlugin = Effect.fn(function* () {
  16. const plugin = yield* Plugin.Service
  17. const aisdk = yield* AISDK.Service
  18. const host = yield* PluginHost.make(plugin)
  19. yield* CloudflareWorkersAIPlugin.effect(host)
  20. })
  21. function required<T>(value: T | undefined): T {
  22. if (value === undefined) throw new Error("Expected value")
  23. return value
  24. }
  25. function withEnv<A, E, R>(vars: Record<string, string | undefined>, effect: () => Effect.Effect<A, E, R>) {
  26. return Effect.acquireUseRelease(
  27. Effect.sync(() => {
  28. const previous = Object.fromEntries(Object.keys(vars).map((key) => [key, process.env[key]]))
  29. Object.entries(vars).forEach(([key, value]) => {
  30. if (value === undefined) delete process.env[key]
  31. else process.env[key] = value
  32. })
  33. return previous
  34. }),
  35. effect,
  36. (previous) =>
  37. Effect.sync(() =>
  38. Object.entries(previous).forEach(([key, value]) => {
  39. if (value === undefined) delete process.env[key]
  40. else process.env[key] = value
  41. }),
  42. ),
  43. )
  44. }
  45. function fakeSelectorSdk(calls: string[]) {
  46. const make = (method: string) => (id: string) => {
  47. calls.push(`${method}:${id}`)
  48. return { modelId: id, provider: method, specificationVersion: "v3" } as unknown as LanguageModelV3
  49. }
  50. return {
  51. responses: make("responses"),
  52. messages: make("messages"),
  53. chat: make("chat"),
  54. languageModel: make("languageModel"),
  55. }
  56. }
  57. function cloudflareLanguage(sdk: unknown, modelID = "@cf/model") {
  58. return (sdk as { languageModel: (id: string) => { config: CloudflareConfig; provider: string } }).languageModel(
  59. modelID,
  60. )
  61. }
  62. type CloudflareConfig = {
  63. url: (input: { path: string; modelId: string }) => string
  64. headers: () => Record<string, string> | Promise<Record<string, string>>
  65. }
  66. function cloudflareURL(sdk: unknown, modelID = "@cf/model") {
  67. return cloudflareLanguage(sdk, modelID).config.url({ path: "/chat/completions", modelId: modelID })
  68. }
  69. function cloudflareHeaders(sdk: unknown, modelID = "@cf/model") {
  70. return cloudflareLanguage(sdk, modelID).config.headers()
  71. }
  72. describe("CloudflareWorkersAIPlugin", () => {
  73. it.effect("prompts for the account ID when the environment does not provide it", () =>
  74. withEnv({ CLOUDFLARE_ACCOUNT_ID: undefined }, () =>
  75. Effect.gen(function* () {
  76. yield* addPlugin()
  77. expect(
  78. (yield* (yield* Integration.Service).get(Integration.ID.make("cloudflare-workers-ai")))?.methods,
  79. ).toContainEqual({
  80. type: "key",
  81. label: "API key",
  82. prompts: [
  83. {
  84. type: "text",
  85. key: "accountId",
  86. message: "Enter your Cloudflare Account ID",
  87. placeholder: "e.g. 1234567890abcdef1234567890abcdef",
  88. },
  89. ],
  90. })
  91. }),
  92. ),
  93. )
  94. it.effect("maps account ID to endpoint URL and creates an OpenAI-compatible SDK", () =>
  95. withEnv({ CLOUDFLARE_ACCOUNT_ID: "acct", CLOUDFLARE_API_KEY: "key" }, () =>
  96. Effect.gen(function* () {
  97. const plugin = yield* Plugin.Service
  98. const aisdk = yield* AISDK.Service
  99. const catalog = yield* Catalog.Service
  100. yield* catalog.transform((catalog) =>
  101. catalog.provider.update(Provider.ID.make("cloudflare-workers-ai"), (provider) => {
  102. provider.package = Provider.aisdk("test-provider")
  103. }),
  104. )
  105. yield* addPlugin()
  106. const provider = required(yield* catalog.provider.get(Provider.ID.make("cloudflare-workers-ai")))
  107. const sdk = yield* aisdk.runSDK({
  108. model: Model.Info.make({
  109. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("@cf/model")),
  110. modelID: Model.ID.make("@cf/model"),
  111. package: provider.package,
  112. settings: provider.settings,
  113. }),
  114. package: "@ai-sdk/openai-compatible",
  115. options: { name: "cloudflare-workers-ai", headers: { custom: "header" } },
  116. })
  117. expect(provider).toMatchObject({
  118. package: "aisdk:test-provider",
  119. settings: { baseURL: "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1" },
  120. })
  121. expect(sdk.sdk).toBeDefined()
  122. }),
  123. ),
  124. )
  125. it.effect("preserves a configured endpoint URL instead of deriving one from account ID", () =>
  126. withEnv({ CLOUDFLARE_ACCOUNT_ID: "acct" }, () =>
  127. Effect.gen(function* () {
  128. const catalog = yield* Catalog.Service
  129. yield* catalog.transform((catalog) =>
  130. catalog.provider.update(Provider.ID.make("cloudflare-workers-ai"), (provider) => {
  131. provider.package = Provider.aisdk("test-provider")
  132. provider.settings = { ...provider.settings, baseURL: "https://proxy.example/v1" }
  133. }),
  134. )
  135. yield* addPlugin()
  136. expect(required(yield* catalog.provider.get(Provider.ID.make("cloudflare-workers-ai")))).toMatchObject({
  137. package: "aisdk:test-provider",
  138. settings: { baseURL: "https://proxy.example/v1" },
  139. })
  140. }),
  141. ),
  142. )
  143. it.effect("allows a configured baseURL without account ID", () =>
  144. withEnv({ CLOUDFLARE_ACCOUNT_ID: undefined, CLOUDFLARE_API_KEY: "key" }, () =>
  145. Effect.gen(function* () {
  146. const plugin = yield* Plugin.Service
  147. const aisdk = yield* AISDK.Service
  148. yield* addPlugin()
  149. const result = yield* aisdk.runSDK({
  150. model: Model.Info.make({
  151. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("@cf/model")),
  152. modelID: Model.ID.make("@cf/model"),
  153. package: "aisdk:@ai-sdk/openai-compatible",
  154. settings: { baseURL: "https://proxy.example/v1" },
  155. }),
  156. package: "@ai-sdk/openai-compatible",
  157. options: { name: "cloudflare-workers-ai", baseURL: "https://proxy.example/v1" },
  158. })
  159. expect(cloudflareURL(result.sdk)).toBe("https://proxy.example/v1/chat/completions")
  160. }),
  161. ),
  162. )
  163. it.effect("uses env account ID over configured account ID", () =>
  164. withEnv({ CLOUDFLARE_ACCOUNT_ID: "env-acct" }, () =>
  165. Effect.gen(function* () {
  166. const catalog = yield* Catalog.Service
  167. yield* catalog.transform((catalog) =>
  168. catalog.provider.update(Provider.ID.make("cloudflare-workers-ai"), (provider) => {
  169. provider.package = Provider.aisdk("test-provider")
  170. provider.settings = { ...provider.settings, accountId: "configured-acct" }
  171. }),
  172. )
  173. yield* addPlugin()
  174. expect(required(yield* catalog.provider.get(Provider.ID.make("cloudflare-workers-ai")))).toMatchObject({
  175. package: "aisdk:test-provider",
  176. settings: { baseURL: "https://api.cloudflare.com/client/v4/accounts/env-acct/ai/v1" },
  177. })
  178. }),
  179. ),
  180. )
  181. it.effect("uses env API key over auth or configured API key and keeps the Cloudflare User-Agent", () =>
  182. withEnv({ CLOUDFLARE_ACCOUNT_ID: "acct", CLOUDFLARE_API_KEY: "env-key" }, () =>
  183. Effect.gen(function* () {
  184. const plugin = yield* Plugin.Service
  185. const aisdk = yield* AISDK.Service
  186. yield* addPlugin()
  187. const result = yield* aisdk.runSDK({
  188. model: Model.Info.make({
  189. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("@cf/model")),
  190. modelID: Model.ID.make("@cf/model"),
  191. package: "aisdk:@ai-sdk/openai-compatible",
  192. settings: { baseURL: "https://proxy.example/v1" },
  193. }),
  194. package: "@ai-sdk/openai-compatible",
  195. options: {
  196. name: "cloudflare-workers-ai",
  197. apiKey: "auth-key",
  198. baseURL: "https://proxy.example/v1",
  199. headers: { custom: "header" },
  200. },
  201. })
  202. const headers = yield* Effect.promise(() => Promise.resolve(cloudflareHeaders(result.sdk)))
  203. expect(headers.authorization).toBe("Bearer env-key")
  204. expect(headers.custom).toBe("header")
  205. expect(headers["user-agent"]).toMatch(/^opencode\/.* cloudflare-workers-ai \(.+\) ai-sdk\/openai-compatible\//)
  206. }),
  207. ),
  208. )
  209. it.effect("expands account ID vars in endpoint URLs", () =>
  210. withEnv({ CLOUDFLARE_ACCOUNT_ID: "acct", CLOUDFLARE_API_KEY: "key" }, () =>
  211. Effect.gen(function* () {
  212. const plugin = yield* Plugin.Service
  213. const aisdk = yield* AISDK.Service
  214. yield* addPlugin()
  215. const result = yield* aisdk.runSDK({
  216. model: Model.Info.make({
  217. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("@cf/model")),
  218. modelID: Model.ID.make("@cf/model"),
  219. package: "aisdk:@ai-sdk/openai-compatible",
  220. settings: { baseURL: "https://api.cloudflare.com/client/v4/accounts/${CLOUDFLARE_ACCOUNT_ID}/ai/v1" },
  221. }),
  222. package: "@ai-sdk/openai-compatible",
  223. options: {
  224. name: "cloudflare-workers-ai",
  225. baseURL: "https://api.cloudflare.com/client/v4/accounts/${CLOUDFLARE_ACCOUNT_ID}/ai/v1",
  226. },
  227. })
  228. expect(cloudflareURL(result.sdk)).toBe(
  229. "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions",
  230. )
  231. }),
  232. ),
  233. )
  234. it.effect("selects languageModel with the API model ID", () =>
  235. Effect.gen(function* () {
  236. const plugin = yield* Plugin.Service
  237. const aisdk = yield* AISDK.Service
  238. const calls: string[] = []
  239. yield* addPlugin()
  240. const result = yield* aisdk.runLanguage({
  241. model: Model.Info.make({
  242. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("alias")),
  243. modelID: Model.ID.make("@cf/api-model"),
  244. package: "aisdk:test-provider",
  245. }),
  246. sdk: fakeSelectorSdk(calls),
  247. options: {},
  248. })
  249. expect(result.language).toBeDefined()
  250. expect(calls).toEqual(["languageModel:@cf/api-model"])
  251. }),
  252. )
  253. it.effect("does not create an SDK for non OpenAI-compatible packages", () =>
  254. withEnv({ CLOUDFLARE_ACCOUNT_ID: "acct", CLOUDFLARE_API_KEY: "key" }, () =>
  255. Effect.gen(function* () {
  256. const plugin = yield* Plugin.Service
  257. const aisdk = yield* AISDK.Service
  258. yield* addPlugin()
  259. const result = yield* aisdk.runSDK({
  260. model: Model.Info.make({
  261. ...Model.Info.default(Provider.ID.make("cloudflare-workers-ai"), Model.ID.make("@cf/model")),
  262. modelID: Model.ID.make("@cf/model"),
  263. package: "aisdk:@ai-sdk/anthropic",
  264. settings: { baseURL: "https://proxy.example/v1" },
  265. }),
  266. package: "@ai-sdk/anthropic",
  267. options: { name: "cloudflare-workers-ai" },
  268. })
  269. expect(result.sdk).toBeUndefined()
  270. }),
  271. ),
  272. )
  273. })