fix: bound model streaming

This commit is contained in:
lofyer
2026-08-13 03:18:55 +08:00
parent 980f3a0c8f
commit 5b579ae100
8 changed files with 1081 additions and 475 deletions
+48
View File
@@ -0,0 +1,48 @@
import { describe, expect, it } from 'vitest'
import { readBoundedResponseText } from './bounded-response'
describe('readBoundedResponseText', () => {
it('cancels an oversized response as soon as it crosses the byte limit', async () => {
const chunk = new Uint8Array(1024 * 1024)
let pulls = 0
const response = new Response(
new ReadableStream<Uint8Array>({
pull(controller) {
pulls += 1
controller.enqueue(chunk)
}
})
)
await expect(
readBoundedResponseText(response, {
maxBytes: 8 * 1024 * 1024,
tooLargeMessage: 'response too large'
})
).rejects.toThrow('response too large')
expect(pulls).toBeLessThan(20)
})
it('rejects an invalid declared response length without reading the body', async () => {
let pulls = 0
const response = new Response(
new ReadableStream<Uint8Array>({
pull(controller) {
pulls += 1
controller.enqueue(new Uint8Array([1]))
}
}),
{
headers: { 'content-length': 'invalid' }
}
)
await expect(
readBoundedResponseText(response, {
maxBytes: 1024,
tooLargeMessage: 'response too large'
})
).rejects.toThrow('response too large')
expect(pulls).toBe(0)
})
})
+54
View File
@@ -0,0 +1,54 @@
export type BoundedResponseTextOptions = {
maxBytes: number
missingBodyMessage?: string
tooLargeMessage: string
}
export async function readBoundedResponseText(
response: Response,
options: BoundedResponseTextOptions
): Promise<string> {
const declaredLength = response.headers.get('content-length')
if (declaredLength !== null) {
const parsedLength = Number(declaredLength)
if (
!Number.isSafeInteger(parsedLength) ||
parsedLength < 0 ||
parsedLength > options.maxBytes
) {
await response.body?.cancel().catch(() => undefined)
throw new Error(options.tooLargeMessage)
}
}
if (!response.body) {
if (options.missingBodyMessage) {
throw new Error(options.missingBodyMessage)
}
return ''
}
const reader = response.body.getReader()
const chunks: Uint8Array[] = []
let completed = false
let total = 0
try {
while (true) {
const { done, value } = await reader.read()
if (done) {
completed = true
break
}
total += value.byteLength
if (total > options.maxBytes) {
throw new Error(options.tooLargeMessage)
}
chunks.push(value)
}
} finally {
if (!completed) {
await reader.cancel().catch(() => undefined)
}
reader.releaseLock()
}
return Buffer.concat(chunks, total).toString('utf8')
}
+59 -2
View File
@@ -7,7 +7,7 @@ import {
writeFile writeFile
} from 'node:fs/promises' } from 'node:fs/promises'
import { existsSync, readFileSync } from 'node:fs' import { existsSync, readFileSync } from 'node:fs'
import { createHash } from 'node:crypto' import { createHash, randomUUID } from 'node:crypto'
import { createServer } from 'node:http' import { createServer } from 'node:http'
import { tmpdir } from 'node:os' import { tmpdir } from 'node:os'
import { join } from 'node:path' import { join } from 'node:path'
@@ -165,8 +165,14 @@ describe('ContinueHostAdapter', () => {
'let r=[eS.join(hu.continueHome,AKt)],o=' 'let r=[eS.join(hu.continueHome,AKt)],o='
) )
expect(bundle).toContain('goodbuddyEvents:[]') expect(bundle).toContain('goodbuddyEvents:[]')
expect(bundle).toContain('goodbuddyEventsBytes:0')
expect(bundle).toContain('goodbuddyEventsBytes+=Buffer.byteLength')
expect(bundle).toContain('goodbuddyEventsBytes<=2097152')
expect(bundle).toContain('l.length<=1e5')
expect(bundle).toContain('goodbuddyEventsOverflow:!1')
expect(bundle).toContain('goodbuddyEventsOverflow=!0')
expect(bundle).toContain('goodbuddyEvents:ce') expect(bundle).toContain('goodbuddyEvents:ce')
expect(bundle).toContain('type:"text",delta:u') expect(bundle).toContain('type:"text",delta:l')
expect(bundle).toContain('onToolStart?.(c.name,c.arguments,c.id)') expect(bundle).toContain('onToolStart?.(c.name,c.arguments,c.id)')
expect(bundle).toContain( expect(bundle).toContain(
'function ZZo(e){let t=[];if(e.allow)' 'function ZZo(e){let t=[];if(e.allow)'
@@ -994,6 +1000,57 @@ describe('ContinueHostAdapter', () => {
expect(killed).toBe(true) expect(killed).toBe(true)
}) })
it('fails when the patched host reports dropped stream events', async () => {
const distribution = await createDistribution()
let stateRequests = 0
vi.stubGlobal(
'fetch',
vi.fn(async (input: string | URL | Request) => {
if (String(input).endsWith('/state')) {
stateRequests += 1
return Response.json({
session: { history: [] },
isProcessing: stateRequests > 1,
messageQueueLength: 0,
pendingPermission: null,
goodbuddyEventsOverflow: stateRequests > 1
})
}
return Response.json({})
})
)
const adapter = new ContinueHostAdapter({
binaryPath: distribution.entryPath,
configPath: '',
workspace: process.cwd(),
cacheRoot: distribution.cacheRoot,
trustedBundleHashes: [distribution.sourceHash],
launchHost: () => ({
exitCode: null,
killed: false,
stderr: null,
once: () => undefined,
kill: () => true
}),
modelProfile: {
id: randomUUID(),
name: 'Local model',
baseUrl: 'http://127.0.0.1:11434/v1',
modelName: 'qwen3',
protocol: 'openai-chat-completions',
authentication: 'none'
}
})
await expect(
adapter.run(
'hello',
new AbortController().signal,
async () => 'deny'
)
).rejects.toThrow('流式事件超过安全限制')
})
it('uses auto mode and returns audit metadata for agent tools', async () => { it('uses auto mode and returns audit metadata for agent tools', async () => {
const distribution = await createDistribution() const distribution = await createDistribution()
let launchArgs: string[] = [] let launchArgs: string[] = []
+59 -23
View File
@@ -38,6 +38,7 @@ import {
safeToolErrorDetail safeToolErrorDetail
} from './approval-summary' } from './approval-summary'
import { stageRuntimeSkillPackages } from './runtime-skill-packages' import { stageRuntimeSkillPackages } from './runtime-skill-packages'
import { readBoundedResponseText } from './bounded-response'
const supportedVersion = '1.5.47' const supportedVersion = '1.5.47'
const supportedBundleHashes = new Set([ const supportedBundleHashes = new Set([
@@ -49,6 +50,8 @@ const maximumMessageBytes = 20 * 1024 * 1024
const maximumConfigBytes = 1024 * 1024 const maximumConfigBytes = 1024 * 1024
const maximumConfiguredMcpServers = 100 const maximumConfiguredMcpServers = 100
const maximumStreamEvents = 5_000 const maximumStreamEvents = 5_000
const maximumStreamEventBytes = 2 * 1024 * 1024
const maximumExecutionMilliseconds = 10 * 60_000
const knowledgeMcpName = 'goodbuddy-knowledge' const knowledgeMcpName = 'goodbuddy-knowledge'
export const continueConfigurationRequiredMessage = export const continueConfigurationRequiredMessage =
'Continue 尚未配置模型连接,请在设置中选择 GoodBuddy 模型连接或指定 Continue 配置文件' 'Continue 尚未配置模型连接,请在设置中选择 GoodBuddy 模型连接或指定 Continue 配置文件'
@@ -116,7 +119,8 @@ const stateSchema = z.object({
goodbuddyEvents: z goodbuddyEvents: z
.array(continueHostStreamEventSchema) .array(continueHostStreamEventSchema)
.max(maximumStreamEvents) .max(maximumStreamEvents)
.optional() .optional(),
goodbuddyEventsOverflow: z.boolean().optional()
}) })
type ContinueHostState = z.infer<typeof stateSchema> type ContinueHostState = z.infer<typeof stateSchema>
@@ -704,17 +708,17 @@ export class ContinueHostAdapter {
patched = replaceExactly( patched = replaceExactly(
patched, patched,
streamCallbacksMarker, streamCallbacksMarker,
'a={onContent:u=>{u&&e.goodbuddyEvents.length<5e3&&e.goodbuddyEvents.push({type:"text",delta:u})},onContentComplete:u=>{},onToolStart:(u,l,c)=>{c&&e.goodbuddyEvents.length<5e3&&e.goodbuddyEvents.push({type:"tool",callId:c,name:u,state:"running",input:(()=>{try{return JSON.stringify(l).slice(0,4e3)}catch{return"[无法序列化]"}})()})},onToolResult:(u,l,c,d)=>{d&&e.goodbuddyEvents.length<5e3&&e.goodbuddyEvents.push({type:"tool",callId:d,name:l,state:c==="done"?"completed":"failed",output:String(u).slice(0,16e3)})},onToolError:(u,l,c)=>{c&&e.goodbuddyEvents.length<5e3&&e.goodbuddyEvents.push({type:"tool",callId:c,name:l??"unknown",state:"failed",error:String(u).slice(0,1e3)})},onToolPermissionRequest:' 'a={onContent:u=>{if(!u)return;let l=String(u);e.goodbuddyEventsBytes+=Buffer.byteLength(l);l.length<=1e5&&e.goodbuddyEvents.length<5e3&&e.goodbuddyEventsBytes<=2097152?e.goodbuddyEvents.push({type:"text",delta:l}):e.goodbuddyEventsOverflow=!0},onContentComplete:u=>{},onToolStart:(u,l,c)=>{if(!c)return;let d={type:"tool",callId:c,name:u,state:"running",input:(()=>{try{return JSON.stringify(l).slice(0,4e3)}catch{return"[无法序列化]"}})()};e.goodbuddyEventsBytes+=Buffer.byteLength(JSON.stringify(d));e.goodbuddyEvents.length<5e3&&e.goodbuddyEventsBytes<=2097152?e.goodbuddyEvents.push(d):e.goodbuddyEventsOverflow=!0},onToolResult:(u,l,c,d)=>{if(!d)return;let p={type:"tool",callId:d,name:l,state:c==="done"?"completed":"failed",output:String(u).slice(0,16e3)};e.goodbuddyEventsBytes+=Buffer.byteLength(JSON.stringify(p));e.goodbuddyEvents.length<5e3&&e.goodbuddyEventsBytes<=2097152?e.goodbuddyEvents.push(p):e.goodbuddyEventsOverflow=!0},onToolError:(u,l,c)=>{if(!c)return;let d={type:"tool",callId:c,name:l??"unknown",state:"failed",error:String(u).slice(0,1e3)};e.goodbuddyEventsBytes+=Buffer.byteLength(JSON.stringify(d));e.goodbuddyEvents.length<5e3&&e.goodbuddyEventsBytes<=2097152?e.goodbuddyEvents.push(d):e.goodbuddyEventsOverflow=!0},onToolPermissionRequest:'
) )
patched = replaceExactly( patched = replaceExactly(
patched, patched,
serverStateMarker, serverStateMarker,
'pendingPermission:null,goodbuddyEvents:[]},B=' 'pendingPermission:null,goodbuddyEvents:[],goodbuddyEventsBytes:0,goodbuddyEventsOverflow:!1},B='
) )
patched = replaceExactly( patched = replaceExactly(
patched, patched,
serverStateEndpointMarker, serverStateEndpointMarker,
'j.get("/state",(we,Te)=>{M.lastActivity=Date.now(),B();let ue=e7e(M.session,M.isProcessing,rS.getQueueLength(),M.pendingPermission),ce=M.goodbuddyEvents.splice(0);Te.json({...ue,goodbuddyEvents:ce})})' 'j.get("/state",(we,Te)=>{M.lastActivity=Date.now(),B();let ue=e7e(M.session,M.isProcessing,rS.getQueueLength(),M.pendingPermission),ce=M.goodbuddyEvents.splice(0),de=M.goodbuddyEventsOverflow;M.goodbuddyEventsBytes=0,M.goodbuddyEventsOverflow=!1;Te.json({...ue,goodbuddyEvents:ce,goodbuddyEventsOverflow:de})})'
) )
patched = replaceExactly( patched = replaceExactly(
patched, patched,
@@ -839,14 +843,10 @@ export class ContinueHostAdapter {
redirect: 'error', redirect: 'error',
signal: init.signal signal: init.signal
}) })
const contentLength = Number(response.headers.get('content-length') ?? 0) const body = await readBoundedResponseText(response, {
if (contentLength > maximumStateBytes) { maxBytes: maximumStateBytes,
throw new Error('Continue 宿主响应超过安全大小限制') tooLargeMessage: 'Continue 宿主响应超过安全大小限制'
} })
const body = await response.text()
if (Buffer.byteLength(body) > maximumStateBytes) {
throw new Error('Continue 宿主响应超过安全大小限制')
}
if (!response.ok) { if (!response.ok) {
throw new Error(`Continue 宿主请求失败(HTTP ${response.status}`) throw new Error(`Continue 宿主请求失败(HTTP ${response.status}`)
} }
@@ -861,8 +861,11 @@ export class ContinueHostAdapter {
signal: AbortSignal signal: AbortSignal
): Promise<ContinueHostState> { ): Promise<ContinueHostState> {
const expiresAt = Date.now() + 30_000 const expiresAt = Date.now() + 30_000
const timeoutSignal = AbortSignal.timeout(30_000)
const startupSignal = AbortSignal.any([signal, timeoutSignal])
try {
while (Date.now() < expiresAt) { while (Date.now() < expiresAt) {
signal.throwIfAborted() startupSignal.throwIfAborted()
const childFailure = getChildFailure() const childFailure = getChildFailure()
if (childFailure) { if (childFailure) {
throw childFailure throw childFailure
@@ -872,12 +875,20 @@ export class ContinueHostAdapter {
} }
try { try {
return stateSchema.parse( return stateSchema.parse(
await this.request(origin, token, '/state', { signal }) await this.request(origin, token, '/state', {
signal: startupSignal
})
) )
} catch { } catch {
await delay(150, signal) await delay(150, startupSignal)
} }
} }
} catch (error) {
if (timeoutSignal.aborted && !signal.aborted) {
throw new Error('Continue 宿主启动超时', { cause: error })
}
throw error
}
throw new Error('Continue 宿主启动超时') throw new Error('Continue 宿主启动超时')
} }
@@ -1150,6 +1161,7 @@ export class ContinueHostAdapter {
let observedTools: ContinueHostTool[] = [] let observedTools: ContinueHostTool[] = []
let streamedText = false let streamedText = false
let executionTimeoutSignal: AbortSignal | undefined
try { try {
const initialState = await this.waitForStartup( const initialState = await this.waitForStartup(
child, child,
@@ -1159,6 +1171,13 @@ export class ContinueHostAdapter {
signal signal
) )
const startIndex = initialState.session.history.length const startIndex = initialState.session.history.length
executionTimeoutSignal = AbortSignal.timeout(
maximumExecutionMilliseconds
)
const executionSignal = AbortSignal.any([
signal,
executionTimeoutSignal
])
const message = const message =
runOptions.images && runOptions.images.length > 0 runOptions.images && runOptions.images.length > 0
? [ ? [
@@ -1178,13 +1197,13 @@ export class ContinueHostAdapter {
await this.request(origin, token, '/message', { await this.request(origin, token, '/message', {
method: 'POST', method: 'POST',
body: messageBody, body: messageBody,
signal signal: executionSignal
}) })
const expiresAt = Date.now() + 10 * 60_000 const expiresAt = Date.now() + maximumExecutionMilliseconds
const handledPermissionIds = new Set<string>() const handledPermissionIds = new Set<string>()
while (Date.now() < expiresAt) { while (Date.now() < expiresAt) {
signal.throwIfAborted() executionSignal.throwIfAborted()
if (childFailure) { if (childFailure) {
throw childFailure throw childFailure
} }
@@ -1194,8 +1213,19 @@ export class ContinueHostAdapter {
) )
} }
const state = stateSchema.parse( const state = stateSchema.parse(
await this.request(origin, token, '/state', { signal }) await this.request(origin, token, '/state', {
signal: executionSignal
})
) )
if (state.goodbuddyEventsOverflow) {
throw new Error('Continue 宿主流式事件超过安全限制')
}
const streamEventBytes = Buffer.byteLength(
JSON.stringify(state.goodbuddyEvents ?? [])
)
if (streamEventBytes > maximumStreamEventBytes) {
throw new Error('Continue 宿主流式事件超过安全限制')
}
observedTools = mergeContinueTools( observedTools = mergeContinueTools(
observedTools, observedTools,
extractContinueTools(state.session.history, startIndex) extractContinueTools(state.session.history, startIndex)
@@ -1265,7 +1295,7 @@ export class ContinueHostAdapter {
requestId: pending.requestId, requestId: pending.requestId,
approved: decision !== 'deny' approved: decision !== 'deny'
}), }),
signal signal: executionSignal
}) })
} }
if ( if (
@@ -1308,16 +1338,22 @@ export class ContinueHostAdapter {
: {}) : {})
} }
} }
await delay(150, signal) await delay(150, executionSignal)
} }
throw new Error('Continue 宿主执行超时') throw new Error('Continue 宿主执行超时')
} catch (error) { } catch (error) {
if (error instanceof ContinueHostRunError) { if (error instanceof ContinueHostRunError) {
throw error throw error
} }
const normalizedError =
executionTimeoutSignal?.aborted && !signal.aborted
? new Error('Continue 宿主执行超时', { cause: error })
: error
throw new ContinueHostRunError( throw new ContinueHostRunError(
error instanceof Error ? error.message : 'Continue 宿主执行失败', normalizedError instanceof Error
{ cause: error, tools: observedTools } ? normalizedError.message
: 'Continue 宿主执行失败',
{ cause: normalizedError, tools: observedTools }
) )
} finally { } finally {
signal.removeEventListener('abort', abort) signal.removeEventListener('abort', abort)
+85
View File
@@ -1,5 +1,6 @@
import { beforeEach, describe, expect, it, vi } from 'vitest' import { beforeEach, describe, expect, it, vi } from 'vitest'
import type { RuntimeEvent } from './runtime' import type { RuntimeEvent } from './runtime'
import { randomUUID } from 'node:crypto'
import { import {
ContinueHostRunError, ContinueHostRunError,
type ContinueHostAdapterOptions type ContinueHostAdapterOptions
@@ -632,6 +633,90 @@ describe('ContinueAgentRuntime', () => {
]) ])
}) })
it('fails instead of silently dropping an overflowing stream queue', async () => {
mocks.runHost.mockImplementation(
async (
_prompt,
_signal,
_authorize,
options
) => {
for (let index = 0; index < 1_001; index += 1) {
options?.onEvent?.({
type: 'text',
delta: String(index)
})
}
return { text: 'done', streamedText: true }
}
)
const stream = createRuntime().run(
{
requestId: randomUUID(),
conversationId: 'overflow-conversation',
prompt: 'test'
},
new AbortController().signal
)
await expect(async () => {
for await (const _event of stream) {
void _event
}
}).rejects.toThrow('流式事件积压超过安全限制')
})
it('aborts the host run when stream consumption ends early', async () => {
let resolveHost: (() => void) | undefined
const hostFinished = new Promise<void>((resolve) => {
resolveHost = resolve
})
let hostSignal: AbortSignal | undefined
mocks.runHost.mockImplementation(
async (
_prompt,
signal,
_authorize,
options
) => {
hostSignal = signal
await options?.onEvent?.({
type: 'text',
delta: 'partial'
})
await new Promise<void>((resolve) => {
signal.addEventListener(
'abort',
() => {
resolve()
resolveHost?.()
},
{ once: true }
)
})
throw signal.reason
}
)
const stream = createRuntime().run(
{
requestId: randomUUID(),
conversationId: 'early-close-conversation',
prompt: 'test'
},
new AbortController().signal
)
await expect(stream.next()).resolves.toMatchObject({
value: { type: 'status' }
})
await expect(stream.next()).resolves.toMatchObject({
value: { type: 'text', delta: 'partial' }
})
await stream.return()
await hostFinished
expect(hostSignal?.aborted).toBe(true)
})
it('emits terminal tool audits before a failed Continue run', async () => { it('emits terminal tool audits before a failed Continue run', async () => {
mocks.runHost.mockRejectedValue( mocks.runHost.mockRejectedValue(
new ContinueHostRunError('Continue failed', { new ContinueHostRunError('Continue failed', {
+26 -3
View File
@@ -51,6 +51,7 @@ export type ContinueRuntimeOptions = {
// The prompt reaches the Continue host through a local HTTP POST body, so no // The prompt reaches the Continue host through a local HTTP POST body, so no
// platform command-line limit applies to it. // platform command-line limit applies to it.
const MAX_CONTINUE_PROMPT_CHARACTERS = 128_000 const MAX_CONTINUE_PROMPT_CHARACTERS = 128_000
const MAX_QUEUED_STREAM_EVENTS = 1_000
const scopedReadToolNameSet = new Set<string>(scopedReadToolNames) const scopedReadToolNameSet = new Set<string>(scopedReadToolNames)
function continueToolFailureMessage(tool: ContinueHostTool): string { function continueToolFailureMessage(tool: ContinueHostTool): string {
@@ -331,7 +332,12 @@ export class ContinueAgentRuntime implements AgentRuntime {
let streamFinished = false let streamFinished = false
let streamResult: ContinueHostRunResult | undefined let streamResult: ContinueHostRunResult | undefined
let streamError: unknown let streamError: unknown
const hostController = new AbortController()
const hostSignal = AbortSignal.any([signal, hostController.signal])
const onEvent = (event: ContinueHostStreamEvent): void => { const onEvent = (event: ContinueHostStreamEvent): void => {
if (queuedEvents.length >= MAX_QUEUED_STREAM_EVENTS) {
throw new Error('Continue 流式事件积压超过安全限制')
}
queuedEvents.push(event) queuedEvents.push(event)
wakeStream?.() wakeStream?.()
wakeStream = undefined wakeStream = undefined
@@ -339,7 +345,7 @@ export class ContinueAgentRuntime implements AgentRuntime {
const hostRun = host const hostRun = host
.run( .run(
conversationContext, conversationContext,
signal, hostSignal,
authorize, authorize,
{ {
workMode: request.workMode, workMode: request.workMode,
@@ -361,7 +367,7 @@ export class ContinueAgentRuntime implements AgentRuntime {
wakeStream?.() wakeStream?.()
wakeStream = undefined wakeStream = undefined
}) })
try {
while (!streamFinished || queuedEvents.length > 0) { while (!streamFinished || queuedEvents.length > 0) {
if (queuedEvents.length === 0) { if (queuedEvents.length === 0) {
await new Promise<void>((resolve) => { await new Promise<void>((resolve) => {
@@ -385,7 +391,12 @@ export class ContinueAgentRuntime implements AgentRuntime {
false false
) )
} }
} finally {
hostController.abort(new Error('Continue 流式消费已结束'))
wakeStream?.()
wakeStream = undefined
await hostRun await hostRun
}
if (streamError) { if (streamError) {
throw streamError throw streamError
} }
@@ -396,7 +407,19 @@ export class ContinueAgentRuntime implements AgentRuntime {
} catch (error) { } catch (error) {
if (error instanceof ContinueHostRunError) { if (error instanceof ContinueHostRunError) {
for (const tool of error.tools) { for (const tool of error.tools) {
yield toContinueToolEvent(request.requestId, tool, true) const terminalEvent = toContinueToolEvent(
request.requestId,
tool,
true
)
const previous = emittedTools.get(tool.callId)
if (
!previous ||
previous.state !== tool.state ||
previous.error !== tool.error
) {
yield terminalEvent
}
} }
} }
throw error throw error
+222
View File
@@ -321,6 +321,228 @@ describe('ModelAgentRuntime', () => {
await expect(consume()).rejects.toThrow('意外中断') await expect(consume()).rejects.toThrow('意外中断')
}) })
it('rejects malformed SSE JSON instead of silently skipping it', async () => {
const runtime = new ModelAgentRuntime({
apiKey: 'test-key',
baseUrl: 'https://bigtoken.ai',
model: 'sonnet-5',
protocol: 'anthropic-messages',
authentication: 'api-key',
fetcher: vi.fn<typeof fetch>(async () =>
new Response('data: {invalid}\n\n', {
status: 200,
headers: { 'content-type': 'text/event-stream' }
})
)
})
const consume = async (): Promise<void> => {
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'test'
},
new AbortController().signal
)) {
void _event
}
}
await expect(consume()).rejects.toThrow('无效的流式 JSON')
})
it('parses CRLF event separators split across response chunks', async () => {
const payload = createEventStream('split CRLF').replaceAll('\n', '\r\n')
const splitAt = payload.indexOf('\r\n\r\n') + 3
const body = new ReadableStream<Uint8Array>({
start(controller) {
controller.enqueue(
new TextEncoder().encode(payload.slice(0, splitAt))
)
controller.enqueue(
new TextEncoder().encode(payload.slice(splitAt))
)
controller.close()
}
})
const runtime = new ModelAgentRuntime({
apiKey: 'test-key',
baseUrl: 'https://bigtoken.ai',
model: 'sonnet-5',
protocol: 'anthropic-messages',
authentication: 'api-key',
fetcher: vi.fn<typeof fetch>(async () =>
new Response(body, {
status: 200,
headers: { 'content-type': 'text/event-stream' }
})
)
})
const events = []
for await (const event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'test'
},
new AbortController().signal
)) {
events.push(event)
}
expect(events).toContainEqual(
expect.objectContaining({
type: 'text',
delta: 'split CRLF'
})
)
})
it('aborts a model request that exceeds the runtime timeout', async () => {
vi.useFakeTimers()
try {
const runtime = new ModelAgentRuntime({
apiKey: 'test-key',
baseUrl: 'https://bigtoken.ai',
model: 'sonnet-5',
protocol: 'anthropic-messages',
authentication: 'api-key',
requestTimeoutMs: 50,
fetcher: vi.fn<typeof fetch>(
async (_input, init) =>
new Promise<Response>((_resolve, reject) => {
init?.signal?.addEventListener(
'abort',
() => reject(init.signal?.reason),
{ once: true }
)
})
)
})
const stream = runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'test'
},
new AbortController().signal
)
await expect(stream.next()).resolves.toMatchObject({
value: { type: 'status' }
})
const result = stream.next()
const assertion = expect(result).rejects.toThrow(
'模型接口请求超时'
)
await vi.advanceTimersByTimeAsync(50)
await assertion
} finally {
vi.useRealTimers()
}
})
it('aborts a stalled response body after headers arrive', async () => {
vi.useFakeTimers()
try {
let responseSignal: AbortSignal | null | undefined
const runtime = new ModelAgentRuntime({
apiKey: 'test-key',
baseUrl: 'https://bigtoken.ai',
model: 'sonnet-5',
protocol: 'anthropic-messages',
authentication: 'api-key',
requestTimeoutMs: 50,
fetcher: vi.fn<typeof fetch>(async (_input, init) => {
responseSignal = init?.signal
return new Response(
new ReadableStream<Uint8Array>({
start(controller) {
init?.signal?.addEventListener(
'abort',
() => controller.error(init.signal?.reason),
{ once: true }
)
}
}),
{
status: 200,
headers: { 'content-type': 'text/event-stream' }
}
)
})
})
const stream = runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'test'
},
new AbortController().signal
)
await expect(stream.next()).resolves.toMatchObject({
value: { type: 'status' }
})
const result = stream.next()
const assertion = expect(result).rejects.toThrow(
'模型接口请求超时'
)
await vi.advanceTimersByTimeAsync(50)
await assertion
expect(responseSignal?.aborted).toBe(true)
} finally {
vi.useRealTimers()
}
})
it('bounds the total ordinary streaming response size', async () => {
const chunk = new TextEncoder().encode(
`data: ${JSON.stringify({
type: 'content_block_delta',
delta: {
type: 'text_delta',
text: 'x'.repeat(65_000)
}
})}\n\n`
)
const body = new ReadableStream<Uint8Array>({
pull(controller) {
controller.enqueue(chunk)
}
})
const runtime = new ModelAgentRuntime({
apiKey: 'test-key',
baseUrl: 'https://bigtoken.ai',
model: 'sonnet-5',
protocol: 'anthropic-messages',
authentication: 'api-key',
fetcher: vi.fn<typeof fetch>(async () =>
new Response(body, {
status: 200,
headers: { 'content-type': 'text/event-stream' }
})
)
})
const consume = async (): Promise<void> => {
for await (const _event of runtime.run(
{
requestId: crypto.randomUUID(),
conversationId: crypto.randomUUID(),
prompt: 'test'
},
new AbortController().signal
)) {
void _event
}
}
await expect(consume()).rejects.toThrow(
'流式响应超过安全限制'
)
})
it('preserves bounded provider error messages', async () => { it('preserves bounded provider error messages', async () => {
const runtime = new ModelAgentRuntime({ const runtime = new ModelAgentRuntime({
apiKey: 'test-key', apiKey: 'test-key',
+229 -148
View File
@@ -37,6 +37,7 @@ import {
boundedToolDetail, boundedToolDetail,
safeToolErrorDetail safeToolErrorDetail
} from './approval-summary' } from './approval-summary'
import { readBoundedResponseText } from './bounded-response'
type ConversationMessage = { type ConversationMessage = {
role: 'user' | 'assistant' role: 'user' | 'assistant'
@@ -98,12 +99,14 @@ type ModelToolResponse = {
const maxGeneratedImageBytes = 3_900_000 const maxGeneratedImageBytes = 3_900_000
const maxImageResponseBytes = 5_300_000 const maxImageResponseBytes = 5_300_000
const maxChatResponseBytes = 2 * 1024 * 1024 const maxChatResponseBytes = 2 * 1024 * 1024
const maxStreamBlockBytes = 1024 * 1024
const maxToolArgumentBytes = 128 * 1024 const maxToolArgumentBytes = 128 * 1024
const maxToolContextBytes = 1024 * 1024 const maxToolContextBytes = 1024 * 1024
const maxToolCallsPerRun = 40 const maxToolCallsPerRun = 40
const maxToolRounds = 24 const maxToolRounds = 24
const maxRepeatedIdenticalCalls = 3 const maxRepeatedIdenticalCalls = 3
const maxIdenticalRoundsWithoutProgress = 2 const maxIdenticalRoundsWithoutProgress = 2
const defaultModelRequestTimeoutMs = 10 * 60_000
function getCurrentTimeInstruction(now = new Date()): string { function getCurrentTimeInstruction(now = new Date()): string {
const systemTime = [ const systemTime = [
@@ -138,6 +141,7 @@ export type ModelRuntimeOptions = {
webSearchEnabled?: boolean webSearchEnabled?: boolean
toolProvider?: ModelToolProviderLike toolProvider?: ModelToolProviderLike
fetcher?: typeof fetch fetcher?: typeof fetch
requestTimeoutMs?: number
} }
function getErrorMessage(value: unknown): string | undefined { function getErrorMessage(value: unknown): string | undefined {
@@ -412,33 +416,48 @@ function createUsageEvent(
} }
} }
async function readBoundedText( function createRequestSignal(
response: Response, signal: AbortSignal,
maxBytes: number timeoutMs: number
): Promise<string> { ): {
if (!response.body) { signal: AbortSignal
throw new Error('模型接口未返回响应内容') clear: () => void
timedOut: () => boolean
} {
const controller = new AbortController()
let timedOut = false
const abortFromCaller = (): void => {
controller.abort(signal.reason)
} }
const reader = response.body.getReader() signal.addEventListener('abort', abortFromCaller, { once: true })
const chunks: Uint8Array[] = [] const timeout = setTimeout(() => {
let total = 0 if (controller.signal.aborted) {
try { return
while (true) {
const { done, value } = await reader.read()
if (done) {
break
} }
total += value.byteLength timedOut = true
if (total > maxBytes) { controller.abort(new Error('模型接口请求超时'))
await reader.cancel().catch(() => undefined) }, timeoutMs)
throw new Error('模型接口响应超过安全限制') if (signal.aborted) {
abortFromCaller()
} }
chunks.push(value) return {
signal: controller.signal,
clear: () => {
clearTimeout(timeout)
signal.removeEventListener('abort', abortFromCaller)
},
timedOut: () => timedOut
} }
} finally {
reader.releaseLock()
} }
return Buffer.concat(chunks, total).toString('utf8')
function normalizeRequestError(
error: unknown,
timedOut: boolean
): never {
if (timedOut) {
throw new Error('模型接口请求超时', { cause: error })
}
throw error
} }
function parseGeneratedImage(value: unknown): { function parseGeneratedImage(value: unknown): {
@@ -898,6 +917,14 @@ function parseModelToolResponse(
} }
} }
function getSseData(block: string): string {
return block
.split('\n')
.filter((line) => line.startsWith('data:'))
.map((line) => line.slice(5).trimStart())
.join('\n')
}
function parseStreamBlock( function parseStreamBlock(
block: string, block: string,
protocol: ModelProtocol protocol: ModelProtocol
@@ -907,11 +934,7 @@ function parseStreamBlock(
stopped: boolean stopped: boolean
usage?: ModelUsageUpdate usage?: ModelUsageUpdate
} { } {
const data = block const data = getSseData(block)
.split('\n')
.filter((line) => line.startsWith('data:'))
.map((line) => line.slice(5).trimStart())
.join('\n')
if (!data) { if (!data) {
return { stopped: false } return { stopped: false }
} }
@@ -925,8 +948,10 @@ function parseStreamBlock(
let event: unknown let event: unknown
try { try {
event = JSON.parse(data) event = JSON.parse(data)
} catch { } catch (error) {
return { stopped: false } throw new Error('模型接口返回了无效的流式 JSON', {
cause: error
})
} }
const error = getErrorMessage(event) const error = getErrorMessage(event)
if (error) { if (error) {
@@ -985,11 +1010,7 @@ function parseStreamBlock(
function parseSseData( function parseSseData(
block: string block: string
): { event?: unknown; stopped: boolean } { ): { event?: unknown; stopped: boolean } {
const data = block const data = getSseData(block)
.split('\n')
.filter((line) => line.startsWith('data:'))
.map((line) => line.slice(5).trimStart())
.join('\n')
if (!data) { if (!data) {
return { stopped: false } return { stopped: false }
} }
@@ -998,8 +1019,56 @@ function parseSseData(
} }
try { try {
return { event: JSON.parse(data), stopped: false } return { event: JSON.parse(data), stopped: false }
} catch { } catch (error) {
return { stopped: false } throw new Error('模型接口返回了无效的流式 JSON', {
cause: error
})
}
}
async function* readBoundedSseBlocks(
response: Response
): AsyncGenerator<string, void, void> {
if (!response.body) {
throw new Error('模型接口未返回流式响应')
}
const reader = response.body.getReader()
const decoder = new TextDecoder()
let buffer = ''
let completed = false
let receivedBytes = 0
try {
while (true) {
const { done, value } = await reader.read()
receivedBytes += value?.byteLength ?? 0
if (receivedBytes > maxChatResponseBytes) {
throw new Error('模型接口流式响应超过安全限制')
}
buffer += decoder.decode(value, { stream: !done })
buffer = buffer.replaceAll('\r\n', '\n')
if (Buffer.byteLength(buffer) > maxStreamBlockBytes) {
throw new Error('模型接口流式响应块超过安全限制')
}
const blocks = buffer.split('\n\n')
buffer = blocks.pop() ?? ''
if (done && buffer.trim()) {
blocks.push(buffer)
buffer = ''
}
for (const block of blocks) {
yield block
}
if (done) {
completed = true
break
}
}
} finally {
if (!completed) {
await reader.cancel().catch(() => undefined)
}
reader.releaseLock()
} }
} }
@@ -1010,9 +1079,18 @@ export class ModelAgentRuntime implements AgentRuntime {
private readonly knownConversationIds = new Set<string>() private readonly knownConversationIds = new Set<string>()
private readonly fetcher: typeof fetch private readonly fetcher: typeof fetch
private readonly toolProvider: ModelToolProviderLike private readonly toolProvider: ModelToolProviderLike
private readonly requestTimeoutMs: number
constructor(private readonly options: ModelRuntimeOptions) { constructor(private readonly options: ModelRuntimeOptions) {
this.fetcher = options.fetcher ?? fetch this.fetcher = options.fetcher ?? fetch
this.requestTimeoutMs =
options.requestTimeoutMs ?? defaultModelRequestTimeoutMs
if (
!Number.isSafeInteger(this.requestTimeoutMs) ||
this.requestTimeoutMs < 1
) {
throw new Error('模型接口请求超时设置无效')
}
this.toolProvider = this.toolProvider =
options.toolProvider ?? options.toolProvider ??
new ModelToolProvider( new ModelToolProvider(
@@ -1073,6 +1151,31 @@ export class ModelAgentRuntime implements AgentRuntime {
return headers return headers
} }
private async fetchWithTimeout(
input: URL,
init: RequestInit,
signal: AbortSignal
): Promise<{
response: Response
clear: () => void
timedOut: () => boolean
}> {
const request = createRequestSignal(signal, this.requestTimeoutMs)
try {
return {
response: await this.fetcher(input, {
...init,
signal: request.signal
}),
clear: request.clear,
timedOut: request.timedOut
}
} catch (error) {
request.clear()
return normalizeRequestError(error, request.timedOut())
}
}
async getStatus(): Promise<AgentRuntimeStatus> { async getStatus(): Promise<AgentRuntimeStatus> {
const imageGeneration = this.capability === 'image-generation' const imageGeneration = this.capability === 'image-generation'
return { return {
@@ -1123,9 +1226,16 @@ export class ModelAgentRuntime implements AgentRuntime {
) )
}) })
if (!response.ok) { if (!response.ok) {
const responseText = await readBoundedResponseText(response, {
maxBytes: 128 * 1024,
missingBodyMessage: '模型接口未返回响应内容',
tooLargeMessage: '模型接口响应超过安全限制'
})
let detail: string | undefined let detail: string | undefined
try { try {
detail = getErrorMessage(await response.json()) detail = getErrorMessage(
responseText.trim() ? JSON.parse(responseText) : undefined
)
} catch { } catch {
detail = undefined detail = undefined
} }
@@ -1276,21 +1386,33 @@ export class ModelAgentRuntime implements AgentRuntime {
model: this.options.model, model: this.options.model,
prompt: request.prompt.slice(0, 100_000), prompt: request.prompt.slice(0, 100_000),
n: 1, n: 1,
quality: quality: this.options.imageGenerationQuality ?? 'auto',
this.options.imageGenerationQuality ??
'auto',
response_format: 'b64_json' response_format: 'b64_json'
} }
const response = await this.fetcher(this.getEndpoint(), { const modelRequest = await this.fetchWithTimeout(
this.getEndpoint(),
{
method: 'POST', method: 'POST',
headers: this.getHeaders(), headers: this.getHeaders(),
body: JSON.stringify(imageRequest), body: JSON.stringify(imageRequest)
},
signal signal
})
const responseText = await readBoundedText(
response,
response.ok ? maxImageResponseBytes : 128 * 1024
) )
const response = modelRequest.response
let responseText: string
try {
responseText = await readBoundedResponseText(response, {
maxBytes: response.ok
? maxImageResponseBytes
: 128 * 1024,
missingBodyMessage: '模型接口未返回响应内容',
tooLargeMessage: '模型接口响应超过安全限制'
})
} catch (error) {
return normalizeRequestError(error, modelRequest.timedOut())
} finally {
modelRequest.clear()
}
if (!response.ok) { if (!response.ok) {
let errorPayload: unknown let errorPayload: unknown
try { try {
@@ -1421,17 +1543,23 @@ export class ModelAgentRuntime implements AgentRuntime {
if (Buffer.byteLength(body) > 2 * 1024 * 1024) { if (Buffer.byteLength(body) > 2 * 1024 * 1024) {
throw new Error('模型工具请求上下文超过 2MB 安全限制') throw new Error('模型工具请求上下文超过 2MB 安全限制')
} }
const response = await this.fetcher(this.getEndpoint(), { const request = await this.fetchWithTimeout(
this.getEndpoint(),
{
method: 'POST', method: 'POST',
headers: this.getHeaders(), headers: this.getHeaders(),
body, body
},
signal signal
})
if (!response.ok) {
const responseText = await readBoundedText(
response,
128 * 1024
) )
const response = request.response
try {
if (!response.ok) {
const responseText = await readBoundedResponseText(response, {
maxBytes: 128 * 1024,
missingBodyMessage: '模型接口未返回响应内容',
tooLargeMessage: '模型接口响应超过安全限制'
})
let detail: string | undefined let detail: string | undefined
try { try {
detail = getErrorMessage( detail = getErrorMessage(
@@ -1443,8 +1571,7 @@ export class ModelAgentRuntime implements AgentRuntime {
detail = undefined detail = undefined
} }
throw new Error( throw new Error(
detail ?? detail ?? `模型接口请求失败(HTTP ${response.status}`
`模型接口请求失败(HTTP ${response.status}`
) )
} }
if ( if (
@@ -1454,11 +1581,6 @@ export class ModelAgentRuntime implements AgentRuntime {
?.toLocaleLowerCase() ?.toLocaleLowerCase()
.includes('text/event-stream') .includes('text/event-stream')
) { ) {
if (!response.body) {
throw new Error('模型接口未返回流式响应')
}
const reader = response.body.getReader()
const decoder = new TextDecoder()
const streamedToolCalls = new Map< const streamedToolCalls = new Map<
number, number,
{ arguments: string; id: string; name: string } { arguments: string; id: string; name: string }
@@ -1468,33 +1590,9 @@ export class ModelAgentRuntime implements AgentRuntime {
} }
let answer = '' let answer = ''
let reasoning = '' let reasoning = ''
let buffer = ''
let receivedStop = false let receivedStop = false
let receivedBytes = 0
let streamEnded = false
try { for await (const block of readBoundedSseBlocks(response)) {
while (!receivedStop) {
const { done, value } = await reader.read()
streamEnded = done
receivedBytes += value?.byteLength ?? 0
if (receivedBytes > maxChatResponseBytes) {
throw new Error('模型接口流式响应超过安全限制')
}
buffer += decoder.decode(value, { stream: !done }).replaceAll(
'\r\n',
'\n'
)
if (Buffer.byteLength(buffer) > maxChatResponseBytes) {
throw new Error('模型接口流式响应块超过安全限制')
}
const blocks = buffer.split('\n\n')
buffer = blocks.pop() ?? ''
if (done && buffer.trim()) {
blocks.push(buffer)
buffer = ''
}
for (const block of blocks) {
const parsed = parseSseData(block) const parsed = parseSseData(block)
if (parsed.stopped) { if (parsed.stopped) {
receivedStop = true receivedStop = true
@@ -1511,9 +1609,7 @@ export class ModelAgentRuntime implements AgentRuntime {
usage, usage,
getUsageUpdate(parsed.event, 'openai') getUsageUpdate(parsed.event, 'openai')
) )
const reasoningDelta = getOpenAIReasoningDelta( const reasoningDelta = getOpenAIReasoningDelta(parsed.event)
parsed.event
)
if (reasoningDelta) { if (reasoningDelta) {
reasoning += reasoningDelta reasoning += reasoningDelta
yield { yield {
@@ -1590,16 +1686,6 @@ export class ModelAgentRuntime implements AgentRuntime {
streamedToolCalls.set(index as number, next) streamedToolCalls.set(index as number, next)
} }
} }
if (done) {
break
}
}
} finally {
if (!streamEnded) {
await reader.cancel().catch(() => undefined)
}
reader.releaseLock()
}
if (!receivedStop) { if (!receivedStop) {
throw new Error('模型接口流式响应意外中断') throw new Error('模型接口流式响应意外中断')
} }
@@ -1640,17 +1726,20 @@ export class ModelAgentRuntime implements AgentRuntime {
streamed: true streamed: true
} }
} }
const responseText = await readBoundedText( const responseText = await readBoundedResponseText(response, {
response, maxBytes: maxChatResponseBytes,
maxChatResponseBytes missingBodyMessage: '模型接口未返回响应内容',
) tooLargeMessage: '模型接口响应超过安全限制'
})
let payload: unknown let payload: unknown
try { try {
payload = responseText.trim() payload = responseText.trim()
? JSON.parse(responseText) ? JSON.parse(responseText)
: undefined : undefined
} catch (error) { } catch (error) {
throw new Error('模型接口返回了无效 JSON', { cause: error }) throw new Error('模型接口返回了无效 JSON', {
cause: error
})
} }
const providerError = getErrorMessage(payload) const providerError = getErrorMessage(payload)
if (providerError) { if (providerError) {
@@ -1664,6 +1753,11 @@ export class ModelAgentRuntime implements AgentRuntime {
? 'anthropic' ? 'anthropic'
: 'openai' : 'openai'
) )
} catch (error) {
return normalizeRequestError(error, request.timedOut())
} finally {
request.clear()
}
} }
private async *runToolExecution( private async *runToolExecution(
@@ -1742,10 +1836,18 @@ export class ModelAgentRuntime implements AgentRuntime {
request.requestId request.requestId
) )
let responseStep = await responseStream.next() let responseStep = await responseStream.next()
try {
while (!responseStep.done) { while (!responseStep.done) {
yield responseStep.value yield responseStep.value
responseStep = await responseStream.next() responseStep = await responseStream.next()
} }
} finally {
if (!responseStep.done) {
await responseStream
.throw(new Error('模型流式消费已结束'))
.catch(() => undefined)
}
}
const response = responseStep.value const response = responseStep.value
const usage = { const usage = {
reported: false reported: false
@@ -2079,7 +2181,9 @@ export class ModelAgentRuntime implements AgentRuntime {
: responses : responses
? this.getResponsesInput(request) ? this.getResponsesInput(request)
: this.getOpenAIMessages(request, system) : this.getOpenAIMessages(request, system)
const response = await this.fetcher(this.getEndpoint(), { const modelRequest = await this.fetchWithTimeout(
this.getEndpoint(),
{
method: 'POST', method: 'POST',
headers: this.getHeaders(), headers: this.getHeaders(),
body: JSON.stringify( body: JSON.stringify(
@@ -2108,14 +2212,25 @@ export class ModelAgentRuntime implements AgentRuntime {
}, },
messages messages
} }
), )
},
signal signal
}) )
const response = modelRequest.response
try {
if (!response.ok) { if (!response.ok) {
const responseText = await readBoundedResponseText(response, {
maxBytes: 128 * 1024,
missingBodyMessage: '模型接口未返回响应内容',
tooLargeMessage: '模型接口响应超过安全限制'
})
let detail: string | undefined let detail: string | undefined
try { try {
detail = getErrorMessage(await response.json()) detail = getErrorMessage(
responseText.trim()
? JSON.parse(responseText)
: undefined
)
} catch { } catch {
detail = undefined detail = undefined
} }
@@ -2124,41 +2239,13 @@ export class ModelAgentRuntime implements AgentRuntime {
) )
} }
if (!response.body) {
throw new Error('模型接口未返回流式响应')
}
const reader = response.body.getReader()
const decoder = new TextDecoder()
let buffer = ''
let answer = '' let answer = ''
let receivedStop = false let receivedStop = false
let streamEnded = false
const usage = { const usage = {
reported: false reported: false
} satisfies ModelUsageAccumulator } satisfies ModelUsageAccumulator
try { for await (const block of readBoundedSseBlocks(response)) {
while (!receivedStop) {
const { done, value } = await reader.read()
streamEnded = done
buffer += decoder.decode(value, { stream: !done }).replaceAll(
'\r\n',
'\n'
)
if (Buffer.byteLength(buffer) > 1024 * 1024) {
throw new Error('模型接口流式响应块超过安全限制')
}
const blocks = buffer.split('\n\n')
buffer = blocks.pop() ?? ''
if (done && buffer.trim()) {
blocks.push(buffer)
buffer = ''
}
for (const block of blocks) {
const parsed = parseStreamBlock(block, this.options.protocol) const parsed = parseStreamBlock(block, this.options.protocol)
if (parsed.usage) { if (parsed.usage) {
applyUsageUpdate(usage, parsed.usage) applyUsageUpdate(usage, parsed.usage)
@@ -2186,17 +2273,6 @@ export class ModelAgentRuntime implements AgentRuntime {
} }
} }
if (done) {
break
}
}
} finally {
if (!streamEnded) {
await reader.cancel().catch(() => undefined)
}
reader.releaseLock()
}
if (!receivedStop) { if (!receivedStop) {
throw new Error('模型接口流式响应意外中断') throw new Error('模型接口流式响应意外中断')
} }
@@ -2225,6 +2301,11 @@ export class ModelAgentRuntime implements AgentRuntime {
requestId: request.requestId, requestId: request.requestId,
type: 'done' type: 'done'
} }
} catch (error) {
return normalizeRequestError(error, modelRequest.timedOut())
} finally {
modelRequest.clear()
}
} }
async dispose(): Promise<void> { async dispose(): Promise<void> {