mirror of
https://github.com/langgenius/dify.git
synced 2026-07-30 16:59:35 +08:00
test: use SQLite sessions in core app apps (#39102)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
3f1036e55a
commit
1593b7cc60
@ -4,6 +4,7 @@ from contextlib import contextmanager
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig
|
||||
from core.app.apps.advanced_chat.generate_task_pipeline import (
|
||||
@ -54,10 +55,11 @@ from core.workflow.nodes.human_input.entities import UserActionConfig
|
||||
from core.workflow.nodes.human_input.pause_reason import DifyHITLEventType
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from graphon.enums import BuiltinNodeTypes
|
||||
from graphon.file import FileTransferMethod, FileType
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.enums import MessageStatus
|
||||
from models.model import AppMode, EndUser
|
||||
from models.model import AppMode, EndUser, Message, MessageFile
|
||||
from tests.workflow_test_utils import build_test_variable_pool
|
||||
|
||||
|
||||
@ -110,6 +112,40 @@ def _make_pipeline():
|
||||
return pipeline
|
||||
|
||||
|
||||
def _persist_message(session: Session, *, message_id: str = "message-id") -> Message:
|
||||
message = Message(
|
||||
app_id="app",
|
||||
model_provider="provider",
|
||||
model_id="model",
|
||||
override_model_configs=None,
|
||||
conversation_id="conv-id",
|
||||
inputs={},
|
||||
query="hello",
|
||||
message="",
|
||||
message_tokens=0,
|
||||
message_unit_price=0,
|
||||
message_price_unit=0,
|
||||
answer="",
|
||||
answer_tokens=0,
|
||||
answer_unit_price=0,
|
||||
answer_price_unit=0,
|
||||
parent_message_id=None,
|
||||
provider_response_latency=0,
|
||||
total_price=0,
|
||||
currency="USD",
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
from_source="api",
|
||||
from_end_user_id="end-user",
|
||||
from_account_id=None,
|
||||
app_mode=AppMode.ADVANCED_CHAT,
|
||||
status=MessageStatus.PAUSED,
|
||||
)
|
||||
message.id = message_id
|
||||
session.add(message)
|
||||
session.commit()
|
||||
return message
|
||||
|
||||
|
||||
class TestAdvancedChatGenerateTaskPipeline:
|
||||
def test_ensure_workflow_initialized_raises(self):
|
||||
pipeline = _make_pipeline()
|
||||
@ -255,53 +291,29 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
pipeline._base_task_pipeline.handle_error = lambda **kwargs: ValueError("boom")
|
||||
pipeline._base_task_pipeline.error_to_stream_response = lambda err: err
|
||||
|
||||
@contextmanager
|
||||
def _fake_session():
|
||||
yield SimpleNamespace()
|
||||
|
||||
pipeline._database_session = _fake_session
|
||||
|
||||
responses = list(pipeline._handle_error_event(QueueErrorEvent(error=ValueError("boom"))))
|
||||
|
||||
assert isinstance(responses[0], ValueError)
|
||||
|
||||
def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_handle_workflow_started_event_sets_run_id(self, sqlite_session: Session):
|
||||
pipeline = _make_pipeline()
|
||||
message = _persist_message(sqlite_session)
|
||||
other_message = _persist_message(sqlite_session, message_id="other-message-id")
|
||||
pipeline._graph_runtime_state = GraphRuntimeState(
|
||||
variable_pool=build_test_variable_pool(variables=build_system_variables(workflow_execution_id="run-id")),
|
||||
start_at=0.0,
|
||||
)
|
||||
pipeline._workflow_response_converter.workflow_start_to_stream_response = lambda **kwargs: "started"
|
||||
|
||||
# Track database operations for verification
|
||||
executed_statements = []
|
||||
|
||||
@contextmanager
|
||||
def _fake_session():
|
||||
sess = SimpleNamespace()
|
||||
|
||||
def _execute(stmt):
|
||||
executed_statements.append(stmt)
|
||||
return SimpleNamespace()
|
||||
|
||||
sess.execute = _execute
|
||||
yield sess
|
||||
|
||||
monkeypatch.setattr(pipeline, "_database_session", _fake_session)
|
||||
monkeypatch.setattr(pipeline, "_get_message", lambda **kwargs: SimpleNamespace())
|
||||
|
||||
responses = list(pipeline._handle_workflow_started_event(QueueWorkflowStartedEvent()))
|
||||
|
||||
assert pipeline._workflow_run_id == "run-id"
|
||||
assert responses == ["started"]
|
||||
|
||||
# Verify database operation was executed
|
||||
assert len(executed_statements) == 1
|
||||
# Verify the UPDATE statement targets the correct message and sets workflow_run_id
|
||||
update_stmt = executed_statements[0]
|
||||
stmt_str = str(update_stmt)
|
||||
assert "UPDATE messages" in stmt_str
|
||||
assert "WHERE messages.id" in stmt_str
|
||||
sqlite_session.refresh(message)
|
||||
sqlite_session.refresh(other_message)
|
||||
assert message.workflow_run_id == "run-id"
|
||||
assert other_message.workflow_run_id is None
|
||||
|
||||
def test_message_end_to_stream_response_strips_annotation_reply(self):
|
||||
pipeline = _make_pipeline()
|
||||
@ -464,7 +476,7 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
assert list(pipeline._handle_loop_next_event(loop_next)) == ["loop_next"]
|
||||
assert list(pipeline._handle_loop_completed_event(loop_done)) == ["loop_done"]
|
||||
|
||||
def test_workflow_finish_handlers(self, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_workflow_finish_handlers(self):
|
||||
pipeline = _make_pipeline()
|
||||
pipeline._workflow_run_id = "run-id"
|
||||
pipeline._graph_runtime_state = GraphRuntimeState(
|
||||
@ -482,12 +494,6 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
pipeline._base_task_pipeline.error_to_stream_response = lambda err: err
|
||||
pipeline._get_message = lambda **kwargs: SimpleNamespace(id="message-id")
|
||||
|
||||
@contextmanager
|
||||
def _fake_session():
|
||||
yield SimpleNamespace(scalar=lambda *args, **kwargs: None)
|
||||
|
||||
monkeypatch.setattr(pipeline, "_database_session", _fake_session)
|
||||
|
||||
succeeded_responses = list(pipeline._handle_workflow_succeeded_event(QueueWorkflowSucceededEvent(outputs={})))
|
||||
assert len(succeeded_responses) == 2
|
||||
assert isinstance(succeeded_responses[0], MessageEndStreamResponse)
|
||||
@ -644,12 +650,12 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
assert list(pipeline._handle_human_input_form_timeout_event(timeout_event)) == ["timeout"]
|
||||
assert persisted == ["saved"]
|
||||
|
||||
def test_save_message_preserves_full_answer_and_sets_usage(self):
|
||||
def test_save_message_preserves_full_answer_and_sets_usage(self, sqlite_session: Session):
|
||||
pipeline = _make_pipeline()
|
||||
pipeline._recorded_files = [
|
||||
{
|
||||
"type": "image",
|
||||
"transfer_method": "remote",
|
||||
"type": FileType.IMAGE,
|
||||
"transfer_method": FileTransferMethod.REMOTE_URL,
|
||||
"remote_url": "http://example.com/file.png",
|
||||
"related_id": "file-id",
|
||||
}
|
||||
@ -660,32 +666,7 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
pipeline._task_state.first_token_time = pipeline._base_task_pipeline.start_at + 0.1
|
||||
pipeline._task_state.last_token_time = pipeline._base_task_pipeline.start_at + 0.2
|
||||
|
||||
message = SimpleNamespace(
|
||||
id="message-id",
|
||||
status=MessageStatus.PAUSED,
|
||||
answer="",
|
||||
updated_at=None,
|
||||
provider_response_latency=None,
|
||||
message_tokens=None,
|
||||
message_unit_price=None,
|
||||
message_price_unit=None,
|
||||
answer_tokens=None,
|
||||
answer_unit_price=None,
|
||||
answer_price_unit=None,
|
||||
total_price=None,
|
||||
currency=None,
|
||||
message_metadata=None,
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
from_account_id=None,
|
||||
from_end_user_id="end-user",
|
||||
)
|
||||
|
||||
class _Session:
|
||||
def scalar(self, *args, **kwargs):
|
||||
return message
|
||||
|
||||
def add_all(self, items):
|
||||
self.items = items
|
||||
message = _persist_message(sqlite_session)
|
||||
|
||||
graph_runtime_state = GraphRuntimeState(
|
||||
variable_pool=VariablePool.from_bootstrap(
|
||||
@ -694,13 +675,15 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
start_at=0.0,
|
||||
)
|
||||
|
||||
pipeline._save_message(session=_Session(), graph_runtime_state=graph_runtime_state)
|
||||
pipeline._save_message(session=sqlite_session, graph_runtime_state=graph_runtime_state)
|
||||
sqlite_session.commit()
|
||||
|
||||
assert message.status == MessageStatus.NORMAL
|
||||
assert message.answer == " hello "
|
||||
assert message.message_metadata
|
||||
assert sqlite_session.query(MessageFile).filter_by(message_id=message.id).count() == 1
|
||||
|
||||
def test_handle_stop_event_saves_message_for_moderation(self, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_handle_stop_event_saves_message_for_moderation(self):
|
||||
pipeline = _make_pipeline()
|
||||
pipeline._message_end_to_stream_response = lambda: "end"
|
||||
saved: list[str] = []
|
||||
@ -710,18 +693,12 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
|
||||
pipeline._save_message = _save_message
|
||||
|
||||
@contextmanager
|
||||
def _fake_session():
|
||||
yield SimpleNamespace()
|
||||
|
||||
monkeypatch.setattr(pipeline, "_database_session", _fake_session)
|
||||
|
||||
responses = list(pipeline._handle_stop_event(QueueStopEvent(stopped_by=QueueStopEvent.StopBy.INPUT_MODERATION)))
|
||||
|
||||
assert responses == ["end"]
|
||||
assert saved == ["saved"]
|
||||
|
||||
def test_handle_message_end_event_applies_output_moderation(self, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_handle_message_end_event_applies_output_moderation(self):
|
||||
pipeline = _make_pipeline()
|
||||
pipeline._graph_runtime_state = GraphRuntimeState(
|
||||
variable_pool=VariablePool.from_bootstrap(
|
||||
@ -740,12 +717,6 @@ class TestAdvancedChatGenerateTaskPipeline:
|
||||
|
||||
pipeline._save_message = _save_message
|
||||
|
||||
@contextmanager
|
||||
def _fake_session():
|
||||
yield SimpleNamespace()
|
||||
|
||||
monkeypatch.setattr(pipeline, "_database_session", _fake_session)
|
||||
|
||||
responses = list(pipeline._handle_advanced_chat_message_end_event(QueueAdvancedChatMessageEndEvent()))
|
||||
|
||||
assert responses == ["replace", "end"]
|
||||
|
||||
Loading…
Reference in New Issue
Block a user