llm.ts 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208
  1. import { Provider } from "@/provider/provider"
  2. import { Log } from "@/util/log"
  3. import {
  4. streamText,
  5. wrapLanguageModel,
  6. type ModelMessage,
  7. type StreamTextResult,
  8. type Tool,
  9. type ToolSet,
  10. extractReasoningMiddleware,
  11. } from "ai"
  12. import { clone, mergeDeep, pipe } from "remeda"
  13. import { ProviderTransform } from "@/provider/transform"
  14. import { Config } from "@/config/config"
  15. import { Instance } from "@/project/instance"
  16. import type { Agent } from "@/agent/agent"
  17. import type { MessageV2 } from "./message-v2"
  18. import { Plugin } from "@/plugin"
  19. import { SystemPrompt } from "./system"
  20. import { Flag } from "@/flag/flag"
  21. import { PermissionNext } from "@/permission/next"
  22. export namespace LLM {
  23. const log = Log.create({ service: "llm" })
  24. export const OUTPUT_TOKEN_MAX = Flag.OPENCODE_EXPERIMENTAL_OUTPUT_TOKEN_MAX || 32_000
  25. export type StreamInput = {
  26. user: MessageV2.User
  27. sessionID: string
  28. model: Provider.Model
  29. agent: Agent.Info
  30. system: string[]
  31. abort: AbortSignal
  32. messages: ModelMessage[]
  33. small?: boolean
  34. tools: Record<string, Tool>
  35. retries?: number
  36. }
  37. export type StreamOutput = StreamTextResult<ToolSet, unknown>
  38. export async function stream(input: StreamInput) {
  39. const l = log
  40. .clone()
  41. .tag("providerID", input.model.providerID)
  42. .tag("modelID", input.model.id)
  43. .tag("sessionID", input.sessionID)
  44. .tag("small", (input.small ?? false).toString())
  45. .tag("agent", input.agent.name)
  46. l.info("stream", {
  47. modelID: input.model.id,
  48. providerID: input.model.providerID,
  49. })
  50. const [language, cfg] = await Promise.all([Provider.getLanguage(input.model), Config.get()])
  51. const system = SystemPrompt.header(input.model.providerID)
  52. system.push(
  53. [
  54. // use agent prompt otherwise provider prompt
  55. ...(input.agent.prompt ? [input.agent.prompt] : SystemPrompt.provider(input.model)),
  56. // any custom prompt passed into this call
  57. ...input.system,
  58. // any custom prompt from last user message
  59. ...(input.user.system ? [input.user.system] : []),
  60. ]
  61. .filter((x) => x)
  62. .join("\n"),
  63. )
  64. const header = system[0]
  65. const original = clone(system)
  66. await Plugin.trigger("experimental.chat.system.transform", {}, { system })
  67. if (system.length === 0) {
  68. system.push(...original)
  69. }
  70. // rejoin to maintain 2-part structure for caching if header unchanged
  71. if (system.length > 2 && system[0] === header) {
  72. const rest = system.slice(1)
  73. system.length = 0
  74. system.push(header, rest.join("\n"))
  75. }
  76. const provider = await Provider.getProvider(input.model.providerID)
  77. const variant =
  78. !input.small && input.model.variants && input.user.variant ? input.model.variants[input.user.variant] : {}
  79. const base = input.small
  80. ? ProviderTransform.smallOptions(input.model)
  81. : ProviderTransform.options(input.model, input.sessionID, provider.options)
  82. const options = pipe(base, mergeDeep(input.model.options), mergeDeep(input.agent.options), mergeDeep(variant))
  83. const params = await Plugin.trigger(
  84. "chat.params",
  85. {
  86. sessionID: input.sessionID,
  87. agent: input.agent,
  88. model: input.model,
  89. provider: Provider.getProvider(input.model.providerID),
  90. message: input.user,
  91. },
  92. {
  93. temperature: input.model.capabilities.temperature
  94. ? (input.agent.temperature ?? ProviderTransform.temperature(input.model))
  95. : undefined,
  96. topP: input.agent.topP ?? ProviderTransform.topP(input.model),
  97. topK: ProviderTransform.topK(input.model),
  98. options,
  99. },
  100. )
  101. l.info("params", {
  102. params,
  103. })
  104. const maxOutputTokens = ProviderTransform.maxOutputTokens(
  105. input.model.api.npm,
  106. params.options,
  107. input.model.limit.output,
  108. OUTPUT_TOKEN_MAX,
  109. )
  110. const tools = await resolveTools(input)
  111. return streamText({
  112. onError(error) {
  113. l.error("stream error", {
  114. error,
  115. })
  116. },
  117. async experimental_repairToolCall(failed) {
  118. const lower = failed.toolCall.toolName.toLowerCase()
  119. if (lower !== failed.toolCall.toolName && tools[lower]) {
  120. l.info("repairing tool call", {
  121. tool: failed.toolCall.toolName,
  122. repaired: lower,
  123. })
  124. return {
  125. ...failed.toolCall,
  126. toolName: lower,
  127. }
  128. }
  129. return {
  130. ...failed.toolCall,
  131. input: JSON.stringify({
  132. tool: failed.toolCall.toolName,
  133. error: failed.error.message,
  134. }),
  135. toolName: "invalid",
  136. }
  137. },
  138. temperature: params.temperature,
  139. topP: params.topP,
  140. topK: params.topK,
  141. providerOptions: ProviderTransform.providerOptions(input.model, params.options),
  142. activeTools: Object.keys(tools).filter((x) => x !== "invalid"),
  143. tools,
  144. maxOutputTokens,
  145. abortSignal: input.abort,
  146. headers: {
  147. ...(input.model.providerID.startsWith("opencode")
  148. ? {
  149. "x-opencode-project": Instance.project.id,
  150. "x-opencode-session": input.sessionID,
  151. "x-opencode-request": input.user.id,
  152. "x-opencode-client": Flag.OPENCODE_CLIENT,
  153. }
  154. : undefined),
  155. ...input.model.headers,
  156. },
  157. maxRetries: input.retries ?? 0,
  158. messages: [
  159. ...system.map(
  160. (x): ModelMessage => ({
  161. role: "system",
  162. content: x,
  163. }),
  164. ),
  165. ...input.messages,
  166. ],
  167. model: wrapLanguageModel({
  168. model: language,
  169. middleware: [
  170. {
  171. async transformParams(args) {
  172. if (args.type === "stream") {
  173. // @ts-expect-error
  174. args.params.prompt = ProviderTransform.message(args.params.prompt, input.model)
  175. }
  176. return args.params
  177. },
  178. },
  179. extractReasoningMiddleware({ tagName: "think", startWithReasoning: false }),
  180. ],
  181. }),
  182. experimental_telemetry: { isEnabled: cfg.experimental?.openTelemetry },
  183. })
  184. }
  185. async function resolveTools(input: Pick<StreamInput, "tools" | "agent" | "user">) {
  186. const disabled = PermissionNext.disabled(Object.keys(input.tools), input.agent.permission)
  187. for (const tool of Object.keys(input.tools)) {
  188. if (input.user.tools?.[tool] === false || disabled.has(tool)) {
  189. delete input.tools[tool]
  190. }
  191. }
  192. return input.tools
  193. }
  194. }