refactor(models): pass session into Conversation.to_dict and Message.to_dict (#42834)

Co-authored-by: Dorian <dorian0326@users.noreply.github.com>
This commit is contained in:
Dorian
2026-09-23 19:47:44 +00:00
committed by GitHub
co-authored by Dorian
parent e43ea0986e
commit 7cc0898f58
5 changed files with 21 additions and 19 deletions
+5 -5
View File
@@ -1000,7 +1000,7 @@ class TraceTask:
message_trace_info = MessageTraceInfo(
trace_id=self.trace_id,
message_id=message_id,
message_data=message_data.to_dict(),
message_data=message_data.to_dict(session=db.session()),
conversation_model=conversation_mode,
message_tokens=message_tokens,
answer_tokens=message_data.answer_tokens,
@@ -1052,7 +1052,7 @@ class TraceTask:
trace_id=self.trace_id,
message_id=workflow_app_log_id or message_id,
inputs=inputs,
message_data=message_data.to_dict(),
message_data=message_data.to_dict(session=db.session()),
flagged=moderation_result.flagged,
action=moderation_result.action,
preset_response=moderation_result.preset_response,
@@ -1094,7 +1094,7 @@ class TraceTask:
suggested_question_trace_info = SuggestedQuestionTraceInfo(
trace_id=self.trace_id,
message_id=workflow_app_log_id or message_id,
message_data=message_data.to_dict(),
message_data=message_data.to_dict(session=db.session()),
inputs=message_data.message,
outputs=message_data.answer,
start_time=timer.get("start"),
@@ -1197,7 +1197,7 @@ class TraceTask:
start_time=timer.get("start"),
end_time=timer.get("end"),
metadata=metadata,
message_data=message_data.to_dict(),
message_data=message_data.to_dict(session=db.session()),
error=kwargs.get("error"),
)
@@ -1262,7 +1262,7 @@ class TraceTask:
tool_trace_info = ToolTraceInfo(
trace_id=self.trace_id,
message_id=message_id,
message_data=message_data.to_dict(),
message_data=message_data.to_dict(session=db.session()),
tool_name=tool_name,
start_time=timer.get("start") if timer else created_time,
end_time=timer.get("end") if timer else end_time,
+4 -4
View File
@@ -1402,7 +1402,7 @@ class Conversation(Base):
def in_debug_mode(self) -> bool:
return self.override_model_configs is not None
def to_dict(self) -> ConversationDict:
def to_dict(self, *, session: Session) -> ConversationDict:
return {
"id": self.id,
"app_id": self.app_id,
@@ -1413,7 +1413,7 @@ class Conversation(Base):
"mode": self.mode,
"name": self.name,
"summary": self.summary,
"inputs": self.inputs_with_session(session=db.session()),
"inputs": self.inputs_with_session(session=session),
"introduction": self.introduction,
"system_instruction": self.system_instruction,
"system_instruction_tokens": self.system_instruction_tokens,
@@ -1775,13 +1775,13 @@ class Message(Base):
return None
def to_dict(self) -> MessageDict:
def to_dict(self, *, session: Session) -> MessageDict:
return {
"id": self.id,
"app_id": self.app_id,
"conversation_id": self.conversation_id,
"model_id": self.model_id,
"inputs": self.inputs_with_session(session=db.session()),
"inputs": self.inputs_with_session(session=session),
"query": self.query,
"total_price": self.total_price,
"message": self.message,
@@ -148,7 +148,7 @@ class ClearFreePlanTenantExpiredLogs:
f"-{time.time()}.json",
json.dumps(
jsonable_encoder(
[message.to_dict() for message in messages],
[message.to_dict(session=session) for message in messages],
),
).encode("utf-8"),
)
@@ -188,7 +188,7 @@ class ClearFreePlanTenantExpiredLogs:
f"-{time.time()}.json",
json.dumps(
jsonable_encoder(
[conversation.to_dict() for conversation in conversations],
[conversation.to_dict(session=session) for conversation in conversations],
),
).encode("utf-8"),
)
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock
import pytest
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceTask
@@ -21,12 +21,12 @@ TABLES = (App, Conversation, Message, MessageFile, WorkflowAppLog, WorkflowNodeE
def _bind_trace_database(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
"""Use real SQLite sessions for ORM lookups without changing trace-domain data."""
monkeypatch.setattr(
"core.ops.ops_trace_manager.db",
SimpleNamespace(engine=sqlite_engine, session=sqlite_session),
SimpleNamespace(engine=sqlite_engine, session=scoped_session(sqlite_session_factory)),
)
@@ -84,7 +84,7 @@ def _make_message_data():
def __init__(self, values):
self.__dict__.update(values)
def to_dict(self):
def to_dict(self, **_kwargs: object):
return dict(self.__dict__)
return _MessageData(data)
@@ -823,7 +823,8 @@ class TestConversationModel:
# Assert
assert result is True
def test_conversation_to_dict_serialization(self):
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
def test_conversation_to_dict_serialization(self, sqlite_session: Session):
"""Test conversation to_dict method."""
# Arrange
app_id = str(uuid4())
@@ -841,7 +842,7 @@ class TestConversationModel:
conversation._inputs = {"query": "test"}
# Act
result = conversation.to_dict()
result = conversation.to_dict(session=sqlite_session)
# Assert
assert result["id"] == conversation.id
@@ -1004,7 +1005,8 @@ class TestMessageModel:
# Assert
assert result == {}
def test_message_to_dict_serialization(self):
@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True)
def test_message_to_dict_serialization(self, sqlite_session: Session):
"""Test message to_dict method."""
# Arrange
app_id = str(uuid4())
@@ -1030,7 +1032,7 @@ class TestMessageModel:
)
# Act
result = message.to_dict()
result = message.to_dict(session=sqlite_session)
# Assert
assert result["id"] == message.id