runner.ts 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563
  1. import { isOutputTruncatedError, streamChat } from "../llm-client"
  2. import type { StreamCallbacks } from "../llm-client"
  3. import { isFunctionCallingEnabled, providerUsesTextToolCalls } from "./config"
  4. import { accumulateToolCalls, parseTextToolCalls } from "./tool-call-parser"
  5. import { toOpenAITools } from "./tools-schema"
  6. import type { ToolRegistry } from "./registry"
  7. import type { AgentConfig, AgentMessage, AgentRunCallbacks, AgentRunRecord, ToolCall, ToolCallDelta } from "./types"
  8. import { DEFAULT_MAX_ROUNDS, TOOL_EXECUTE_TIMEOUT_MS } from "./types"
  9. import type { TaskBreakpoint } from "./task-breakpoint"
  10. import {
  11. clearTaskBreakpoint,
  12. createTaskBreakpoint,
  13. saveTaskBreakpoint,
  14. updateBreakpointStage,
  15. } from "./task-breakpoint"
  16. import { getEffectiveMaxContextSize, type ChatMessage } from "../llm-providers"
  17. import { isReasoningDisabled, isReasoningOnlyResponseError, withReasoningDisabled } from "../reasoning-retry"
  18. import { addLlmUsage, mergeLlmUsageSnapshot, type LlmUsage } from "../llm-usage"
  19. import { trimChatMessagesToTokenBudget } from "../chat-request-budget"
  20. import { logReasoningReplay } from "../reasoning-replay-debug"
  21. import { ToolEvidenceLedger } from "./tool-evidence-ledger"
  22. import {
  23. RequiredToolsNotCalledError,
  24. buildRequiredToolNudgeMessage,
  25. missingRequiredToolsOnce,
  26. } from "./required-tools-gate"
  27. import { isToolErrorResult } from "./tool-result"
  28. export class ModelDoesNotSupportToolsError extends Error {
  29. constructor() {
  30. super("当前模型不支持工具调用")
  31. this.name = "ModelDoesNotSupportToolsError"
  32. }
  33. }
  34. function messageContentText(content: AgentMessage["content"]): string {
  35. if (typeof content === "string") return content
  36. return content
  37. .map((block) => (block.type === "text" ? block.text : ""))
  38. .join("")
  39. }
  40. function withToolTimeout<T>(operation: Promise<T>, timeoutMs: number | undefined): Promise<T> {
  41. const resolvedTimeoutMs = timeoutMs ?? TOOL_EXECUTE_TIMEOUT_MS
  42. if (resolvedTimeoutMs <= 0) return operation
  43. return Promise.race([
  44. operation,
  45. new Promise<never>((_, reject) =>
  46. setTimeout(() => reject(new Error("工具执行超时")), resolvedTimeoutMs),
  47. ),
  48. ])
  49. }
  50. export class AgentRunner {
  51. async run(
  52. config: AgentConfig,
  53. registry: ToolRegistry,
  54. messages: AgentMessage[],
  55. callbacks: AgentRunCallbacks,
  56. signal?: AbortSignal,
  57. ): Promise<AgentRunRecord> {
  58. const record: AgentRunRecord = { toolCalls: [], roundsUsed: 0, finalText: "" }
  59. const workingMessages = [...messages]
  60. let finalText = ""
  61. const maxRounds = config.maxRounds || DEFAULT_MAX_ROUNDS
  62. const projectPath = config.projectPath
  63. const taskGoal =
  64. config.taskGoal ||
  65. messageContentText([...messages].reverse().find((m) => m.role === "user")?.content ?? "") ||
  66. "未命名任务"
  67. const taskContract = `## 任务契约\n初始任务目标:${taskGoal.slice(0, 1800)}\n执行过程中不得因历史裁剪丢失该目标;当前用户新要求优先。`
  68. const contractInsertIndex = workingMessages.findIndex((message) => message.role !== "system")
  69. workingMessages.splice(contractInsertIndex < 0 ? workingMessages.length : contractInsertIndex, 0, {
  70. role: "system",
  71. content: taskContract,
  72. })
  73. const evidenceLedger = new ToolEvidenceLedger(config.toolResultContextLimit ?? 6000)
  74. let taskBreakpoint: TaskBreakpoint | null = projectPath
  75. ? createTaskBreakpoint({
  76. taskGoal,
  77. currentStage: "agent_round_1",
  78. })
  79. : null
  80. const persistTaskBreakpoint = async () => {
  81. if (!projectPath || !taskBreakpoint) return
  82. try {
  83. await saveTaskBreakpoint(projectPath, taskBreakpoint)
  84. } catch {
  85. // 断点保存失败不应中断当前 AI 会话
  86. }
  87. }
  88. const clearPersistedBreakpoint = async () => {
  89. if (!projectPath) return
  90. try {
  91. await clearTaskBreakpoint(projectPath)
  92. } catch {
  93. // clearTaskBreakpoint 内部已吞掉错误,这里保持双保险
  94. }
  95. }
  96. if (taskBreakpoint) {
  97. await persistTaskBreakpoint()
  98. }
  99. for (let round = 0; round < maxRounds; round++) {
  100. record.roundsUsed = round + 1
  101. if (signal?.aborted) {
  102. for (const tc of record.toolCalls) {
  103. if (tc.status === "running") {
  104. tc.status = "cancelled"
  105. tc.finishedAt = Date.now()
  106. callbacks.onToolEvent?.({
  107. type: "cancelled",
  108. callId: tc.id,
  109. name: tc.name,
  110. params: tc.params,
  111. timestamp: tc.finishedAt,
  112. })
  113. }
  114. }
  115. callbacks.onError(new Error("操作已取消"))
  116. return record
  117. }
  118. const toolCallDeltas: ToolCallDelta[] = []
  119. let roundText = ""
  120. let roundReasoningContent = ""
  121. let streamError: Error | undefined
  122. let roundUsage: LlmUsage | undefined
  123. const streamCallbacks: StreamCallbacks = {
  124. onToken: (t: string) => {
  125. roundText += t
  126. },
  127. onReasoningToken: (t: string) => {
  128. roundReasoningContent += t
  129. callbacks.onReasoningToken?.(t)
  130. },
  131. onToolCallDelta: (delta: ToolCallDelta) => {
  132. toolCallDeltas.push(delta)
  133. },
  134. onUsage: (usage) => {
  135. roundUsage = mergeLlmUsageSnapshot(roundUsage, usage)
  136. if (roundUsage) callbacks.onUsage?.(roundUsage)
  137. },
  138. onUserMemoryDecision: (decision) => {
  139. if (record.userMemoryDecision === undefined) {
  140. record.userMemoryDecision = decision
  141. callbacks.onUserMemoryDecision?.(decision)
  142. }
  143. },
  144. onDone: () => {
  145. // stream finished
  146. },
  147. onError: (err: Error) => {
  148. streamError = err
  149. },
  150. }
  151. const toolsAllowed = isFunctionCallingEnabled(config.llmConfig) && config.tools.length > 0
  152. let openaiTools = toolsAllowed ? toOpenAITools(config.tools) : undefined
  153. let attemptedToolsFallback = false
  154. const buildRequestOverrides = (baseOverrides = config.requestOverrides) =>
  155. openaiTools
  156. ? { ...baseOverrides, tools: openaiTools as any, toolChoice: "auto" as const }
  157. : baseOverrides
  158. let requestOverrides = buildRequestOverrides()
  159. const isToolUnsupportedError = (err: unknown) => {
  160. const msg = err instanceof Error ? err.message : String(err)
  161. return /function[\s_.-]*call|tool_choice|tools?\s+(?:is|are)\s+not\s+supported|does\s+not\s+support\s+(?:function|tools?)|unsupported\s+(?:function|tools?|tool_choice)|不支持\s*(?:工具|function\s*call|FunctionCall)/i.test(msg)
  162. }
  163. const failToolsUnsupported = () => {
  164. callbacks.onError(new ModelDoesNotSupportToolsError())
  165. return record
  166. }
  167. const streamRound = async () => {
  168. // maxContextSize is already a token count; the remaining quarter of the
  169. // window covers the response and prompt scaffolding.
  170. const effectiveContext = getEffectiveMaxContextSize(config.llmConfig)
  171. const internalBudget = Math.max(1, Math.floor(effectiveContext * 0.75))
  172. let compacted: AgentMessage[]
  173. try {
  174. compacted = trimChatMessagesToTokenBudget(
  175. workingMessages as ChatMessage[],
  176. internalBudget,
  177. ) as AgentMessage[]
  178. } catch {
  179. // streamChat retries with a 512-token output floor before giving up;
  180. // surface a readable reason instead of the bare budget error.
  181. throw new Error(
  182. "模型上下文不足:当前对话即使压缩后仍放不下系统提示与最新请求。请缩短输入,或在设置中调高该模型的上下文窗口。",
  183. )
  184. }
  185. workingMessages.splice(0, workingMessages.length, ...compacted)
  186. await streamChat(
  187. config.llmConfig,
  188. workingMessages as ChatMessage[],
  189. streamCallbacks,
  190. signal,
  191. requestOverrides,
  192. )
  193. }
  194. const retryWithoutTools = async () => {
  195. attemptedToolsFallback = true
  196. openaiTools = undefined
  197. roundText = ""
  198. roundReasoningContent = ""
  199. toolCallDeltas.length = 0
  200. streamError = undefined
  201. roundUsage = undefined
  202. requestOverrides = buildRequestOverrides(config.requestOverrides)
  203. await streamRound()
  204. }
  205. try {
  206. await streamRound()
  207. } catch (err) {
  208. if (openaiTools && isToolUnsupportedError(err)) {
  209. try {
  210. await retryWithoutTools()
  211. } catch {
  212. return failToolsUnsupported()
  213. }
  214. } else {
  215. callbacks.onError(err instanceof Error ? err : new Error(String(err)))
  216. return record
  217. }
  218. }
  219. if (
  220. streamError &&
  221. openaiTools &&
  222. isToolUnsupportedError(streamError)
  223. ) {
  224. try {
  225. await retryWithoutTools()
  226. } catch {
  227. return failToolsUnsupported()
  228. }
  229. }
  230. if (
  231. streamError &&
  232. isReasoningOnlyResponseError(streamError) &&
  233. !isReasoningDisabled(config.llmConfig, requestOverrides)
  234. ) {
  235. roundText = ""
  236. roundReasoningContent = ""
  237. toolCallDeltas.length = 0
  238. streamError = undefined
  239. roundUsage = undefined
  240. requestOverrides = buildRequestOverrides(withReasoningDisabled(config.requestOverrides))
  241. try {
  242. await streamRound()
  243. } catch (err) {
  244. callbacks.onError(err instanceof Error ? err : new Error(String(err)))
  245. return record
  246. }
  247. }
  248. if (roundUsage) {
  249. record.lastRequestUsage = { ...roundUsage }
  250. record.usage = addLlmUsage(record.usage, roundUsage)
  251. }
  252. if (streamError) {
  253. if (attemptedToolsFallback) {
  254. return failToolsUnsupported()
  255. }
  256. // Token-limit truncation: keep the partial round text so callers
  257. // can show it and offer continuation, instead of dropping the
  258. // whole round on the floor.
  259. if (
  260. isOutputTruncatedError(streamError) &&
  261. toolCallDeltas.length === 0 &&
  262. roundText.trim()
  263. ) {
  264. finalText = roundText
  265. record.finalText = finalText
  266. callbacks.onText(roundText)
  267. }
  268. callbacks.onError(streamError)
  269. return record
  270. }
  271. // Check for tool calls (native deltas, or text JSON for cursor-cli bridge)
  272. let toolCalls = accumulateToolCalls(toolCallDeltas)
  273. if (
  274. toolCalls.length === 0 &&
  275. openaiTools &&
  276. providerUsesTextToolCalls(config.llmConfig.provider)
  277. ) {
  278. const parsed = parseTextToolCalls(
  279. roundText,
  280. new Set(config.tools.map((tool) => tool.name)),
  281. )
  282. if (parsed.toolCalls.length > 0) {
  283. toolCalls = parsed.toolCalls
  284. roundText = parsed.residualText
  285. }
  286. }
  287. if (toolCalls.length === 0) {
  288. const missingRequired = missingRequiredToolsOnce({
  289. requiredToolsOnce: config.requiredToolsOnce,
  290. availableToolNames: config.tools.map((tool) => tool.name),
  291. calledToolNames: record.toolCalls.map((call) => call.name),
  292. toolsEnabled: Boolean(openaiTools),
  293. })
  294. if (missingRequired.length > 0) {
  295. if (roundText.trim() || roundReasoningContent) {
  296. workingMessages.push({
  297. role: "assistant",
  298. content: roundText || "",
  299. reasoning_content: roundReasoningContent,
  300. })
  301. }
  302. workingMessages.push({
  303. role: "system",
  304. content: buildRequiredToolNudgeMessage(missingRequired),
  305. })
  306. const isLastRound = round >= maxRounds - 1
  307. if (isLastRound) {
  308. await clearPersistedBreakpoint()
  309. callbacks.onError(new RequiredToolsNotCalledError(missingRequired))
  310. return record
  311. }
  312. continue
  313. }
  314. finalText = roundText
  315. record.finalText = finalText
  316. if (roundText) callbacks.onText(roundText)
  317. await clearPersistedBreakpoint()
  318. callbacks.onDone()
  319. return record
  320. }
  321. // Add assistant message with tool calls.
  322. // DeepSeek/Kimi thinking mode requires reasoning_content on every
  323. // tool-call assistant message in subsequent rounds — even "".
  324. const assistantMsg: AgentMessage = {
  325. role: "assistant",
  326. content: roundText || "",
  327. tool_calls: toolCalls,
  328. reasoning_content: roundReasoningContent,
  329. }
  330. logReasoningReplay("agent.round.tool_assistant", {
  331. round: round + 1,
  332. contentLen: (roundText || "").length,
  333. reasoningLen: roundReasoningContent.length,
  334. toolNames: toolCalls.map((call) => call.function.name),
  335. workingMessageCount: workingMessages.length + 1,
  336. })
  337. workingMessages.push(assistantMsg)
  338. // Execute each tool call
  339. for (const tc of toolCalls) {
  340. const toolName = tc.function.name
  341. const tool = registry.get(toolName)
  342. const saveToolProgress = async () => {
  343. if (!taskBreakpoint) return
  344. const usedTools = taskBreakpoint.usedTools.includes(toolName)
  345. ? taskBreakpoint.usedTools
  346. : [...taskBreakpoint.usedTools, toolName]
  347. taskBreakpoint = updateBreakpointStage(
  348. { ...taskBreakpoint, usedTools },
  349. `agent_round_${round + 1}`,
  350. `tool:${toolName}`,
  351. )
  352. await persistTaskBreakpoint()
  353. }
  354. const params = (() => {
  355. try { return JSON.parse(tc.function.arguments || "{}") }
  356. catch { return {} }
  357. })()
  358. const toolCallRecord: AgentRunRecord["toolCalls"][number] = {
  359. id: tc.id,
  360. name: toolName,
  361. params,
  362. result: "",
  363. status: "running",
  364. startedAt: Date.now(),
  365. finishedAt: Date.now(),
  366. }
  367. const callbackToolCall: ToolCall = { id: tc.id, name: toolName, arguments: params }
  368. const executionContext = {
  369. callId: tc.id,
  370. toolName,
  371. onToolEvent: callbacks.onToolEvent,
  372. onActivityEvent: callbacks.onActivityEvent,
  373. }
  374. callbacks.onToolCall(callbackToolCall)
  375. callbacks.onToolEvent?.({
  376. type: "call_started",
  377. callId: tc.id,
  378. name: toolName,
  379. params,
  380. timestamp: toolCallRecord.startedAt,
  381. })
  382. if (!tool) {
  383. const errorMsg = `错误: 未知工具 ${toolName}`
  384. callbacks.onToolError(tc.id, errorMsg)
  385. toolCallRecord.status = "error"
  386. toolCallRecord.result = errorMsg
  387. toolCallRecord.finishedAt = Date.now()
  388. record.toolCalls.push(toolCallRecord)
  389. callbacks.onToolEvent?.({
  390. type: "error",
  391. callId: tc.id,
  392. name: toolName,
  393. params,
  394. result: errorMsg,
  395. timestamp: toolCallRecord.finishedAt,
  396. })
  397. workingMessages.push({
  398. role: "tool",
  399. content: evidenceLedger.format(toolName, params, toolCallRecord.result),
  400. tool_call_id: tc.id,
  401. name: toolName,
  402. })
  403. await saveToolProgress()
  404. continue
  405. }
  406. const permission = tool.permission ?? (tool.category === "write" ? "confirm" : "auto")
  407. if (permission === "confirm") {
  408. let preview = ""
  409. try {
  410. const previewFn = tool.generatePreview ?? tool.execute
  411. preview = await withToolTimeout(previewFn(params, signal, executionContext), tool.executeTimeoutMs)
  412. } catch (e) {
  413. const errorMsg = `预览生成失败:${e instanceof Error ? e.message : String(e)}`
  414. toolCallRecord.status = "error"
  415. toolCallRecord.result = errorMsg
  416. toolCallRecord.finishedAt = Date.now()
  417. record.toolCalls.push(toolCallRecord)
  418. callbacks.onToolError(tc.id, errorMsg)
  419. callbacks.onToolEvent?.({
  420. type: "error",
  421. callId: tc.id,
  422. name: toolName,
  423. params,
  424. result: errorMsg,
  425. timestamp: toolCallRecord.finishedAt,
  426. })
  427. workingMessages.push({
  428. role: "tool",
  429. content: evidenceLedger.format(toolName, params, errorMsg),
  430. tool_call_id: tc.id,
  431. name: toolName,
  432. })
  433. await saveToolProgress()
  434. continue
  435. }
  436. toolCallRecord.status = "approval_required"
  437. ;(toolCallRecord as any).preview = preview
  438. toolCallRecord.result = preview
  439. toolCallRecord.finishedAt = Date.now()
  440. record.toolCalls.push(toolCallRecord)
  441. callbacks.onToolEvent?.({
  442. type: "approval_required",
  443. callId: tc.id,
  444. name: toolName,
  445. params,
  446. result: preview,
  447. preview,
  448. timestamp: toolCallRecord.finishedAt,
  449. })
  450. workingMessages.push({
  451. role: "tool",
  452. content: evidenceLedger.format(toolName, params, preview),
  453. tool_call_id: tc.id,
  454. name: toolName,
  455. })
  456. await saveToolProgress()
  457. continue
  458. }
  459. try {
  460. const result = await withToolTimeout(tool.execute(params, signal, executionContext), tool.executeTimeoutMs)
  461. toolCallRecord.result = result
  462. toolCallRecord.finishedAt = Date.now()
  463. if (isToolErrorResult(result)) {
  464. toolCallRecord.status = "error"
  465. callbacks.onToolError(tc.id, result)
  466. callbacks.onToolEvent?.({
  467. type: "error",
  468. callId: tc.id,
  469. name: toolName,
  470. params,
  471. result,
  472. timestamp: toolCallRecord.finishedAt,
  473. })
  474. } else {
  475. toolCallRecord.status = "done"
  476. callbacks.onToolResult(tc.id, result)
  477. callbacks.onToolEvent?.({
  478. type: "result",
  479. callId: tc.id,
  480. name: toolName,
  481. params,
  482. result,
  483. timestamp: toolCallRecord.finishedAt,
  484. })
  485. }
  486. } catch (err) {
  487. toolCallRecord.status = "error"
  488. toolCallRecord.result = `错误: ${err instanceof Error ? err.message : String(err)}`
  489. toolCallRecord.finishedAt = Date.now()
  490. callbacks.onToolError(tc.id, toolCallRecord.result)
  491. callbacks.onToolEvent?.({
  492. type: "error",
  493. callId: tc.id,
  494. name: toolName,
  495. params,
  496. result: toolCallRecord.result,
  497. timestamp: toolCallRecord.finishedAt,
  498. })
  499. }
  500. record.toolCalls.push(toolCallRecord)
  501. await saveToolProgress()
  502. workingMessages.push({
  503. role: "tool",
  504. content: evidenceLedger.format(toolName, params, toolCallRecord.result),
  505. tool_call_id: tc.id,
  506. name: toolName,
  507. })
  508. }
  509. // Continue loop
  510. if (signal?.aborted) {
  511. for (const tc of record.toolCalls) {
  512. if (tc.status === "running") {
  513. tc.status = "cancelled"
  514. tc.finishedAt = Date.now()
  515. callbacks.onToolEvent?.({
  516. type: "cancelled",
  517. callId: tc.id,
  518. name: tc.name,
  519. params: tc.params,
  520. timestamp: tc.finishedAt,
  521. })
  522. }
  523. }
  524. callbacks.onError(new Error("操作已取消"))
  525. return record
  526. }
  527. }
  528. // Exceeded max rounds
  529. callbacks.onError(new Error(`Agent 已达到最大调用轮次(${maxRounds}),请尝试减少引用内容或拆分任务`))
  530. return record
  531. }
  532. }