fix(responses): preserve signed thinking when replaying streams

This commit is contained in:
jinwukong 2026-09-10 18:25:01 +08:00
parent 238f434153
commit 14bbd67ddd
6 changed files with 167 additions and 7 deletions

View file

@ -58,6 +58,7 @@ class _ThinkingBlockFragment(TypedDict, total=False):
class _ThinkingDelta(TypedDict, total=False):
thinking_blocks: Sequence[_ThinkingBlockFragment]
provider_specific_fields: ReadOnly[Mapping[str, object] | None]
class _ThinkingChoice(TypedDict, total=False):
@ -687,7 +688,7 @@ class ChunkProcessor:
def _flush_thinking_block() -> None:
nonlocal current_thinking_text_parts, current_signature
if len(current_thinking_text_parts) > 0 and current_signature:
if current_signature:
thinking_blocks.append(
ChatCompletionThinkingBlock(
type="thinking",
@ -717,10 +718,19 @@ class ChunkProcessor:
)
)
else:
thinking_text = thinking_block.get("thinking", None)
thinking_text, signature, provider_fields = (
thinking_block.get("thinking"),
thinking_block.get("signature"),
delta.get("provider_specific_fields"),
)
if (
signature
and isinstance(provider_fields, Mapping)
and provider_fields.get("thinking_blocks") == thinking
):
current_thinking_text_parts.clear()
if thinking_text:
current_thinking_text_parts.append(thinking_text)
signature = thinking_block.get("signature", None)
if signature:
current_signature = signature
_flush_thinking_block()

View file

@ -649,6 +649,16 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
)
return response
def _encoded_thinking_blocks(self) -> str | None:
response: Final = (
self.litellm_model_response
if isinstance(self.litellm_model_response, ModelResponse)
else self.create_litellm_model_response()
)
if response is None:
return None
return LiteLLMCompletionResponsesConfig.encode_thinking_blocks(response.choices[0].message)
@staticmethod
def _snapshot_chunk_for_stream_chunk_builder(
chunk: ModelResponseStream,
@ -839,6 +849,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
**{
"id": reasoning_item_id,
"type": "reasoning",
"encrypted_content": self._encoded_thinking_blocks(),
"summary": [
{
"type": "summary_text",

View file

@ -1490,7 +1490,7 @@ class LiteLLMCompletionResponsesConfig:
input_item: Mapping[str, object],
) -> tuple[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, ...] | None:
"""
Decode ``encrypted_content`` written by ``_encode_thinking_blocks`` back
Decode ``encrypted_content`` written by ``encode_thinking_blocks`` back
into the signed thinking blocks it serialized.
LiteLLM writes this field itself for providers whose reasoning is signed
@ -2571,7 +2571,7 @@ class LiteLLMCompletionResponsesConfig:
return output_items
@staticmethod
def _encode_thinking_blocks(message: Message) -> str | None:
def encode_thinking_blocks(message: Message) -> str | None:
thinking_blocks: Final[Sequence[Mapping[str, object]]] = getattr(message, "thinking_blocks", None) or ()
preserved: Final = tuple(block for block in thinking_blocks if block.get("signature") or block.get("data"))
return json.dumps(preserved, separators=(",", ":")) if preserved else None
@ -2585,7 +2585,7 @@ class LiteLLMCompletionResponsesConfig:
if hasattr(choice, "message") and choice.message:
message = choice.message
reasoning_content: str = getattr(message, "reasoning_content", None) or ""
encrypted_content = LiteLLMCompletionResponsesConfig._encode_thinking_blocks(message)
encrypted_content = LiteLLMCompletionResponsesConfig.encode_thinking_blocks(message)
if reasoning_content or encrypted_content:
# Only check the first choice for reasoning content
return [

View file

@ -9,6 +9,7 @@ from litellm import ChatCompletionUsageBlock, stream_chunk_builder
from litellm.types.utils import GenericStreamingChunk
from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
from litellm.types.llms.openai import ChatCompletionThinkingBlock
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
@ -236,6 +237,30 @@ def test_get_combined_thinking_content_preserves_interleaved_blocks():
assert result[2]["signature"] == "sig_block2"
@pytest.mark.parametrize("snapshot", [True, False], ids=["provider-snapshot", "genuine-final-delta"])
def test_stream_chunk_builder_distinguishes_thinking_snapshots_from_repeated_deltas(snapshot: bool) -> None:
signed: Final = ChatCompletionThinkingBlock(type="thinking", thinking="echo", signature="test-signature")
deltas: Final = (
Delta(thinking_blocks=[ChatCompletionThinkingBlock(type="thinking", thinking="echo")]),
Delta(thinking_blocks=[signed], provider_specific_fields={"thinking_blocks": [signed]} if snapshot else None),
)
chunks: Final = [
ModelResponseStream(
id="chatcmpl-thinking",
model="claude-opus-5",
choices=[StreamingChoices(index=0, delta=delta, finish_reason="stop" if index == 1 else None)],
)
for index, delta in enumerate(deltas)
]
response: Final = stream_chunk_builder(chunks=chunks)
assert response is not None
assert response.choices[0].message.thinking_blocks == [
{"type": "thinking", "thinking": "echo" if snapshot else "echoecho", "signature": "test-signature"}
]
def test_cache_read_input_tokens_retained():
chunk1 = ModelResponseStream(
id="chatcmpl-95aabb85-c39f-443d-ae96-0370c404d70c",

View file

@ -167,7 +167,7 @@ class TestEncryptedReasoningRoundTrip:
{"type": "redacted_thinking", "data": "redacted-payload"},
]
message = Message(role="assistant", content="answer", thinking_blocks=blocks)
encoded = LiteLLMCompletionResponsesConfig._encode_thinking_blocks(message)
encoded = LiteLLMCompletionResponsesConfig.encode_thinking_blocks(message)
decoded = LiteLLMCompletionResponsesConfig._decode_thinking_blocks_from_input_item(
{"type": "reasoning", "encrypted_content": encoded}
)

View file

@ -19,9 +19,15 @@ import pytest
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
BaseLiteLLMOpenAIResponseObject,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ResponseCompletedEvent,
ResponsesAPIStreamEvents,
)
from litellm.types.responses.main import build_web_search_call
@ -1131,3 +1137,111 @@ async def test_plain_text_stream_announces_exactly_one_message_item(sync_mode: b
ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
):
assert event.item_id == message_item_adds[0].item.id
def _signed_thinking_chunks(cumulative: bool) -> tuple[ModelResponseStream, ...]:
blocks: Final = [
ChatCompletionThinkingBlock(type="thinking", thinking="One plus one equals two.", signature="test-signature"),
ChatCompletionRedactedThinkingBlock(type="redacted_thinking", data="test-redacted-data"),
]
deltas: Final = (
Delta(
reasoning_content="One plus ",
thinking_blocks=[ChatCompletionThinkingBlock(type="thinking", thinking="One plus ")],
),
Delta(
reasoning_content="one equals two.",
thinking_blocks=[ChatCompletionThinkingBlock(type="thinking", thinking="one equals two.")],
),
Delta(
thinking_blocks=blocks
if cumulative
else [
ChatCompletionThinkingBlock(type="thinking", thinking="", signature="test-signature"),
ChatCompletionRedactedThinkingBlock(type="redacted_thinking", data="test-redacted-data"),
],
provider_specific_fields={"thinking_blocks": blocks} if cumulative else None,
),
Delta(content="2"),
)
return tuple(
ModelResponseStream(
id=CHAT_COMPLETION_ID,
model="test-model",
choices=[StreamingChoices(index=0, delta=delta, finish_reason="stop" if index == 3 else None)],
)
for index, delta in enumerate(deltas)
)
@pytest.mark.parametrize("cumulative", [True, False], ids=["cumulative-provider-blocks", "delta-blocks"])
@pytest.mark.parametrize("asynchronous", [True, False], ids=["async", "sync"])
async def test_completed_response_replays_signed_thinking_unchanged(cumulative: bool, asynchronous: bool) -> None:
iterator: Final = _build_iterator(_signed_thinking_chunks(cumulative))
events: Final = [event async for event in iterator] if asynchronous else list(iterator)
completed: Final = next(event for event in events if isinstance(event, ResponseCompletedEvent))
reasoning: Final = next(item for item in completed.response.output if item.type == "reasoning")
assert reasoning.encrypted_content is not None
assert json.loads(reasoning.encrypted_content) == [
{"type": "thinking", "thinking": "One plus one equals two.", "signature": "test-signature"},
{"type": "redacted_thinking", "data": "test-redacted-data"},
]
messages: Final = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message(
input_item=reasoning.model_dump(exclude_none=True), replay_reasoning=True
)
assert len(messages) == 1
assert messages[0]["thinking_blocks"] == [
{"type": "thinking", "thinking": "One plus one equals two.", "signature": "test-signature"},
{"type": "redacted_thinking", "data": "test-redacted-data"},
]
@pytest.mark.parametrize("cumulative", [True, False], ids=["cumulative-provider-blocks", "delta-blocks"])
async def test_reasoning_done_preserves_the_replay_payload(cumulative: bool) -> None:
iterator: Final = _build_iterator(_signed_thinking_chunks(cumulative))
events: Final = [event async for event in iterator]
done: Final = next(
event for event in events if event.type == "response.output_item.done" and event.item.type == "reasoning"
)
completed: Final = next(event for event in events if isinstance(event, ResponseCompletedEvent))
reasoning: Final = next(item for item in completed.response.output if item.type == "reasoning")
payload: Final = done.item.model_dump().get("encrypted_content")
assert payload is not None
assert json.loads(payload) == [
{"type": "thinking", "thinking": "One plus one equals two.", "signature": "test-signature"},
{"type": "redacted_thinking", "data": "test-redacted-data"},
]
assert payload == reasoning.encrypted_content
@pytest.mark.parametrize("asynchronous", [True, False], ids=["async", "sync"])
async def test_streamed_signature_only_thinking_is_replayable(asynchronous: bool) -> None:
block: Final = ChatCompletionThinkingBlock(type="thinking", thinking="", signature="opaque-signature")
chunks: Final = [
ModelResponseStream(
id=CHAT_COMPLETION_ID,
model="test-model",
choices=[
StreamingChoices(
index=0,
delta=Delta(thinking_blocks=[block], provider_specific_fields={"thinking_blocks": [block]}),
)
],
),
_tool_call_chunk(finish_reason="tool_calls"),
]
iterator: Final = _build_iterator(chunks)
events: Final = [event async for event in iterator] if asynchronous else list(iterator)
completed: Final = next(event for event in events if isinstance(event, ResponseCompletedEvent))
reasoning: Final = next(item for item in completed.response.output if item.type == "reasoning")
assert json.loads(reasoning.encrypted_content) == [block]
messages: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=[item.model_dump(exclude_none=True) for item in completed.response.output],
responses_api_request={},
replay_reasoning=True,
)
tool_message: Final = next(message for message in messages if message.get("tool_calls"))
assert tool_message["thinking_blocks"] == [block]