Skip to content

Commit 4b2a0d6

Browse files
authored
Support Stream chat completion response (#1276)
1 parent 3578888 commit 4b2a0d6

2 files changed

Lines changed: 699 additions & 106 deletions

File tree

‎packages/server/server_tests/memmachine_server/common/language_model/test_openai_chat_completions_language_model.py‎

Lines changed: 367 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
"""Unit tests for OpenAIChatCompletionsLanguageModel."""
22

3+
from dataclasses import dataclass
4+
from typing import Any, cast
35
from unittest.mock import AsyncMock, MagicMock, patch
46

57
import httpx
@@ -9,7 +11,7 @@
911
from openai.types.chat.chat_completion_message_function_tool_call import (
1012
Function as ToolCallFunction,
1113
)
12-
from pydantic import ValidationError
14+
from pydantic import BaseModel, ValidationError
1315

1416
from memmachine_server.common.data_types import ExternalServiceAPIError
1517
from memmachine_server.common.language_model.openai_chat_completions_language_model import (
@@ -19,6 +21,38 @@
1921
from memmachine_server.common.metrics_factory import MetricsFactory
2022

2123

24+
class FakeAsyncStream:
25+
def __init__(self, chunks):
26+
self._chunks = iter(chunks)
27+
28+
def __aiter__(self):
29+
return self
30+
31+
async def __anext__(self):
32+
try:
33+
return next(self._chunks)
34+
except StopIteration as exc:
35+
raise StopAsyncIteration from exc
36+
37+
38+
@dataclass
39+
class FakeReasoningChunk:
40+
type: str = "response.reasoning.delta"
41+
42+
43+
class ParsedResponse(BaseModel):
44+
answer: str
45+
confidence: float
46+
47+
48+
def make_chat_completion(data: dict[str, Any]) -> openai_chat.ChatCompletion:
49+
return openai_chat.ChatCompletion.model_validate(data)
50+
51+
52+
def make_chat_completion_chunk(data: dict[str, Any]) -> openai_chat.ChatCompletionChunk:
53+
return openai_chat.ChatCompletionChunk.model_validate(data)
54+
55+
2256
@pytest.fixture
2357
def mock_metrics_factory():
2458
"""Fixture for a mocked MetricsFactory."""
@@ -216,6 +250,338 @@ async def test_generate_response_success(mock_async_openai, minimal_config):
216250
assert call_args.kwargs["store"] is False
217251

218252

253+
@pytest.mark.asyncio
254+
async def test_generate_parsed_response_success(mock_async_openai, minimal_config):
255+
"""Test a non-streamed structured response is parsed correctly."""
256+
mock_response = make_chat_completion(
257+
{
258+
"id": "completion-1",
259+
"object": "chat.completion",
260+
"created": 1,
261+
"model": "test-model",
262+
"choices": [
263+
{
264+
"index": 0,
265+
"finish_reason": "stop",
266+
"message": {
267+
"role": "assistant",
268+
"content": '{"answer":"ok","confidence":0.9}',
269+
"refusal": None,
270+
},
271+
}
272+
],
273+
"usage": {
274+
"prompt_tokens": 7,
275+
"completion_tokens": 5,
276+
"total_tokens": 12,
277+
},
278+
}
279+
)
280+
281+
mock_client = mock_async_openai.return_value
282+
mock_client.chat.completions.create.return_value = mock_response
283+
284+
lm = OpenAIChatCompletionsLanguageModel(minimal_config)
285+
parsed = await lm.generate_parsed_response(
286+
ParsedResponse,
287+
system_prompt="System prompt",
288+
user_prompt="User prompt",
289+
)
290+
291+
assert parsed == ParsedResponse(answer="ok", confidence=0.9)
292+
call_args = mock_client.chat.completions.create.call_args
293+
assert call_args.kwargs["model"] == "test-model"
294+
assert call_args.kwargs["messages"] == [
295+
{"role": "system", "content": "System prompt"},
296+
{"role": "user", "content": "User prompt"},
297+
]
298+
assert call_args.kwargs["response_format"]["type"] == "json_schema"
299+
300+
301+
@pytest.mark.asyncio
302+
async def test_generate_parsed_response_from_stream(
303+
mock_async_openai,
304+
minimal_config,
305+
):
306+
"""Test a streamed structured response is aggregated and validated."""
307+
streamed_response = FakeAsyncStream(
308+
[
309+
FakeReasoningChunk(),
310+
make_chat_completion_chunk(
311+
{
312+
"id": "chunk-1",
313+
"object": "chat.completion.chunk",
314+
"created": 1,
315+
"model": "test-model",
316+
"choices": [
317+
{
318+
"index": 0,
319+
"delta": {
320+
"role": "assistant",
321+
"content": '{"answer":"ok",',
322+
},
323+
"finish_reason": None,
324+
"logprobs": None,
325+
}
326+
],
327+
"usage": None,
328+
}
329+
),
330+
make_chat_completion_chunk(
331+
{
332+
"id": "chunk-2",
333+
"object": "chat.completion.chunk",
334+
"created": 2,
335+
"model": "test-model",
336+
"choices": [
337+
{
338+
"index": 0,
339+
"delta": {"content": '"confidence":0.9}'},
340+
"finish_reason": "stop",
341+
"logprobs": None,
342+
}
343+
],
344+
"usage": None,
345+
}
346+
),
347+
]
348+
)
349+
350+
mock_client = mock_async_openai.return_value
351+
mock_client.chat.completions.create.return_value = streamed_response
352+
353+
lm = OpenAIChatCompletionsLanguageModel(minimal_config)
354+
parsed = await lm.generate_parsed_response(ParsedResponse)
355+
356+
assert parsed == ParsedResponse(answer="ok", confidence=0.9)
357+
358+
359+
@pytest.mark.asyncio
360+
async def test_generate_parsed_response_from_stream_discards_reasoning_delta_chunks(
361+
mock_async_openai,
362+
minimal_config,
363+
):
364+
"""Test streamed structured responses ignore chunks with reasoning_content."""
365+
reasoning_chunk = make_chat_completion_chunk(
366+
{
367+
"id": "chunk-0",
368+
"object": "chat.completion.chunk",
369+
"created": 0,
370+
"model": "test-model",
371+
"choices": [
372+
{
373+
"index": 0,
374+
"delta": {"content": None},
375+
"finish_reason": None,
376+
"logprobs": None,
377+
}
378+
],
379+
"usage": None,
380+
}
381+
)
382+
cast(Any, reasoning_chunk.choices[0].delta).reasoning_content = "chain-of-thought"
383+
384+
streamed_response = FakeAsyncStream(
385+
[
386+
reasoning_chunk,
387+
make_chat_completion_chunk(
388+
{
389+
"id": "chunk-1",
390+
"object": "chat.completion.chunk",
391+
"created": 1,
392+
"model": "test-model",
393+
"choices": [
394+
{
395+
"index": 0,
396+
"delta": {
397+
"role": "assistant",
398+
"content": '{"answer":"ok",',
399+
},
400+
"finish_reason": None,
401+
"logprobs": None,
402+
}
403+
],
404+
"usage": None,
405+
}
406+
),
407+
make_chat_completion_chunk(
408+
{
409+
"id": "chunk-2",
410+
"object": "chat.completion.chunk",
411+
"created": 2,
412+
"model": "test-model",
413+
"choices": [
414+
{
415+
"index": 0,
416+
"delta": {"content": '"confidence":0.9}'},
417+
"finish_reason": "stop",
418+
"logprobs": None,
419+
}
420+
],
421+
"usage": None,
422+
}
423+
),
424+
]
425+
)
426+
427+
mock_client = mock_async_openai.return_value
428+
mock_client.chat.completions.create.return_value = streamed_response
429+
430+
lm = OpenAIChatCompletionsLanguageModel(minimal_config)
431+
parsed = await lm.generate_parsed_response(ParsedResponse)
432+
433+
assert parsed == ParsedResponse(answer="ok", confidence=0.9)
434+
435+
436+
@pytest.mark.asyncio
437+
async def test_generate_response_streamed_chat_completion_chunks(
438+
mock_async_openai,
439+
minimal_config,
440+
):
441+
"""Test a streamed chat completion is aggregated into a single response."""
442+
streamed_response = FakeAsyncStream(
443+
[
444+
make_chat_completion_chunk(
445+
{
446+
"id": "chunk-1",
447+
"object": "chat.completion.chunk",
448+
"created": 1,
449+
"model": "test-model",
450+
"choices": [
451+
{
452+
"index": 0,
453+
"delta": {"role": "assistant", "content": "Hello"},
454+
"finish_reason": None,
455+
"logprobs": None,
456+
}
457+
],
458+
"usage": None,
459+
}
460+
),
461+
make_chat_completion_chunk(
462+
{
463+
"id": "chunk-2",
464+
"object": "chat.completion.chunk",
465+
"created": 2,
466+
"model": "test-model",
467+
"choices": [
468+
{
469+
"index": 0,
470+
"delta": {"content": ", world!"},
471+
"finish_reason": "stop",
472+
"logprobs": None,
473+
}
474+
],
475+
"usage": {
476+
"prompt_tokens": 10,
477+
"completion_tokens": 4,
478+
"total_tokens": 14,
479+
},
480+
}
481+
),
482+
]
483+
)
484+
485+
mock_client = mock_async_openai.return_value
486+
mock_client.chat.completions.create.return_value = streamed_response
487+
488+
lm = OpenAIChatCompletionsLanguageModel(minimal_config)
489+
(
490+
content,
491+
tool_calls,
492+
input_tokens,
493+
output_tokens,
494+
) = await lm.generate_response_with_token_usage()
495+
496+
assert content == "Hello, world!"
497+
assert tool_calls == []
498+
assert input_tokens == 10
499+
assert output_tokens == 4
500+
501+
502+
@pytest.mark.asyncio
503+
async def test_generate_response_streamed_tool_calls_and_discards_reasoning_chunks(
504+
mock_async_openai,
505+
minimal_config,
506+
):
507+
"""Test streamed tool calls are reconstructed while reasoning chunks are ignored."""
508+
streamed_response = FakeAsyncStream(
509+
[
510+
FakeReasoningChunk(),
511+
make_chat_completion_chunk(
512+
{
513+
"id": "chunk-1",
514+
"object": "chat.completion.chunk",
515+
"created": 1,
516+
"model": "test-model",
517+
"choices": [
518+
{
519+
"index": 0,
520+
"delta": {
521+
"tool_calls": [
522+
{
523+
"index": 0,
524+
"id": "call_123",
525+
"type": "function",
526+
"function": {
527+
"name": "get_weather",
528+
"arguments": '{"location"',
529+
},
530+
}
531+
]
532+
},
533+
"finish_reason": None,
534+
"logprobs": None,
535+
}
536+
],
537+
"usage": None,
538+
}
539+
),
540+
make_chat_completion_chunk(
541+
{
542+
"id": "chunk-2",
543+
"object": "chat.completion.chunk",
544+
"created": 2,
545+
"model": "test-model",
546+
"choices": [
547+
{
548+
"index": 0,
549+
"delta": {
550+
"tool_calls": [
551+
{
552+
"index": 0,
553+
"function": {"arguments": ': "Boston"}'},
554+
}
555+
]
556+
},
557+
"finish_reason": "tool_calls",
558+
"logprobs": None,
559+
}
560+
],
561+
"usage": None,
562+
}
563+
),
564+
]
565+
)
566+
567+
mock_client = mock_async_openai.return_value
568+
mock_client.chat.completions.create.return_value = streamed_response
569+
570+
lm = OpenAIChatCompletionsLanguageModel(minimal_config)
571+
content, tool_calls = await lm.generate_response()
572+
573+
assert content == ""
574+
assert tool_calls == [
575+
{
576+
"call_id": "call_123",
577+
"function": {
578+
"name": "get_weather",
579+
"arguments": {"location": "Boston"},
580+
},
581+
}
582+
]
583+
584+
219585
@pytest.mark.asyncio
220586
async def test_generate_response_with_tool_calls(
221587
mock_async_openai,

0 commit comments

Comments
 (0)