diff --git a/src/graphon/model_runtime/entities/message_entities.py b/src/graphon/model_runtime/entities/message_entities.py index 15e31c5..a022a93 100644 --- a/src/graphon/model_runtime/entities/message_entities.py +++ b/src/graphon/model_runtime/entities/message_entities.py @@ -5,7 +5,7 @@ from enum import StrEnum, auto from typing import Annotated, Any, Literal, Self -from pydantic import BaseModel, Field, field_serializer, field_validator +from pydantic import BaseModel, Field, JsonValue, field_serializer, field_validator class PromptMessageRole(StrEnum): @@ -59,6 +59,7 @@ class PromptMessageContent(ABC, BaseModel): """Model class for prompt message content.""" type: PromptMessageContentType + opaque_body: JsonValue | None = None class TextPromptMessageContent(PromptMessageContent): @@ -258,6 +259,7 @@ def transform_id_to_str(cls, value: Any) -> str: role: PromptMessageRole = PromptMessageRole.ASSISTANT tool_calls: list[ToolCall] = [] + opaque_body: JsonValue | None = None def is_empty(self) -> bool: """Check whether the assistant message has no content or tool calls.""" diff --git a/src/graphon/model_runtime/model_providers/base/large_language_model.py b/src/graphon/model_runtime/model_providers/base/large_language_model.py index deae9fb..7603624 100644 --- a/src/graphon/model_runtime/model_providers/base/large_language_model.py +++ b/src/graphon/model_runtime/model_providers/base/large_language_model.py @@ -4,6 +4,8 @@ from collections.abc import Callable, Generator, Iterator, Mapping, Sequence from dataclasses import dataclass, field +from pydantic import JsonValue + from graphon.model_runtime.callbacks.base_callback import Callback from graphon.model_runtime.callbacks.logging_callback import LoggingCallback from graphon.model_runtime.entities.llm_entities import ( @@ -122,6 +124,7 @@ class _LLMChunkAccumulator: usage: LLMUsage = field(default_factory=LLMUsage.empty_usage) system_fingerprint: str | None = None tool_calls: list[AssistantPromptMessage.ToolCall] = field(default_factory=list) + opaque_body: JsonValue | None = None def consume_all(self, chunks: Iterator[LLMResultChunk]) -> None: for chunk in chunks: @@ -135,6 +138,8 @@ def consume(self, chunk: LLMResultChunk) -> None: self.usage = chunk.delta.usage if chunk.system_fingerprint: self.system_fingerprint = chunk.system_fingerprint + if chunk.delta.message.opaque_body is not None: + self.opaque_body = chunk.delta.message.opaque_body def _consume_content(self, chunk: LLMResultChunk) -> None: content = chunk.delta.message.content @@ -155,6 +160,7 @@ def to_result( message=AssistantPromptMessage( content=self.content or self.content_list, tool_calls=self.tool_calls, + opaque_body=self.opaque_body, ), usage=self.usage, system_fingerprint=self.system_fingerprint, @@ -167,6 +173,7 @@ class _StreamingInvokeAccumulator: message_content: list[PromptMessageContentUnionTypes] = field(default_factory=list) usage: LLMUsage | None = None system_fingerprint: str | None = None + opaque_body: JsonValue | None = None def consume(self, chunk: LLMResultChunk) -> None: self._consume_content(chunk.delta.message.content) @@ -175,6 +182,8 @@ def consume(self, chunk: LLMResultChunk) -> None: self.usage = chunk.delta.usage if chunk.system_fingerprint: self.system_fingerprint = chunk.system_fingerprint + if chunk.delta.message.opaque_body is not None: + self.opaque_body = chunk.delta.message.opaque_body def _consume_content( self, @@ -196,7 +205,10 @@ def to_result( return LLMResult( model=self.real_model, prompt_messages=prompt_messages, - message=AssistantPromptMessage(content=self.message_content), + message=AssistantPromptMessage( + content=self.message_content, + opaque_body=self.opaque_body, + ), usage=self.usage or LLMUsage.empty_usage(), system_fingerprint=self.system_fingerprint, ) diff --git a/tests/model_runtime/test_large_language_model_non_stream_result.py b/tests/model_runtime/test_large_language_model_non_stream_result.py index ddd3664..341a21a 100644 --- a/tests/model_runtime/test_large_language_model_non_stream_result.py +++ b/tests/model_runtime/test_large_language_model_non_stream_result.py @@ -1,4 +1,5 @@ from collections.abc import Iterator +from typing import Any from graphon.model_runtime.entities.llm_entities import ( LLMResult, @@ -24,8 +25,13 @@ def _make_chunk( tool_calls: list[AssistantPromptMessage.ToolCall] | None = None, usage: LLMUsage | None = None, system_fingerprint: str | None = None, + opaque_body: Any | None = None, ) -> LLMResultChunk: - message = AssistantPromptMessage(content=content, tool_calls=tool_calls or []) + message = AssistantPromptMessage( + content=content, + tool_calls=tool_calls or [], + opaque_body=opaque_body, + ) delta = LLMResultChunkDelta(index=0, message=message, usage=usage) return LLMResultChunk( model=model, @@ -151,6 +157,63 @@ def test__normalize_non_stream_runtime_result__empty_iterator_defaults() -> None assert result.system_fingerprint is None +def test_non_stream_result_preserves_opaque_body() -> None: + prompt_messages = [UserPromptMessage(content="hi")] + opaque_body = { + "assistant_blocks": [{"type": "thinking", "signature": "sig-1"}], + } + chunk = _make_chunk( + content="hello", + usage=LLMUsage.empty_usage(), + opaque_body=opaque_body, + ) + + result = normalize_non_stream_runtime_result( + model="test-model", + prompt_messages=prompt_messages, + result=iter([chunk]), + ) + + assert result.message.opaque_body == opaque_body + + +def test_non_stream_result_opaque_body_last_non_none_wins() -> None: + prompt_messages = [UserPromptMessage(content="hi")] + chunks = iter([ + _make_chunk( + content="a", + usage=LLMUsage.empty_usage(), + opaque_body={"snapshot": 1}, + ), + _make_chunk(content="b", usage=LLMUsage.empty_usage()), + _make_chunk( + content="c", + usage=LLMUsage.empty_usage(), + opaque_body={"snapshot": 2}, + ), + ]) + + result = normalize_non_stream_runtime_result( + model="test-model", + prompt_messages=prompt_messages, + result=chunks, + ) + + assert result.message.opaque_body == {"snapshot": 2} + + +def test_non_stream_result_opaque_body_defaults_to_none() -> None: + prompt_messages = [UserPromptMessage(content="hi")] + + result = normalize_non_stream_runtime_result( + model="test-model", + prompt_messages=prompt_messages, + result=iter([_make_chunk(content="hello", usage=LLMUsage.empty_usage())]), + ) + + assert result.message.opaque_body is None + + def test__normalize_non_stream_runtime_result__accumulates_all_chunks() -> None: prompt_messages = [UserPromptMessage(content="hi")] diff --git a/tests/model_runtime/test_message_entities.py b/tests/model_runtime/test_message_entities.py index 293477e..de49299 100644 --- a/tests/model_runtime/test_message_entities.py +++ b/tests/model_runtime/test_message_entities.py @@ -1,4 +1,10 @@ +from graphon.model_runtime.entities.llm_entities import ( + LLMResultChunk, + LLMResultChunkDelta, + LLMUsage, +) from graphon.model_runtime.entities.message_entities import ( + AssistantPromptMessage, ImagePromptMessageContent, TextPromptMessageContent, UserPromptMessage, @@ -29,5 +35,76 @@ def test_prompt_message_normalizes_dict_content_items_for_serialization() -> Non assert isinstance(message.content, list) assert isinstance(message.content[0], TextPromptMessageContent) assert message.model_dump(mode="json")["content"] == [ - {"type": "text", "data": "hello"}, + {"type": "text", "data": "hello", "opaque_body": None}, ] + + +def test_assistant_prompt_message_opaque_body_defaults_to_none() -> None: + message = AssistantPromptMessage(content="ok") + + assert message.opaque_body is None + + +def test_assistant_prompt_message_opaque_body_survives_json_round_trip() -> None: + opaque_body = { + "assistant_blocks": [{"type": "thinking", "signature": "sig-1"}], + } + message = AssistantPromptMessage(content="ok", opaque_body=opaque_body) + + restored = AssistantPromptMessage.model_validate_json(message.model_dump_json()) + + assert restored.opaque_body == opaque_body + + +def test_assistant_prompt_message_accepts_payload_without_opaque_body() -> None: + message = AssistantPromptMessage.model_validate( + {"role": "assistant", "content": "ok", "tool_calls": []}, + ) + + assert message.opaque_body is None + + +def test_prompt_message_content_opaque_body_survives_json_round_trip() -> None: + content = TextPromptMessageContent( + data="hello", + opaque_body={"thought_signature": "sig-1"}, + ) + message = AssistantPromptMessage(content=[content]) + + restored = AssistantPromptMessage.model_validate_json(message.model_dump_json()) + + assert isinstance(restored.content, list) + assert isinstance(restored.content[0], TextPromptMessageContent) + assert restored.content[0].opaque_body == {"thought_signature": "sig-1"} + + +def test_prompt_message_content_accepts_dict_with_opaque_body() -> None: + message = UserPromptMessage.model_validate( + {"content": [{"type": "text", "data": "hello", "opaque_body": {"key": "v"}}]}, + ) + + assert isinstance(message.content, list) + assert isinstance(message.content[0], TextPromptMessageContent) + assert message.content[0].opaque_body == {"key": "v"} + + +def test_llm_result_chunk_json_round_trip_preserves_opaque_body() -> None: + """Simulate the plugin daemon -> core deserialization boundary.""" + opaque_body = { + "assistant_blocks": [ + {"type": "thinking", "thinking": "...", "signature": "sig-1"}, + {"type": "redacted_thinking", "data": "enc-1"}, + ], + } + chunk = LLMResultChunk( + model="test-model", + delta=LLMResultChunkDelta( + index=0, + message=AssistantPromptMessage(content="ok", opaque_body=opaque_body), + usage=LLMUsage.empty_usage(), + ), + ) + + restored = LLMResultChunk.model_validate_json(chunk.model_dump_json()) + + assert restored.delta.message.opaque_body == opaque_body diff --git a/tests/model_runtime/test_streaming_invoke_accumulator.py b/tests/model_runtime/test_streaming_invoke_accumulator.py new file mode 100644 index 0000000..4366ae0 --- /dev/null +++ b/tests/model_runtime/test_streaming_invoke_accumulator.py @@ -0,0 +1,67 @@ +from typing import Any + +from graphon.model_runtime.entities.llm_entities import ( + LLMResultChunk, + LLMResultChunkDelta, + LLMUsage, +) +from graphon.model_runtime.entities.message_entities import ( + AssistantPromptMessage, + UserPromptMessage, +) +from graphon.model_runtime.model_providers.base.large_language_model import ( + _StreamingInvokeAccumulator, +) + + +def _make_chunk( + *, + content: str | None, + opaque_body: Any | None = None, + usage: LLMUsage | None = None, +) -> LLMResultChunk: + message = AssistantPromptMessage( + content=content, + tool_calls=[], + opaque_body=opaque_body, + ) + delta = LLMResultChunkDelta(index=0, message=message, usage=usage) + return LLMResultChunk(model="test-model", delta=delta) + + +def test_streaming_invoke_accumulator_preserves_opaque_body() -> None: + accumulator = _StreamingInvokeAccumulator(real_model="test-model") + opaque_body = { + "assistant_blocks": [{"type": "thinking", "signature": "sig-1"}], + } + + accumulator.consume(_make_chunk(content="hello")) + accumulator.consume(_make_chunk(content=" world", opaque_body=opaque_body)) + + result = accumulator.to_result(prompt_messages=[UserPromptMessage(content="hi")]) + + assert result.message.opaque_body == opaque_body + + +def test_streaming_invoke_accumulator_opaque_body_defaults_to_none() -> None: + accumulator = _StreamingInvokeAccumulator(real_model="test-model") + + accumulator.consume(_make_chunk(content="hello")) + + result = accumulator.to_result(prompt_messages=[UserPromptMessage(content="hi")]) + + assert result.message.opaque_body is None + + +def test_streaming_invoke_accumulator_none_chunk_does_not_clobber_snapshot() -> None: + accumulator = _StreamingInvokeAccumulator(real_model="test-model") + opaque_body = { + "assistant_blocks": [{"type": "redacted_thinking", "data": "enc-1"}], + } + + accumulator.consume(_make_chunk(content="a", opaque_body=opaque_body)) + accumulator.consume(_make_chunk(content="b")) + + result = accumulator.to_result(prompt_messages=[UserPromptMessage(content="hi")]) + + assert result.message.opaque_body == opaque_body