mirror of
https://github.com/langgenius/dify.git
synced 2026-07-26 14:18:35 +08:00
test: use sqlite3 session in test_workflow_execute_task (#38704)
This commit is contained in:
parent
8caa1d1379
commit
ca1782b105
@ -4,18 +4,23 @@ import json
|
||||
import logging
|
||||
import uuid
|
||||
from contextlib import nullcontext
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import App, AppMode, Conversation
|
||||
from models.workflow import Workflow, WorkflowRun
|
||||
from models.base import TypeBase
|
||||
from models.enums import ConversationFromSource, CreatorUserRole, WorkflowRunTriggeredFrom
|
||||
from models.model import App, AppMode, Conversation, Message
|
||||
from models.workflow import Workflow, WorkflowRun, WorkflowType
|
||||
from repositories.sqlalchemy_api_workflow_run_repository import _WorkflowRunError
|
||||
from tasks.app_generate import workflow_execute_task as workflow_execute_task_module
|
||||
from tasks.app_generate.workflow_execute_task import (
|
||||
@ -28,17 +33,6 @@ from tasks.app_generate.workflow_execute_task import (
|
||||
)
|
||||
|
||||
|
||||
class _FakeSessionContext:
|
||||
def __init__(self, session: MagicMock):
|
||||
self._session = session
|
||||
|
||||
def __enter__(self) -> MagicMock:
|
||||
return self._session
|
||||
|
||||
def __exit__(self, exc_type, exc, tb) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
class _StreamEventModel(BaseModel):
|
||||
event: object | None = None
|
||||
task_id: object | None = None
|
||||
@ -81,6 +75,147 @@ def _published_payloads(topic: MagicMock) -> list[dict[str, object] | str]:
|
||||
return [_decode_published_payload(call.args[0]) for call in topic.publish.call_args_list]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]:
|
||||
tables = [
|
||||
TypeBase.metadata.tables[model.__tablename__] for model in (App, Workflow, WorkflowRun, Conversation, Message)
|
||||
]
|
||||
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
|
||||
return sessionmaker(bind=sqlite_engine, expire_on_commit=False)
|
||||
|
||||
|
||||
def _persist_app_and_workflow(session_factory: sessionmaker[Session]) -> None:
|
||||
app = App(
|
||||
id="app-id",
|
||||
tenant_id="tenant-id",
|
||||
name="Test App",
|
||||
mode=AppMode.ADVANCED_CHAT,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
)
|
||||
workflow = Workflow(
|
||||
id="workflow-id",
|
||||
tenant_id="tenant-id",
|
||||
app_id=app.id,
|
||||
type=WorkflowType.CHAT,
|
||||
version=Workflow.VERSION_DRAFT,
|
||||
graph="{}",
|
||||
features="{}",
|
||||
created_by="workflow-owner",
|
||||
environment_variables=[],
|
||||
conversation_variables=[],
|
||||
rag_pipeline_variables=[],
|
||||
)
|
||||
with session_factory.begin() as session:
|
||||
session.add_all([app, workflow])
|
||||
|
||||
|
||||
def _persist_resumption_models(
|
||||
session_factory: sessionmaker[Session],
|
||||
*,
|
||||
workflow_run_id: str,
|
||||
conversation_id: str | None = None,
|
||||
) -> None:
|
||||
app = App(
|
||||
id="app-id",
|
||||
tenant_id="tenant-id",
|
||||
name="Test App",
|
||||
mode=AppMode.ADVANCED_CHAT,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
)
|
||||
workflow = Workflow(
|
||||
id="wf-id",
|
||||
tenant_id="tenant-id",
|
||||
app_id=app.id,
|
||||
type=WorkflowType.CHAT,
|
||||
version=Workflow.VERSION_DRAFT,
|
||||
graph="{}",
|
||||
features="{}",
|
||||
created_by="workflow-owner",
|
||||
environment_variables=[],
|
||||
conversation_variables=[],
|
||||
rag_pipeline_variables=[],
|
||||
)
|
||||
workflow_run = WorkflowRun(
|
||||
id=workflow_run_id,
|
||||
tenant_id="tenant-id",
|
||||
app_id=app.id,
|
||||
workflow_id=workflow.id,
|
||||
type=WorkflowType.CHAT,
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
version=workflow.version,
|
||||
graph="{}",
|
||||
inputs="{}",
|
||||
status=WorkflowExecutionStatus.RUNNING,
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="account-id",
|
||||
)
|
||||
|
||||
with session_factory.begin() as session:
|
||||
session.add_all([app, workflow, workflow_run])
|
||||
if conversation_id is None:
|
||||
return
|
||||
|
||||
conversation = Conversation(
|
||||
id=conversation_id,
|
||||
app_id=app.id,
|
||||
mode=AppMode.ADVANCED_CHAT,
|
||||
name="Test Conversation",
|
||||
inputs={},
|
||||
from_source=ConversationFromSource.API,
|
||||
)
|
||||
messages = [
|
||||
Message(
|
||||
id="older-message-id",
|
||||
app_id=app.id,
|
||||
conversation_id=conversation.id,
|
||||
inputs={},
|
||||
query="older matching message",
|
||||
message={"role": "user", "content": "older"},
|
||||
answer="older",
|
||||
message_unit_price=Decimal(0),
|
||||
answer_unit_price=Decimal(0),
|
||||
currency="USD",
|
||||
from_source=ConversationFromSource.API,
|
||||
workflow_run_id=workflow_run_id,
|
||||
created_at=datetime(2025, 1, 1),
|
||||
),
|
||||
Message(
|
||||
id="expected-message-id",
|
||||
app_id=app.id,
|
||||
conversation_id=conversation.id,
|
||||
inputs={},
|
||||
query="newer matching message",
|
||||
message={"role": "user", "content": "expected"},
|
||||
answer="expected",
|
||||
message_unit_price=Decimal(0),
|
||||
answer_unit_price=Decimal(0),
|
||||
currency="USD",
|
||||
from_source=ConversationFromSource.API,
|
||||
workflow_run_id=workflow_run_id,
|
||||
created_at=datetime(2025, 1, 2),
|
||||
),
|
||||
Message(
|
||||
id="other-run-message-id",
|
||||
app_id=app.id,
|
||||
conversation_id=conversation.id,
|
||||
inputs={},
|
||||
query="newest message from another run",
|
||||
message={"role": "user", "content": "other run"},
|
||||
answer="other run",
|
||||
message_unit_price=Decimal(0),
|
||||
answer_unit_price=Decimal(0),
|
||||
currency="USD",
|
||||
from_source=ConversationFromSource.API,
|
||||
workflow_run_id="other-run-id",
|
||||
created_at=datetime(2025, 1, 3),
|
||||
),
|
||||
]
|
||||
session.add(conversation)
|
||||
session.add_all(messages)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("event", "expected"),
|
||||
[
|
||||
@ -397,7 +532,9 @@ def test_publish_streaming_response_does_not_publish_synthetic_failure_after_ter
|
||||
|
||||
|
||||
def test_app_runner_streaming_failure_publishes_started_then_failed_workflow_finished(
|
||||
mock_topic: MagicMock, monkeypatch
|
||||
mock_topic: MagicMock,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
exec_params = AppExecutionParams(
|
||||
app_id="app-id",
|
||||
@ -410,14 +547,9 @@ def test_app_runner_streaming_failure_publishes_started_then_failed_workflow_fin
|
||||
streaming=True,
|
||||
workflow_run_id="workflow-run-id",
|
||||
)
|
||||
runner = _AppRunner(session_factory=MagicMock(), exec_params=exec_params)
|
||||
_persist_app_and_workflow(sqlite_session_factory)
|
||||
runner = _AppRunner(session_factory=sqlite_session_factory, exec_params=exec_params)
|
||||
|
||||
workflow = SimpleNamespace(id="workflow-id", app_id="app-id", created_by="workflow-owner")
|
||||
app = SimpleNamespace(id="app-id")
|
||||
fake_session = MagicMock()
|
||||
fake_session.get.side_effect = [workflow, app]
|
||||
|
||||
monkeypatch.setattr(runner, "_session", lambda: nullcontext(fake_session))
|
||||
monkeypatch.setattr(runner, "_resolve_user", lambda: MagicMock())
|
||||
monkeypatch.setattr(runner, "_setup_flask_context", lambda _user: nullcontext())
|
||||
monkeypatch.setattr(runner, "_run_app", lambda **_kwargs: (_ for _ in ()).throw(ValueError("Invalid upload file")))
|
||||
@ -494,6 +626,7 @@ def test_resolve_account_for_run_without_switching_tenant():
|
||||
def test_app_runner_streaming_failure_keeps_existing_pre_runtime_helper_behavior(
|
||||
mock_topic: MagicMock,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
exec_params = AppExecutionParams(
|
||||
app_id="app-id",
|
||||
@ -506,14 +639,9 @@ def test_app_runner_streaming_failure_keeps_existing_pre_runtime_helper_behavior
|
||||
streaming=True,
|
||||
workflow_run_id="workflow-run-id",
|
||||
)
|
||||
runner = _AppRunner(session_factory=MagicMock(), exec_params=exec_params)
|
||||
_persist_app_and_workflow(sqlite_session_factory)
|
||||
runner = _AppRunner(session_factory=sqlite_session_factory, exec_params=exec_params)
|
||||
|
||||
workflow = SimpleNamespace(id="workflow-id", app_id="app-id", created_by="workflow-owner")
|
||||
app = SimpleNamespace(id="app-id")
|
||||
fake_session = MagicMock()
|
||||
fake_session.get.side_effect = [workflow, app]
|
||||
|
||||
monkeypatch.setattr(runner, "_session", lambda: nullcontext(fake_session))
|
||||
monkeypatch.setattr(runner, "_resolve_user", lambda: MagicMock())
|
||||
monkeypatch.setattr(runner, "_setup_flask_context", lambda _user: nullcontext())
|
||||
monkeypatch.setattr(runner, "_run_app", lambda **_kwargs: (_ for _ in ()).throw(ValueError("Invalid upload file")))
|
||||
@ -532,6 +660,7 @@ def test_app_runner_streaming_failure_keeps_existing_pre_runtime_helper_behavior
|
||||
|
||||
def test_app_runner_streaming_success_calls_publish_streaming_response_with_full_signature(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
exec_params = AppExecutionParams(
|
||||
app_id="app-id",
|
||||
@ -544,16 +673,11 @@ def test_app_runner_streaming_success_calls_publish_streaming_response_with_full
|
||||
streaming=True,
|
||||
workflow_run_id="workflow-run-id",
|
||||
)
|
||||
runner = _AppRunner(session_factory=MagicMock(), exec_params=exec_params)
|
||||
|
||||
workflow = SimpleNamespace(id="workflow-id", app_id="app-id", created_by="workflow-owner")
|
||||
app = SimpleNamespace(id="app-id")
|
||||
fake_session = MagicMock()
|
||||
fake_session.get.side_effect = [workflow, app]
|
||||
_persist_app_and_workflow(sqlite_session_factory)
|
||||
runner = _AppRunner(session_factory=sqlite_session_factory, exec_params=exec_params)
|
||||
response_stream = _single_event_generator({"event": "message"})
|
||||
publish_streaming_response = MagicMock()
|
||||
|
||||
monkeypatch.setattr(runner, "_session", lambda: nullcontext(fake_session))
|
||||
monkeypatch.setattr(runner, "_resolve_user", lambda: MagicMock())
|
||||
monkeypatch.setattr(runner, "_setup_flask_context", lambda _user: nullcontext())
|
||||
monkeypatch.setattr(runner, "_run_app", lambda **_kwargs: response_stream)
|
||||
@ -574,12 +698,20 @@ def test_app_runner_streaming_success_calls_publish_streaming_response_with_full
|
||||
)
|
||||
|
||||
|
||||
def test_resume_app_execution_queries_message_by_conversation_and_workflow_run(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_resume_app_execution_queries_message_by_conversation_and_workflow_run(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_engine: Engine,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
workflow_run_id = "run-id"
|
||||
conversation_id = "conversation-id"
|
||||
message = MagicMock()
|
||||
_persist_resumption_models(
|
||||
sqlite_session_factory,
|
||||
workflow_run_id=workflow_run_id,
|
||||
conversation_id=conversation_id,
|
||||
)
|
||||
|
||||
monkeypatch.setattr("tasks.app_generate.workflow_execute_task.db", SimpleNamespace(engine=object()))
|
||||
monkeypatch.setattr("tasks.app_generate.workflow_execute_task.db", SimpleNamespace(engine=sqlite_engine))
|
||||
|
||||
pause_entity = MagicMock()
|
||||
pause_entity.get_state.return_value = b"state"
|
||||
@ -604,36 +736,6 @@ def test_resume_app_execution_queries_message_by_conversation_and_workflow_run(m
|
||||
lambda *_args, **_kwargs: MagicMock(),
|
||||
)
|
||||
|
||||
workflow_run = SimpleNamespace(
|
||||
workflow_id="wf-id",
|
||||
app_id="app-id",
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="account-id",
|
||||
tenant_id="tenant-id",
|
||||
)
|
||||
workflow = SimpleNamespace(created_by="workflow-owner")
|
||||
app_model = SimpleNamespace(id="app-id")
|
||||
conversation = SimpleNamespace(id=conversation_id)
|
||||
|
||||
session = MagicMock()
|
||||
|
||||
def _session_get(model, key):
|
||||
if model is WorkflowRun:
|
||||
return workflow_run
|
||||
if model is Workflow:
|
||||
return workflow
|
||||
if model is App:
|
||||
return app_model
|
||||
if model is Conversation:
|
||||
return conversation
|
||||
return None
|
||||
|
||||
session.get.side_effect = _session_get
|
||||
session.scalar.return_value = message
|
||||
|
||||
monkeypatch.setattr(
|
||||
"tasks.app_generate.workflow_execute_task.Session", lambda *_args, **_kwargs: _FakeSessionContext(session)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"tasks.app_generate.workflow_execute_task._resolve_user_for_run", lambda *_args, **_kwargs: MagicMock()
|
||||
)
|
||||
@ -643,29 +745,21 @@ def test_resume_app_execution_queries_message_by_conversation_and_workflow_run(m
|
||||
|
||||
_resume_app_execution({"workflow_run_id": workflow_run_id})
|
||||
|
||||
stmt = session.scalar.call_args.args[0]
|
||||
stmt_text = str(stmt)
|
||||
assert "messages.conversation_id = :conversation_id_1" in stmt_text
|
||||
assert "messages.workflow_run_id = :workflow_run_id_1" in stmt_text
|
||||
assert "ORDER BY messages.created_at DESC" in stmt_text
|
||||
assert " LIMIT " in stmt_text
|
||||
|
||||
compiled_params = stmt.compile().params
|
||||
assert conversation_id in compiled_params.values()
|
||||
assert workflow_run_id in compiled_params.values()
|
||||
|
||||
workflow_run_repo.resume_workflow_pause.assert_called_once_with(workflow_run_id, pause_entity)
|
||||
resume_advanced_chat.assert_called_once()
|
||||
assert resume_advanced_chat.call_args.kwargs["conversation"] is conversation
|
||||
assert resume_advanced_chat.call_args.kwargs["message"] is message
|
||||
assert resume_advanced_chat.call_args.kwargs["conversation"].id == conversation_id
|
||||
assert resume_advanced_chat.call_args.kwargs["message"].id == "expected-message-id"
|
||||
|
||||
|
||||
def test_resume_app_execution_returns_early_when_advanced_chat_missing_conversation_id(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_engine: Engine,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
workflow_run_id = "run-id"
|
||||
_persist_resumption_models(sqlite_session_factory, workflow_run_id=workflow_run_id)
|
||||
|
||||
monkeypatch.setattr("tasks.app_generate.workflow_execute_task.db", SimpleNamespace(engine=object()))
|
||||
monkeypatch.setattr("tasks.app_generate.workflow_execute_task.db", SimpleNamespace(engine=sqlite_engine))
|
||||
|
||||
pause_entity = MagicMock()
|
||||
pause_entity.get_state.return_value = b"state"
|
||||
@ -690,32 +784,6 @@ def test_resume_app_execution_returns_early_when_advanced_chat_missing_conversat
|
||||
lambda *_args, **_kwargs: MagicMock(),
|
||||
)
|
||||
|
||||
workflow_run = SimpleNamespace(
|
||||
workflow_id="wf-id",
|
||||
app_id="app-id",
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="account-id",
|
||||
tenant_id="tenant-id",
|
||||
)
|
||||
workflow = SimpleNamespace(created_by="workflow-owner")
|
||||
app_model = SimpleNamespace(id="app-id")
|
||||
|
||||
session = MagicMock()
|
||||
|
||||
def _session_get(model, key):
|
||||
if model is WorkflowRun:
|
||||
return workflow_run
|
||||
if model is Workflow:
|
||||
return workflow
|
||||
if model is App:
|
||||
return app_model
|
||||
return None
|
||||
|
||||
session.get.side_effect = _session_get
|
||||
|
||||
monkeypatch.setattr(
|
||||
"tasks.app_generate.workflow_execute_task.Session", lambda *_args, **_kwargs: _FakeSessionContext(session)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"tasks.app_generate.workflow_execute_task._resolve_user_for_run", lambda *_args, **_kwargs: MagicMock()
|
||||
)
|
||||
@ -724,12 +792,14 @@ def test_resume_app_execution_returns_early_when_advanced_chat_missing_conversat
|
||||
|
||||
_resume_app_execution({"workflow_run_id": workflow_run_id})
|
||||
|
||||
session.scalar.assert_not_called()
|
||||
workflow_run_repo.resume_workflow_pause.assert_not_called()
|
||||
resume_advanced_chat.assert_not_called()
|
||||
|
||||
|
||||
def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
generate_entity = _build_advanced_chat_generate_entity(conversation_id="conversation-id")
|
||||
generate_entity.stream = False
|
||||
workflow = SimpleNamespace(id="workflow-id", created_by="workflow-owner")
|
||||
@ -765,7 +835,7 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monk
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
session_factory=sqlite_session_factory,
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
workflow_run=SimpleNamespace(triggered_from="app_run"),
|
||||
@ -785,7 +855,10 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monk
|
||||
)
|
||||
|
||||
|
||||
def test_resume_workflow_publishes_events_for_originally_blocking_runs(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_resume_workflow_publishes_events_for_originally_blocking_runs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
generate_entity = _build_workflow_generate_entity(stream=False)
|
||||
workflow = SimpleNamespace(id="workflow-id", created_by="workflow-owner")
|
||||
|
||||
@ -819,7 +892,7 @@ def test_resume_workflow_publishes_events_for_originally_blocking_runs(monkeypat
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
session_factory=sqlite_session_factory,
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
workflow_run=SimpleNamespace(triggered_from="app_run"),
|
||||
@ -840,7 +913,10 @@ def test_resume_workflow_publishes_events_for_originally_blocking_runs(monkeypat
|
||||
workflow_run_repo.delete_workflow_pause.assert_called_once_with(pause_entity)
|
||||
|
||||
|
||||
def test_resume_workflow_ignores_missing_old_pause_after_repause(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_resume_workflow_ignores_missing_old_pause_after_repause(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
):
|
||||
generate_entity = _build_workflow_generate_entity(stream=False)
|
||||
workflow = SimpleNamespace(id="workflow-id", created_by="workflow-owner")
|
||||
|
||||
@ -875,7 +951,7 @@ def test_resume_workflow_ignores_missing_old_pause_after_repause(monkeypatch: py
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
session_factory=sqlite_session_factory,
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
workflow_run=SimpleNamespace(triggered_from="app_run"),
|
||||
|
||||
Loading…
Reference in New Issue
Block a user