provider.ts 8.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323
  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 { LspDiagnosticTool } from "../tool/lsp-diagnostics"
  15. import { LspHoverTool } from "../tool/lsp-hover"
  16. import { PatchTool } from "../tool/patch"
  17. import { ReadTool } from "../tool/read"
  18. import type { Tool } from "../tool/tool"
  19. import { WriteTool } from "../tool/write"
  20. import { TodoReadTool, TodoWriteTool } from "../tool/todo"
  21. import { AuthAnthropic } from "../auth/anthropic"
  22. import { ModelsDev } from "./models"
  23. import { NamedError } from "../util/error"
  24. import { Auth } from "../auth"
  25. import { TaskTool } from "../tool/task"
  26. export namespace Provider {
  27. const log = Log.create({ service: "provider" })
  28. type CustomLoader = (
  29. provider: ModelsDev.Provider,
  30. ) => Promise<Record<string, any> | false>
  31. type Source = "env" | "config" | "custom" | "api"
  32. const CUSTOM_LOADERS: Record<string, CustomLoader> = {
  33. async anthropic(provider) {
  34. const access = await AuthAnthropic.access()
  35. if (!access) return false
  36. for (const model of Object.values(provider.models)) {
  37. model.cost = {
  38. input: 0,
  39. inputCached: 0,
  40. output: 0,
  41. outputCached: 0,
  42. }
  43. }
  44. return {
  45. apiKey: "",
  46. headers: {
  47. authorization: `Bearer ${access}`,
  48. "anthropic-beta": "oauth-2025-04-20",
  49. },
  50. }
  51. },
  52. "amazon-bedrock": async () => {
  53. if (!process.env["AWS_PROFILE"]) return false
  54. const { fromNodeProviderChain } = await import(
  55. await BunProc.install("@aws-sdk/credential-providers")
  56. )
  57. return {
  58. region: process.env["AWS_REGION"] ?? "us-east-1",
  59. credentialProvider: fromNodeProviderChain(),
  60. }
  61. },
  62. }
  63. const state = App.state("provider", async () => {
  64. const config = await Config.get()
  65. const database = await ModelsDev.get()
  66. const providers: {
  67. [providerID: string]: {
  68. source: Source
  69. info: ModelsDev.Provider
  70. options: Record<string, any>
  71. }
  72. } = {}
  73. const models = new Map<
  74. string,
  75. { info: ModelsDev.Model; language: LanguageModel }
  76. >()
  77. const sdk = new Map<string, SDK>()
  78. log.info("init")
  79. function mergeProvider(
  80. id: string,
  81. options: Record<string, any>,
  82. source: Source,
  83. ) {
  84. const provider = providers[id]
  85. if (!provider) {
  86. providers[id] = {
  87. source,
  88. info: database[id],
  89. options,
  90. }
  91. return
  92. }
  93. provider.options = mergeDeep(provider.options, options)
  94. provider.source = source
  95. }
  96. for (const [providerID, provider] of Object.entries(
  97. config.provider ?? {},
  98. )) {
  99. const existing = database[providerID]
  100. const parsed: ModelsDev.Provider = {
  101. id: providerID,
  102. name: provider.name ?? existing?.name ?? providerID,
  103. env: provider.env ?? existing?.env ?? [],
  104. models: existing?.models ?? {},
  105. }
  106. for (const [modelID, model] of Object.entries(provider.models ?? {})) {
  107. const existing = parsed.models[modelID]
  108. const parsedModel: ModelsDev.Model = {
  109. id: modelID,
  110. name: model.name ?? existing?.name ?? modelID,
  111. attachment: model.attachment ?? existing?.attachment ?? false,
  112. reasoning: model.reasoning ?? existing?.reasoning ?? false,
  113. temperature: model.temperature ?? existing?.temperature ?? false,
  114. cost: model.cost ??
  115. existing?.cost ?? {
  116. input: 0,
  117. output: 0,
  118. inputCached: 0,
  119. outputCached: 0,
  120. },
  121. limit: model.limit ??
  122. existing?.limit ?? {
  123. context: 0,
  124. output: 0,
  125. },
  126. }
  127. parsed.models[modelID] = parsedModel
  128. }
  129. database[providerID] = parsed
  130. }
  131. // load env
  132. for (const [providerID, provider] of Object.entries(database)) {
  133. if (provider.env.some((item) => process.env[item])) {
  134. mergeProvider(providerID, {}, "env")
  135. }
  136. }
  137. // load apikeys
  138. for (const [providerID, provider] of Object.entries(await Auth.all())) {
  139. if (provider.type === "api") {
  140. mergeProvider(providerID, { apiKey: provider.key }, "api")
  141. }
  142. }
  143. // load custom
  144. for (const [providerID, fn] of Object.entries(CUSTOM_LOADERS)) {
  145. const result = await fn(database[providerID])
  146. if (result) mergeProvider(providerID, result, "custom")
  147. }
  148. // load config
  149. for (const [providerID, provider] of Object.entries(
  150. config.provider ?? {},
  151. )) {
  152. mergeProvider(providerID, provider.options ?? {}, "config")
  153. }
  154. for (const providerID of Object.keys(providers)) {
  155. log.info("found", { providerID })
  156. }
  157. return {
  158. models,
  159. providers,
  160. sdk,
  161. }
  162. })
  163. export async function list() {
  164. return state().then((state) => state.providers)
  165. }
  166. async function getSDK(providerID: string) {
  167. return (async () => {
  168. using _ = log.time("getSDK", {
  169. providerID,
  170. })
  171. const s = await state()
  172. const existing = s.sdk.get(providerID)
  173. if (existing) return existing
  174. const [pkg, version] = await ModelsDev.pkg(providerID)
  175. const mod = await import(await BunProc.install(pkg, version))
  176. const fn = mod[Object.keys(mod).find((key) => key.startsWith("create"))!]
  177. const loaded = fn(s.providers[providerID]?.options)
  178. s.sdk.set(providerID, loaded)
  179. return loaded as SDK
  180. })().catch((e) => {
  181. throw new InitError({ providerID: providerID }, { cause: e })
  182. })
  183. }
  184. export async function getModel(providerID: string, modelID: string) {
  185. const key = `${providerID}/${modelID}`
  186. const s = await state()
  187. if (s.models.has(key)) return s.models.get(key)!
  188. log.info("getModel", {
  189. providerID,
  190. modelID,
  191. })
  192. const provider = s.providers[providerID]
  193. if (!provider) throw new ModelNotFoundError({ providerID, modelID })
  194. const info = provider.info.models[modelID]
  195. if (!info) throw new ModelNotFoundError({ providerID, modelID })
  196. const sdk = await getSDK(providerID)
  197. try {
  198. const language =
  199. // @ts-expect-error
  200. "responses" in sdk ? sdk.responses(modelID) : sdk.languageModel(modelID)
  201. log.info("found", { providerID, modelID })
  202. s.models.set(key, {
  203. info,
  204. language,
  205. })
  206. return {
  207. info,
  208. language,
  209. }
  210. } catch (e) {
  211. if (e instanceof NoSuchModelError)
  212. throw new ModelNotFoundError(
  213. {
  214. modelID: modelID,
  215. providerID,
  216. },
  217. { cause: e },
  218. )
  219. throw e
  220. }
  221. }
  222. const priority = ["gemini-2.5-pro-preview", "codex-mini", "claude-sonnet-4"]
  223. export function sort(models: ModelsDev.Model[]) {
  224. return sortBy(
  225. models,
  226. [
  227. (model) => priority.findIndex((filter) => model.id.includes(filter)),
  228. "desc",
  229. ],
  230. [(model) => (model.id.includes("latest") ? 0 : 1), "asc"],
  231. [(model) => model.id, "desc"],
  232. )
  233. }
  234. export async function defaultModel() {
  235. const [provider] = await list().then((val) => Object.values(val))
  236. if (!provider) throw new Error("no providers found")
  237. const [model] = sort(Object.values(provider.info.models))
  238. if (!model) throw new Error("no models found")
  239. return {
  240. providerID: provider.info.id,
  241. modelID: model.id,
  242. }
  243. }
  244. const TOOLS = [
  245. EditTool,
  246. WebFetchTool,
  247. GlobTool,
  248. GrepTool,
  249. ListTool,
  250. LspDiagnosticTool,
  251. LspHoverTool,
  252. PatchTool,
  253. ReadTool,
  254. EditTool,
  255. // MultiEditTool,
  256. WriteTool,
  257. TaskTool,
  258. ]
  259. const TOOL_MAPPING: Record<string, Tool.Info[]> = {
  260. anthropic: TOOLS.filter((t) => t.id !== "opencode.patch"),
  261. openai: TOOLS,
  262. google: TOOLS,
  263. }
  264. export async function tools(providerID: string) {
  265. /*
  266. const cfg = await Config.get()
  267. if (cfg.tool?.provider?.[providerID])
  268. return cfg.tool.provider[providerID].map(
  269. (id) => TOOLS.find((t) => t.id === id)!,
  270. )
  271. */
  272. return TOOL_MAPPING[providerID] ?? TOOLS
  273. }
  274. export const ModelNotFoundError = NamedError.create(
  275. "ProviderModelNotFoundError",
  276. z.object({
  277. providerID: z.string(),
  278. modelID: z.string(),
  279. }),
  280. )
  281. export const InitError = NamedError.create(
  282. "ProviderInitError",
  283. z.object({
  284. providerID: z.string(),
  285. }),
  286. )
  287. export const AuthError = NamedError.create(
  288. "ProviderAuthError",
  289. z.object({
  290. providerID: z.string(),
  291. message: z.string(),
  292. }),
  293. )
  294. }