| |
| |
|
|
| from unittest.mock import MagicMock, patch |
|
|
| import pytest |
| from openai_harmony import Author, Message, Role, StreamState, TextContent |
|
|
| from vllm.entrypoints.openai.responses.context import ( |
| HarmonyContext, |
| SimpleContext, |
| StreamingHarmonyContext, |
| TurnMetrics, |
| ) |
| from vllm.outputs import CompletionOutput, RequestOutput |
|
|
|
|
| def create_mock_request_output( |
| prompt_token_ids=None, |
| output_token_ids=None, |
| num_cached_tokens=0, |
| finished=True, |
| ): |
| """Helper function to create a mock RequestOutput object for testing.""" |
| outputs = [] |
| token_ids = output_token_ids if output_token_ids is not None else [] |
| outputs = [ |
| CompletionOutput( |
| index=0, |
| text="Test output", |
| token_ids=token_ids, |
| cumulative_logprob=0.0, |
| logprobs=None, |
| finish_reason=None, |
| stop_reason=None, |
| ) |
| ] |
|
|
| return RequestOutput( |
| request_id="test-id", |
| prompt="Test prompt", |
| prompt_token_ids=prompt_token_ids, |
| prompt_logprobs=None, |
| outputs=outputs, |
| finished=finished, |
| num_cached_tokens=num_cached_tokens, |
| ) |
|
|
|
|
| async def generate_mock_outputs( |
| num_turns, prompt_token_counts, output_token_counts, cached_token_counts=None |
| ): |
| """Generate a sequence of mock RequestOutput objects to simulate multiple |
| turns.""" |
| if cached_token_counts is None: |
| cached_token_counts = [0] * num_turns |
|
|
| for i in range(num_turns): |
| |
| prompt_token_ids = list(range(1, prompt_token_counts[i] + 1)) |
| output_token_ids = list(range(1, output_token_counts[i] + 1)) |
|
|
| |
| yield create_mock_request_output( |
| prompt_token_ids=prompt_token_ids, |
| output_token_ids=output_token_ids, |
| num_cached_tokens=cached_token_counts[i], |
| ) |
|
|
|
|
| @pytest.fixture |
| def mock_parser(): |
| """Set up a mock parser for tests.""" |
| with patch( |
| "vllm.entrypoints.openai.responses.context.get_streamable_parser_for_assistant" |
| ) as mock_parser_factory: |
| |
| parser = MagicMock() |
| parser.messages = [] |
| parser.current_channel = None |
| parser.state = StreamState.EXPECT_START |
| mock_parser_factory.return_value = parser |
| yield parser |
|
|
|
|
| def test_single_turn_token_counting(): |
| """Test token counting behavior for a single turn.""" |
| |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| |
| mock_output = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3, 4, 5], |
| output_token_ids=[6, 7, 8], |
| num_cached_tokens=2, |
| ) |
|
|
| |
| context.append_output(mock_output) |
|
|
| |
| assert context.num_prompt_tokens == 5 |
| assert context.num_output_tokens == 3 |
| assert context.num_cached_tokens == 2 |
| assert context.num_tool_output_tokens == 0 |
|
|
| |
| assert not context.is_first_turn |
| assert len(context.all_turn_metrics) == 1 |
| previous_turn = context.all_turn_metrics[0] |
| assert previous_turn.input_tokens == 5 |
| assert previous_turn.output_tokens == 3 |
| assert previous_turn.cached_input_tokens == 2 |
| assert previous_turn.tool_output_tokens == 0 |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_multi_turn_token_counting(): |
| """Test token counting behavior across multiple turns with tool output.""" |
| |
| context = HarmonyContext(messages=[], available_tools=["browser"]) |
|
|
| |
| |
| |
| |
| prompt_token_counts = [5, 15, 20] |
| output_token_counts = [3, 4, 5] |
| cached_token_counts = [0, 5, 15] |
| mock_generator = generate_mock_outputs( |
| 3, prompt_token_counts, output_token_counts, cached_token_counts |
| ) |
|
|
| |
| mock_output1 = await anext(mock_generator) |
| context.append_output(mock_output1) |
|
|
| |
| assert context.num_prompt_tokens == 5 |
| assert context.num_output_tokens == 3 |
| assert context.num_tool_output_tokens == 0 |
|
|
| |
| mock_output2 = await anext(mock_generator) |
| context.append_output(mock_output2) |
| |
| |
| expected_tool_output = 7 |
|
|
| assert context.num_prompt_tokens == 5 + 15 |
| assert context.num_output_tokens == 3 + 4 |
| assert context.num_tool_output_tokens == expected_tool_output |
| assert context.num_cached_tokens == 5 |
|
|
| |
| mock_output3 = await anext(mock_generator) |
| context.append_output(mock_output3) |
| |
| |
| |
| expected_tool_output = 7 + 1 |
|
|
| assert context.num_prompt_tokens == 5 + 15 + 20 |
| assert context.num_output_tokens == 3 + 4 + 5 |
| assert context.num_tool_output_tokens == expected_tool_output |
| assert context.num_cached_tokens == 5 + 15 |
|
|
| |
| assert len(context.all_turn_metrics) == 3 |
| for i, turn in enumerate(context.all_turn_metrics): |
| assert turn.input_tokens == prompt_token_counts[i] |
| assert turn.output_tokens == output_token_counts[i] |
| assert turn.cached_input_tokens == cached_token_counts[i] |
| assert context.all_turn_metrics[1].tool_output_tokens == 7 |
| assert context.all_turn_metrics[2].tool_output_tokens == 1 |
|
|
|
|
| def test_empty_output_tokens(): |
| """Test behavior when RequestOutput has empty output tokens.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| |
| mock_output = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], |
| output_token_ids=[], |
| num_cached_tokens=1, |
| ) |
|
|
| context.append_output(mock_output) |
|
|
| |
| assert context.num_prompt_tokens == 3 |
| assert context.num_output_tokens == 0 |
| assert context.num_cached_tokens == 1 |
| assert context.num_tool_output_tokens == 0 |
|
|
|
|
| def test_missing_prompt_token_ids(): |
| """Test behavior when RequestOutput has None prompt_token_ids.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| mock_output = create_mock_request_output( |
| prompt_token_ids=None, |
| output_token_ids=[1, 2], |
| num_cached_tokens=0, |
| ) |
|
|
| |
| |
| context.append_output(mock_output) |
|
|
| |
| assert context.num_prompt_tokens == 0 |
| assert context.num_output_tokens == 2 |
| assert context.num_cached_tokens == 0 |
| assert context.num_tool_output_tokens == 0 |
|
|
|
|
| def test_reasoning_tokens_counting(mock_parser): |
| """Test that reasoning tokens are counted correctly.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| |
| mock_parser.current_channel = "analysis" |
|
|
| mock_output = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], |
| output_token_ids=[4, 5, 6, 7], |
| num_cached_tokens=0, |
| ) |
|
|
| context.append_output(mock_output) |
|
|
| |
| assert context.num_reasoning_tokens == 4 |
| assert context.num_output_tokens == 4 |
|
|
|
|
| def test_preamble_tokens_not_counted_as_reasoning(mock_parser): |
| """Preambles (commentary with no recipient) are visible user text, |
| not hidden reasoning. They must NOT inflate num_reasoning_tokens.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| mock_parser.current_channel = "commentary" |
| mock_parser.current_recipient = None |
|
|
| mock_output = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], |
| output_token_ids=[4, 5, 6], |
| num_cached_tokens=0, |
| ) |
| context.append_output(mock_output) |
|
|
| assert context.num_reasoning_tokens == 0 |
| assert context.num_output_tokens == 3 |
|
|
|
|
| def test_commentary_with_recipient_counted_as_reasoning(mock_parser): |
| """Commentary directed at a tool (recipient != None) is hidden from |
| the user, so it should still count as reasoning tokens.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| mock_parser.current_channel = "commentary" |
| mock_parser.current_recipient = "python" |
|
|
| mock_output = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], |
| output_token_ids=[4, 5, 6], |
| num_cached_tokens=0, |
| ) |
| context.append_output(mock_output) |
|
|
| assert context.num_reasoning_tokens == 3 |
| assert context.num_output_tokens == 3 |
|
|
|
|
| def test_zero_tokens_edge_case(): |
| """Test behavior with all zero token counts.""" |
| context = HarmonyContext(messages=[], available_tools=[]) |
|
|
| |
| |
| mock_output = create_mock_request_output( |
| prompt_token_ids=[], |
| output_token_ids=[], |
| num_cached_tokens=0, |
| ) |
|
|
| context.append_output(mock_output) |
|
|
| |
| assert context.num_prompt_tokens == 0 |
| assert context.num_output_tokens == 0 |
| assert context.num_cached_tokens == 0 |
| assert context.num_tool_output_tokens == 0 |
| assert context.num_reasoning_tokens == 0 |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_single_turn_no_tool_output(): |
| """Test that first turn never generates tool output tokens.""" |
| context = HarmonyContext( |
| messages=[], |
| available_tools=["browser"], |
| ) |
|
|
| |
| mock_output = create_mock_request_output( |
| prompt_token_ids=list(range(100)), |
| output_token_ids=[1, 2, 3], |
| num_cached_tokens=0, |
| ) |
|
|
| context.append_output(mock_output) |
|
|
| |
| assert context.num_tool_output_tokens == 0 |
| assert context.is_first_turn is False |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_negative_tool_tokens_edge_case(): |
| """Test edge case where calculation could result in negative tool |
| tokens. We should log an error and clamp the value to 0.""" |
| |
| with patch("vllm.entrypoints.openai.responses.context.logger.error") as mock_log: |
| context = HarmonyContext(messages=[], available_tools=["browser"]) |
|
|
| |
| mock_output1 = create_mock_request_output( |
| prompt_token_ids=list(range(10)), |
| output_token_ids=[1, 2, 3, 4, 5], |
| ) |
| context.append_output(mock_output1) |
|
|
| |
| |
| mock_output2 = create_mock_request_output( |
| prompt_token_ids=list(range(12)), |
| output_token_ids=[6, 7], |
| ) |
| context.append_output(mock_output2) |
|
|
| |
| |
| assert context.num_tool_output_tokens == 0 |
| assert context.num_prompt_tokens == 10 + 12 |
| assert context.num_output_tokens == 5 + 2 |
|
|
| |
| mock_log.assert_called_once() |
|
|
| |
| args, _ = mock_log.call_args |
| log_message = args[0] |
|
|
| |
| assert "Negative tool output tokens calculated" in log_message |
| assert "-3" in str(args) |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_streaming_multi_turn_token_counting(mock_parser): |
| """Test token counting for streaming multi-turn conversations. |
| |
| This test focuses on how StreamingHarmonyContext counts tokens in a |
| multi-turn conversation with streaming (token-by-token) outputs and |
| message boundaries. |
| """ |
| |
| context = StreamingHarmonyContext(messages=[], available_tools=["browser"]) |
|
|
| num_prompt_tokens = [3, 8, 13] |
| num_output_tokens = [3, 3, 2] |
| num_cached_tokens = [0, 3, 8] |
|
|
| |
| |
| |
| |
|
|
| |
| |
| context.append_output( |
| create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], |
| output_token_ids=[101], |
| num_cached_tokens=num_cached_tokens[0], |
| finished=False, |
| ) |
| ) |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| output_token_ids=[102], |
| finished=False, |
| ) |
| ) |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| output_token_ids=[103], |
| finished=True, |
| ) |
| ) |
|
|
| |
| assert context.num_prompt_tokens == 3 |
| assert context.num_output_tokens == 3 |
| assert context.num_cached_tokens == 0 |
| assert context.num_tool_output_tokens == 0 |
| assert context.first_tok_of_message is True |
|
|
| |
| mock_parser.current_channel = "analysis" |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| prompt_token_ids=[ |
| 1, |
| 2, |
| 3, |
| 101, |
| 102, |
| 103, |
| 4, |
| 5, |
| ], |
| output_token_ids=[201], |
| num_cached_tokens=num_cached_tokens[1], |
| finished=False, |
| ) |
| ) |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| output_token_ids=[202], |
| finished=False, |
| ) |
| ) |
|
|
| context.append_output( |
| create_mock_request_output( |
| output_token_ids=[203], |
| finished=True, |
| ) |
| ) |
|
|
| |
| assert context.num_prompt_tokens == 3 + 8 |
| assert context.num_output_tokens == 3 + 3 |
| assert context.num_reasoning_tokens == 3 |
| assert context.num_cached_tokens == 3 |
|
|
| |
| expected_tool_tokens = 8 - 3 - 3 |
| assert context.num_tool_output_tokens == expected_tool_tokens |
|
|
| |
| mock_parser.current_channel = "final" |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| prompt_token_ids=[ |
| 1, |
| 2, |
| 3, |
| 101, |
| 102, |
| 103, |
| 4, |
| 5, |
| 201, |
| 202, |
| 203, |
| 6, |
| 7, |
| ], |
| output_token_ids=[301], |
| num_cached_tokens=num_cached_tokens[2], |
| finished=False, |
| ) |
| ) |
|
|
| context.append_output( |
| create_mock_request_output( |
| output_token_ids=[302], |
| finished=True, |
| ) |
| ) |
|
|
| |
| assert context.num_prompt_tokens == sum(num_prompt_tokens) |
| assert context.num_output_tokens == sum(num_output_tokens) |
| assert context.num_reasoning_tokens == 3 |
| assert context.num_cached_tokens == sum( |
| num_cached_tokens |
| ) |
|
|
| |
| |
| additional_tool_tokens = 13 - 8 - 3 |
| assert ( |
| context.num_tool_output_tokens == expected_tool_tokens + additional_tool_tokens |
| ) |
|
|
| |
| assert len(context.all_turn_metrics) == 3 |
| for i, turn in enumerate(context.all_turn_metrics): |
| assert turn.input_tokens == num_prompt_tokens[i] |
| assert turn.output_tokens == num_output_tokens[i] |
| assert turn.cached_input_tokens == num_cached_tokens[i] |
| assert context.all_turn_metrics[1].tool_output_tokens == 2 |
| assert context.all_turn_metrics[2].tool_output_tokens == 2 |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_streaming_message_synchronization(mock_parser): |
| """Test message synchronization logic from lines 413-417 in context.py. |
| |
| This test verifies that when parser.messages contains more messages than |
| the context's _messages (minus initial messages), the context properly |
| extends its message list with the new parser messages. |
| """ |
|
|
| |
| initial_messages = [ |
| Message( |
| author=Author(role=Role.USER, name="user"), |
| content=[TextContent(text="Hello")], |
| recipient=Role.ASSISTANT, |
| ) |
| ] |
| context = StreamingHarmonyContext(messages=initial_messages, available_tools=[]) |
|
|
| |
| assert len(context._messages) == 1 |
| assert context.num_init_messages == 1 |
|
|
| |
| |
| mock_parser.messages = [ |
| Message( |
| author=Author(role=Role.ASSISTANT, name="assistant"), |
| content=[TextContent(text="Response 1")], |
| recipient=Role.USER, |
| ), |
| ] |
|
|
| |
| context.append_output( |
| create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], output_token_ids=[101], finished=False |
| ) |
| ) |
|
|
| |
| assert len(context._messages) == 2 |
|
|
| |
| assert context._messages[1].content[0].text == "Response 1" |
|
|
| |
| |
| messages_minus_init = len(context._messages) - context.num_init_messages |
| parser_messages_count = len(mock_parser.messages) |
|
|
| |
| assert messages_minus_init == parser_messages_count |
|
|
| |
| mock_parser.messages.append( |
| Message( |
| author=Author(role=Role.ASSISTANT, name="assistant"), |
| content=[TextContent(text="Response 4")], |
| recipient=Role.USER, |
| ) |
| ) |
|
|
| |
| mock_output2 = create_mock_request_output( |
| prompt_token_ids=[1, 2, 3], output_token_ids=[102], finished=True |
| ) |
|
|
| context.append_output(mock_output2) |
|
|
| |
| assert len(context._messages) == 3 |
| assert context.num_init_messages == 1 |
| assert context._messages[2].content[0].text == "Response 4" |
|
|
|
|
| def test_turn_metrics_copy_and_reset(): |
| """Test TurnMetrics copy and reset methods work correctly.""" |
| |
| original_metrics = TurnMetrics( |
| input_tokens=10, |
| output_tokens=20, |
| cached_input_tokens=5, |
| tool_output_tokens=3, |
| ) |
|
|
| |
| copied_metrics = original_metrics.copy() |
|
|
| |
| assert copied_metrics.input_tokens == 10 |
| assert copied_metrics.output_tokens == 20 |
| assert copied_metrics.cached_input_tokens == 5 |
| assert copied_metrics.tool_output_tokens == 3 |
|
|
| |
| assert copied_metrics is not original_metrics |
|
|
| |
| copied_metrics.input_tokens = 999 |
| assert original_metrics.input_tokens == 10 |
| assert copied_metrics.input_tokens == 999 |
|
|
| |
| original_metrics.reset() |
|
|
| |
| assert original_metrics.input_tokens == 0 |
| assert original_metrics.output_tokens == 0 |
| assert original_metrics.cached_input_tokens == 0 |
| assert original_metrics.tool_output_tokens == 0 |
|
|
| |
| assert copied_metrics.input_tokens == 999 |
| assert copied_metrics.output_tokens == 20 |
| assert copied_metrics.cached_input_tokens == 5 |
| assert copied_metrics.tool_output_tokens == 3 |
|
|
|
|
| |
|
|
|
|
| def create_simple_context_output( |
| text="", |
| token_ids=None, |
| prompt="Test prompt", |
| prompt_token_ids=None, |
| num_cached_tokens=0, |
| logprobs=None, |
| finished=True, |
| ): |
| """Helper to create a RequestOutput with customizable text for |
| SimpleContext tests.""" |
| if token_ids is None: |
| token_ids = [] |
| return RequestOutput( |
| request_id="test-id", |
| prompt=prompt, |
| prompt_token_ids=prompt_token_ids, |
| prompt_logprobs=None, |
| outputs=[ |
| CompletionOutput( |
| index=0, |
| text=text, |
| token_ids=token_ids, |
| cumulative_logprob=0.0, |
| logprobs=logprobs, |
| finish_reason=None, |
| stop_reason=None, |
| ) |
| ], |
| finished=finished, |
| num_cached_tokens=num_cached_tokens, |
| ) |
|
|
|
|
| def test_simple_context_output_messages_empty(): |
| """output_messages should be empty before any output is appended.""" |
| context = SimpleContext() |
| assert context.output_messages == [] |
|
|
|
|
| def test_simple_context_output_messages_single_call(): |
| """Non-streaming: single append_output produces a single output message.""" |
| context = SimpleContext() |
| output = create_simple_context_output( |
| text="Hello world", |
| token_ids=[10, 20, 30], |
| prompt_token_ids=[1, 2, 3], |
| ) |
| context.append_output(output) |
|
|
| messages = context.output_messages |
| assert len(messages) == 1 |
| assert messages[0].message == "Hello world" |
| assert messages[0].tokens == [10, 20, 30] |
| assert messages[0].type == "raw_message_tokens" |
|
|
|
|
| def test_simple_context_output_messages_streaming_consolidation(): |
| """Streaming: multiple append_output calls consolidate into one message.""" |
| context = SimpleContext() |
|
|
| |
| context.append_output( |
| create_simple_context_output( |
| text="Hello", |
| token_ids=[10], |
| prompt_token_ids=[1, 2, 3], |
| ) |
| ) |
| context.append_output( |
| create_simple_context_output( |
| text=" world", |
| token_ids=[20], |
| prompt_token_ids=[1, 2, 3], |
| ) |
| ) |
| context.append_output( |
| create_simple_context_output( |
| text="!", |
| token_ids=[30], |
| prompt_token_ids=[1, 2, 3], |
| ) |
| ) |
|
|
| messages = context.output_messages |
| assert len(messages) == 1 |
| assert messages[0].message == "Hello world!" |
| assert messages[0].tokens == [10, 20, 30] |
|
|
|
|
| def test_simple_context_output_messages_many_deltas(): |
| """Streaming with many small deltas still produces a single message.""" |
| context = SimpleContext() |
|
|
| words = ["The", " quick", " brown", " fox", " jumps"] |
| for i, word in enumerate(words): |
| context.append_output( |
| create_simple_context_output( |
| text=word, |
| token_ids=[100 + i], |
| prompt_token_ids=[1, 2], |
| ) |
| ) |
|
|
| messages = context.output_messages |
| assert len(messages) == 1 |
| assert messages[0].message == "The quick brown fox jumps" |
| assert messages[0].tokens == [100, 101, 102, 103, 104] |
|
|
|
|
| def test_simple_context_input_messages(): |
| """input_messages is populated on the first append_output call.""" |
| context = SimpleContext() |
| assert context.input_messages == [] |
|
|
| context.append_output( |
| create_simple_context_output( |
| text="Hi", |
| token_ids=[10], |
| prompt="My prompt text", |
| prompt_token_ids=[1, 2, 3], |
| ) |
| ) |
|
|
| assert len(context.input_messages) == 1 |
| assert context.input_messages[0].message == "My prompt text" |
| assert context.input_messages[0].tokens == [1, 2, 3] |
|
|
| |
| context.append_output( |
| create_simple_context_output( |
| text=" there", |
| token_ids=[20], |
| prompt="My prompt text", |
| prompt_token_ids=[1, 2, 3], |
| ) |
| ) |
|
|
| assert len(context.input_messages) == 1 |
|
|
|
|
| def test_simple_context_token_counting(): |
| """Token counting accumulates across streaming deltas.""" |
| context = SimpleContext() |
|
|
| context.append_output( |
| create_simple_context_output( |
| text="a", |
| token_ids=[10, 11], |
| prompt_token_ids=[1, 2, 3, 4, 5], |
| num_cached_tokens=2, |
| ) |
| ) |
| context.append_output( |
| create_simple_context_output( |
| text="b", |
| token_ids=[12], |
| prompt_token_ids=[1, 2, 3, 4, 5], |
| num_cached_tokens=2, |
| ) |
| ) |
|
|
| assert context.num_prompt_tokens == 5 |
| assert context.num_output_tokens == 3 |
| assert context.num_cached_tokens == 2 |
|
|
|
|
| def test_simple_context_final_output(): |
| """final_output reconstructs accumulated text and token_ids.""" |
| context = SimpleContext() |
|
|
| context.append_output( |
| create_simple_context_output( |
| text="foo", |
| token_ids=[1, 2], |
| prompt_token_ids=[10], |
| ) |
| ) |
| context.append_output( |
| create_simple_context_output( |
| text="bar", |
| token_ids=[3], |
| prompt_token_ids=[10], |
| ) |
| ) |
|
|
| final = context.final_output |
| assert final is not None |
| assert final.outputs[0].text == "foobar" |
| assert final.outputs[0].token_ids == (1, 2, 3) |
|
|
|
|
| def test_simple_context_output_messages_empty_text_with_tokens(): |
| """output_messages should be returned when tokens exist even if text is |
| empty (e.g. special tokens).""" |
| context = SimpleContext() |
| context.append_output( |
| create_simple_context_output( |
| text="", |
| token_ids=[99], |
| prompt_token_ids=[1], |
| ) |
| ) |
|
|
| messages = context.output_messages |
| assert len(messages) == 1 |
| assert messages[0].message == "" |
| assert messages[0].tokens == [99] |
|
|
|
|
| def test_simple_context_output_messages_no_mutation(): |
| """Each call to output_messages returns a fresh list; callers can't |
| corrupt internal state.""" |
| context = SimpleContext() |
| context.append_output( |
| create_simple_context_output( |
| text="hello", |
| token_ids=[1], |
| prompt_token_ids=[10], |
| ) |
| ) |
|
|
| msgs1 = context.output_messages |
| msgs2 = context.output_messages |
| assert msgs1 is not msgs2 |
| assert msgs1[0].message == msgs2[0].message |
|
|
| |
| context.append_output( |
| create_simple_context_output( |
| text=" world", |
| token_ids=[2], |
| prompt_token_ids=[10], |
| ) |
| ) |
|
|
| msgs3 = context.output_messages |
| assert len(msgs3) == 1 |
| assert msgs3[0].message == "hello world" |
| assert msgs3[0].tokens == [1, 2] |
|
|