provider.ts 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570
  1. import z from "zod"
  2. import { App } from "../app/app"
  3. import { Config } from "../config/config"
  4. import { mergeDeep, sortBy } from "remeda"
  5. import { NoSuchModelError, type LanguageModel, type Provider as SDK } from "ai"
  6. import { Log } from "../util/log"
  7. import { BunProc } from "../bun"
  8. import { BashTool } from "../tool/bash"
  9. import { EditTool } from "../tool/edit"
  10. import { WebFetchTool } from "../tool/webfetch"
  11. import { GlobTool } from "../tool/glob"
  12. import { GrepTool } from "../tool/grep"
  13. import { ListTool } from "../tool/ls"
  14. import { PatchTool } from "../tool/patch"
  15. import { ReadTool } from "../tool/read"
  16. import type { Tool } from "../tool/tool"
  17. import { WriteTool } from "../tool/write"
  18. import { TodoReadTool, TodoWriteTool } from "../tool/todo"
  19. import { AuthAnthropic } from "../auth/anthropic"
  20. import { AuthCopilot } from "../auth/copilot"
  21. import { ModelsDev } from "./models"
  22. import { NamedError } from "../util/error"
  23. import { Auth } from "../auth"
  24. // import { TaskTool } from "../tool/task"
  25. export namespace Provider {
  26. const log = Log.create({ service: "provider" })
  27. type CustomLoader = (
  28. provider: ModelsDev.Provider,
  29. api?: string,
  30. ) => Promise<{
  31. autoload: boolean
  32. getModel?: (sdk: any, modelID: string) => Promise<any>
  33. options?: Record<string, any>
  34. }>
  35. type Source = "env" | "config" | "custom" | "api"
  36. const CUSTOM_LOADERS: Record<string, CustomLoader> = {
  37. async anthropic(provider) {
  38. const access = await AuthAnthropic.access()
  39. if (!access) return { autoload: false }
  40. for (const model of Object.values(provider.models)) {
  41. model.cost = {
  42. input: 0,
  43. output: 0,
  44. }
  45. }
  46. return {
  47. autoload: true,
  48. options: {
  49. apiKey: "",
  50. async fetch(input: any, init: any) {
  51. const access = await AuthAnthropic.access()
  52. const headers = {
  53. ...init.headers,
  54. authorization: `Bearer ${access}`,
  55. "anthropic-beta": "oauth-2025-04-20",
  56. }
  57. delete headers["x-api-key"]
  58. return fetch(input, {
  59. ...init,
  60. headers,
  61. })
  62. },
  63. },
  64. }
  65. },
  66. "github-copilot": async (provider) => {
  67. const copilot = await AuthCopilot()
  68. if (!copilot) return { autoload: false }
  69. let info = await Auth.get("github-copilot")
  70. if (!info || info.type !== "oauth") return { autoload: false }
  71. if (provider && provider.models) {
  72. for (const model of Object.values(provider.models)) {
  73. model.cost = {
  74. input: 0,
  75. output: 0,
  76. }
  77. }
  78. }
  79. return {
  80. autoload: true,
  81. options: {
  82. apiKey: "",
  83. async fetch(input: any, init: any) {
  84. const info = await Auth.get("github-copilot")
  85. if (!info || info.type !== "oauth") return
  86. if (!info.access || info.expires < Date.now()) {
  87. const tokens = await copilot.access(info.refresh)
  88. if (!tokens)
  89. throw new Error("GitHub Copilot authentication expired")
  90. await Auth.set("github-copilot", {
  91. type: "oauth",
  92. ...tokens,
  93. })
  94. info.access = tokens.access
  95. }
  96. let isAgentCall = false
  97. try {
  98. const body =
  99. typeof init.body === "string"
  100. ? JSON.parse(init.body)
  101. : init.body
  102. if (body?.messages) {
  103. isAgentCall = body.messages.some(
  104. (msg: any) =>
  105. msg.role && ["tool", "assistant"].includes(msg.role),
  106. )
  107. }
  108. } catch {}
  109. const headers = {
  110. ...init.headers,
  111. ...copilot.HEADERS,
  112. Authorization: `Bearer ${info.access}`,
  113. "Openai-Intent": "conversation-edits",
  114. "X-Initiator": isAgentCall ? "agent" : "user",
  115. }
  116. delete headers["x-api-key"]
  117. return fetch(input, {
  118. ...init,
  119. headers,
  120. })
  121. },
  122. },
  123. }
  124. },
  125. openai: async () => {
  126. return {
  127. autoload: false,
  128. async getModel(sdk: any, modelID: string) {
  129. return sdk.responses(modelID)
  130. },
  131. options: {},
  132. }
  133. },
  134. "amazon-bedrock": async () => {
  135. if (!process.env["AWS_PROFILE"] && !process.env["AWS_ACCESS_KEY_ID"])
  136. return { autoload: false }
  137. const region = process.env["AWS_REGION"] ?? "us-east-1"
  138. const { fromNodeProviderChain } = await import(
  139. await BunProc.install("@aws-sdk/credential-providers")
  140. )
  141. return {
  142. autoload: true,
  143. options: {
  144. region,
  145. credentialProvider: fromNodeProviderChain(),
  146. },
  147. async getModel(sdk: any, modelID: string) {
  148. let regionPrefix = region.split("-")[0]
  149. switch (regionPrefix) {
  150. case "us": {
  151. const modelRequiresPrefix = ["claude", "deepseek"].some((m) =>
  152. modelID.includes(m),
  153. )
  154. if (modelRequiresPrefix) {
  155. modelID = `${regionPrefix}.${modelID}`
  156. }
  157. break
  158. }
  159. case "eu": {
  160. const regionRequiresPrefix = [
  161. "eu-west-1",
  162. "eu-west-3",
  163. "eu-north-1",
  164. "eu-central-1",
  165. "eu-south-1",
  166. "eu-south-2",
  167. ].some((r) => region.includes(r))
  168. const modelRequiresPrefix = [
  169. "claude",
  170. "nova-lite",
  171. "nova-micro",
  172. "llama3",
  173. "pixtral",
  174. ].some((m) => modelID.includes(m))
  175. if (regionRequiresPrefix && modelRequiresPrefix) {
  176. modelID = `${regionPrefix}.${modelID}`
  177. }
  178. break
  179. }
  180. case "ap": {
  181. const modelRequiresPrefix = [
  182. "claude",
  183. "nova-lite",
  184. "nova-micro",
  185. "nova-pro",
  186. ].some((m) => modelID.includes(m))
  187. if (modelRequiresPrefix) {
  188. regionPrefix = "apac"
  189. modelID = `${regionPrefix}.${modelID}`
  190. }
  191. break
  192. }
  193. }
  194. return sdk.languageModel(modelID)
  195. },
  196. }
  197. },
  198. openrouter: async (provider) => {
  199. return {
  200. autoload: false,
  201. options: {
  202. headers: {
  203. "HTTP-Referer": "https://opencode.ai/",
  204. "X-Title": "opencode",
  205. },
  206. },
  207. }
  208. },
  209. }
  210. const state = App.state("provider", async () => {
  211. const config = await Config.get()
  212. const database = await ModelsDev.get()
  213. const providers: {
  214. [providerID: string]: {
  215. source: Source
  216. info: ModelsDev.Provider
  217. getModel?: (sdk: any, modelID: string) => Promise<any>
  218. options: Record<string, any>
  219. }
  220. } = {}
  221. const models = new Map<
  222. string,
  223. { info: ModelsDev.Model; language: LanguageModel }
  224. >()
  225. const sdk = new Map<string, SDK>()
  226. log.info("init")
  227. function mergeProvider(
  228. id: string,
  229. options: Record<string, any>,
  230. source: Source,
  231. getModel?: (sdk: any, modelID: string) => Promise<any>,
  232. ) {
  233. const provider = providers[id]
  234. if (!provider) {
  235. const info = database[id]
  236. if (!info) return
  237. if (info.api) options["baseURL"] = info.api
  238. providers[id] = {
  239. source,
  240. info,
  241. options,
  242. getModel,
  243. }
  244. return
  245. }
  246. provider.options = mergeDeep(provider.options, options)
  247. provider.source = source
  248. provider.getModel = getModel ?? provider.getModel
  249. }
  250. const configProviders = Object.entries(config.provider ?? {})
  251. for (const [providerID, provider] of configProviders) {
  252. const existing = database[providerID]
  253. const parsed: ModelsDev.Provider = {
  254. id: providerID,
  255. npm: provider.npm ?? existing?.npm,
  256. name: provider.name ?? existing?.name ?? providerID,
  257. env: provider.env ?? existing?.env ?? [],
  258. api: provider.api ?? existing?.api,
  259. models: existing?.models ?? {},
  260. }
  261. for (const [modelID, model] of Object.entries(provider.models ?? {})) {
  262. const existing = parsed.models[modelID]
  263. const parsedModel: ModelsDev.Model = {
  264. id: modelID,
  265. name: model.name ?? existing?.name ?? modelID,
  266. release_date: model.release_date ?? existing?.release_date,
  267. attachment: model.attachment ?? existing?.attachment ?? false,
  268. reasoning: model.reasoning ?? existing?.reasoning ?? false,
  269. temperature: model.temperature ?? existing?.temperature ?? false,
  270. tool_call: model.tool_call ?? existing?.tool_call ?? true,
  271. cost: {
  272. ...existing?.cost,
  273. ...model.cost,
  274. input: 0,
  275. output: 0,
  276. cache_read: 0,
  277. cache_write: 0,
  278. },
  279. options: {
  280. ...existing?.options,
  281. ...model.options,
  282. },
  283. limit: model.limit ??
  284. existing?.limit ?? {
  285. context: 0,
  286. output: 0,
  287. },
  288. }
  289. parsed.models[modelID] = parsedModel
  290. }
  291. database[providerID] = parsed
  292. }
  293. const disabled = await Config.get().then(
  294. (cfg) => new Set(cfg.disabled_providers ?? []),
  295. )
  296. // load env
  297. for (const [providerID, provider] of Object.entries(database)) {
  298. if (disabled.has(providerID)) continue
  299. const apiKey = provider.env.map((item) => process.env[item]).at(0)
  300. if (!apiKey) continue
  301. mergeProvider(
  302. providerID,
  303. // only include apiKey if there's only one potential option
  304. provider.env.length === 1 ? { apiKey } : {},
  305. "env",
  306. )
  307. }
  308. // load apikeys
  309. for (const [providerID, provider] of Object.entries(await Auth.all())) {
  310. if (disabled.has(providerID)) continue
  311. if (provider.type === "api") {
  312. mergeProvider(providerID, { apiKey: provider.key }, "api")
  313. }
  314. }
  315. // load custom
  316. for (const [providerID, fn] of Object.entries(CUSTOM_LOADERS)) {
  317. if (disabled.has(providerID)) continue
  318. const result = await fn(database[providerID])
  319. if (result && (result.autoload || providers[providerID])) {
  320. mergeProvider(
  321. providerID,
  322. result.options ?? {},
  323. "custom",
  324. result.getModel,
  325. )
  326. }
  327. }
  328. // load config
  329. for (const [providerID, provider] of configProviders) {
  330. mergeProvider(providerID, provider.options ?? {}, "config")
  331. }
  332. for (const [providerID, provider] of Object.entries(providers)) {
  333. if (Object.keys(provider.info.models).length === 0) {
  334. delete providers[providerID]
  335. continue
  336. }
  337. log.info("found", { providerID })
  338. }
  339. return {
  340. models,
  341. providers,
  342. sdk,
  343. }
  344. })
  345. export async function list() {
  346. return state().then((state) => state.providers)
  347. }
  348. async function getSDK(provider: ModelsDev.Provider) {
  349. return (async () => {
  350. using _ = log.time("getSDK", {
  351. providerID: provider.id,
  352. })
  353. const s = await state()
  354. const existing = s.sdk.get(provider.id)
  355. if (existing) return existing
  356. const pkg = provider.npm ?? provider.id
  357. const mod = await import(await BunProc.install(pkg, "latest"))
  358. const fn = mod[Object.keys(mod).find((key) => key.startsWith("create"))!]
  359. const loaded = fn(s.providers[provider.id]?.options)
  360. s.sdk.set(provider.id, loaded)
  361. return loaded as SDK
  362. })().catch((e) => {
  363. throw new InitError({ providerID: provider.id }, { cause: e })
  364. })
  365. }
  366. export async function getModel(providerID: string, modelID: string) {
  367. const key = `${providerID}/${modelID}`
  368. const s = await state()
  369. if (s.models.has(key)) return s.models.get(key)!
  370. log.info("getModel", {
  371. providerID,
  372. modelID,
  373. })
  374. const provider = s.providers[providerID]
  375. if (!provider) throw new ModelNotFoundError({ providerID, modelID })
  376. const info = provider.info.models[modelID]
  377. if (!info) throw new ModelNotFoundError({ providerID, modelID })
  378. const sdk = await getSDK(provider.info)
  379. try {
  380. const language = provider.getModel
  381. ? await provider.getModel(sdk, modelID)
  382. : sdk.languageModel(modelID)
  383. log.info("found", { providerID, modelID })
  384. s.models.set(key, {
  385. info,
  386. language,
  387. })
  388. return {
  389. info,
  390. language,
  391. }
  392. } catch (e) {
  393. if (e instanceof NoSuchModelError)
  394. throw new ModelNotFoundError(
  395. {
  396. modelID: modelID,
  397. providerID,
  398. },
  399. { cause: e },
  400. )
  401. throw e
  402. }
  403. }
  404. const priority = ["gemini-2.5-pro-preview", "codex-mini", "claude-sonnet-4"]
  405. export function sort(models: ModelsDev.Model[]) {
  406. return sortBy(
  407. models,
  408. [
  409. (model) => priority.findIndex((filter) => model.id.includes(filter)),
  410. "desc",
  411. ],
  412. [(model) => (model.id.includes("latest") ? 0 : 1), "asc"],
  413. [(model) => model.id, "desc"],
  414. )
  415. }
  416. export async function defaultModel() {
  417. const cfg = await Config.get()
  418. if (cfg.model) return parseModel(cfg.model)
  419. const provider = await list()
  420. .then((val) => Object.values(val))
  421. .then((x) =>
  422. x.find(
  423. (p) => !cfg.provider || Object.keys(cfg.provider).includes(p.info.id),
  424. ),
  425. )
  426. if (!provider) throw new Error("no providers found")
  427. const [model] = sort(Object.values(provider.info.models))
  428. if (!model) throw new Error("no models found")
  429. return {
  430. providerID: provider.info.id,
  431. modelID: model.id,
  432. }
  433. }
  434. export function parseModel(model: string) {
  435. const [providerID, ...rest] = model.split("/")
  436. return {
  437. providerID: providerID,
  438. modelID: rest.join("/"),
  439. }
  440. }
  441. const TOOLS = [
  442. BashTool,
  443. EditTool,
  444. WebFetchTool,
  445. GlobTool,
  446. GrepTool,
  447. ListTool,
  448. // LspDiagnosticTool,
  449. // LspHoverTool,
  450. PatchTool,
  451. ReadTool,
  452. // MultiEditTool,
  453. WriteTool,
  454. TodoWriteTool,
  455. TodoReadTool,
  456. // TaskTool,
  457. ]
  458. const TOOL_MAPPING: Record<string, Tool.Info[]> = {
  459. anthropic: TOOLS.filter((t) => t.id !== "patch"),
  460. openai: TOOLS.map((t) => ({
  461. ...t,
  462. parameters: optionalToNullable(t.parameters),
  463. })),
  464. azure: TOOLS.map((t) => ({
  465. ...t,
  466. parameters: optionalToNullable(t.parameters),
  467. })),
  468. google: TOOLS,
  469. }
  470. export async function tools(providerID: string) {
  471. /*
  472. const cfg = await Config.get()
  473. if (cfg.tool?.provider?.[providerID])
  474. return cfg.tool.provider[providerID].map(
  475. (id) => TOOLS.find((t) => t.id === id)!,
  476. )
  477. */
  478. return TOOL_MAPPING[providerID] ?? TOOLS
  479. }
  480. function optionalToNullable(schema: z.ZodTypeAny): z.ZodTypeAny {
  481. if (schema instanceof z.ZodObject) {
  482. const shape = schema.shape
  483. const newShape: Record<string, z.ZodTypeAny> = {}
  484. for (const [key, value] of Object.entries(shape)) {
  485. const zodValue = value as z.ZodTypeAny
  486. if (zodValue instanceof z.ZodOptional) {
  487. newShape[key] = zodValue.unwrap().nullable()
  488. } else {
  489. newShape[key] = optionalToNullable(zodValue)
  490. }
  491. }
  492. return z.object(newShape)
  493. }
  494. if (schema instanceof z.ZodArray) {
  495. return z.array(optionalToNullable(schema.element))
  496. }
  497. if (schema instanceof z.ZodUnion) {
  498. return z.union(
  499. schema.options.map((option: z.ZodTypeAny) =>
  500. optionalToNullable(option),
  501. ) as [z.ZodTypeAny, z.ZodTypeAny, ...z.ZodTypeAny[]],
  502. )
  503. }
  504. return schema
  505. }
  506. export const ModelNotFoundError = NamedError.create(
  507. "ProviderModelNotFoundError",
  508. z.object({
  509. providerID: z.string(),
  510. modelID: z.string(),
  511. }),
  512. )
  513. export const InitError = NamedError.create(
  514. "ProviderInitError",
  515. z.object({
  516. providerID: z.string(),
  517. }),
  518. )
  519. export const AuthError = NamedError.create(
  520. "ProviderAuthError",
  521. z.object({
  522. providerID: z.string(),
  523. message: z.string(),
  524. }),
  525. )
  526. }