Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion src/graphon/model_runtime/entities/message_entities.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -59,6 +59,7 @@ class PromptMessageContent(ABC, BaseModel):
"""Model class for prompt message content."""

type: PromptMessageContentType
opaque_body: JsonValue | None = None


class TextPromptMessageContent(PromptMessageContent):
Expand Down Expand Up @@ -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."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand All @@ -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,
)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from collections.abc import Iterator
from typing import Any

from graphon.model_runtime.entities.llm_entities import (
LLMResult,
Expand All @@ -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,
Expand Down Expand Up @@ -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")]

Expand Down
79 changes: 78 additions & 1 deletion tests/model_runtime/test_message_entities.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down Expand Up @@ -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
67 changes: 67 additions & 0 deletions tests/model_runtime/test_streaming_invoke_accumulator.py
Original file line number Diff line number Diff line change
@@ -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
Loading