Просмотр исходного кода

fix(agent): 支持 Cursor CLI 文本协议工具调用

cursor-api-proxy 无法返回原生 tool_calls,改为解析模型 JSON 文本,
解除写作 Agent 调度模型被误拦的问题。

Co-authored-by: Cursor <cursoragent@cursor.com>
darknessomi 2 месяцев назад
Родитель
Сommit
dffdb03ea2

+ 10 - 3
src/lib/agent/config.ts

@@ -11,15 +11,18 @@ export const TOOL_UNSUPPORTED_MODEL_PREFIXES: string[] = [
   "deepseek-reasoner",
   "claude-code",
   "codex-cli",
-  "cursor-cli",
 ]
 
 const TOOL_UNSUPPORTED_PROVIDERS = new Set<LlmConfig["provider"]>([
   "claude-code",
   "codex-cli",
-  "cursor-cli",
 ])
 
+/** cursor-api-proxy 无法返回原生 tool_calls delta,需从文本中解析工具调用。 */
+export function providerUsesTextToolCalls(provider: LlmConfig["provider"]): boolean {
+  return provider === "cursor-cli"
+}
+
 export interface BuildAgentConfigOptions extends ToolFactoryOptions {
   llmConfig: LlmConfig
   requestOverrides?: AgentConfig["requestOverrides"]
@@ -51,10 +54,14 @@ export function buildAgentConfig(
   registry.clear()
   registerAllBuiltInTools(registry, options)
 
+  const prompt = providerUsesTextToolCalls(options.llmConfig.provider)
+    ? `${systemPrompt}\n\n当需要调用工具时,请只输出一个 JSON 对象,格式为 {"name":"工具名","arguments":{...}},不要附加其他说明文字。收到工具结果后继续推理;若无需工具则直接回答。`
+    : systemPrompt
+
   return {
     maxRounds: DEFAULT_MAX_ROUNDS,
     tools: registry.list(),
-    systemPrompt,
+    systemPrompt: prompt,
     llmConfig: options.llmConfig,
     modelId,
     projectPath: options.projectPath,

+ 49 - 0
src/lib/agent/runner.spec.ts

@@ -188,6 +188,55 @@ describe("AgentRunner", () => {
     expect(result.roundsUsed).toBe(2)
   })
 
+  it("parses text JSON tool calls for cursor-cli providers", async () => {
+    const tool: Tool = {
+      name: "read_chapter",
+      description: "read",
+      category: "read",
+      parameters: { name: { type: "string", description: "name" } },
+      execute: vi.fn().mockResolvedValue("Chapter content"),
+    }
+    registry.register(tool)
+
+    let callCount = 0
+    mockStreamChat.mockImplementation(async (_config: unknown, _msgs: unknown[], cb: StreamCallbacks) => {
+      callCount++
+      if (callCount === 1) {
+        cb.onToken('{"name":"read_chapter","arguments":{"name":"ch1"}}')
+        cb.onDone()
+      } else {
+        cb.onToken("完成")
+        cb.onDone()
+      }
+    })
+
+    const callbacks = {
+      onText: vi.fn(),
+      onToolCall: vi.fn(),
+      onToolResult: vi.fn(),
+      onToolError: vi.fn(),
+      onDone: vi.fn(),
+      onError: vi.fn(),
+    }
+
+    const config: AgentConfig = {
+      maxRounds: 3,
+      tools: [tool],
+      systemPrompt: "You are helpful",
+      llmConfig: { ...mockLlmConfig, provider: "cursor-cli" },
+    }
+    const result = await runner.run(config, registry, [systemMsg, userMsg], callbacks, undefined)
+
+    expect(tool.execute).toHaveBeenCalledWith(
+      { name: "ch1" },
+      undefined,
+      expect.objectContaining({ toolName: "read_chapter" }),
+    )
+    expect(callbacks.onToolCall).toHaveBeenCalledOnce()
+    expect(result.finalText).toBe("完成")
+    expect(result.roundsUsed).toBe(2)
+  })
+
   it("passes the real tool call id into tool execution context", async () => {
     const tool: Tool = {
       name: "read_chapter",

+ 18 - 3
src/lib/agent/runner.ts

@@ -1,6 +1,7 @@
 import { streamChat } from "../llm-client"
 import type { StreamCallbacks } from "../llm-client"
-import { accumulateToolCalls } from "./tool-call-parser"
+import { providerUsesTextToolCalls } from "./config"
+import { accumulateToolCalls, parseTextToolCalls } from "./tool-call-parser"
 import { toOpenAITools } from "./tools-schema"
 import type { ToolRegistry } from "./registry"
 import type { AgentConfig, AgentMessage, AgentRunCallbacks, AgentRunRecord, ToolCall, ToolCallDelta } from "./types"
@@ -192,8 +193,22 @@ export class AgentRunner {
         return record
       }
 
-      // Check for tool calls
-      const toolCalls = accumulateToolCalls(toolCallDeltas)
+      // Check for tool calls (native deltas, or text JSON for cursor-cli bridge)
+      let toolCalls = accumulateToolCalls(toolCallDeltas)
+      if (
+        toolCalls.length === 0 &&
+        openaiTools &&
+        providerUsesTextToolCalls(config.llmConfig.provider)
+      ) {
+        const parsed = parseTextToolCalls(
+          roundText,
+          new Set(config.tools.map((tool) => tool.name)),
+        )
+        if (parsed.toolCalls.length > 0) {
+          toolCalls = parsed.toolCalls
+          roundText = parsed.residualText
+        }
+      }
 
       if (toolCalls.length === 0) {
         finalText = roundText

+ 48 - 1
src/lib/agent/tool-call-parser.spec.ts

@@ -1,5 +1,5 @@
 import { describe, expect, it } from "vitest"
-import { accumulateToolCalls } from "./tool-call-parser"
+import { accumulateToolCalls, parseTextToolCalls } from "./tool-call-parser"
 import type { ToolCallDelta } from "./types"
 
 describe("accumulateToolCalls", () => {
@@ -49,3 +49,50 @@ describe("accumulateToolCalls", () => {
     expect(result[0].function.arguments).toBe("not json")
   })
 })
+
+describe("parseTextToolCalls", () => {
+  const allowed = new Set(["read_chapter", "route_task"])
+
+  it("parses name/arguments JSON", () => {
+    const result = parseTextToolCalls(
+      '{"name":"read_chapter","arguments":{"name":"第1章"}}',
+      allowed,
+    )
+    expect(result.toolCalls).toHaveLength(1)
+    expect(result.toolCalls[0].function.name).toBe("read_chapter")
+    expect(result.toolCalls[0].function.arguments).toBe('{"name":"第1章"}')
+    expect(result.residualText).toBe("")
+  })
+
+  it("parses fenced JSON and keeps residual prose", () => {
+    const result = parseTextToolCalls(
+      '先查一下章节。\n```json\n{"name":"read_chapter","parameters":{"name":"第1章"}}\n```\n',
+      allowed,
+    )
+    expect(result.toolCalls[0].function.name).toBe("read_chapter")
+    expect(result.residualText).toContain("先查一下章节")
+  })
+
+  it("parses OpenAI-shaped tool_calls wrapper", () => {
+    const result = parseTextToolCalls(
+      JSON.stringify({
+        tool_calls: [
+          {
+            id: "call_1",
+            type: "function",
+            function: { name: "route_task", arguments: '{"intent":"write_chapter"}' },
+          },
+        ],
+      }),
+      allowed,
+    )
+    expect(result.toolCalls[0].id).toBe("call_1")
+    expect(result.toolCalls[0].function.name).toBe("route_task")
+  })
+
+  it("ignores JSON that is not an allowed tool", () => {
+    const result = parseTextToolCalls('{"name":"unknown_tool","arguments":{}}', allowed)
+    expect(result.toolCalls).toEqual([])
+    expect(result.residualText).toContain("unknown_tool")
+  })
+})

+ 150 - 0
src/lib/agent/tool-call-parser.ts

@@ -24,3 +24,153 @@ export function accumulateToolCalls(deltas: ToolCallDelta[]): ToolCall[] {
     }
   })
 }
+
+function extractJsonCandidates(text: string): string[] {
+  const trimmed = text.trim()
+  if (!trimmed) return []
+
+  const candidates: string[] = []
+  const fenced = trimmed.match(/```(?:json)?\s*([\s\S]*?)```/i)
+  if (fenced?.[1]?.trim()) {
+    candidates.push(fenced[1].trim())
+  }
+
+  const firstObj = trimmed.indexOf("{")
+  const lastObj = trimmed.lastIndexOf("}")
+  if (firstObj >= 0 && lastObj > firstObj) {
+    candidates.push(trimmed.slice(firstObj, lastObj + 1))
+  }
+
+  const firstArr = trimmed.indexOf("[")
+  const lastArr = trimmed.lastIndexOf("]")
+  if (firstArr >= 0 && lastArr > firstArr) {
+    candidates.push(trimmed.slice(firstArr, lastArr + 1))
+  }
+
+  if (trimmed.startsWith("{") || trimmed.startsWith("[")) {
+    candidates.unshift(trimmed)
+  }
+
+  return [...new Set(candidates)]
+}
+
+function stringifyArguments(value: unknown): string {
+  if (typeof value === "string") {
+    const trimmed = value.trim()
+    if (!trimmed) return "{}"
+    try {
+      JSON.parse(trimmed)
+      return trimmed
+    } catch {
+      return JSON.stringify(value)
+    }
+  }
+  if (value && typeof value === "object") {
+    return JSON.stringify(value)
+  }
+  return "{}"
+}
+
+function toolCallFromObject(
+  raw: Record<string, unknown>,
+  allowedToolNames: ReadonlySet<string>,
+  index: number,
+): ToolCall | null {
+  const nestedFunction =
+    raw.function && typeof raw.function === "object"
+      ? (raw.function as Record<string, unknown>)
+      : null
+
+  const nameCandidate = [
+    raw.name,
+    raw.tool,
+    raw.tool_name,
+    nestedFunction?.name,
+  ].find((value): value is string => typeof value === "string" && value.trim().length > 0)
+
+  if (!nameCandidate) return null
+  const name = nameCandidate.trim()
+  if (!allowedToolNames.has(name)) return null
+
+  const args =
+    raw.arguments ??
+    raw.parameters ??
+    raw.input ??
+    nestedFunction?.arguments ??
+    nestedFunction?.parameters ??
+    {}
+
+  const id =
+    (typeof raw.id === "string" && raw.id.trim()) ||
+    (typeof nestedFunction?.id === "string" && nestedFunction.id.trim()) ||
+    `text_call_${index}_${Date.now()}`
+
+  return {
+    id,
+    type: "function",
+    function: {
+      name,
+      arguments: stringifyArguments(args),
+    },
+  }
+}
+
+function collectToolCallsFromParsed(
+  parsed: unknown,
+  allowedToolNames: ReadonlySet<string>,
+): ToolCall[] {
+  if (Array.isArray(parsed)) {
+    return parsed.flatMap((item, index) => {
+      if (!item || typeof item !== "object") return []
+      const call = toolCallFromObject(item as Record<string, unknown>, allowedToolNames, index)
+      return call ? [call] : []
+    })
+  }
+
+  if (!parsed || typeof parsed !== "object") return []
+  const obj = parsed as Record<string, unknown>
+
+  if (Array.isArray(obj.tool_calls)) {
+    return collectToolCallsFromParsed(obj.tool_calls, allowedToolNames)
+  }
+
+  if (Array.isArray(obj.tools)) {
+    return collectToolCallsFromParsed(obj.tools, allowedToolNames)
+  }
+
+  const single = toolCallFromObject(obj, allowedToolNames, 0)
+  return single ? [single] : []
+}
+
+/**
+ * cursor-api-proxy 等桥接层只能把 tools schema 注入 prompt,无法返回原生
+ * tool_calls delta。从模型文本里解析 JSON 工具调用。
+ */
+export function parseTextToolCalls(
+  text: string,
+  allowedToolNames: ReadonlySet<string>,
+): { toolCalls: ToolCall[]; residualText: string } {
+  if (!text.trim() || allowedToolNames.size === 0) {
+    return { toolCalls: [], residualText: text }
+  }
+
+  for (const candidate of extractJsonCandidates(text)) {
+    let parsed: unknown
+    try {
+      parsed = JSON.parse(candidate)
+    } catch {
+      continue
+    }
+
+    const toolCalls = collectToolCallsFromParsed(parsed, allowedToolNames)
+    if (toolCalls.length === 0) continue
+
+    const residualText = text.includes(candidate)
+      ? text.replace(candidate, "").replace(/```(?:json)?\s*```/gi, "").trim()
+      : ""
+
+    return { toolCalls, residualText }
+  }
+
+  return { toolCalls: [], residualText: text }
+}

+ 2 - 2
src/lib/cursor-cli-provider.spec.ts

@@ -22,8 +22,8 @@ describe("cursor-cli provider", () => {
     expect(hasUsableLlm(base, { "cursor-cli": { enabled: true } })).toBe(true)
   })
 
-  it("does not support agent tools", () => {
-    expect(modelSupportsTools("composer-2-fast", "cursor-cli")).toBe(false)
+  it("supports agent tools via text tool-call parsing", () => {
+    expect(modelSupportsTools("composer-2-fast", "cursor-cli")).toBe(true)
     expect(modelSupportsTools("gpt-4o", "openai")).toBe(true)
   })