feat: add computer control and managed browser
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user