aisdk.ts 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235
  1. export * as AISDK from "./aisdk"
  2. import { makeLocationNode } from "./effect/app-node"
  3. import type { LanguageModelV3 } from "@ai-sdk/provider"
  4. import { Cause, Context, Effect, Layer, Schema, Scope } from "effect"
  5. import { ModelV2 } from "./model"
  6. import { ProviderV2 } from "./provider"
  7. import { State } from "./state"
  8. type SDK = any
  9. export interface SDKEvent {
  10. readonly model: ModelV2.Info
  11. readonly package: string
  12. readonly options: Record<string, any>
  13. sdk?: SDK
  14. }
  15. export interface LanguageEvent {
  16. readonly model: ModelV2.Info
  17. readonly sdk: SDK
  18. readonly options: Record<string, any>
  19. language?: LanguageModelV3
  20. }
  21. function wrapSSE(res: Response, ms: number, ctl: AbortController) {
  22. if (typeof ms !== "number" || ms <= 0) return res
  23. if (!res.body) return res
  24. if (!res.headers.get("content-type")?.includes("text/event-stream")) return res
  25. const reader = res.body.getReader()
  26. const body = new ReadableStream<Uint8Array>({
  27. async pull(ctrl) {
  28. const part = await new Promise<Awaited<ReturnType<typeof reader.read>>>((resolve, reject) => {
  29. const id = setTimeout(() => {
  30. const err = new Error("SSE read timed out")
  31. ctl.abort(err)
  32. void reader.cancel(err)
  33. reject(err)
  34. }, ms)
  35. reader.read().then(
  36. (part) => {
  37. clearTimeout(id)
  38. resolve(part)
  39. },
  40. (err) => {
  41. clearTimeout(id)
  42. reject(err)
  43. },
  44. )
  45. })
  46. if (part.done) {
  47. ctrl.close()
  48. return
  49. }
  50. ctrl.enqueue(part.value)
  51. },
  52. async cancel(reason) {
  53. ctl.abort(reason)
  54. await reader.cancel(reason)
  55. },
  56. })
  57. return new Response(body, {
  58. headers: new Headers(res.headers),
  59. status: res.status,
  60. statusText: res.statusText,
  61. })
  62. }
  63. function prepareOptions(model: ModelV2.Info, pkg: string) {
  64. const options: Record<string, any> = {
  65. name: model.providerID,
  66. ...(model.api.type === "aisdk" ? (model.api.settings ?? {}) : {}),
  67. ...model.request.body,
  68. }
  69. if (model.api.type === "aisdk" && model.api.url) options.baseURL = model.api.url
  70. const customFetch = options.fetch
  71. const chunkTimeout = options.chunkTimeout
  72. delete options.chunkTimeout
  73. options.fetch = async (input: Parameters<typeof fetch>[0], init?: RequestInit) => {
  74. const opts = { ...(init ?? {}) }
  75. const signals = [
  76. opts.signal,
  77. typeof chunkTimeout === "number" && chunkTimeout > 0 ? new AbortController() : undefined,
  78. options.timeout !== undefined && options.timeout !== null && options.timeout !== false
  79. ? AbortSignal.timeout(options.timeout)
  80. : undefined,
  81. ].filter((item): item is AbortSignal | AbortController => Boolean(item))
  82. const chunkAbortCtl = signals.find((item): item is AbortController => item instanceof AbortController)
  83. const abortSignals = signals.map((item) => (item instanceof AbortController ? item.signal : item))
  84. if (abortSignals.length === 1) opts.signal = abortSignals[0]
  85. if (abortSignals.length > 1) opts.signal = AbortSignal.any(abortSignals)
  86. if (
  87. (pkg === "@ai-sdk/openai" || pkg === "@ai-sdk/azure" || pkg === "@ai-sdk/amazon-bedrock/mantle") &&
  88. opts.body &&
  89. opts.method === "POST"
  90. ) {
  91. const body = JSON.parse(opts.body as string)
  92. if (body.store !== true && Array.isArray(body.input)) {
  93. for (const item of body.input) {
  94. if ("id" in item) delete item.id
  95. }
  96. opts.body = JSON.stringify(body)
  97. }
  98. }
  99. const res = await (typeof customFetch === "function" ? customFetch : fetch)(input, {
  100. ...opts,
  101. timeout: false,
  102. })
  103. if (!chunkAbortCtl || typeof chunkTimeout !== "number") return res
  104. return wrapSSE(res, chunkTimeout, chunkAbortCtl)
  105. }
  106. return options
  107. }
  108. export class InitError extends Schema.TaggedErrorClass<InitError>()("AISDK.InitError", {
  109. providerID: ProviderV2.ID,
  110. cause: Schema.Defect(),
  111. }) {}
  112. function initError(providerID: ProviderV2.ID) {
  113. return Effect.catchCause((cause) => Effect.fail(new InitError({ providerID, cause: Cause.squash(cause) })))
  114. }
  115. export interface Interface {
  116. readonly hook: {
  117. readonly sdk: (
  118. callback: (event: SDKEvent) => Effect.Effect<void> | void,
  119. ) => Effect.Effect<State.Registration, never, Scope.Scope>
  120. readonly language: (
  121. callback: (event: LanguageEvent) => Effect.Effect<void> | void,
  122. ) => Effect.Effect<State.Registration, never, Scope.Scope>
  123. }
  124. readonly runSDK: (event: SDKEvent) => Effect.Effect<SDKEvent>
  125. readonly runLanguage: (event: LanguageEvent) => Effect.Effect<LanguageEvent>
  126. readonly language: (model: ModelV2.Info) => Effect.Effect<LanguageModelV3, InitError>
  127. }
  128. export class Service extends Context.Service<Service, Interface>()("@kirincode/v2/AISDK") {}
  129. export const locationLayer = Layer.effect(
  130. Service,
  131. Effect.gen(function* () {
  132. let sdkHooks: ((event: SDKEvent) => Effect.Effect<void> | void)[] = []
  133. let languageHooks: ((event: LanguageEvent) => Effect.Effect<void> | void)[] = []
  134. const languages = new Map<string, LanguageModelV3>()
  135. const sdks = new Map<string, SDK>()
  136. const register = <Event>(
  137. hooks: () => ((event: Event) => Effect.Effect<void> | void)[],
  138. update: (hooks: ((event: Event) => Effect.Effect<void> | void)[]) => void,
  139. ) =>
  140. Effect.fn("AISDK.hook")(function* (callback: (event: Event) => Effect.Effect<void> | void) {
  141. const scope = yield* Scope.Scope
  142. let active = true
  143. update([...hooks(), callback])
  144. const dispose = Effect.sync(() => {
  145. if (!active) return
  146. active = false
  147. update(hooks().filter((item) => item !== callback))
  148. })
  149. yield* Scope.addFinalizer(scope, dispose)
  150. return { dispose }
  151. })
  152. const run = Effect.fnUntraced(function* <Event>(
  153. hooks: readonly ((event: Event) => Effect.Effect<void> | void)[],
  154. event: Event,
  155. ) {
  156. for (const hook of hooks) {
  157. const result = hook(event)
  158. if (Effect.isEffect(result)) yield* result
  159. }
  160. return event
  161. })
  162. const service = Service.of({
  163. hook: {
  164. sdk: register(
  165. () => sdkHooks,
  166. (next) => (sdkHooks = next),
  167. ),
  168. language: register(
  169. () => languageHooks,
  170. (next) => (languageHooks = next),
  171. ),
  172. },
  173. runSDK: (event) => run(sdkHooks, event),
  174. runLanguage: (event) => run(languageHooks, event),
  175. language: Effect.fn("AISDK.language")(function* (model) {
  176. const key = `${model.providerID}/${model.id}/${model.request.variant ?? "default"}`
  177. const existing = languages.get(key)
  178. if (existing) return existing
  179. if (model.api.type !== "aisdk")
  180. return yield* new InitError({
  181. providerID: model.providerID,
  182. cause: new Error(`Unsupported api ${model.api.type}`),
  183. })
  184. const options = prepareOptions(model, model.api.package)
  185. const sdkKey = JSON.stringify({
  186. providerID: model.providerID,
  187. api: model.api,
  188. options,
  189. })
  190. const sdk =
  191. sdks.get(sdkKey) ??
  192. (yield* service.runSDK({ model, package: model.api.package, options }).pipe(initError(model.providerID))).sdk
  193. if (!sdk)
  194. return yield* new InitError({
  195. providerID: model.providerID,
  196. cause: new Error("No AISDK provider plugin returned an SDK"),
  197. })
  198. sdks.set(sdkKey, sdk)
  199. const result = yield* service.runLanguage({ model, sdk, options }).pipe(initError(model.providerID))
  200. const language = yield* Effect.sync(() => result.language ?? sdk.languageModel(model.api.id)).pipe(
  201. initError(model.providerID),
  202. )
  203. languages.set(key, language)
  204. return language
  205. }),
  206. })
  207. return service
  208. }),
  209. )
  210. export const node = makeLocationNode({ service: Service, layer: locationLayer, deps: [] })