fix: parse production learning agent streams
This commit is contained in:
1 parent
01bee3188b
commit
aa57dff677
2 files changed
+162
-15
No files matched your search
@@ -38,6 +38,63 @@ function record(value: unknown): Record<string, unknown> {
|
|||||||
return value && typeof value === 'object' && !Array.isArray(value) ? value as Record<string, unknown> : {};
|
return value && typeof value === 'object' && !Array.isArray(value) ? value as Record<string, unknown> : {};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function createSseDataParser() {
|
||||||
|
let line = '';
|
||||||
|
let pendingCarriageReturn = false;
|
||||||
|
let dataLines: string[] = [];
|
||||||
|
|
||||||
|
const commitLine = (events: string[]) => {
|
||||||
|
if (line === '') {
|
||||||
|
if (dataLines.length > 0) events.push(dataLines.join('\n'));
|
||||||
|
dataLines = [];
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if (!line.startsWith(':')) {
|
||||||
|
const separator = line.indexOf(':');
|
||||||
|
const field = separator === -1 ? line : line.slice(0, separator);
|
||||||
|
let value = separator === -1 ? '' : line.slice(separator + 1);
|
||||||
|
if (value.startsWith(' ')) value = value.slice(1);
|
||||||
|
if (field === 'data') dataLines.push(value);
|
||||||
|
}
|
||||||
|
line = '';
|
||||||
|
};
|
||||||
|
|
||||||
|
const push = (chunk: string): string[] => {
|
||||||
|
const events: string[] = [];
|
||||||
|
for (const character of chunk) {
|
||||||
|
if (pendingCarriageReturn) {
|
||||||
|
pendingCarriageReturn = false;
|
||||||
|
commitLine(events);
|
||||||
|
if (character === '\n') continue;
|
||||||
|
}
|
||||||
|
if (character === '\r') {
|
||||||
|
pendingCarriageReturn = true;
|
||||||
|
} else if (character === '\n') {
|
||||||
|
commitLine(events);
|
||||||
|
} else {
|
||||||
|
line += character;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return events;
|
||||||
|
};
|
||||||
|
|
||||||
|
const finish = (): string[] => {
|
||||||
|
const events: string[] = [];
|
||||||
|
if (pendingCarriageReturn) {
|
||||||
|
pendingCarriageReturn = false;
|
||||||
|
commitLine(events);
|
||||||
|
}
|
||||||
|
if (line !== '') commitLine(events);
|
||||||
|
if (dataLines.length > 0) {
|
||||||
|
events.push(dataLines.join('\n'));
|
||||||
|
dataLines = [];
|
||||||
|
}
|
||||||
|
return events;
|
||||||
|
};
|
||||||
|
|
||||||
|
return { push, finish };
|
||||||
|
}
|
||||||
|
|
||||||
export function createLearningAgentClient(dependencies: Dependencies = {}) {
|
export function createLearningAgentClient(dependencies: Dependencies = {}) {
|
||||||
const fetchImpl = dependencies.fetchImpl ?? proxyAwareFetch;
|
const fetchImpl = dependencies.fetchImpl ?? proxyAwareFetch;
|
||||||
const getAccessToken = dependencies.getAccessToken ?? getValidWorksSquareAccessToken;
|
const getAccessToken = dependencies.getAccessToken ?? getValidWorksSquareAccessToken;
|
||||||
@@ -129,28 +186,38 @@ export function createLearningAgentClient(dependencies: Dependencies = {}) {
|
|||||||
}
|
}
|
||||||
const reader = response.body.getReader();
|
const reader = response.body.getReader();
|
||||||
const decoder = new TextDecoder();
|
const decoder = new TextDecoder();
|
||||||
let buffer = '';
|
const parser = createSseDataParser();
|
||||||
let text = '';
|
let text = '';
|
||||||
|
const consume = (data: string): { text?: string; completed?: true } => {
|
||||||
|
if (!data) return {};
|
||||||
|
const envelope = record(JSON.parse(data));
|
||||||
|
if (envelope.run_id !== runId) return {};
|
||||||
|
const payload = record(envelope.payload);
|
||||||
|
if (envelope.type === 'learning.assistant.delta' && typeof payload.delta === 'string') {
|
||||||
|
text += payload.delta;
|
||||||
|
}
|
||||||
|
if (envelope.type === 'learning.assistant.failed') {
|
||||||
|
throw new Error(typeof payload.message === 'string' ? payload.message : '助教回答失败');
|
||||||
|
}
|
||||||
|
if (envelope.type === 'learning.assistant.completed') {
|
||||||
|
return { text: text.trim(), completed: true };
|
||||||
|
}
|
||||||
|
return {};
|
||||||
|
};
|
||||||
try {
|
try {
|
||||||
while (true) {
|
while (true) {
|
||||||
const chunk = await reader.read();
|
const chunk = await reader.read();
|
||||||
if (chunk.done) break;
|
const events = chunk.done
|
||||||
buffer += decoder.decode(chunk.value, { stream: true });
|
? [...parser.push(decoder.decode()), ...parser.finish()]
|
||||||
const frames = buffer.split('\n\n');
|
: parser.push(decoder.decode(chunk.value, { stream: true }));
|
||||||
buffer = frames.pop() || '';
|
for (const data of events) {
|
||||||
for (const frame of frames) {
|
const result = consume(data);
|
||||||
const data = frame.split('\n').find((line) => line.startsWith('data:'))?.slice(5).trim();
|
if (result.completed) {
|
||||||
if (!data) continue;
|
|
||||||
const envelope = record(JSON.parse(data));
|
|
||||||
if (envelope.run_id !== runId) continue;
|
|
||||||
const payload = record(envelope.payload);
|
|
||||||
if (envelope.type === 'learning.assistant.delta' && typeof payload.delta === 'string') text += payload.delta;
|
|
||||||
if (envelope.type === 'learning.assistant.failed') throw new Error(typeof payload.message === 'string' ? payload.message : '助教回答失败');
|
|
||||||
if (envelope.type === 'learning.assistant.completed') {
|
|
||||||
controller.abort();
|
controller.abort();
|
||||||
return { text: text.trim() };
|
return { text: result.text ?? '' };
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if (chunk.done) break;
|
||||||
}
|
}
|
||||||
} finally {
|
} finally {
|
||||||
clearTimeout(timeout);
|
clearTimeout(timeout);
|
||||||
|
|||||||
@@ -6,6 +6,32 @@ describe('Learning Agent client', () => {
|
|||||||
const fetchImpl = vi.fn<typeof fetch>();
|
const fetchImpl = vi.fn<typeof fetch>();
|
||||||
const getAccessToken = vi.fn();
|
const getAccessToken = vi.fn();
|
||||||
|
|
||||||
|
const chunkedStream = (...chunks: string[]) => new Response(new ReadableStream<Uint8Array>({
|
||||||
|
start(controller) {
|
||||||
|
const encoder = new TextEncoder();
|
||||||
|
for (const chunk of chunks) controller.enqueue(encoder.encode(chunk));
|
||||||
|
controller.close();
|
||||||
|
},
|
||||||
|
}), { status: 200, headers: { 'Content-Type': 'text/event-stream' } });
|
||||||
|
|
||||||
|
const mockRun = (stream: Response) => {
|
||||||
|
fetchImpl
|
||||||
|
.mockResolvedValueOnce(new Response(JSON.stringify({ session_id: 'session-1' }), { status: 201 }))
|
||||||
|
.mockResolvedValueOnce(new Response(JSON.stringify({ run_id: 'run-1' }), { status: 202 }))
|
||||||
|
.mockResolvedValueOnce(new Response(JSON.stringify({ stream_url: '/api/agents/sessions/session-1/events?ticket=ticket-1' }), { status: 200 }))
|
||||||
|
.mockResolvedValueOnce(stream);
|
||||||
|
};
|
||||||
|
|
||||||
|
const ask = () => createLearningAgentClient({
|
||||||
|
fetchImpl,
|
||||||
|
getAccessToken,
|
||||||
|
apiBaseUrl: 'https://square.example',
|
||||||
|
}).ask({
|
||||||
|
courseId: 'course-1',
|
||||||
|
contentHash: 'a'.repeat(64),
|
||||||
|
message: '为什么要学 Python?',
|
||||||
|
});
|
||||||
|
|
||||||
beforeEach(() => {
|
beforeEach(() => {
|
||||||
fetchImpl.mockReset();
|
fetchImpl.mockReset();
|
||||||
getAccessToken.mockReset();
|
getAccessToken.mockReset();
|
||||||
@@ -68,6 +94,60 @@ describe('Learning Agent client', () => {
|
|||||||
expect(fetchImpl.mock.calls.some((call) => call[1]?.method === 'DELETE')).toBe(false);
|
expect(fetchImpl.mock.calls.some((call) => call[1]?.method === 'DELETE')).toBe(false);
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it('parses production CRLF-delimited SSE frames', async () => {
|
||||||
|
const event = (type: string, payload: Record<string, unknown>) => (
|
||||||
|
`id: 1\r\nevent: agent.event\r\ndata: ${JSON.stringify({ run_id: 'run-1', type, payload })}\r\n\r\n`
|
||||||
|
);
|
||||||
|
mockRun(chunkedStream(
|
||||||
|
event('learning.assistant.delta', { delta: 'CRLF 正常。' }),
|
||||||
|
event('learning.assistant.completed', {}),
|
||||||
|
));
|
||||||
|
|
||||||
|
await expect(ask()).resolves.toEqual({ text: 'CRLF 正常。' });
|
||||||
|
});
|
||||||
|
|
||||||
|
it('parses a CRLF frame delimiter split across stream chunks', async () => {
|
||||||
|
const delta = JSON.stringify({
|
||||||
|
run_id: 'run-1',
|
||||||
|
type: 'learning.assistant.delta',
|
||||||
|
payload: { delta: '跨块正常。' },
|
||||||
|
});
|
||||||
|
const completed = JSON.stringify({
|
||||||
|
run_id: 'run-1',
|
||||||
|
type: 'learning.assistant.completed',
|
||||||
|
payload: {},
|
||||||
|
});
|
||||||
|
mockRun(chunkedStream(
|
||||||
|
`data: ${delta}\r\n\r`,
|
||||||
|
`\ndata: ${completed}\r`,
|
||||||
|
'\n\r',
|
||||||
|
'\n',
|
||||||
|
));
|
||||||
|
|
||||||
|
await expect(ask()).resolves.toEqual({ text: '跨块正常。' });
|
||||||
|
});
|
||||||
|
|
||||||
|
it('joins multiple data lines and ignores other SSE fields and comments', async () => {
|
||||||
|
const delta = JSON.stringify({
|
||||||
|
run_id: 'run-1',
|
||||||
|
type: 'learning.assistant.delta',
|
||||||
|
payload: { delta: '多行正常。' },
|
||||||
|
});
|
||||||
|
const split = delta.indexOf(',"type"');
|
||||||
|
const completed = JSON.stringify({
|
||||||
|
run_id: 'run-1',
|
||||||
|
type: 'learning.assistant.completed',
|
||||||
|
payload: {},
|
||||||
|
});
|
||||||
|
mockRun(chunkedStream(
|
||||||
|
': keep-alive\r\nid: 9\r\nevent: agent.event\r\nretry: 3000\r\n',
|
||||||
|
`data: ${delta.slice(0, split)},\r\ndata: ${delta.slice(split + 1)}\r\n\r\n`,
|
||||||
|
`data: ${completed}\r\n\r\n`,
|
||||||
|
));
|
||||||
|
|
||||||
|
await expect(ask()).resolves.toEqual({ text: '多行正常。' });
|
||||||
|
});
|
||||||
|
|
||||||
it('stops before the network for an invalid package hash', async () => {
|
it('stops before the network for an invalid package hash', async () => {
|
||||||
const client = createLearningAgentClient({ fetchImpl, getAccessToken });
|
const client = createLearningAgentClient({ fetchImpl, getAccessToken });
|
||||||
await expect(client.ask({ courseId: 'course-1', contentHash: 'bad', message: '你好' })).rejects.toThrow('助教请求无效');
|
await expect(client.ask({ courseId: 'course-1', contentHash: 'bad', message: '你好' })).rejects.toThrow('助教请求无效');
|
||||||
|
|||||||
Reference in new issue
Block a user