|
1 | 1 | """Unit tests for OpenAIChatCompletionsLanguageModel.""" |
2 | 2 |
|
| 3 | +from dataclasses import dataclass |
| 4 | +from typing import Any, cast |
3 | 5 | from unittest.mock import AsyncMock, MagicMock, patch |
4 | 6 |
|
5 | 7 | import httpx |
|
9 | 11 | from openai.types.chat.chat_completion_message_function_tool_call import ( |
10 | 12 | Function as ToolCallFunction, |
11 | 13 | ) |
12 | | -from pydantic import ValidationError |
| 14 | +from pydantic import BaseModel, ValidationError |
13 | 15 |
|
14 | 16 | from memmachine_server.common.data_types import ExternalServiceAPIError |
15 | 17 | from memmachine_server.common.language_model.openai_chat_completions_language_model import ( |
|
19 | 21 | from memmachine_server.common.metrics_factory import MetricsFactory |
20 | 22 |
|
21 | 23 |
|
| 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 | + |
22 | 56 | @pytest.fixture |
23 | 57 | def mock_metrics_factory(): |
24 | 58 | """Fixture for a mocked MetricsFactory.""" |
@@ -216,6 +250,338 @@ async def test_generate_response_success(mock_async_openai, minimal_config): |
216 | 250 | assert call_args.kwargs["store"] is False |
217 | 251 |
|
218 | 252 |
|
| 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 | + |
219 | 585 | @pytest.mark.asyncio |
220 | 586 | async def test_generate_response_with_tool_calls( |
221 | 587 | mock_async_openai, |
|
0 commit comments