From 21293ca401b9ae78bb8d47fd60ea12273a0f9c41 Mon Sep 17 00:00:00 2001 From: carey-bk <180496115+carey-bk@users.noreply.github.com> Date: Sun, 4 Oct 2026 20:54:31 +0800 Subject: [PATCH] fix: preserve reasoning from compatible Chat Completions providers --- .../generators/openaichatgenerator.mdx | 2 + .../generators/openaichatgenerator.mdx | 2 + haystack/components/generators/chat/openai.py | 212 +++++---- ...penai-chat-reasoning-7e25f15ac5ff3d32.yaml | 9 + .../components/generators/chat/test_openai.py | 427 +++++++++++++++++- 5 files changed, 549 insertions(+), 103 deletions(-) create mode 100644 releasenotes/notes/openai-chat-reasoning-7e25f15ac5ff3d32.yaml diff --git a/docs-website/docs/pipeline-components/generators/openaichatgenerator.mdx b/docs-website/docs/pipeline-components/generators/openaichatgenerator.mdx index c4bcb24e3d2..97187f8bcea 100644 --- a/docs-website/docs/pipeline-components/generators/openaichatgenerator.mdx +++ b/docs-website/docs/pipeline-components/generators/openaichatgenerator.mdx @@ -93,6 +93,8 @@ print(response["replies"][0].text) ### Streaming +For OpenAI-compatible providers that return reasoning text, non-empty string `reasoning` or `reasoning_content` fields populate `ChatMessage.reasoning` and `StreamingChunk.reasoning`, separately from the answer text. `reasoning` takes precedence when both contain text; empty or non-string values are ignored. A provider delta containing both reasoning and text produces separate callback chunks. Reasoning tokens in usage metadata alone do not provide reasoning text. + You can stream output as it’s generated. Pass a callback to `streaming_callback`. Use the built-in `print_streaming_chunk` to print text tokens and tool events (tool calls and tool results). ```python diff --git a/docs-website/versioned_docs/version-3.3/pipeline-components/generators/openaichatgenerator.mdx b/docs-website/versioned_docs/version-3.3/pipeline-components/generators/openaichatgenerator.mdx index c4bcb24e3d2..97187f8bcea 100644 --- a/docs-website/versioned_docs/version-3.3/pipeline-components/generators/openaichatgenerator.mdx +++ b/docs-website/versioned_docs/version-3.3/pipeline-components/generators/openaichatgenerator.mdx @@ -93,6 +93,8 @@ print(response["replies"][0].text) ### Streaming +For OpenAI-compatible providers that return reasoning text, non-empty string `reasoning` or `reasoning_content` fields populate `ChatMessage.reasoning` and `StreamingChunk.reasoning`, separately from the answer text. `reasoning` takes precedence when both contain text; empty or non-string values are ignored. A provider delta containing both reasoning and text produces separate callback chunks. Reasoning tokens in usage metadata alone do not provide reasoning text. + You can stream output as it’s generated. Pass a callback to `streaming_callback`. Use the built-in `print_streaming_chunk` to print text tokens and tool events (tool calls and tool results). ```python diff --git a/haystack/components/generators/chat/openai.py b/haystack/components/generators/chat/openai.py index e649e2d2e27..02da19e4644 100644 --- a/haystack/components/generators/chat/openai.py +++ b/haystack/components/generators/chat/openai.py @@ -5,6 +5,7 @@ import asyncio import json import os +from dataclasses import replace from datetime import datetime from typing import Any, ClassVar @@ -32,6 +33,7 @@ ChatMessage, ComponentInfo, FinishReason, + ReasoningContent, StreamingCallbackT, StreamingChunk, SyncStreamingCallbackT, @@ -63,6 +65,11 @@ class OpenAIChatGenerator: from OpenAI API. It uses [ChatMessage](https://docs.haystack.deepset.ai/docs/chatmessage) format in input and output. + For OpenAI-compatible providers, non-empty string `reasoning` or `reasoning_content` response fields + populate `ChatMessage.reasoning` and `StreamingChunk.reasoning`, separately from text. `reasoning` takes + precedence; empty or non-string values are ignored. A delta containing both reasoning and text produces + separate streaming chunks. + You can customize how the text is generated by passing parameters to the OpenAI API. Use the `**generation_kwargs` argument when you initialize the component or when you run it. Any parameter that works with @@ -566,13 +573,15 @@ def _prepare_api_call( # noqa: PLR0913 def _handle_stream_response(self, chat_completion: Stream, callback: SyncStreamingCallbackT) -> list[ChatMessage]: component_info = ComponentInfo.from_component(self) chunks: list[StreamingChunk] = [] + block_indices: dict[str | int, int] = {} for chunk in chat_completion: assert len(chunk.choices) <= 1, "Streaming responses should have at most one choice." - chunk_delta = _convert_chat_completion_chunk_to_streaming_chunk( - chunk=chunk, previous_chunks=chunks, component_info=component_info + chunk_deltas = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=chunk, block_indices=block_indices, previous_chunks=chunks, component_info=component_info ) - chunks.append(chunk_delta) - callback(chunk_delta) + for chunk_delta in chunk_deltas: + chunks.append(chunk_delta) + callback(chunk_delta) return [_convert_streaming_chunks_to_chat_message(chunks=chunks)] async def _handle_async_stream_response( @@ -580,14 +589,16 @@ async def _handle_async_stream_response( ) -> list[ChatMessage]: component_info = ComponentInfo.from_component(self) chunks: list[StreamingChunk] = [] + block_indices: dict[str | int, int] = {} try: async for chunk in chat_completion: assert len(chunk.choices) <= 1, "Streaming responses should have at most one choice." - chunk_delta = _convert_chat_completion_chunk_to_streaming_chunk( - chunk=chunk, previous_chunks=chunks, component_info=component_info + chunk_deltas = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=chunk, block_indices=block_indices, previous_chunks=chunks, component_info=component_info ) - chunks.append(chunk_delta) - await _invoke_streaming_callback(callback, chunk_delta) + for chunk_delta in chunk_deltas: + chunks.append(chunk_delta) + await _invoke_streaming_callback(callback, chunk_delta) except asyncio.CancelledError: await asyncio.shield(chat_completion.close()) @@ -648,6 +659,16 @@ def _check_finish_reason(meta: dict[str, Any]) -> None: ) +def _get_reasoning_text(message: BaseModel | None) -> str | None: + # Compatible providers expose these non-standard fields as SDK model extras. + extra = (message.model_extra or {}) if message is not None else {} + for field in ("reasoning", "reasoning_content"): + value = extra.get(field) + if isinstance(value, str) and value: + return value + return None + + def _convert_chat_completion_to_chat_message( completion: ChatCompletion | ParsedChatCompletion, choice: Choice ) -> ChatMessage: @@ -690,22 +711,26 @@ def _convert_chat_completion_to_chat_message( if logprobs: meta["logprobs"] = logprobs - return ChatMessage.from_assistant(text=text, tool_calls=tool_calls, meta=meta) + return ChatMessage.from_assistant( + text=text, tool_calls=tool_calls, meta=meta, reasoning=_get_reasoning_text(message) + ) -def _convert_chat_completion_chunk_to_streaming_chunk( - chunk: ChatCompletionChunk, previous_chunks: list[StreamingChunk], component_info: ComponentInfo | None = None -) -> StreamingChunk: +def _convert_chat_completion_chunk_to_streaming_chunks( + chunk: ChatCompletionChunk, + block_indices: dict[str | int, int], + previous_chunks: list[StreamingChunk], + component_info: ComponentInfo | None = None, +) -> list[StreamingChunk]: """ - Converts the streaming response chunk from the OpenAI API to a StreamingChunk. + Convert an SDK delta, splitting reasoning and other content into separate blocks. :param chunk: The chunk returned by the OpenAI API. - :param previous_chunks: A list of previously received StreamingChunks. - :param component_info: An optional `ComponentInfo` object containing information about the component that - generated the chunk, such as the component name and type. - - :returns: - A StreamingChunk object representing the content of the chunk from the OpenAI API. + :param block_indices: Per-response mapping of text, reasoning, and provider tool indices to content block indices. + Updated as new blocks arrive. + :param previous_chunks: Previously emitted chunks, used to preserve the non-reasoning start markers. + :param component_info: Information about the component that generated the chunk. + :returns: StreamingChunks representing the delta, with completion metadata on the last chunk. """ finish_reason_mapping: dict[str, FinishReason] = { "stop": "stop", @@ -714,93 +739,94 @@ def _convert_chat_completion_chunk_to_streaming_chunk( "tool_calls": "tool_calls", "function_call": "tool_calls", } - # On very first chunk so len(previous_chunks) == 0, the Choices field only provides role info (e.g. "assistant") - # Choices is empty if include_usage is set to True where the usage information is returned. - if len(chunk.choices) == 0: - return StreamingChunk( - content="", - component_info=component_info, - # Index is None since it's only set to an int when a content block is present - index=None, - finish_reason=None, - meta={ - "model": chunk.model, - "received_at": datetime.now().isoformat(), - "usage": _serialize_object(chunk.usage), - }, - ) + meta = {"model": chunk.model, "received_at": datetime.now().isoformat(), "usage": _serialize_object(chunk.usage)} + if not chunk.choices: + return [StreamingChunk(content="", component_info=component_info, meta=meta)] choice: ChunkChoice = chunk.choices[0] + delta = choice.delta + meta.update( + index=choice.index, + tool_calls=delta.tool_calls if delta and delta.tool_calls else None, + finish_reason=choice.finish_reason, + ) + finish_reason = finish_reason_mapping.get(choice.finish_reason) if choice.finish_reason else None + reasoning_text = _get_reasoning_text(delta) + has_reasoning = reasoning_text is not None or "reasoning" in block_indices + + result = [] + if reasoning_text is not None: + start = "reasoning" not in block_indices + index = block_indices.setdefault("reasoning", max(block_indices.values(), default=-1) + 1) + result.append( + StreamingChunk( + content="", + reasoning=ReasoningContent(reasoning_text=reasoning_text), + index=index, + start=start, + component_info=component_info, + meta={**meta, "tool_calls": None}, + ) + ) + + if delta and delta.content and (has_reasoning or not delta.tool_calls): + start = "text" not in block_indices if has_reasoning else len(previous_chunks) == 1 + text_index = ( + None + if not has_reasoning and delta.role is not None + else block_indices.setdefault("text", max(block_indices.values(), default=-1) + 1 if has_reasoning else 0) + ) + text_meta = {**meta, "tool_calls": None} + if choice.logprobs: + text_meta["logprobs"] = _serialize_object(choice.logprobs) + result.append( + StreamingChunk( + content=delta.content, index=text_index, start=start, component_info=component_info, meta=text_meta + ) + ) - # create a list of ToolCallDelta objects from the tool calls - if choice.delta and choice.delta.tool_calls: + if delta and delta.tool_calls: tool_calls_deltas = [] - for tool_call in choice.delta.tool_calls: + for tool_call in delta.tool_calls: + block_indices.setdefault( + tool_call.index, max(block_indices.values(), default=-1) + 1 if has_reasoning else tool_call.index + ) function = tool_call.function tool_calls_deltas.append( ToolCallDelta( + # Reconstruction groups and sorts tools by this provider index, not by arrival order. index=tool_call.index, id=tool_call.id, tool_name=function.name if function else None, arguments=function.arguments if function and function.arguments else None, ) ) - return StreamingChunk( - content=choice.delta.content or "", - component_info=component_info, - # We adopt the first tool_calls_deltas.index as the overall index of the chunk. - index=tool_calls_deltas[0].index, - tool_calls=tool_calls_deltas, - start=tool_calls_deltas[0].tool_name is not None, - finish_reason=finish_reason_mapping.get(choice.finish_reason) if choice.finish_reason else None, - meta={ - "model": chunk.model, - "index": choice.index, - "tool_calls": choice.delta.tool_calls, - "finish_reason": choice.finish_reason, - "received_at": datetime.now().isoformat(), - "usage": _serialize_object(chunk.usage), - }, + result.append( + StreamingChunk( + content="" if has_reasoning else delta.content or "", + index=block_indices[tool_calls_deltas[0].index], + tool_calls=tool_calls_deltas, + start=tool_calls_deltas[0].tool_name is not None, + component_info=component_info, + meta=meta, + ) ) - # On very first chunk the choice field only provides role info (e.g. "assistant") so we set index to None - # We set all chunks missing the content field to index of None. E.g. can happen if chunk only contains finish - # reason. - if choice.delta and (choice.delta.content is None or choice.delta.role is not None): - resolved_index = None - else: - # We set the index to be 0 since if text content is being streamed then no tool calls are being streamed - # NOTE: We may need to revisit this if OpenAI allows planning/thinking content before tool calls like - # Anthropic Claude - resolved_index = 0 - - # Initialize meta dictionary - meta = { - "model": chunk.model, - "index": choice.index, - "tool_calls": choice.delta.tool_calls if choice.delta and choice.delta.tool_calls else None, - "finish_reason": choice.finish_reason, - "received_at": datetime.now().isoformat(), - "usage": _serialize_object(chunk.usage), - } - - # check if logprobs are present - # logprobs are returned only for text content - logprobs = _serialize_object(choice.logprobs) if choice.logprobs else None - if logprobs: - meta["logprobs"] = logprobs + if not result: + if choice.logprobs: + meta["logprobs"] = _serialize_object(choice.logprobs) + result.append( + StreamingChunk( + content="", + component_info=component_info, + meta=meta, + # Preserve legacy empty/null delta indices and start markers when reasoning is absent. + index=None if has_reasoning or (delta and (delta.content is None or delta.role is not None)) else 0, + start=not has_reasoning and len(previous_chunks) == 1, + ) + ) - content = "" - if choice.delta and choice.delta.content: - content = choice.delta.content - - return StreamingChunk( - content=content, - component_info=component_info, - index=resolved_index, - # The first chunk is always a start message chunk that only contains role information, so if we reach here - # and previous_chunks is length 1 then this is the start of text content. - start=len(previous_chunks) == 1, - finish_reason=finish_reason_mapping.get(choice.finish_reason) if choice.finish_reason else None, - meta=meta, - ) + # A single SDK delta can contain reasoning, text, and a finish reason. Signal completion only after all its content. + return [ + replace(converted, meta={**converted.meta, "finish_reason": None, "usage": None}) for converted in result[:-1] + ] + [replace(result[-1], finish_reason=finish_reason)] diff --git a/releasenotes/notes/openai-chat-reasoning-7e25f15ac5ff3d32.yaml b/releasenotes/notes/openai-chat-reasoning-7e25f15ac5ff3d32.yaml new file mode 100644 index 00000000000..031b98c2728 --- /dev/null +++ b/releasenotes/notes/openai-chat-reasoning-7e25f15ac5ff3d32.yaml @@ -0,0 +1,9 @@ +--- +fixes: + - | + ``OpenAIChatGenerator`` now preserves non-empty string ``reasoning`` and + ``reasoning_content`` fields returned by OpenAI-compatible Chat Completions + providers in ``ChatMessage.reasoning`` and ``StreamingChunk.reasoning``. + Previously these fields were discarded. Reasoning stays separate from answer + text, including when both arrive in the same streaming delta. The + ``reasoning`` field takes precedence; empty or non-string values are ignored. diff --git a/test/components/generators/chat/test_openai.py b/test/components/generators/chat/test_openai.py index df3ca86f322..f85bfdfd1f6 100644 --- a/test/components/generators/chat/test_openai.py +++ b/test/components/generators/chat/test_openai.py @@ -36,7 +36,7 @@ from haystack.components.generators.chat.openai import ( OpenAIChatGenerator, _check_finish_reason, - _convert_chat_completion_chunk_to_streaming_chunk, + _convert_chat_completion_chunk_to_streaming_chunks, _make_schema_strict, ) from haystack.components.generators.utils import print_streaming_chunk @@ -45,6 +45,7 @@ ChatRole, FileContent, ImageContent, + ReasoningContent, StreamingChunk, ToolCall, ToolCallDelta, @@ -1911,17 +1912,18 @@ def handler(request: httpx.Request) -> httpx.Response: class TestChatCompletionChunkConversion: - def test_convert_chat_completion_chunk_to_streaming_chunk( + def test_convert_chat_completion_chunk_to_streaming_chunks( self, chat_completion_chunks: MagicMock, streaming_chunks: Any ) -> None: + block_indices: dict[str | int, int] = {} previous_chunks: list[StreamingChunk] = [] for openai_chunk, haystack_chunk in zip(chat_completion_chunks, streaming_chunks, strict=True): - stream_chunk = _convert_chat_completion_chunk_to_streaming_chunk( - chunk=openai_chunk, previous_chunks=previous_chunks + stream_chunks = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=openai_chunk, block_indices=block_indices, previous_chunks=previous_chunks ) - assert stream_chunk == haystack_chunk - previous_chunks.append(stream_chunk) + assert stream_chunks == [haystack_chunk] + previous_chunks.extend(stream_chunks) def test_convert_chat_completion_chunk_with_empty_tool_calls(self) -> None: @@ -1941,7 +1943,9 @@ def test_convert_chat_completion_chunk_with_empty_tool_calls(self) -> None: model="gpt-5-mini", object="chat.completion.chunk", ) - result = _convert_chat_completion_chunk_to_streaming_chunk(chunk=chunk, previous_chunks=[]) + (result,) = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=chunk, block_indices={}, previous_chunks=[] + ) assert result.content == "" assert result.start is False assert result.tool_calls == [ToolCallDelta(index=0)] @@ -1956,8 +1960,8 @@ def test_convert_chat_completion_chunk_with_delta_none(self, chat_completion_chu This should not happen, but some OpenAI-compatible providers sometimes return a delta set to None. """ - result = _convert_chat_completion_chunk_to_streaming_chunk( - chunk=chat_completion_chunk_delta_none, previous_chunks=[] + (result,) = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=chat_completion_chunk_delta_none, block_indices={}, previous_chunks=[] ) assert result.content == "" @@ -2023,7 +2027,9 @@ def test_convert_usage_chunk_to_streaming_chunk(self) -> None: prompt_tokens_details=PromptTokensDetails(audio_tokens=0, cached_tokens=0), ), ) - result = _convert_chat_completion_chunk_to_streaming_chunk(chunk=usage_chunk, previous_chunks=[]) + (result,) = _convert_chat_completion_chunk_to_streaming_chunks( + chunk=usage_chunk, block_indices={}, previous_chunks=[] + ) assert result.content == "" assert result.start is False assert result.tool_calls is None @@ -2340,3 +2346,404 @@ def test_live_run_strict_nested_tool(self) -> None: assert "address" in tool_call.arguments assert "street" in tool_call.arguments["address"] assert "city" in tool_call.arguments["address"] + + +@pytest.fixture +def reasoning_response(): + async def run_response(payload, *, run_mode, streaming=False): + chunks = [] + requests = [] + + def respond(request): + requests.append(json.loads(request.content)) + assert request.url.path == "/v1/chat/completions" + if streaming: + data = "".join(f"data: {json.dumps(chunk)}\n\n" for chunk in payload) + return httpx.Response( + 200, headers={"content-type": "text/event-stream"}, text=data + "data: [DONE]\n\n" + ) + return httpx.Response(200, json=payload) + + async def async_callback(chunk): + chunks.append(chunk) + + generator = OpenAIChatGenerator( + api_key=Secret.from_token("offline-test-key"), + api_base_url="https://provider.invalid/v1", + model="compatible-reasoner", + http_client_kwargs={"transport": httpx.MockTransport(respond)}, + max_retries=0, + ) + callback = async_callback if run_mode == "async" else chunks.append + try: + if run_mode == "sync": + result = generator.run("Test", streaming_callback=callback if streaming else None) + else: + result = await generator.run_async("Test", streaming_callback=callback if streaming else None) + finally: + if run_mode == "sync": + generator.close() + else: + await generator.close_async() + assert len(requests) == 1 + assert requests[0]["stream"] is streaming + assert requests[0]["messages"] == [{"role": "user", "content": "Test"}] + assert set(result) == {"replies"} + assert len(result["replies"]) == 1 + return result["replies"][0], chunks + + return run_response + + +def reasoning_api_chunk(delta=None, *, finish_reason=None, usage=None): + return { + "id": "offline-completion", + "object": "chat.completion.chunk", + "created": 1, + "model": "compatible-reasoner", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}] if delta is not None else [], + "usage": usage, + } + + +class TestOpenAIReasoning: + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + @pytest.mark.parametrize("streaming", [False, True]) + @pytest.mark.parametrize( + "extra, expected", + [ + ({"reasoning": "Think."}, "Think."), + ({"reasoning_content": "Think."}, "Think."), + ({"reasoning": "New", "reasoning_content": "Old"}, "New"), + ({"reasoning": "", "reasoning_content": "Fallback"}, "Fallback"), + ({"reasoning": None, "reasoning_content": "Fallback"}, "Fallback"), + ({"reasoning": {"text": "unsupported"}, "reasoning_content": "Fallback"}, "Fallback"), + ({"reasoning": " "}, " "), + ({}, None), + ({"reasoning": "", "reasoning_content": None}, None), + ({"reasoning": 42, "reasoning_content": ["unsupported"]}, None), + ], + ) + async def test_reasoning_fields(self, reasoning_response, run_mode, streaming, extra, expected): + payload: dict[str, Any] | list[dict[str, Any]] + if streaming: + payload = [ + reasoning_api_chunk({"role": "assistant", "content": ""}), + reasoning_api_chunk({"content": "Answer", **extra}), + reasoning_api_chunk({}, finish_reason="stop"), + ] + else: + payload = { + "id": "offline-completion", + "object": "chat.completion", + "created": 1, + "model": "compatible-reasoner", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Answer", **extra}, + "finish_reason": "stop", + } + ], + } + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=streaming) + meta: dict[str, Any] = {"model": "compatible-reasoner", "index": 0, "finish_reason": "stop", "usage": None} + if streaming: + meta["completion_start_time"] = ANY + assert "".join(chunk.content for chunk in chunks) == "Answer" + assert [chunk.reasoning.reasoning_text for chunk in chunks if chunk.reasoning] == ( + [expected] if expected is not None else [] + ) + if expected is None: + assert [chunk.index for chunk in chunks] == [None, 0, None] + assert [chunk.start for chunk in chunks] == [False, True, False] + assert reply.to_dict() == ChatMessage.from_assistant(text="Answer", reasoning=expected, meta=meta).to_dict() + + @pytest.mark.parametrize("run_mode", ["sync", "async", "async_sync_callback"]) + @pytest.mark.parametrize("field", ["reasoning", "reasoning_content"]) + @pytest.mark.parametrize("combined", [False, True]) + async def test_streaming_reasoning_and_text_blocks(self, reasoning_response, run_mode, field, combined): + deltas = [{"role": "assistant", "content": ""}, {field: "Think"}] + deltas += [{field: " more", "content": "An"}] if combined else [{field: " more"}, {"content": "An"}] + deltas += [{"content": "swer"}] + usage = {"prompt_tokens": 2, "completion_tokens": 4, "total_tokens": 6} + payload = [reasoning_api_chunk(delta) for delta in deltas] + payload += [reasoning_api_chunk({}, finish_reason="stop"), reasoning_api_chunk(usage=usage)] + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=True) + expected = [ + ("", None, None, False, None), + ("", "Think", 0, True, None), + ("", " more", 0, False, None), + ("An", None, 1, True, None), + ("swer", None, 1, False, None), + ("", None, None, False, "stop"), + ("", None, None, False, None), + ] + expected_usage = CompletionUsage.model_validate(usage).model_dump() + assert len(chunks) == len(expected) + for position, (chunk, (content, reasoning, index, start, finish)) in enumerate( + zip(chunks, expected, strict=True) + ): + meta: dict[str, Any] = {"model": "compatible-reasoner", "received_at": ANY, "usage": None} + if position == len(expected) - 1: + meta["usage"] = expected_usage + else: + meta.update(index=0, tool_calls=None, finish_reason=finish) + assert chunk.to_dict() == { + "content": content, + "reasoning": {"reasoning_text": reasoning, "extra": {}} if reasoning is not None else None, + "index": index, + "start": start, + "finish_reason": finish, + "meta": meta, + "tool_calls": None, + "tool_call_result": None, + "component_info": { + "type": "haystack.components.generators.chat.openai.OpenAIChatGenerator", + "name": None, + }, + } + assert ( + reply.to_dict() + == ChatMessage.from_assistant( + text="Answer", + reasoning="Think more", + meta={ + "model": "compatible-reasoner", + "index": 0, + "finish_reason": "stop", + "usage": expected_usage, + "completion_start_time": ANY, + }, + ).to_dict() + ) + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + @pytest.mark.parametrize("streaming", [False, True]) + @pytest.mark.parametrize("with_tools", [False, True]) + async def test_reasoning_only_and_tool_calls(self, reasoning_response, run_mode, streaming, with_tools): + payload: dict[str, Any] | list[dict[str, Any]] + tool_calls = [ + {"id": "call-0", "type": "function", "function": {"name": "weather", "arguments": '{"city":"Paris"}'}}, + {"id": "call-1", "type": "function", "function": {"name": "weather", "arguments": '{"city":"Berlin"}'}}, + ] + finish = "tool_calls" if with_tools else "stop" + message: dict[str, Any] = {"role": "assistant", "content": None, "reasoning": "Think."} + if with_tools: + message["tool_calls"] = tool_calls + if streaming: + if with_tools: + message["tool_calls"] = [ + {**tc, "index": i, "function": {"name": "weather", "arguments": '{"city":'}} + for i, tc in enumerate(tool_calls) + ] + payload = [ + reasoning_api_chunk(message), + reasoning_api_chunk( + { + "tool_calls": [ + {"index": 0, "function": {"arguments": '"Paris"}'}}, + {"index": 1, "function": {"arguments": '"Berlin"}'}}, + ] + }, + finish_reason=finish, + ), + ] + else: + payload = [reasoning_api_chunk(message, finish_reason=finish)] + else: + payload = { + "id": "offline-completion", + "object": "chat.completion", + "created": 1, + "model": "compatible-reasoner", + "choices": [{"index": 0, "message": message, "finish_reason": finish}], + } + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=streaming) + meta = {"model": "compatible-reasoner", "index": 0, "finish_reason": finish, "usage": None} + if streaming: + meta["completion_start_time"] = ANY + assert chunks[0].reasoning == ReasoningContent(reasoning_text="Think.") + assert chunks[0].index == 0 + assert chunks[0].start is True + assert [chunk.finish_reason for chunk in chunks if chunk.finish_reason] == [finish] + assert all(not chunk.content for chunk in chunks) + if with_tools: + tool_chunks = [chunk for chunk in chunks if chunk.tool_calls] + assert [chunk.index for chunk in tool_chunks] == [1, 1] + assert [chunk.start for chunk in tool_chunks] == [True, False] + assert [[tc.index for tc in chunk.tool_calls] for chunk in tool_chunks] == [[0, 1], [0, 1]] + expected_tools = ( + [ + ToolCall(id="call-0", tool_name="weather", arguments={"city": "Paris"}), + ToolCall(id="call-1", tool_name="weather", arguments={"city": "Berlin"}), + ] + if with_tools + else None + ) + assert ( + reply.to_dict() + == ChatMessage.from_assistant(reasoning="Think.", tool_calls=expected_tools, meta=meta).to_dict() + ) + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + async def test_finish_metadata_after_mixed_reasoning_delta(self, reasoning_response, run_mode): + usage = {"prompt_tokens": 2, "completion_tokens": 4, "total_tokens": 6} + logprobs: dict[str, Any] = { + "content": [{"token": "Answer", "bytes": None, "logprob": -0.1, "top_logprobs": []}], + "refusal": None, + } + payload = reasoning_api_chunk( + {"role": "assistant", "reasoning_content": "Think.", "content": "Answer"}, + finish_reason="length", + usage=usage, + ) + payload["choices"][0]["logprobs"] = logprobs + reply, chunks = await reasoning_response([payload], run_mode=run_mode, streaming=True) + assert [(chunk.index, chunk.start, chunk.finish_reason) for chunk in chunks] == [ + (0, True, None), + (1, True, "length"), + ] + assert chunks[0].meta["usage"] is None + assert chunks[0].meta["finish_reason"] is None + assert "logprobs" not in chunks[0].meta + assert chunks[1].meta["usage"] == CompletionUsage.model_validate(usage).model_dump() + assert chunks[1].meta["finish_reason"] == "length" + assert chunks[1].meta["logprobs"] == logprobs + assert reply.reasoning == ReasoningContent(reasoning_text="Think.") + assert reply.text == "Answer" + assert reply.meta["finish_reason"] == "length" + assert reply.meta["logprobs"] == [logprobs] + assert reply.meta["usage"] == CompletionUsage.model_validate(usage).model_dump() + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + async def test_reasoning_between_text_and_tool_deltas(self, reasoning_response, run_mode): + payload = [ + reasoning_api_chunk({"content": "An"}), + reasoning_api_chunk({"reasoning": "Think"}), + reasoning_api_chunk( + { + "tool_calls": [ + { + "index": 0, + "id": "call-0", + "type": "function", + "function": {"name": "weather", "arguments": '{"city":'}, + } + ] + } + ), + reasoning_api_chunk({"reasoning": " more", "content": "swer"}), + reasoning_api_chunk( + {"tool_calls": [{"index": 0, "function": {"arguments": '"Paris"}'}}]}, finish_reason="tool_calls" + ), + ] + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=True) + assert [(chunk.index, chunk.start) for chunk in chunks] == [ + (0, False), + (1, True), + (2, True), + (1, False), + (0, False), + (2, False), + ] + assert reply.text == "Answer" + assert reply.reasoning == ReasoningContent(reasoning_text="Think more") + assert reply.tool_calls == [ToolCall(id="call-0", tool_name="weather", arguments={"city": "Paris"})] + assert reply.meta["finish_reason"] == "tool_calls" + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + @pytest.mark.parametrize("tool_indices", [[0, 1], [2, 5], [5, 2]]) + @pytest.mark.parametrize("prefix_text", [False, True]) + async def test_streaming_without_reasoning_preserves_tool_indices_and_order( + self, reasoning_response, run_mode, tool_indices, prefix_text + ): + payload = [reasoning_api_chunk({"role": "assistant", "content": ""})] + if prefix_text: + payload.append(reasoning_api_chunk({"content": "Plan"})) + for index in tool_indices: + payload.append( + reasoning_api_chunk( + { + "tool_calls": [ + { + "index": index, + "id": f"call-{index}", + "type": "function", + "function": {"name": "weather", "arguments": '{"city":'}, + } + ] + } + ) + ) + for index in reversed(tool_indices): + payload.append( + reasoning_api_chunk({"tool_calls": [{"index": index, "function": {"arguments": f'"City {index}"}}'}}]}) + ) + payload.append(reasoning_api_chunk({}, finish_reason="tool_calls")) + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=True) + expected_indices = tool_indices + list(reversed(tool_indices)) + tool_chunks = [chunk for chunk in chunks if chunk.tool_calls] + # ToolCallDelta.index identifies the tool in the provider's list. Reconstruction sorts by that index. + assert [chunk.index for chunk in tool_chunks] == expected_indices + assert [chunk.tool_calls[0].index for chunk in tool_chunks] == expected_indices + assert [chunk.start for chunk in tool_chunks] == [True, True, False, False] + assert reply.tool_calls == [ + ToolCall(id=f"call-{index}", tool_name="weather", arguments={"city": f"City {index}"}) + for index in sorted(tool_indices) + ] + assert reply.text == ("Plan" if prefix_text else None) + assert reply.reasoning is None + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + @pytest.mark.parametrize( + "deltas, expected", + [ + ([{"content": "Answer"}], [(0, False)]), + ([{"role": "assistant", "content": "Answer"}], [(None, False)]), + ([{"role": "assistant", "content": ""}, {"content": "Answer"}], [(None, False), (0, True)]), + ([{"role": "assistant"}, {}, {"content": "Answer"}], [(None, False), (None, True), (0, False)]), + ([{"role": "assistant"}, {"content": "Answer"}, {"content": ""}], [(None, False), (0, True), (0, False)]), + ], + ) + async def test_streaming_without_reasoning_preserves_legacy_text_markers( + self, reasoning_response, run_mode, deltas, expected + ): + reply, chunks = await reasoning_response( + [reasoning_api_chunk(delta) for delta in deltas], run_mode=run_mode, streaming=True + ) + assert [(chunk.index, chunk.start) for chunk in chunks] == expected + assert reply.text == "Answer" + assert reply.reasoning is None + + @pytest.mark.parametrize("run_mode", ["sync", "async"]) + async def test_reasoning_preserves_provider_tool_order(self, reasoning_response, run_mode): + payload = [reasoning_api_chunk({"reasoning": "Think"})] + for index in [5, 2]: + payload.append( + reasoning_api_chunk( + { + "tool_calls": [ + { + "index": index, + "id": f"call-{index}", + "type": "function", + "function": {"name": "weather", "arguments": '{"city":'}, + } + ] + } + ) + ) + for index in [2, 5]: + payload.append( + reasoning_api_chunk({"tool_calls": [{"index": index, "function": {"arguments": f'"City {index}"}}'}}]}) + ) + reply, chunks = await reasoning_response(payload, run_mode=run_mode, streaming=True) + assert [chunk.index for chunk in chunks] == [0, 1, 2, 2, 1] + assert [chunk.tool_calls[0].index for chunk in chunks if chunk.tool_calls] == [5, 2, 2, 5] + assert reply.tool_calls == [ + ToolCall(id="call-2", tool_name="weather", arguments={"city": "City 2"}), + ToolCall(id="call-5", tool_name="weather", arguments={"city": "City 5"}), + ] + assert reply.reasoning == ReasoningContent(reasoning_text="Think") + assert reply.text is None