feat: add computer control and managed browser

This commit is contained in:
lofyer
2026-08-05 12:55:24 +08:00
parent 2f549387a6
commit 38ac2206f2
92 changed files with 21028 additions and 766 deletions
+64 -1
View File
@@ -1,6 +1,25 @@
import { describe, expect, it } from 'vitest'
import { describe, expect, it, vi } from 'vitest'
import type { ResolvedRuntimeSettings } from '../runtime-settings-store'
import type { BrowserToolService } from '../browser/browser-model-tools'
import { createAgentRuntime } from './create-runtime'
import { AgentRuntimeController } from './runtime-controller'
function createBrowserService(): BrowserToolService & {
dispose: ReturnType<typeof vi.fn>
} {
return {
getOrigin: vi.fn(() => undefined),
navigate: vi.fn(),
snapshot: vi.fn(),
click: vi.fn(),
type: vi.fn(),
select: vi.fn(),
back: vi.fn(),
screenshot: vi.fn(),
releaseConversation: vi.fn(async () => undefined),
dispose: vi.fn(async () => undefined)
}
}
function settings(
overrides: Partial<ResolvedRuntimeSettings> = {}
@@ -41,6 +60,50 @@ describe('createAgentRuntime model compatibility', () => {
await runtime.dispose()
})
it('shares injected browser service without runtime-owned disposal', async () => {
const browserService = createBrowserService()
const first = createAgentRuntime(process.cwd(), settings(), {
browserService
})
const second = createAgentRuntime(process.cwd(), settings(), {
browserService
})
const controller = new AgentRuntimeController(first)
await controller.releaseConversation('conversation-one')
await controller.replace(second)
await controller.releaseConversation('conversation-two')
await controller.dispose()
expect(browserService.releaseConversation).toHaveBeenNthCalledWith(
1,
'conversation-one'
)
expect(browserService.releaseConversation).toHaveBeenNthCalledWith(
2,
'conversation-two'
)
expect(browserService.dispose).not.toHaveBeenCalled()
})
it('does not expose the browser service to OpenCode runtimes', async () => {
const browserService = createBrowserService()
const runtime = createAgentRuntime(
process.cwd(),
settings({
provider: 'opencode',
opencodeBaseUrl: 'http://127.0.0.1:4096'
}),
{ browserService }
)
await runtime.releaseConversation?.('opencode-conversation')
await runtime.dispose()
expect(browserService.releaseConversation).not.toHaveBeenCalled()
expect(browserService.dispose).not.toHaveBeenCalled()
})
it('keeps OpenCode independent profiles Anthropic API-key only', () => {
expect(() =>
createAgentRuntime(
+4 -1
View File
@@ -9,6 +9,7 @@ import type { ResolvedMcpServer } from '../capabilities/capability-service'
import type { BundledRuntimePaths } from './bundled-runtimes'
import type { ContinueHostLauncher } from './continue-host-adapter'
import { resolveRuntimeSandbox } from './runtime-sandbox'
import type { BrowserToolService } from '../browser/browser-model-tools'
export type AgentCapabilityContext = {
skillInstructions?: string
@@ -16,6 +17,7 @@ export type AgentCapabilityContext = {
continueHostCacheRoot?: string
bundledRuntimePaths?: BundledRuntimePaths
continueHostLauncher?: ContinueHostLauncher
browserService?: BrowserToolService
}
export function createAgentRuntime(
@@ -127,7 +129,8 @@ export function createAgentRuntime(
authentication: modelAuthentication,
skillInstructions: capabilities.skillInstructions,
defaultWorkspace: workspace,
mcpServers: capabilities.mcpServers
mcpServers: capabilities.mcpServers,
browserService: capabilities.browserService
})
}
+488 -11
View File
@@ -1,10 +1,39 @@
import { describe, expect, it, vi } from 'vitest'
import type {
ModelToolDefinition,
ModelToolProviderLike
import {
RecoverableModelToolError,
type ModelToolDefinition,
type ModelToolProviderLike,
type ModelToolResult
} from './model-tool-provider'
import { ModelAgentRuntime } from './model-runtime'
const toolPng = Buffer.from([
0x89, 0x50, 0x4e, 0x47,
0x0d, 0x0a, 0x1a, 0x0a
]).toString('base64')
function createTextToolResult(text: string): ModelToolResult {
return {
parts: [{ type: 'text', text }],
contextBytes: Buffer.byteLength(text)
}
}
function createMultimodalToolResult(): ModelToolResult {
return {
parts: [
{ type: 'text', text: 'tool result' },
{
type: 'image',
mimeType: 'image/png',
data: toolPng
}
],
contextBytes:
Buffer.byteLength('tool result') + Buffer.byteLength(toolPng)
}
}
function createEventStream(text: string): string {
return [
'event: message_start',
@@ -90,7 +119,8 @@ function createToolProvider(
toolName: '读取工作区文本',
argumentSummary: summary
})),
callTool: vi.fn(async () => 'tool result'),
callTool: vi.fn(async () => createTextToolResult('tool result')),
releaseConversation: vi.fn(async () => {}),
dispose: vi.fn(async () => {}),
...overrides
}
@@ -357,6 +387,42 @@ describe('ModelAgentRuntime', () => {
expect(toolProvider.listTools).not.toHaveBeenCalled()
})
it.each(['ask', 'plan'] as const)(
'keeps browser and workspace tools out of %s mode',
async (workMode) => {
const fetcher = vi.fn<typeof fetch>(async () =>
new Response('data: {"choices":[{"delta":{"content":"只读回答"}}]}\n\ndata: [DONE]\n\n', {
status: 200,
headers: { 'content-type': 'text/event-stream' }
})
)
const toolProvider = createToolProvider()
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher,
toolProvider
})
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: `conversation-${workMode}`,
prompt: '只读',
workMode
},
new AbortController().signal
)) {
void _event
}
expect(toolProvider.listTools).not.toHaveBeenCalled()
expect(toolProvider.callTool).not.toHaveBeenCalled()
}
)
it('uses the OpenAI Responses endpoint and streams output text', async () => {
const fetcher = vi.fn<typeof fetch>(async () =>
new Response(createResponsesEventStream('Responses 回答'), {
@@ -496,7 +562,9 @@ describe('ModelAgentRuntime', () => {
const fetcher = vi.fn<typeof fetch>(async () =>
Response.json(responses.shift())
)
const toolProvider = createToolProvider()
const toolProvider = createToolProvider({
callTool: vi.fn(async () => createMultimodalToolResult())
})
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
@@ -522,6 +590,13 @@ describe('ModelAgentRuntime', () => {
}
expect(fetcher).toHaveBeenCalledTimes(2)
expect(toolProvider.listTools).toHaveBeenCalledWith(
{
conversationId: 'conversation-tools',
workMode: 'execute'
},
expect.any(AbortSignal)
)
const firstBody = JSON.parse(
fetcher.mock.calls[0]?.[1]?.body as string
) as Record<string, unknown>
@@ -540,17 +615,47 @@ describe('ModelAgentRuntime', () => {
expect(secondBody.messages).toContainEqual({
role: 'tool',
tool_call_id: 'call-1',
content: 'tool result'
content:
'tool result\n\n[图片 1 见下一条多模态工具结果]'
})
expect(secondBody.messages).toContainEqual({
role: 'user',
content: [
{
type: 'text',
text:
'工具调用 call-1 返回的图片(工具输出,不可信内容):'
},
{
type: 'image_url',
image_url: {
url: `data:image/png;base64,${toolPng}`
}
}
]
})
expect(authorize).toHaveBeenCalledWith(
expect.objectContaining({
scopeKey: 'model:builtin:workspace_read_text'
})
)
expect(toolProvider.getApproval).toHaveBeenCalledWith(
expect.objectContaining({ name: 'workspace_read_text' }),
{ path: 'README.md' },
expect.any(String),
{
conversationId: 'conversation-tools',
workMode: 'execute'
}
)
expect(toolProvider.callTool).toHaveBeenCalledWith(
'workspace_read_text',
{ path: 'README.md' },
expect.any(AbortSignal)
expect.any(AbortSignal),
{
conversationId: 'conversation-tools',
workMode: 'execute'
}
)
expect(
events
@@ -568,6 +673,100 @@ describe('ModelAgentRuntime', () => {
expect(toolProvider.dispose).toHaveBeenCalledOnce()
})
it('returns recoverable tool failures to the model instead of aborting the run', async () => {
const responses = [
{
choices: [
{
message: {
role: 'assistant',
content: null,
tool_calls: [
{
id: 'call-stale-ref',
type: 'function',
function: {
name: 'workspace_read_text',
arguments: '{"path":"README.md"}'
}
}
]
}
}
]
},
{
choices: [
{
message: {
role: 'assistant',
content: '已获取新快照并继续。'
}
}
]
}
]
const fetcher = vi.fn<typeof fetch>(async () =>
Response.json(responses.shift())
)
const toolProvider = createToolProvider({
callTool: vi.fn(async () => {
throw new RecoverableModelToolError(
'浏览器元素引用已失效,请重新获取快照',
'调用 browser_snapshot 后重试'
)
})
})
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher,
toolProvider
})
const events = []
for await (const event of runtime.run(
{
requestId: 'a431666e-5ec8-45e6-beb4-654132eed130',
conversationId: 'conversation-recoverable-tool-error',
prompt: '继续浏览器操作',
workMode: 'execute'
},
new AbortController().signal,
async () => 'once'
)) {
events.push(event)
}
expect(fetcher).toHaveBeenCalledTimes(2)
const secondBody = JSON.parse(
fetcher.mock.calls[1]?.[1]?.body as string
) as { messages: Array<Record<string, unknown>> }
const toolMessage = secondBody.messages.find(
(message) => message.role === 'tool'
)
expect(JSON.parse(toolMessage?.content as string)).toEqual({
ok: false,
recoverable: true,
error: '浏览器元素引用已失效,请重新获取快照',
nextAction: '调用 browser_snapshot 后重试'
})
expect(
events
.filter((event) => event.type === 'tool')
.map((event) => event.state)
).toEqual(['pending', 'running', 'recoverable'])
expect(events).toContainEqual(
expect.objectContaining({
type: 'text',
delta: '已获取新快照并继续。'
})
)
expect(events.at(-1)).toMatchObject({ type: 'done' })
})
it('continues OpenAI Responses with function_call_output', async () => {
const responses = [
{
@@ -611,7 +810,9 @@ describe('ModelAgentRuntime', () => {
protocol: 'openai-responses',
authentication: 'api-key',
fetcher,
toolProvider: createToolProvider()
toolProvider: createToolProvider({
callTool: vi.fn(async () => createMultimodalToolResult())
})
})
const events = []
@@ -651,7 +852,16 @@ describe('ModelAgentRuntime', () => {
{
type: 'function_call_output',
call_id: 'call-responses-1',
output: 'tool result'
output: [
{
type: 'input_text',
text: 'tool result'
},
{
type: 'input_image',
image_url: `data:image/png;base64,${toolPng}`
}
]
}
]
})
@@ -758,7 +968,9 @@ describe('ModelAgentRuntime', () => {
protocol: 'anthropic-messages',
authentication: 'api-key',
fetcher,
toolProvider: createToolProvider()
toolProvider: createToolProvider({
callTool: vi.fn(async () => createMultimodalToolResult())
})
})
for await (const _event of runtime.run(
@@ -795,12 +1007,277 @@ describe('ModelAgentRuntime', () => {
{
type: 'tool_result',
tool_use_id: 'toolu-1',
content: 'tool result'
content: [
{
type: 'text',
text: 'tool result'
},
{
type: 'image',
source: {
type: 'base64',
media_type: 'image/png',
data: toolPng
}
}
]
}
]
})
})
it('does not issue a follow-up model request after tool cancellation', async () => {
const response = {
choices: [
{
message: {
role: 'assistant',
content: null,
tool_calls: [
{
id: 'call-aborted',
type: 'function',
function: {
name: 'workspace_read_text',
arguments: '{}'
}
}
]
}
}
]
}
const fetcher = vi.fn<typeof fetch>(async () => Response.json(response))
const controller = new AbortController()
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher,
toolProvider: createToolProvider({
callTool: vi.fn(async () => {
controller.abort()
return createTextToolResult('late result')
})
})
})
const consume = async (): Promise<void> => {
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'run',
workMode: 'execute'
},
controller.signal,
async () => 'once'
)) {
void _event
}
}
await expect(consume()).rejects.toThrow()
expect(fetcher).toHaveBeenCalledOnce()
})
it('terminates repeated identical tool rounds without exhausting hard limits', async () => {
let callId = 0
const fetcher = vi.fn<typeof fetch>(async () => {
callId += 1
return Response.json({
choices: [
{
message: {
role: 'assistant',
content: null,
tool_calls: [
{
id: `call-repeat-${callId}`,
type: 'function',
function: {
name: 'workspace_read_text',
arguments: '{"path":"README.md"}'
}
}
]
}
}
]
})
})
const toolProvider = createToolProvider()
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher,
toolProvider
})
const consume = async (): Promise<void> => {
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: 'conversation-repeat',
prompt: 'repeat',
workMode: 'execute'
},
new AbortController().signal,
async () => 'once'
)) {
void _event
}
}
await expect(consume()).rejects.toThrow('没有取得进展')
expect(fetcher).toHaveBeenCalledTimes(3)
expect(toolProvider.callTool).toHaveBeenCalledTimes(2)
})
it('releases provider state for only the requested conversation', async () => {
const toolProvider = createToolProvider()
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
toolProvider
})
await runtime.releaseConversation('conversation-release')
expect(toolProvider.releaseConversation).toHaveBeenCalledOnce()
expect(toolProvider.releaseConversation).toHaveBeenCalledWith(
'conversation-release'
)
})
it('releases known conversations before provider disposal and permits replacement reuse', async () => {
const lifecycle: string[] = []
const released = new Set<string>()
const createProvider = (): ModelToolProviderLike =>
createToolProvider({
releaseConversation: vi.fn(async (conversationId) => {
lifecycle.push(`release:${conversationId}`)
released.add(conversationId)
}),
dispose: vi.fn(async () => {
lifecycle.push('dispose')
})
})
const createRuntime = (toolProvider: ModelToolProviderLike) =>
new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher: vi.fn<typeof fetch>(async () =>
new Response(
'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n',
{
status: 200,
headers: { 'content-type': 'text/event-stream' }
}
)
),
toolProvider
})
const request = {
requestId: crypto.randomUUID(),
conversationId: 'conversation-replacement',
prompt: 'hello',
workMode: 'ask' as const
}
const firstProvider = createProvider()
const firstRuntime = createRuntime(firstProvider)
for await (const _event of firstRuntime.run(
request,
new AbortController().signal
)) {
void _event
}
await firstRuntime.dispose()
expect(lifecycle).toEqual([
'release:conversation-replacement',
'dispose'
])
expect(released).toContain('conversation-replacement')
const replacement = createRuntime(createProvider())
const replacementEvents = []
for await (const event of replacement.run(
{ ...request, requestId: crypto.randomUUID() },
new AbortController().signal
)) {
replacementEvents.push(event)
}
expect(replacementEvents.at(-1)).toMatchObject({ type: 'done' })
await replacement.dispose()
})
it('counts image base64 data against the aggregate tool context limit', async () => {
const response = {
choices: [
{
message: {
role: 'assistant',
content: null,
tool_calls: [
{
id: 'call-large-image',
type: 'function',
function: {
name: 'workspace_read_text',
arguments: '{}'
}
}
]
}
}
]
}
const imageData = Buffer.alloc(1024 * 1024 + 1).toString('base64')
const fetcher = vi.fn<typeof fetch>(async () => Response.json(response))
const runtime = new ModelAgentRuntime({
baseUrl: 'http://127.0.0.1:11434/v1',
model: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none',
fetcher,
toolProvider: createToolProvider({
callTool: vi.fn(async () => ({
parts: [
{
type: 'image' as const,
mimeType: 'image/png' as const,
data: imageData
}
],
contextBytes: Buffer.byteLength(imageData)
}))
})
})
const consume = async (): Promise<void> => {
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'run',
workMode: 'execute'
},
new AbortController().signal,
async () => 'once'
)) {
void _event
}
}
await expect(consume()).rejects.toThrow('结果总量超过 1MB')
expect(fetcher).toHaveBeenCalledOnce()
})
it('generates a bounded image through the BigToken-compatible endpoint', async () => {
const png = Buffer.from([
0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
+254 -28
View File
@@ -5,11 +5,16 @@ import type {
ModelProtocol
} from '../../shared/contracts'
import type { ResolvedMcpServer } from '../capabilities/capability-service'
import type { BrowserToolService } from '../browser/browser-model-tools'
import { createAnthropicMessagesUrl } from './anthropic-endpoint'
import {
ModelToolProvider,
RecoverableModelToolError,
type ModelToolCallContext,
type ModelToolDefinition,
type ModelToolProviderLike
type ModelToolProviderLike,
type ModelToolResult,
type ModelToolResultPart
} from './model-tool-provider'
import {
createOpenAIChatCompletionsUrl,
@@ -86,8 +91,10 @@ const maxImageResponseBytes = 5_300_000
const maxChatResponseBytes = 2 * 1024 * 1024
const maxToolArgumentBytes = 128 * 1024
const maxToolContextBytes = 1024 * 1024
const maxToolCallsPerRun = 12
const maxToolRounds = 8
const maxToolCallsPerRun = 40
const maxToolRounds = 24
const maxRepeatedIdenticalCalls = 3
const maxIdenticalRoundsWithoutProgress = 2
export type ModelRuntimeOptions = {
apiKey?: string
@@ -98,6 +105,7 @@ export type ModelRuntimeOptions = {
skillInstructions?: string
defaultWorkspace?: string
mcpServers?: ResolvedMcpServer[]
browserService?: BrowserToolService
toolProvider?: ModelToolProviderLike
fetcher?: typeof fetch
}
@@ -459,6 +467,152 @@ function parseToolArguments(value: unknown): Record<string, unknown> {
return parsed as Record<string, unknown>
}
function canonicalizeToolArguments(value: unknown): unknown {
if (Array.isArray(value)) {
return value.map(canonicalizeToolArguments)
}
if (value && typeof value === 'object') {
return Object.fromEntries(
Object.entries(value as Record<string, unknown>)
.sort(([left], [right]) => left.localeCompare(right))
.map(([key, item]) => [key, canonicalizeToolArguments(item)])
)
}
return value
}
function getToolCallFingerprint(call: ModelToolCall): string {
return `${call.name}:${JSON.stringify(
canonicalizeToolArguments(call.arguments)
)}`
}
function validateToolResult(result: ModelToolResult): number {
if (
!Array.isArray(result.parts) ||
result.parts.length === 0 ||
!Number.isSafeInteger(result.contextBytes) ||
result.contextBytes < 0
) {
throw new Error('直连模型工具返回了无效结果')
}
let contextBytes = 0
for (const part of result.parts) {
if (part.type === 'text') {
if (typeof part.text !== 'string') {
throw new Error('直连模型工具返回了无效文本结果')
}
contextBytes += Buffer.byteLength(part.text)
} else if (
part.type === 'image' &&
(part.mimeType === 'image/png' ||
part.mimeType === 'image/jpeg' ||
part.mimeType === 'image/webp') &&
typeof part.data === 'string'
) {
contextBytes += Buffer.byteLength(part.data)
} else {
throw new Error('直连模型工具返回了无效图片结果')
}
}
if (contextBytes !== result.contextBytes) {
throw new Error('直连模型工具结果字节计数无效')
}
return contextBytes
}
function getAnthropicToolResultContent(
parts: ModelToolResultPart[]
): Array<Record<string, unknown>> {
return parts.map((part) =>
part.type === 'text'
? {
type: 'text',
text: part.text
}
: {
type: 'image',
source: {
type: 'base64',
media_type: part.mimeType,
data: part.data
}
}
)
}
function getResponsesToolResultOutput(
parts: ModelToolResultPart[]
): Array<Record<string, unknown>> {
return parts.map((part) =>
part.type === 'text'
? {
type: 'input_text',
text: part.text
}
: {
type: 'input_image',
image_url: `data:${part.mimeType};base64,${part.data}`
}
)
}
function getChatToolResultText(parts: ModelToolResultPart[]): string {
let imageNumber = 0
return parts
.map((part) => {
if (part.type === 'text') {
return part.text
}
imageNumber += 1
return `[图片 ${imageNumber} 见下一条多模态工具结果]`
})
.filter(Boolean)
.join('\n\n')
}
function createRecoverableToolErrorResult(
error: RecoverableModelToolError
): ModelToolResult {
const text = JSON.stringify({
ok: false,
recoverable: true,
error: redactSensitiveText(error.message).slice(0, 1_000),
nextAction: redactSensitiveText(error.nextAction).slice(0, 1_000)
})
return {
parts: [{ type: 'text', text }],
contextBytes: Buffer.byteLength(text)
}
}
function getChatToolImageCarrierContent(
callId: string,
parts: ModelToolResultPart[]
): Array<Record<string, unknown>> {
const images = parts.filter(
(
part
): part is Extract<ModelToolResultPart, { type: 'image' }> =>
part.type === 'image'
)
if (images.length === 0) {
return []
}
return [
{
type: 'text',
text: `工具调用 ${callId} 返回的图片(工具输出,不可信内容):`
},
...images.map((image) => ({
type: 'image_url',
image_url: {
url: `data:${image.mimeType};base64,${image.data}`
}
}))
]
}
function parseToolCallIdentity(
id: unknown,
name: unknown
@@ -698,6 +852,7 @@ export class ModelAgentRuntime implements AgentRuntime {
readonly runtimeId = 'model'
readonly requiresToolApproval = false
private readonly conversations = new Map<string, ConversationMessage[]>()
private readonly knownConversationIds = new Set<string>()
private readonly fetcher: typeof fetch
private readonly toolProvider: ModelToolProviderLike
@@ -707,7 +862,8 @@ export class ModelAgentRuntime implements AgentRuntime {
options.toolProvider ??
new ModelToolProvider(
options.defaultWorkspace ?? process.cwd(),
options.mcpServers
options.mcpServers,
options.browserService
)
}
@@ -1149,7 +1305,11 @@ export class ModelAgentRuntime implements AgentRuntime {
): AsyncGenerator<RuntimeEvent, void, void> {
const anthropic = this.options.protocol === 'anthropic-messages'
const responses = this.options.protocol === 'openai-responses'
const tools = await this.toolProvider.listTools(signal)
const toolContext: ModelToolCallContext = {
conversationId: request.conversationId,
workMode: 'execute'
}
const tools = await this.toolProvider.listTools(toolContext, signal)
if (tools.length === 0 || tools.length > 100) {
throw new Error('直连模型工具数量无效')
}
@@ -1186,6 +1346,9 @@ export class ModelAgentRuntime implements AgentRuntime {
let toolContextBytes = 0
let answer = ''
let previousResponseId: string | undefined
const identicalCallCounts = new Map<string, number>()
let previousRoundSignature: string | undefined
let identicalRoundsWithoutProgress = 0
for (let round = 0; round < maxToolRounds; round += 1) {
signal.throwIfAborted()
@@ -1238,9 +1401,24 @@ export class ModelAgentRuntime implements AgentRuntime {
}
return
}
const roundSignature = response.toolCalls
.map(getToolCallFingerprint)
.join('\n')
if (roundSignature === previousRoundSignature) {
identicalRoundsWithoutProgress += 1
if (
identicalRoundsWithoutProgress >=
maxIdenticalRoundsWithoutProgress
) {
throw new Error('直连模型重复了相同工具调用且没有取得进展')
}
} else {
previousRoundSignature = roundSignature
identicalRoundsWithoutProgress = 0
}
totalToolCalls += response.toolCalls.length
if (totalToolCalls > maxToolCallsPerRun) {
throw new Error('直连模型单次运行的工具调用超过 12 个')
throw new Error('直连模型单次运行的工具调用超过 40 个')
}
if (responses) {
if (!response.responseId) {
@@ -1254,8 +1432,16 @@ export class ModelAgentRuntime implements AgentRuntime {
}
const anthropicResults: Array<Record<string, unknown>> = []
const responsesResults: Array<Record<string, unknown>> = []
const chatImageCarrierContent: Array<Record<string, unknown>> = []
for (const call of response.toolCalls) {
signal.throwIfAborted()
const callFingerprint = getToolCallFingerprint(call)
const identicalCallCount =
(identicalCallCounts.get(callFingerprint) ?? 0) + 1
identicalCallCounts.set(callFingerprint, identicalCallCount)
if (identicalCallCount > maxRepeatedIdenticalCalls) {
throw new Error('直连模型重复请求了完全相同的工具调用')
}
if (seenCallIds.has(call.id)) {
throw new Error('模型重复使用了工具调用 ID')
}
@@ -1291,7 +1477,8 @@ export class ModelAgentRuntime implements AgentRuntime {
this.toolProvider.getApproval(
tool,
call.arguments,
safeToolArgumentSummary(call.arguments)
safeToolArgumentSummary(call.arguments),
toolContext
)
)
} catch (error) {
@@ -1316,6 +1503,7 @@ export class ModelAgentRuntime implements AgentRuntime {
}
throw new Error(`用户拒绝了工具「${displayName}`)
}
signal.throwIfAborted()
yield {
requestId: request.requestId,
type: 'tool',
@@ -1325,27 +1513,38 @@ export class ModelAgentRuntime implements AgentRuntime {
summary: `正在执行直连模型工具:${displayName}`
}
let result: string
let result: ModelToolResult
let toolFailed = false
try {
result = await this.toolProvider.callTool(
tool.name,
call.arguments,
signal
signal,
toolContext
)
} catch (error) {
const recoverable = error instanceof RecoverableModelToolError
yield {
requestId: request.requestId,
type: 'tool',
callId: call.id,
name: displayName,
state: 'failed',
summary: `直连模型工具执行失败:${displayName}`
state: recoverable ? 'recoverable' : 'failed',
summary:
recoverable
? `直连模型工具需要刷新后重试:${displayName}`
: `直连模型工具执行失败:${displayName}`
}
if (recoverable) {
result = createRecoverableToolErrorResult(error)
toolFailed = true
} else {
throw new Error(`工具「${displayName}」执行失败`, {
cause: error
})
}
throw new Error(`工具「${displayName}」执行失败`, {
cause: error
})
}
toolContextBytes += Buffer.byteLength(result)
toolContextBytes += validateToolResult(result)
if (toolContextBytes > maxToolContextBytes) {
yield {
requestId: request.requestId,
@@ -1361,28 +1560,34 @@ export class ModelAgentRuntime implements AgentRuntime {
responsesResults.push({
type: 'function_call_output',
call_id: call.id,
output: result
output: getResponsesToolResultOutput(result.parts)
})
} else if (anthropic) {
anthropicResults.push({
type: 'tool_result',
tool_use_id: call.id,
content: result
content: getAnthropicToolResultContent(result.parts),
...(toolFailed ? { is_error: true } : {})
})
} else {
messages.push({
role: 'tool',
tool_call_id: call.id,
content: result
content: getChatToolResultText(result.parts)
})
chatImageCarrierContent.push(
...getChatToolImageCarrierContent(call.id, result.parts)
)
}
yield {
requestId: request.requestId,
type: 'tool',
callId: call.id,
name: displayName,
state: 'completed',
summary: `直连模型工具已完成:${displayName}`
if (!toolFailed) {
yield {
requestId: request.requestId,
type: 'tool',
callId: call.id,
name: displayName,
state: 'completed',
summary: `直连模型工具已完成:${displayName}`
}
}
}
if (anthropic) {
@@ -1392,9 +1597,15 @@ export class ModelAgentRuntime implements AgentRuntime {
})
} else if (responses) {
messages.splice(0, messages.length, ...responsesResults)
} else if (chatImageCarrierContent.length > 0) {
messages.push({
role: 'user',
content: chatImageCarrierContent
})
}
signal.throwIfAborted()
}
throw new Error('直连模型工具调用轮次超过 8 轮')
throw new Error('直连模型工具调用轮次超过 24 轮')
}
async *run(
@@ -1402,6 +1613,7 @@ export class ModelAgentRuntime implements AgentRuntime {
signal: AbortSignal,
authorize?: RuntimeAuthorizer
): AsyncGenerator<RuntimeEvent, void, void> {
this.knownConversationIds.add(request.conversationId)
if (!this.isConfigured()) {
throw new Error('请先在设置中配置模型接口 API Key')
}
@@ -1575,12 +1787,26 @@ export class ModelAgentRuntime implements AgentRuntime {
}
async dispose(): Promise<void> {
const conversationIds = new Set([
...this.knownConversationIds,
...this.conversations.keys()
])
await Promise.allSettled(
[...conversationIds].map((conversationId) =>
this.toolProvider.releaseConversation(conversationId)
)
)
this.knownConversationIds.clear()
this.conversations.clear()
await this.toolProvider.dispose()
}
releaseConversation(conversationId: string): Promise<void> {
async releaseConversation(conversationId: string): Promise<void> {
this.conversations.delete(conversationId)
return Promise.resolve()
try {
await this.toolProvider.releaseConversation(conversationId)
} finally {
this.knownConversationIds.delete(conversationId)
}
}
}
+513 -35
View File
@@ -9,16 +9,24 @@ import { tmpdir } from 'node:os'
import { join } from 'node:path'
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import type { ResolvedMcpServer } from '../capabilities/capability-service'
import type { BrowserToolService } from '../browser/browser-model-tools'
import { BrowserStaleReferenceError } from '../browser/cdp-browser-driver'
const mocks = vi.hoisted(() => {
const tasks = {
callToolStream: vi.fn(),
cancelTask: vi.fn()
}
const client = {
connect: vi.fn(),
listTools: vi.fn(),
callTool: vi.fn(),
experimental: { tasks },
close: vi.fn()
}
return {
client,
tasks,
Client: vi.fn(function Client() {
return client
}),
@@ -33,9 +41,63 @@ vi.mock('../capabilities/mcp-client-transport', () => ({
createMcpTransport: mocks.createMcpTransport
}))
import { ModelToolProvider } from './model-tool-provider'
import {
ModelToolProvider,
type ModelToolCallContext
} from './model-tool-provider'
const temporaryDirectories: string[] = []
const png = Buffer.from([
0x89, 0x50, 0x4e, 0x47,
0x0d, 0x0a, 0x1a, 0x0a
]).toString('base64')
const toolContext = {
conversationId: 'provider-test-conversation',
workMode: 'execute'
} satisfies ModelToolCallContext
function createBrowserService(): BrowserToolService {
return {
getOrigin: vi.fn(() => 'https://example.com'),
navigate: vi.fn(async (_conversationId, url) => ({
url,
origin: 'https://example.com'
})),
snapshot: vi.fn(async () => ({
url: 'https://example.com/',
title: 'Example',
nodes: [],
truncated: false
})),
click: vi.fn(async () => undefined),
type: vi.fn(async () => undefined),
select: vi.fn(async () => undefined),
back: vi.fn(async () => ({
url: 'https://previous.example/',
origin: 'https://previous.example'
})),
screenshot: vi.fn(async () => ({
type: 'image' as const,
mimeType: 'image/png' as const,
data: png
})),
releaseConversation: vi.fn(async () => undefined)
}
}
function createMcpServer(): ResolvedMcpServer {
return {
id: 'd2ef774b-146c-4467-a909-6feb112a9c2c',
name: 'Search MCP',
description: '',
enabled: true,
assignments: ['model'],
secretConfigured: false,
transport: 'stdio',
command: 'node',
args: ['server.js']
}
}
async function createWorkspace(): Promise<string> {
const directory = await mkdtemp(join(tmpdir(), 'goodbuddy-tools-'))
@@ -51,6 +113,13 @@ describe('ModelToolProvider', () => {
mocks.client.callTool.mockResolvedValue({
content: [{ type: 'text', text: 'MCP result' }]
})
mocks.tasks.callToolStream.mockImplementation(async function* () {
yield {
type: 'result',
result: { content: [{ type: 'text', text: 'MCP task result' }] }
}
})
mocks.tasks.cancelTask.mockResolvedValue({})
mocks.client.close.mockResolvedValue(undefined)
})
@@ -71,7 +140,7 @@ describe('ModelToolProvider', () => {
const provider = new ModelToolProvider(workspace)
const signal = new AbortController().signal
await expect(provider.listTools(signal)).resolves.toEqual(
await expect(provider.listTools(toolContext, signal)).resolves.toEqual(
expect.arrayContaining([
expect.objectContaining({ name: 'workspace_read_text' }),
expect.objectContaining({ name: 'workspace_list_directory' }),
@@ -82,23 +151,37 @@ describe('ModelToolProvider', () => {
provider.callTool(
'workspace_read_text',
{ path: 'docs/note.txt' },
signal
signal,
toolContext
)
).resolves.toBe('hello')
await expect(
provider.callTool(
'workspace_list_directory',
{ path: 'docs' },
signal
)
).resolves.toContain('"note.txt"')
await expect(
provider.callTool(
'workspace_write_text',
{ path: 'docs/output.txt', content: 'saved' },
signal
)
).resolves.toContain('"bytesWritten":5')
).resolves.toEqual({
parts: [{ type: 'text', text: 'hello' }],
contextBytes: 5
})
const listing = await provider.callTool(
'workspace_list_directory',
{ path: 'docs' },
signal,
toolContext
)
expect(listing.parts).toEqual([
expect.objectContaining({
type: 'text',
text: expect.stringContaining('"note.txt"')
})
])
const written = await provider.callTool(
'workspace_write_text',
{ path: 'docs/output.txt', content: 'saved' },
signal,
toolContext
)
expect(written.parts).toEqual([
expect.objectContaining({
type: 'text',
text: expect.stringContaining('"bytesWritten":5')
})
])
await expect(
readFile(join(workspace, 'docs', 'output.txt'), 'utf8')
).resolves.toBe('saved')
@@ -112,11 +195,116 @@ describe('ModelToolProvider', () => {
provider.callTool(
'workspace_read_text',
{ path: '../outside.txt' },
new AbortController().signal
new AbortController().signal,
toolContext
)
).rejects.toThrow('不能超出工作区')
})
it('delegates browser tools with per-call conversation context', async () => {
const workspace = await createWorkspace()
const browserService = createBrowserService()
const provider = new ModelToolProvider(workspace, [], browserService)
const firstContext = {
conversationId: 'browser-conversation-one',
workMode: 'execute'
} satisfies ModelToolCallContext
const secondContext = {
conversationId: 'browser-conversation-two',
workMode: 'execute'
} satisfies ModelToolCallContext
const signal = new AbortController().signal
for (const workMode of ['ask', 'plan'] as const) {
const readOnlyContext = {
conversationId: `browser-${workMode}`,
workMode
} satisfies ModelToolCallContext
await expect(
provider.listTools(readOnlyContext, signal)
).resolves.not.toEqual(
expect.arrayContaining([
expect.objectContaining({ name: 'browser_screenshot' })
])
)
await expect(
provider.callTool(
'browser_screenshot',
{},
signal,
readOnlyContext
)
).rejects.toThrow('未知工具')
}
expect(browserService.screenshot).not.toHaveBeenCalled()
const tools = await provider.listTools(firstContext, signal)
expect(
tools.filter((tool) => tool.name.startsWith('browser_'))
).toHaveLength(7)
const navigate = tools.find((tool) => tool.name === 'browser_navigate')
expect(
provider.getApproval(
navigate!,
{ url: 'https://example.com/path?secret=value' },
'runtime summary',
firstContext
)
).toMatchObject({
scopeKey: 'model:browser:navigate:https://example.com',
argumentSummary: 'https://example.com/path?[查询参数已隐藏]',
allowPermanent: false
})
await expect(
provider.callTool('browser_screenshot', {}, signal, firstContext)
).resolves.toEqual({
parts: [{ type: 'image', mimeType: 'image/png', data: png }],
contextBytes: Buffer.byteLength(png)
})
await provider.callTool('browser_screenshot', {}, signal, secondContext)
expect(browserService.screenshot).toHaveBeenNthCalledWith(
1,
firstContext.conversationId,
signal
)
expect(browserService.screenshot).toHaveBeenNthCalledWith(
2,
secondContext.conversationId,
signal
)
await provider.releaseConversation(firstContext.conversationId)
expect(browserService.releaseConversation).toHaveBeenCalledWith(
firstContext.conversationId
)
expect(browserService.releaseConversation).not.toHaveBeenCalledWith(
secondContext.conversationId
)
})
it('marks stale browser references as recoverable model tool errors', async () => {
const workspace = await createWorkspace()
const browserService = createBrowserService()
vi.mocked(browserService.click).mockRejectedValue(
new BrowserStaleReferenceError()
)
const provider = new ModelToolProvider(workspace, [], browserService)
await expect(
provider.callTool(
'browser_click',
{ ref: 'b_currentReference' },
new AbortController().signal,
toolContext
)
).rejects.toMatchObject({
name: 'RecoverableModelToolError',
message: '浏览器元素引用已失效,请重新获取快照',
nextAction: expect.stringContaining('browser_snapshot')
})
})
it('loads and invokes configured MCP tools through provider-safe names', async () => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
@@ -132,21 +320,10 @@ describe('ModelToolProvider', () => {
}
]
})
const server = {
id: 'd2ef774b-146c-4467-a909-6feb112a9c2c',
name: 'Search MCP',
description: '',
enabled: true,
assignments: ['model'],
secretConfigured: false,
transport: 'stdio',
command: 'node',
args: ['server.js']
} satisfies ResolvedMcpServer
const provider = new ModelToolProvider(workspace, [server])
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const signal = new AbortController().signal
const tools = await provider.listTools(signal)
const tools = await provider.listTools(toolContext, signal)
const mcpTool = tools.find((tool) => tool.source === 'mcp')
expect(mcpTool).toMatchObject({
displayName: 'Search MCP / search-web',
@@ -157,9 +334,13 @@ describe('ModelToolProvider', () => {
provider.callTool(
mcpTool?.name ?? '',
{ query: 'GoodBuddy' },
signal
signal,
toolContext
)
).resolves.toBe('MCP result')
).resolves.toEqual({
parts: [{ type: 'text', text: 'MCP result' }],
contextBytes: 10
})
expect(mocks.client.callTool).toHaveBeenCalledWith(
{
name: 'search-web',
@@ -168,11 +349,308 @@ describe('ModelToolProvider', () => {
undefined,
expect.objectContaining({
timeout: 30_000,
signal
signal,
resetTimeoutOnProgress: true,
maxTotalTimeout: 300_000,
onprogress: expect.any(Function)
})
)
await provider.dispose()
expect(mocks.client.close).toHaveBeenCalledOnce()
})
it('preserves ordered bounded MCP text, image, and unsupported audio parts', async () => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [
{
name: 'capture',
inputSchema: { type: 'object' }
}
]
})
mocks.client.callTool.mockResolvedValue({
content: [
{ type: 'text', text: 'before' },
{ type: 'image', mimeType: 'image/png', data: png },
{ type: 'audio', mimeType: 'audio/wav', data: 'ignored' },
{ type: 'text', text: 'after' }
]
})
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, new AbortController().signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).resolves.toEqual({
parts: [
{ type: 'text', text: 'before' },
{
type: 'image',
mimeType: 'image/png',
data: png
},
{ type: 'text', text: '[audio result unsupported]' },
{ type: 'text', text: 'after' }
],
contextBytes:
Buffer.byteLength('before') +
Buffer.byteLength(png) +
Buffer.byteLength('[audio result unsupported]') +
Buffer.byteLength('after')
})
})
it.each([
{
mimeType: 'image/jpeg',
data: Buffer.from([0xff, 0xd8, 0xff]).toString('base64')
},
{
mimeType: 'image/webp',
data: Buffer.from([
0x52, 0x49, 0x46, 0x46,
0x00, 0x00, 0x00, 0x00,
0x57, 0x45, 0x42, 0x50
]).toString('base64')
}
])('accepts a valid $mimeType signature', async ({ mimeType, data }) => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [{ name: 'capture', inputSchema: { type: 'object' } }]
})
mocks.client.callTool.mockResolvedValue({
content: [{ type: 'image', mimeType, data }]
})
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, new AbortController().signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).resolves.toEqual({
parts: [{ type: 'image', mimeType, data }],
contextBytes: Buffer.byteLength(data)
})
})
it.each([
{
name: 'malformed base64',
image: {
type: 'image',
mimeType: 'image/png',
data: `${png.slice(0, -1)}!`
},
message: '无效的 base64'
},
{
name: 'MIME signature mismatch',
image: {
type: 'image',
mimeType: 'image/jpeg',
data: png
},
message: 'MIME 类型与文件签名不匹配'
},
{
name: 'unsupported MIME type',
image: {
type: 'image',
mimeType: 'image/gif',
data: png
},
message: '不支持的图片格式'
}
])('rejects $name in MCP image blocks', async ({ image, message }) => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [{ name: 'capture', inputSchema: { type: 'object' } }]
})
mocks.client.callTool.mockResolvedValue({ content: [image] })
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, new AbortController().signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).rejects.toThrow(message)
})
it('counts encoded and decoded image data against the MCP result budget', async () => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [{ name: 'capture', inputSchema: { type: 'object' } }]
})
const encodedContextOversizedPng = Buffer.concat([
Buffer.from([
0x89, 0x50, 0x4e, 0x47,
0x0d, 0x0a, 0x1a, 0x0a
]),
Buffer.alloc(200 * 1024)
]).toString('base64')
mocks.client.callTool.mockResolvedValue({
content: [
{
type: 'image',
mimeType: 'image/png',
data: encodedContextOversizedPng
}
]
})
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, new AbortController().signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).rejects.toThrow('工具结果超过 256KB')
const decodedOversizedPng = Buffer.concat([
Buffer.from([
0x89, 0x50, 0x4e, 0x47,
0x0d, 0x0a, 0x1a, 0x0a
]),
Buffer.alloc(256 * 1024)
]).toString('base64')
mocks.client.callTool.mockResolvedValue({
content: [
{
type: 'image',
mimeType: 'image/png',
data: decodedOversizedPng
}
]
})
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).rejects.toThrow('过大的 base64 图片')
})
it('bounds MCP content block and image counts', async () => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [{ name: 'capture', inputSchema: { type: 'object' } }]
})
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, new AbortController().signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
mocks.client.callTool.mockResolvedValue({
content: Array.from({ length: 101 }, () => ({
type: 'text',
text: 'x'
}))
})
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).rejects.toThrow('内容块数量超过安全限制')
mocks.client.callTool.mockResolvedValue({
content: Array.from({ length: 9 }, () => ({
type: 'image',
mimeType: 'image/png',
data: png
}))
})
await expect(
provider.callTool(
tool?.name ?? '',
{},
new AbortController().signal,
toolContext
)
).rejects.toThrow('图片数量超过安全限制')
})
it('streams required task tools and best-effort cancels their MCP task', async () => {
const workspace = await createWorkspace()
mocks.client.listTools.mockResolvedValue({
tools: [
{
name: 'long-job',
inputSchema: { type: 'object' },
execution: { taskSupport: 'required' }
}
]
})
const controller = new AbortController()
mocks.tasks.callToolStream.mockImplementation(async function* (
_params,
_schema,
options
) {
yield {
type: 'taskCreated',
task: { taskId: 'task-1', status: 'working' }
}
controller.abort()
throw options.signal.reason
})
const provider = new ModelToolProvider(workspace, [createMcpServer()])
const tools = await provider.listTools(toolContext, controller.signal)
const tool = tools.find((candidate) => candidate.source === 'mcp')
expect(tool?.taskSupport).toBe('required')
await expect(
provider.callTool(
tool?.name ?? '',
{},
controller.signal,
toolContext
)
).rejects.toThrow()
expect(mocks.tasks.callToolStream).toHaveBeenCalledWith(
{
name: 'long-job',
arguments: {}
},
undefined,
expect.objectContaining({
timeout: 30_000,
signal: controller.signal,
resetTimeoutOnProgress: true,
maxTotalTimeout: 300_000
})
)
expect(mocks.tasks.cancelTask).toHaveBeenCalledWith(
'task-1',
{
timeout: 5_000,
maxTotalTimeout: 5_000
}
)
})
})
+350 -86
View File
@@ -14,6 +14,7 @@ import {
resolve
} from 'node:path'
import { z } from 'zod'
import { builtinModelTools } from '../../shared/builtin-model-tools'
import type { ResolvedMcpServer } from '../capabilities/capability-service'
import { createMcpTransport } from '../capabilities/mcp-client-transport'
import {
@@ -23,6 +24,11 @@ import {
readBoundedUtf8File
} from '../workspace-file-access'
import type { RuntimeApprovalRequest } from './runtime'
import {
BrowserModelTools,
type BrowserToolService
} from '../browser/browser-model-tools'
import { BrowserStaleReferenceError } from '../browser/cdp-browser-driver'
const MAX_MODEL_TOOLS = 100
const MAX_MCP_SERVERS = 16
@@ -31,6 +37,15 @@ const MAX_TOOL_RESULT_BYTES = 256 * 1024
const MAX_READ_BYTES = 256 * 1024
const MAX_WRITE_BYTES = 512 * 1024
const MCP_TIMEOUT_MS = 30_000
const MCP_CALL_MAX_TOTAL_TIMEOUT_MS = 5 * 60_000
const MCP_TASK_CANCEL_TIMEOUT_MS = 5_000
const MAX_MCP_CONTENT_BLOCKS = 100
const MAX_MCP_IMAGES = 8
const [
workspaceReadTextTool,
workspaceListDirectoryTool,
workspaceWriteTextTool
] = builtinModelTools
const workspacePathSchema = z
.string()
@@ -66,20 +81,62 @@ export type ModelToolDefinition = {
inputSchema: Record<string, unknown>
source: 'builtin' | 'mcp'
serverName?: string
taskSupport?: 'forbidden' | 'optional' | 'required'
}
export type ModelToolResultPart =
| {
type: 'text'
text: string
}
| {
type: 'image'
mimeType: 'image/png' | 'image/jpeg' | 'image/webp'
data: string
}
export type ModelToolResult = {
parts: ModelToolResultPart[]
contextBytes: number
}
export type ModelToolCallContext = {
conversationId: string
workMode: 'ask' | 'plan' | 'execute'
}
export class RecoverableModelToolError extends Error {
readonly nextAction: string
constructor(
message: string,
nextAction: string,
options?: ErrorOptions
) {
super(message, options)
this.name = 'RecoverableModelToolError'
this.nextAction = nextAction
}
}
export interface ModelToolProviderLike {
listTools(signal: AbortSignal): Promise<ModelToolDefinition[]>
listTools(
context: ModelToolCallContext,
signal: AbortSignal
): Promise<ModelToolDefinition[]>
getApproval(
tool: ModelToolDefinition,
argumentsValue: Record<string, unknown>,
argumentSummary: string
argumentSummary: string,
context: ModelToolCallContext
): RuntimeApprovalRequest
callTool(
name: string,
argumentsValue: Record<string, unknown>,
signal: AbortSignal
): Promise<string>
signal: AbortSignal,
context: ModelToolCallContext
): Promise<ModelToolResult>
releaseConversation(conversationId: string): Promise<void>
dispose(): Promise<void>
}
@@ -151,64 +208,177 @@ function createMcpToolName(serverId: string, originalName: string): string {
return `mcp_${serverHash}_${toolHash}_${readable}`.slice(0, 64)
}
function getMcpResultText(result: unknown): string {
function createTextToolResult(text: string): ModelToolResult {
const contextBytes = Buffer.byteLength(text)
if (contextBytes > MAX_TOOL_RESULT_BYTES) {
throw new Error('工具结果超过 256KB 安全限制')
}
return {
parts: [{ type: 'text', text }],
contextBytes
}
}
function parseMcpImage(
content: Record<string, unknown>
): Extract<ModelToolResultPart, { type: 'image' }> {
const mimeType = content.mimeType
if (
mimeType !== 'image/png' &&
mimeType !== 'image/jpeg' &&
mimeType !== 'image/webp'
) {
throw new Error('MCP 工具返回了不支持的图片格式')
}
if (
typeof content.data !== 'string' ||
content.data.length === 0 ||
content.data.length % 4 !== 0 ||
!/^(?:[A-Za-z0-9+/]{4})*(?:[A-Za-z0-9+/]{2}==|[A-Za-z0-9+/]{3}=)?$/u.test(
content.data
)
) {
throw new Error('MCP 工具返回了无效的 base64 图片')
}
const decoded = Buffer.from(content.data, 'base64')
if (
decoded.length === 0 ||
decoded.length > MAX_TOOL_RESULT_BYTES ||
decoded.toString('base64') !== content.data
) {
throw new Error('MCP 工具返回了无效或过大的 base64 图片')
}
const signatureMatches =
mimeType === 'image/png'
? decoded.length >= 8 &&
decoded.subarray(0, 8).equals(
Buffer.from([
0x89, 0x50, 0x4e, 0x47,
0x0d, 0x0a, 0x1a, 0x0a
])
)
: mimeType === 'image/jpeg'
? decoded.length >= 3 &&
decoded[0] === 0xff &&
decoded[1] === 0xd8 &&
decoded[2] === 0xff
: decoded.length >= 12 &&
decoded.subarray(0, 4).toString('ascii') === 'RIFF' &&
decoded.subarray(8, 12).toString('ascii') === 'WEBP'
if (!signatureMatches) {
throw new Error('MCP 工具图片的 MIME 类型与文件签名不匹配')
}
return {
type: 'image',
mimeType,
data: content.data
}
}
function normalizeMcpResult(result: unknown): ModelToolResult {
if (!result || typeof result !== 'object') {
return boundedJson(result, 'MCP 工具结果无法序列化')
return createTextToolResult(
boundedJson(result, 'MCP 工具结果无法序列化')
)
}
const record = result as Record<string, unknown>
if (record.isError === true) {
throw new Error('MCP Server 报告工具执行失败')
}
if ('toolResult' in record) {
return boundedJson(record.toolResult, 'MCP 工具结果无法序列化')
if (
record.toolResult &&
typeof record.toolResult === 'object' &&
(
Array.isArray(
(record.toolResult as Record<string, unknown>).content
) ||
'structuredContent' in
(record.toolResult as Record<string, unknown>) ||
'isError' in (record.toolResult as Record<string, unknown>)
)
) {
return normalizeMcpResult(record.toolResult)
}
return createTextToolResult(
boundedJson(record.toolResult, 'MCP 工具结果无法序列化')
)
}
const sections: string[] = []
const parts: ModelToolResultPart[] = []
if (
record.structuredContent &&
typeof record.structuredContent === 'object'
) {
sections.push(
boundedJson(
parts.push({
type: 'text',
text: boundedJson(
record.structuredContent,
'MCP 结构化工具结果无法序列化'
)
)
})
}
if (Array.isArray(record.content)) {
for (const item of record.content.slice(0, 100)) {
if (record.content.length > MAX_MCP_CONTENT_BLOCKS) {
throw new Error('MCP 工具结果内容块数量超过安全限制')
}
let imageCount = 0
for (const item of record.content) {
if (!item || typeof item !== 'object') {
continue
}
const content = item as Record<string, unknown>
if (content.type === 'text' && typeof content.text === 'string') {
sections.push(content.text)
parts.push({ type: 'text', text: content.text })
} else if (
content.type === 'resource' &&
content.resource &&
typeof content.resource === 'object' &&
typeof (content.resource as Record<string, unknown>).text === 'string'
) {
sections.push(
(content.resource as Record<string, unknown>).text as string
)
parts.push({
type: 'text',
text: (content.resource as Record<string, unknown>).text as string
})
} else if (content.type === 'resource_link') {
sections.push(
boundedJson(content, 'MCP 资源链接无法序列化')
)
} else if (content.type === 'image' || content.type === 'audio') {
sections.push(`[${String(content.type)} result omitted]`)
parts.push({
type: 'text',
text: boundedJson(content, 'MCP 资源链接无法序列化')
})
} else if (content.type === 'image') {
imageCount += 1
if (imageCount > MAX_MCP_IMAGES) {
throw new Error('MCP 工具结果图片数量超过安全限制')
}
parts.push(parseMcpImage(content))
} else if (content.type === 'audio') {
parts.push({
type: 'text',
text: '[audio result unsupported]'
})
}
}
}
const text = sections.join('\n\n').trim()
if (!text) {
return '{}'
if (parts.length === 0) {
return createTextToolResult('{}')
}
if (Buffer.byteLength(text) > MAX_TOOL_RESULT_BYTES) {
let contextBytes = 0
let decodedImageBytes = 0
for (const part of parts) {
contextBytes += Buffer.byteLength(
part.type === 'text' ? part.text : part.data
)
if (part.type === 'image') {
decodedImageBytes += Buffer.from(part.data, 'base64').length
}
}
if (
contextBytes > MAX_TOOL_RESULT_BYTES ||
decodedImageBytes > MAX_TOOL_RESULT_BYTES
) {
throw new Error('工具结果超过 256KB 安全限制')
}
return text
return { parts, contextBytes }
}
export class ModelToolProvider implements ModelToolProviderLike {
@@ -218,9 +388,21 @@ export class ModelToolProvider implements ModelToolProviderLike {
constructor(
private readonly workspace: string,
private readonly mcpServers: ResolvedMcpServer[] = []
private readonly mcpServers: ResolvedMcpServer[] = [],
private readonly browserService?: BrowserToolService
) {}
private getBrowserTools(
context: ModelToolCallContext
): BrowserModelTools | undefined {
return this.browserService && context.workMode === 'execute'
? new BrowserModelTools({
service: this.browserService,
conversationId: context.conversationId
})
: undefined
}
private async getWorkspace(): Promise<string> {
this.canonicalWorkspace ??= getCanonicalWorkspace(
this.workspace,
@@ -289,10 +471,9 @@ export class ModelToolProvider implements ModelToolProviderLike {
private getBuiltinTools(): ModelToolDefinition[] {
return [
{
name: 'workspace_read_text',
displayName: '读取工作区文本',
description:
'读取当前工作区内一个不超过 256KB 的 UTF-8 文本文件。',
name: workspaceReadTextTool.name,
displayName: workspaceReadTextTool.displayName,
description: workspaceReadTextTool.description,
inputSchema: {
type: 'object',
properties: {
@@ -307,10 +488,9 @@ export class ModelToolProvider implements ModelToolProviderLike {
source: 'builtin'
},
{
name: 'workspace_list_directory',
displayName: '列出工作区目录',
description:
'列出当前工作区内目录的直属内容,最多返回 200 项。',
name: workspaceListDirectoryTool.name,
displayName: workspaceListDirectoryTool.displayName,
description: workspaceListDirectoryTool.description,
inputSchema: {
type: 'object',
properties: {
@@ -324,10 +504,9 @@ export class ModelToolProvider implements ModelToolProviderLike {
source: 'builtin'
},
{
name: 'workspace_write_text',
displayName: '写入工作区文本',
description:
'在当前工作区内新建或覆盖一个不超过 512KB 的 UTF-8 文本文件;父目录必须已存在。',
name: workspaceWriteTextTool.name,
displayName: workspaceWriteTextTool.displayName,
description: workspaceWriteTextTool.description,
inputSchema: {
type: 'object',
properties: {
@@ -366,7 +545,9 @@ export class ModelToolProvider implements ModelToolProviderLike {
timeout: MCP_TIMEOUT_MS,
signal
})
if (result.tools.length > MAX_MODEL_TOOLS - 3) {
const builtinToolCount =
this.getBuiltinTools().length + (this.browserService ? 7 : 0)
if (result.tools.length > MAX_MODEL_TOOLS - builtinToolCount) {
throw new Error(
`MCP Server「${server.name}」提供的工具数量超过安全限制`
)
@@ -386,7 +567,8 @@ export class ModelToolProvider implements ModelToolProviderLike {
.slice(0, 1_000),
inputSchema: normalizeToolSchema(tool.inputSchema),
source: 'mcp',
serverName: server.name
serverName: server.name,
taskSupport: tool.execution?.taskSupport
}
}))
if (
@@ -423,9 +605,11 @@ export class ModelToolProvider implements ModelToolProviderLike {
)
.then((connections) => {
const bindings = new Map<string, McpToolBinding>()
const builtinToolCount =
this.getBuiltinTools().length + (this.browserService ? 7 : 0)
for (const connection of connections) {
for (const binding of connection.tools) {
if (bindings.size + 3 >= MAX_MODEL_TOOLS) {
if (bindings.size + builtinToolCount >= MAX_MODEL_TOOLS) {
throw new Error('直连模型工具总数超过 100 个安全限制')
}
if (bindings.has(binding.definition.name)) {
@@ -448,11 +632,16 @@ export class ModelToolProvider implements ModelToolProviderLike {
return this.mcpBindings
}
async listTools(signal: AbortSignal): Promise<ModelToolDefinition[]> {
async listTools(
context: ModelToolCallContext,
signal: AbortSignal
): Promise<ModelToolDefinition[]> {
signal.throwIfAborted()
const bindings = await this.getMcpBindings(signal)
const browserTools = this.getBrowserTools(context)
return [
...this.getBuiltinTools(),
...(browserTools?.listTools() ?? []),
...[...bindings.values()].map((binding) => binding.definition)
]
}
@@ -460,8 +649,17 @@ export class ModelToolProvider implements ModelToolProviderLike {
getApproval(
tool: ModelToolDefinition,
argumentsValue: Record<string, unknown>,
argumentSummary: string
argumentSummary: string,
context: ModelToolCallContext
): RuntimeApprovalRequest {
const browserTools = this.getBrowserTools(context)
if (browserTools?.ownsTool(tool.name)) {
return browserTools.getApproval(
tool,
argumentsValue,
argumentSummary
)
}
const path =
typeof argumentsValue.path === 'string'
? argumentsValue.path.slice(0, 500)
@@ -490,20 +688,38 @@ export class ModelToolProvider implements ModelToolProviderLike {
async callTool(
name: string,
argumentsValue: Record<string, unknown>,
signal: AbortSignal
): Promise<string> {
signal: AbortSignal,
context: ModelToolCallContext
): Promise<ModelToolResult> {
signal.throwIfAborted()
const browserTools = this.getBrowserTools(context)
if (browserTools?.ownsTool(name)) {
try {
return await browserTools.callTool(name, argumentsValue, signal)
} catch (error) {
if (error instanceof BrowserStaleReferenceError) {
throw new RecoverableModelToolError(
error.message,
'调用 browser_snapshot 获取新快照,然后用新引用重试刚才的操作',
{ cause: error }
)
}
throw error
}
}
if (name === 'workspace_read_text') {
const input = readInputSchema.parse(argumentsValue)
const filePath = await this.resolveExistingPath(input.path, 'file')
return (
await readBoundedUtf8File(
filePath,
MAX_READ_BYTES,
'工作区文本文件超过 256KB 安全限制',
'工作区读取目标不是有效 UTF-8 文本'
)
).content
return createTextToolResult(
(
await readBoundedUtf8File(
filePath,
MAX_READ_BYTES,
'工作区文本文件超过 256KB 安全限制',
'工作区读取目标不是有效 UTF-8 文本'
)
).content
)
}
if (name === 'workspace_list_directory') {
const input = listInputSchema.parse(argumentsValue)
@@ -515,21 +731,23 @@ export class ModelToolProvider implements ModelToolProviderLike {
directoryPath,
200
)
return boundedJson(
{
entries: listing.entries
.sort((left, right) => left.name.localeCompare(right.name))
.map((entry) => ({
name: entry.name,
type: entry.isDirectory()
? 'directory'
: entry.isFile()
? 'file'
: 'other'
})),
truncated: listing.truncated
},
'工作区目录结果无法序列化'
return createTextToolResult(
boundedJson(
{
entries: listing.entries
.sort((left, right) => left.name.localeCompare(right.name))
.map((entry) => ({
name: entry.name,
type: entry.isDirectory()
? 'directory'
: entry.isFile()
? 'file'
: 'other'
})),
truncated: listing.truncated
},
'工作区目录结果无法序列化'
)
)
}
if (name === 'workspace_write_text') {
@@ -551,12 +769,14 @@ export class ModelToolProvider implements ModelToolProviderLike {
await rm(temporaryPath, { force: true }).catch(() => undefined)
throw new Error('无法安全写入工作区文件', { cause: error })
}
return boundedJson(
{
path: input.path,
bytesWritten: Buffer.byteLength(input.content)
},
'工作区写入结果无法序列化'
return createTextToolResult(
boundedJson(
{
path: input.path,
bytesWritten: Buffer.byteLength(input.content)
},
'工作区写入结果无法序列化'
)
)
}
@@ -564,18 +784,52 @@ export class ModelToolProvider implements ModelToolProviderLike {
if (!binding) {
throw new Error('模型请求了未知工具')
}
const result = await binding.client.callTool(
{
name: binding.originalName,
arguments: argumentsValue
},
undefined,
{
timeout: MCP_TIMEOUT_MS,
signal
const params = {
name: binding.originalName,
arguments: argumentsValue
}
const options = {
timeout: MCP_TIMEOUT_MS,
signal,
onprogress: () => undefined,
resetTimeoutOnProgress: true,
maxTotalTimeout: MCP_CALL_MAX_TOTAL_TIMEOUT_MS
}
if (binding.definition.taskSupport !== 'required') {
return normalizeMcpResult(
await binding.client.callTool(params, undefined, options)
)
}
let taskId: string | undefined
try {
for await (const message of binding.client.experimental.tasks.callToolStream(
params,
undefined,
options
)) {
if (
(message.type === 'taskCreated' ||
message.type === 'taskStatus') &&
typeof message.task.taskId === 'string'
) {
taskId = message.task.taskId
} else if (message.type === 'result') {
return normalizeMcpResult(message.result)
} else if (message.type === 'error') {
throw message.error
}
}
)
return getMcpResultText(result)
throw new Error('MCP 任务工具未返回最终结果')
} catch (error) {
if (taskId) {
await binding.client.experimental.tasks.cancelTask(taskId, {
timeout: MCP_TASK_CANCEL_TIMEOUT_MS,
maxTotalTimeout: MCP_TASK_CANCEL_TIMEOUT_MS
}).catch(() => undefined)
}
throw error
}
}
async dispose(): Promise<void> {
@@ -584,4 +838,14 @@ export class ModelToolProvider implements ModelToolProviderLike {
this.mcpBindings = undefined
await Promise.allSettled(clients.map((client) => client.close()))
}
async releaseConversation(conversationId: string): Promise<void> {
if (!this.browserService) {
return
}
await new BrowserModelTools({
service: this.browserService,
conversationId
}).release()
}
}