mirror of
https://github.com/langgenius/dify.git
synced 2026-07-23 20:18:40 +08:00
test: use sqlite3 session in test_async_workflow_service (#38692)
This commit is contained in:
parent
82c8741c73
commit
fd703739b5
@ -1,12 +1,18 @@
|
||||
import json
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import services.async_workflow_service as async_workflow_service_module
|
||||
from models.enums import AppTriggerType, CreatorUserRole, WorkflowRunTriggeredFrom, WorkflowTriggerStatus
|
||||
from models.model import App, AppMode
|
||||
from models.trigger import WorkflowTriggerLog
|
||||
from services.async_workflow_service import AsyncWorkflowService
|
||||
from services.errors.app import QuotaExceededError, WorkflowNotFoundError
|
||||
from services.workflow.entities import AsyncTriggerResponse, TriggerData
|
||||
@ -37,25 +43,66 @@ class AsyncWorkflowServiceTestDataFactory:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def create_trigger_log_with_data(trigger_data: TriggerData, retry_count: int = 0) -> MagicMock:
|
||||
"""Create a mock trigger log with serialized trigger data."""
|
||||
trigger_log = MagicMock()
|
||||
trigger_log.id = "trigger-log-123"
|
||||
trigger_log.trigger_data = trigger_data.model_dump_json()
|
||||
trigger_log.retry_count = retry_count
|
||||
trigger_log.error = "previous-error"
|
||||
trigger_log.status = WorkflowTriggerStatus.FAILED
|
||||
trigger_log.to_dict.return_value = {"id": trigger_log.id}
|
||||
def create_app(app_id: str = "app-123", tenant_id: str = "tenant-123") -> App:
|
||||
"""Create an app that can be persisted for trigger lookup tests."""
|
||||
return App(
|
||||
id=app_id,
|
||||
tenant_id=tenant_id,
|
||||
name="Async workflow app",
|
||||
description="",
|
||||
mode=AppMode.WORKFLOW,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
max_active_requests=0,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def create_trigger_log_with_data(
|
||||
trigger_data: TriggerData,
|
||||
*,
|
||||
trigger_log_id: str = "trigger-log-123",
|
||||
retry_count: int = 0,
|
||||
status: WorkflowTriggerStatus = WorkflowTriggerStatus.FAILED,
|
||||
created_at: datetime | None = None,
|
||||
) -> WorkflowTriggerLog:
|
||||
"""Create a persistent trigger log with serialized trigger data."""
|
||||
trigger_log = WorkflowTriggerLog(
|
||||
tenant_id=trigger_data.tenant_id,
|
||||
app_id=trigger_data.app_id,
|
||||
workflow_id=trigger_data.workflow_id or "workflow-123",
|
||||
workflow_run_id=None,
|
||||
root_node_id=trigger_data.root_node_id,
|
||||
trigger_metadata="{}",
|
||||
trigger_type=trigger_data.trigger_type,
|
||||
trigger_data=trigger_data.model_dump_json(),
|
||||
inputs=json.dumps(dict(trigger_data.inputs)),
|
||||
outputs=None,
|
||||
status=status,
|
||||
error="previous-error",
|
||||
queue_name=QueuePriority.SANDBOX,
|
||||
celery_task_id=None,
|
||||
created_by_role=CreatorUserRole.END_USER,
|
||||
created_by="end-user-123",
|
||||
retry_count=retry_count,
|
||||
elapsed_time=None,
|
||||
total_tokens=None,
|
||||
triggered_at=None,
|
||||
finished_at=None,
|
||||
)
|
||||
trigger_log.id = trigger_log_id
|
||||
if created_at is not None:
|
||||
trigger_log.created_at = created_at
|
||||
return trigger_log
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("sqlite_session")
|
||||
@pytest.mark.parametrize("sqlite_session", [(App, WorkflowTriggerLog)], indirect=True)
|
||||
class TestAsyncWorkflowService:
|
||||
@pytest.fixture
|
||||
def async_workflow_trigger_mocks(self):
|
||||
"""Shared fixture for async workflow trigger tests.
|
||||
|
||||
Yields mocks for:
|
||||
- repo: SQLAlchemyWorkflowTriggerLogRepository
|
||||
- dispatcher_manager_class: QueueDispatcherManager class
|
||||
- dispatcher: dispatcher instance
|
||||
- quota_service: QuotaService mock
|
||||
@ -64,23 +111,10 @@ class TestAsyncWorkflowService:
|
||||
- team_task: execute_workflow_team
|
||||
- sandbox_task: execute_workflow_sandbox
|
||||
"""
|
||||
mock_repo = MagicMock()
|
||||
|
||||
def _create_side_effect(new_log):
|
||||
new_log.id = "trigger-log-123"
|
||||
return new_log
|
||||
|
||||
mock_repo.create.side_effect = _create_side_effect
|
||||
|
||||
mock_dispatcher = MagicMock()
|
||||
mock_quota_service = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
async_workflow_service_module,
|
||||
"SQLAlchemyWorkflowTriggerLogRepository",
|
||||
return_value=mock_repo,
|
||||
),
|
||||
patch.object(async_workflow_service_module, "QueueDispatcherManager") as mock_dispatcher_manager_class,
|
||||
patch.object(async_workflow_service_module, "WorkflowService"),
|
||||
patch.object(
|
||||
@ -100,7 +134,6 @@ class TestAsyncWorkflowService:
|
||||
mock_dispatcher_manager_class.return_value.get_dispatcher.return_value = mock_dispatcher
|
||||
|
||||
yield {
|
||||
"repo": mock_repo,
|
||||
"dispatcher_manager_class": mock_dispatcher_manager_class,
|
||||
"dispatcher": mock_dispatcher,
|
||||
"quota_service": mock_quota_service,
|
||||
@ -119,15 +152,16 @@ class TestAsyncWorkflowService:
|
||||
],
|
||||
)
|
||||
def test_should_dispatch_to_matching_celery_task_when_triggering_workflow(
|
||||
self, queue_name, selected_task_attr, async_workflow_trigger_mocks
|
||||
self,
|
||||
queue_name,
|
||||
selected_task_attr,
|
||||
async_workflow_trigger_mocks,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
"""Test queue-based task routing and successful async trigger response."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
session.commit = MagicMock()
|
||||
app_model = MagicMock()
|
||||
app_model.id = "app-123"
|
||||
session.scalar.return_value = app_model
|
||||
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
|
||||
sqlite_session.commit()
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
|
||||
workflow = MagicMock()
|
||||
workflow.id = "workflow-123"
|
||||
@ -153,20 +187,25 @@ class TestAsyncWorkflowService:
|
||||
user = DummyAccount("account-123")
|
||||
|
||||
# Act
|
||||
result = AsyncWorkflowService.trigger_workflow_async(session=session, user=user, trigger_data=trigger_data)
|
||||
result = AsyncWorkflowService.trigger_workflow_async(
|
||||
session=sqlite_session, user=user, trigger_data=trigger_data
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, AsyncTriggerResponse)
|
||||
assert result.workflow_trigger_log_id == "trigger-log-123"
|
||||
assert result.workflow_trigger_log_id
|
||||
assert result.task_id == "task-123"
|
||||
assert result.status == "queued"
|
||||
assert result.queue == queue_name
|
||||
|
||||
mocks["quota_service"].reserve.assert_called_once()
|
||||
quota_charge_mock.commit.assert_called_once()
|
||||
assert session.commit.call_count == 3
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
created_log = mocks["repo"].create.call_args[0][0]
|
||||
created_log = sqlite_session.scalar(
|
||||
select(WorkflowTriggerLog).where(WorkflowTriggerLog.id == result.workflow_trigger_log_id)
|
||||
)
|
||||
assert created_log is not None
|
||||
assert created_log.status == WorkflowTriggerStatus.QUEUED
|
||||
assert created_log.queue_name == queue_name
|
||||
assert created_log.created_by_role == CreatorUserRole.ACCOUNT
|
||||
@ -182,18 +221,17 @@ class TestAsyncWorkflowService:
|
||||
}
|
||||
for task_attr, task_mock in task_mocks.items():
|
||||
if task_attr == selected_task_attr:
|
||||
task_mock.delay.assert_called_once_with({"workflow_trigger_log_id": "trigger-log-123"})
|
||||
task_mock.delay.assert_called_once_with({"workflow_trigger_log_id": result.workflow_trigger_log_id})
|
||||
else:
|
||||
task_mock.delay.assert_not_called()
|
||||
|
||||
def test_should_set_end_user_role_when_triggered_by_end_user(self, async_workflow_trigger_mocks):
|
||||
def test_should_set_end_user_role_when_triggered_by_end_user(
|
||||
self, async_workflow_trigger_mocks, sqlite_session: Session
|
||||
):
|
||||
"""Test that non-account users are tracked as END_USER in trigger logs."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
session.commit = MagicMock()
|
||||
app_model = MagicMock()
|
||||
app_model.id = "app-123"
|
||||
session.scalar.return_value = app_model
|
||||
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
|
||||
sqlite_session.commit()
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
|
||||
workflow = MagicMock()
|
||||
workflow.id = "workflow-123"
|
||||
@ -208,43 +246,43 @@ class TestAsyncWorkflowService:
|
||||
user = SimpleNamespace(id="end-user-123")
|
||||
|
||||
# Act
|
||||
AsyncWorkflowService.trigger_workflow_async(session=session, user=user, trigger_data=trigger_data)
|
||||
response = AsyncWorkflowService.trigger_workflow_async(
|
||||
session=sqlite_session, user=user, trigger_data=trigger_data
|
||||
)
|
||||
|
||||
# Assert
|
||||
created_log = mocks["repo"].create.call_args[0][0]
|
||||
created_log = sqlite_session.get(WorkflowTriggerLog, response.workflow_trigger_log_id)
|
||||
assert created_log is not None
|
||||
assert created_log.created_by_role == CreatorUserRole.END_USER
|
||||
assert created_log.created_by == "end-user-123"
|
||||
|
||||
def test_should_raise_workflow_not_found_when_app_does_not_exist(self):
|
||||
def test_should_raise_workflow_not_found_when_app_does_not_exist(self, sqlite_session: Session):
|
||||
"""Test trigger failure when app lookup returns no result."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = None
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data(app_id="missing-app")
|
||||
|
||||
with (
|
||||
patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository"),
|
||||
patch.object(async_workflow_service_module, "QueueDispatcherManager"),
|
||||
patch.object(async_workflow_service_module, "WorkflowService"),
|
||||
):
|
||||
# Act / Assert
|
||||
with pytest.raises(WorkflowNotFoundError, match="App not found: missing-app"):
|
||||
AsyncWorkflowService.trigger_workflow_async(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
user=SimpleNamespace(id="user-123"),
|
||||
trigger_data=trigger_data,
|
||||
)
|
||||
|
||||
def test_should_mark_log_rate_limited_and_reraise_when_quota_exceeded(
|
||||
self, async_workflow_trigger_mocks, caplog: pytest.LogCaptureFixture
|
||||
self,
|
||||
async_workflow_trigger_mocks,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
"""Test quota-exceeded path updates trigger log and preserves the quota exception."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
session.commit = MagicMock()
|
||||
app_model = MagicMock()
|
||||
app_model.id = "app-123"
|
||||
session.scalar.return_value = app_model
|
||||
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_app())
|
||||
sqlite_session.commit()
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
|
||||
workflow = MagicMock()
|
||||
workflow.id = "workflow-123"
|
||||
@ -262,7 +300,7 @@ class TestAsyncWorkflowService:
|
||||
# Act / Assert
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
AsyncWorkflowService.trigger_workflow_async(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
user=SimpleNamespace(id="user-123"),
|
||||
trigger_data=trigger_data,
|
||||
)
|
||||
@ -270,42 +308,37 @@ class TestAsyncWorkflowService:
|
||||
assert exc_info.value.feature == "workflow"
|
||||
assert exc_info.value.tenant_id == "tenant-123"
|
||||
assert exc_info.value.required == 1
|
||||
assert session.commit.call_count == 3
|
||||
updated_log = mocks["repo"].update.call_args[0][0]
|
||||
assert not sqlite_session.in_transaction()
|
||||
updated_log = sqlite_session.scalar(select(WorkflowTriggerLog))
|
||||
assert updated_log is not None
|
||||
assert updated_log.status == WorkflowTriggerStatus.RATE_LIMITED
|
||||
assert "Quota limit reached" in updated_log.error
|
||||
assert (
|
||||
"Workflow quota exceeded for tenant tenant-123, app app-123, workflow workflow-123, "
|
||||
"trigger log trigger-log-123"
|
||||
f"trigger log {updated_log.id}"
|
||||
) in caplog.messages
|
||||
mocks["professional_task"].delay.assert_not_called()
|
||||
mocks["team_task"].delay.assert_not_called()
|
||||
mocks["sandbox_task"].delay.assert_not_called()
|
||||
|
||||
def test_should_raise_when_reinvoke_target_log_does_not_exist(self):
|
||||
def test_should_raise_when_reinvoke_target_log_does_not_exist(self, sqlite_session: Session):
|
||||
"""Test reinvoke_trigger error path when original trigger log is missing."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
repo = MagicMock()
|
||||
repo.get_by_id.return_value = None
|
||||
# Act / Assert
|
||||
with pytest.raises(ValueError, match="Trigger log not found: missing-log"):
|
||||
AsyncWorkflowService.reinvoke_trigger(
|
||||
session=sqlite_session,
|
||||
user=SimpleNamespace(id="user-123"),
|
||||
workflow_trigger_log_id="missing-log",
|
||||
)
|
||||
|
||||
with patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository", return_value=repo):
|
||||
# Act / Assert
|
||||
with pytest.raises(ValueError, match="Trigger log not found: missing-log"):
|
||||
AsyncWorkflowService.reinvoke_trigger(
|
||||
session=session,
|
||||
user=SimpleNamespace(id="user-123"),
|
||||
workflow_trigger_log_id="missing-log",
|
||||
)
|
||||
|
||||
def test_should_update_original_log_and_requeue_when_reinvoking(self):
|
||||
def test_should_update_original_log_and_requeue_when_reinvoking(self, sqlite_session: Session):
|
||||
"""Test reinvoke flow updates original log state and triggers a new async run."""
|
||||
# Arrange
|
||||
session = MagicMock()
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
|
||||
trigger_log = AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(trigger_data, retry_count=1)
|
||||
repo = MagicMock()
|
||||
repo.get_by_id.return_value = trigger_log
|
||||
sqlite_session.add(trigger_log)
|
||||
sqlite_session.commit()
|
||||
|
||||
expected_response = AsyncTriggerResponse(
|
||||
workflow_trigger_log_id="new-trigger-log-456",
|
||||
@ -315,7 +348,6 @@ class TestAsyncWorkflowService:
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(async_workflow_service_module, "SQLAlchemyWorkflowTriggerLogRepository", return_value=repo),
|
||||
patch.object(
|
||||
async_workflow_service_module.AsyncWorkflowService,
|
||||
"trigger_workflow_async",
|
||||
@ -326,145 +358,142 @@ class TestAsyncWorkflowService:
|
||||
|
||||
# Act
|
||||
response = AsyncWorkflowService.reinvoke_trigger(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
user=user,
|
||||
workflow_trigger_log_id="trigger-log-123",
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert response == expected_response
|
||||
assert not sqlite_session.in_transaction()
|
||||
sqlite_session.refresh(trigger_log)
|
||||
assert trigger_log.status == WorkflowTriggerStatus.RETRYING
|
||||
assert trigger_log.retry_count == 2
|
||||
assert trigger_log.error is None
|
||||
assert trigger_log.triggered_at is not None
|
||||
repo.update.assert_called_once_with(trigger_log)
|
||||
session.commit.assert_called_once()
|
||||
called_trigger_data = mock_trigger_workflow_async.call_args.args[1]
|
||||
assert isinstance(called_trigger_data, TriggerData)
|
||||
assert called_trigger_data.app_id == "app-123"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("repo_result", "expected"),
|
||||
("lookup_id", "tenant_id", "expected_id"),
|
||||
[
|
||||
(None, None),
|
||||
(MagicMock(), {"id": "trigger-log-123"}),
|
||||
("missing-log", "tenant-123", None),
|
||||
("trigger-log-123", "tenant-123", "trigger-log-123"),
|
||||
("trigger-log-123", "other-tenant", None),
|
||||
],
|
||||
)
|
||||
def test_should_return_trigger_log_dict_or_none(self, repo_result, expected):
|
||||
"""Test get_trigger_log returns serialized log data or None."""
|
||||
def test_should_return_trigger_log_dict_or_none(
|
||||
self,
|
||||
lookup_id: str,
|
||||
tenant_id: str,
|
||||
expected_id: str | None,
|
||||
sqlite_session: Session,
|
||||
sqlite_engine: Engine,
|
||||
):
|
||||
"""Test get_trigger_log returns persisted data with tenant isolation."""
|
||||
# Arrange
|
||||
mock_session = MagicMock()
|
||||
mock_repo = MagicMock()
|
||||
fake_engine = MagicMock()
|
||||
mock_repo.get_by_id.return_value = repo_result
|
||||
if repo_result:
|
||||
repo_result.to_dict.return_value = expected
|
||||
trigger_data = AsyncWorkflowServiceTestDataFactory.create_trigger_data()
|
||||
sqlite_session.add(AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(trigger_data))
|
||||
sqlite_session.commit()
|
||||
|
||||
mock_session_context = MagicMock()
|
||||
mock_session_context.__enter__.return_value = mock_session
|
||||
mock_session_context.__exit__.return_value = None
|
||||
|
||||
mock_sessionmaker = MagicMock()
|
||||
mock_sessionmaker.return_value.begin.return_value = mock_session_context
|
||||
|
||||
with (
|
||||
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=fake_engine)),
|
||||
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
|
||||
patch.object(
|
||||
async_workflow_service_module,
|
||||
"SQLAlchemyWorkflowTriggerLogRepository",
|
||||
return_value=mock_repo,
|
||||
),
|
||||
):
|
||||
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
|
||||
# Act
|
||||
result = AsyncWorkflowService.get_trigger_log("trigger-log-123", tenant_id="tenant-123")
|
||||
result = AsyncWorkflowService.get_trigger_log(lookup_id, tenant_id=tenant_id)
|
||||
|
||||
# Assert
|
||||
assert result == expected
|
||||
mock_sessionmaker.assert_called_once_with(fake_engine)
|
||||
mock_repo.get_by_id.assert_called_once_with("trigger-log-123", "tenant-123")
|
||||
assert (result["id"] if result else None) == expected_id
|
||||
|
||||
def test_should_return_recent_logs_as_dict_list(self):
|
||||
"""Test get_recent_logs converts repository models into dictionaries."""
|
||||
def test_should_return_recent_logs_as_dict_list(self, sqlite_session: Session, sqlite_engine: Engine):
|
||||
"""Test recent logs are ordered, paginated, and tenant/app scoped."""
|
||||
# Arrange
|
||||
mock_session = MagicMock()
|
||||
mock_repo = MagicMock()
|
||||
log1 = MagicMock()
|
||||
log1.to_dict.return_value = {"id": "log-1"}
|
||||
log2 = MagicMock()
|
||||
log2.to_dict.return_value = {"id": "log-2"}
|
||||
mock_repo.get_recent_logs.return_value = [log1, log2]
|
||||
now = datetime.now(UTC)
|
||||
logs = [
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
|
||||
trigger_log_id=f"log-{index}",
|
||||
created_at=now - timedelta(minutes=index),
|
||||
)
|
||||
for index in range(1, 4)
|
||||
]
|
||||
logs.extend(
|
||||
[
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(tenant_id="other-tenant"),
|
||||
trigger_log_id="other-tenant-log",
|
||||
created_at=now,
|
||||
),
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(app_id="other-app"),
|
||||
trigger_log_id="other-app-log",
|
||||
created_at=now,
|
||||
),
|
||||
]
|
||||
)
|
||||
sqlite_session.add_all(logs)
|
||||
sqlite_session.commit()
|
||||
|
||||
mock_session_context = MagicMock()
|
||||
mock_session_context.__enter__.return_value = mock_session
|
||||
mock_session_context.__exit__.return_value = None
|
||||
|
||||
mock_sessionmaker = MagicMock()
|
||||
mock_sessionmaker.return_value.begin.return_value = mock_session_context
|
||||
|
||||
with (
|
||||
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=MagicMock())),
|
||||
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
|
||||
patch.object(
|
||||
async_workflow_service_module,
|
||||
"SQLAlchemyWorkflowTriggerLogRepository",
|
||||
return_value=mock_repo,
|
||||
),
|
||||
):
|
||||
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
|
||||
# Act
|
||||
result = AsyncWorkflowService.get_recent_logs(
|
||||
tenant_id="tenant-123",
|
||||
app_id="app-123",
|
||||
hours=12,
|
||||
limit=50,
|
||||
offset=10,
|
||||
limit=2,
|
||||
offset=1,
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result == [{"id": "log-1"}, {"id": "log-2"}]
|
||||
mock_repo.get_recent_logs.assert_called_once_with(
|
||||
tenant_id="tenant-123",
|
||||
app_id="app-123",
|
||||
hours=12,
|
||||
limit=50,
|
||||
offset=10,
|
||||
)
|
||||
assert [log["id"] for log in result] == ["log-2", "log-3"]
|
||||
|
||||
def test_should_return_failed_logs_for_retry_as_dict_list(self):
|
||||
"""Test get_failed_logs_for_retry serializes repository logs into dicts."""
|
||||
def test_should_return_failed_logs_for_retry_as_dict_list(self, sqlite_session: Session, sqlite_engine: Engine):
|
||||
"""Test retry candidates are status, retry-count, and tenant scoped."""
|
||||
# Arrange
|
||||
mock_session = MagicMock()
|
||||
mock_repo = MagicMock()
|
||||
log = MagicMock()
|
||||
log.to_dict.return_value = {"id": "failed-log-1"}
|
||||
mock_repo.get_failed_for_retry.return_value = [log]
|
||||
|
||||
mock_session_context = MagicMock()
|
||||
mock_session_context.__enter__.return_value = mock_session
|
||||
mock_session_context.__exit__.return_value = None
|
||||
|
||||
mock_sessionmaker = MagicMock()
|
||||
mock_sessionmaker.return_value.begin.return_value = mock_session_context
|
||||
|
||||
with (
|
||||
patch.object(async_workflow_service_module, "db", new=SimpleNamespace(engine=MagicMock())),
|
||||
patch.object(async_workflow_service_module, "sessionmaker", mock_sessionmaker),
|
||||
patch.object(
|
||||
async_workflow_service_module,
|
||||
"SQLAlchemyWorkflowTriggerLogRepository",
|
||||
return_value=mock_repo,
|
||||
now = datetime.now(UTC)
|
||||
candidates = [
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
|
||||
trigger_log_id="failed-log-1",
|
||||
retry_count=1,
|
||||
created_at=now - timedelta(minutes=2),
|
||||
),
|
||||
):
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
|
||||
trigger_log_id="rate-limited-log",
|
||||
retry_count=2,
|
||||
status=WorkflowTriggerStatus.RATE_LIMITED,
|
||||
created_at=now - timedelta(minutes=1),
|
||||
),
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
|
||||
trigger_log_id="retry-limit-log",
|
||||
retry_count=4,
|
||||
),
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(tenant_id="other-tenant"),
|
||||
trigger_log_id="other-tenant-log",
|
||||
),
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_log_with_data(
|
||||
AsyncWorkflowServiceTestDataFactory.create_trigger_data(),
|
||||
trigger_log_id="queued-log",
|
||||
status=WorkflowTriggerStatus.QUEUED,
|
||||
),
|
||||
]
|
||||
sqlite_session.add_all(candidates)
|
||||
sqlite_session.commit()
|
||||
|
||||
with patch.object(async_workflow_service_module, "db", SimpleNamespace(engine=sqlite_engine)):
|
||||
# Act
|
||||
result = AsyncWorkflowService.get_failed_logs_for_retry(tenant_id="tenant-123", max_retry_count=4, limit=20)
|
||||
|
||||
# Assert
|
||||
assert result == [{"id": "failed-log-1"}]
|
||||
mock_repo.get_failed_for_retry.assert_called_once_with(tenant_id="tenant-123", max_retry_count=4, limit=20)
|
||||
assert [log["id"] for log in result] == ["failed-log-1", "rate-limited-log"]
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("sqlite_session")
|
||||
@pytest.mark.parametrize("sqlite_session", [(App, WorkflowTriggerLog)], indirect=True)
|
||||
class TestAsyncWorkflowServiceGetWorkflow:
|
||||
def test_should_return_specific_workflow_when_workflow_id_exists(self):
|
||||
def test_should_return_specific_workflow_when_workflow_id_exists(self, sqlite_session: Session):
|
||||
"""Test _get_workflow returns published workflow by id when provided."""
|
||||
# Arrange
|
||||
workflow_service = MagicMock()
|
||||
@ -473,19 +502,18 @@ class TestAsyncWorkflowServiceGetWorkflow:
|
||||
workflow_service.get_published_workflow_by_id.return_value = workflow
|
||||
|
||||
# Act
|
||||
session = MagicMock()
|
||||
result = AsyncWorkflowService._get_workflow(
|
||||
workflow_service, app_model, workflow_id="workflow-123", session=session
|
||||
workflow_service, app_model, workflow_id="workflow-123", session=sqlite_session
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result == workflow
|
||||
workflow_service.get_published_workflow_by_id.assert_called_once_with(
|
||||
app_model, "workflow-123", session=session
|
||||
app_model, "workflow-123", session=sqlite_session
|
||||
)
|
||||
workflow_service.get_published_workflow.assert_not_called()
|
||||
|
||||
def test_should_raise_when_specific_workflow_id_not_found(self):
|
||||
def test_should_raise_when_specific_workflow_id_not_found(self, sqlite_session: Session):
|
||||
"""Test _get_workflow raises WorkflowNotFoundError for unknown workflow id."""
|
||||
# Arrange
|
||||
workflow_service = MagicMock()
|
||||
@ -495,10 +523,10 @@ class TestAsyncWorkflowServiceGetWorkflow:
|
||||
# Act / Assert
|
||||
with pytest.raises(WorkflowNotFoundError, match="Published workflow not found: workflow-404"):
|
||||
AsyncWorkflowService._get_workflow(
|
||||
workflow_service, app_model, workflow_id="workflow-404", session=MagicMock()
|
||||
workflow_service, app_model, workflow_id="workflow-404", session=sqlite_session
|
||||
)
|
||||
|
||||
def test_should_return_default_published_workflow_when_workflow_id_not_provided(self):
|
||||
def test_should_return_default_published_workflow_when_workflow_id_not_provided(self, sqlite_session: Session):
|
||||
"""Test _get_workflow returns default published workflow when no id is provided."""
|
||||
# Arrange
|
||||
workflow_service = MagicMock()
|
||||
@ -508,15 +536,14 @@ class TestAsyncWorkflowServiceGetWorkflow:
|
||||
workflow_service.get_published_workflow.return_value = workflow
|
||||
|
||||
# Act
|
||||
session = MagicMock()
|
||||
result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=session)
|
||||
result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
assert result == workflow
|
||||
workflow_service.get_published_workflow.assert_called_once_with(app_model, session=session)
|
||||
workflow_service.get_published_workflow.assert_called_once_with(app_model, session=sqlite_session)
|
||||
workflow_service.get_published_workflow_by_id.assert_not_called()
|
||||
|
||||
def test_should_raise_when_default_published_workflow_not_found(self):
|
||||
def test_should_raise_when_default_published_workflow_not_found(self, sqlite_session: Session):
|
||||
"""Test _get_workflow raises WorkflowNotFoundError when app has no published workflow."""
|
||||
# Arrange
|
||||
workflow_service = MagicMock()
|
||||
@ -526,4 +553,4 @@ class TestAsyncWorkflowServiceGetWorkflow:
|
||||
|
||||
# Act / Assert
|
||||
with pytest.raises(WorkflowNotFoundError, match="No published workflow found for app: app-123"):
|
||||
AsyncWorkflowService._get_workflow(workflow_service, app_model, session=MagicMock())
|
||||
AsyncWorkflowService._get_workflow(workflow_service, app_model, session=sqlite_session)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user