mirror of
https://github.com/langgenius/dify.git
synced 2026-09-29 17:07:38 +08:00
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:
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user