mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
test: use real ORM models in workflow repository tests (#40019)
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -12,7 +13,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerat
|
||||
from core.workflow.system_variables import default_system_variables
|
||||
from graphon.entities.graph_config import NodeConfigDictAdapter
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from models.workflow import Workflow
|
||||
from models.workflow import Workflow, WorkflowKind
|
||||
|
||||
|
||||
def _make_graph_state():
|
||||
@@ -55,13 +56,14 @@ def test_run_uses_single_node_execution_branch(
|
||||
app_generate_entity.single_iteration_run = single_iteration_run
|
||||
app_generate_entity.single_loop_run = single_loop_run
|
||||
|
||||
workflow = MagicMock(spec=Workflow)
|
||||
workflow.tenant_id = "tenant"
|
||||
workflow.app_id = "app"
|
||||
workflow.id = "workflow"
|
||||
workflow.type = "workflow"
|
||||
workflow.version = "v1"
|
||||
workflow.graph_dict = {"nodes": [], "edges": []}
|
||||
workflow = Workflow(
|
||||
tenant_id="tenant",
|
||||
app_id="app",
|
||||
id="workflow",
|
||||
type="workflow",
|
||||
version="v1",
|
||||
graph=json.dumps({"nodes": [], "edges": []}),
|
||||
)
|
||||
workflow.environment_variables = []
|
||||
|
||||
runner = WorkflowAppRunner(
|
||||
@@ -119,25 +121,28 @@ def test_single_node_run_validates_target_node_config(monkeypatch: pytest.Monkey
|
||||
app_id="app",
|
||||
)
|
||||
|
||||
workflow = MagicMock(spec=Workflow)
|
||||
workflow.id = "workflow"
|
||||
workflow.tenant_id = "tenant"
|
||||
workflow.graph_dict = {
|
||||
"nodes": [
|
||||
workflow = Workflow(
|
||||
id="workflow",
|
||||
tenant_id="tenant",
|
||||
graph=json.dumps(
|
||||
{
|
||||
"id": "loop-node",
|
||||
"data": {
|
||||
"type": "loop",
|
||||
"title": "Loop",
|
||||
"loop_count": 1,
|
||||
"start_node_id": "loop-start",
|
||||
"break_conditions": [],
|
||||
"logical_operator": "and",
|
||||
},
|
||||
"nodes": [
|
||||
{
|
||||
"id": "loop-node",
|
||||
"data": {
|
||||
"type": "loop",
|
||||
"title": "Loop",
|
||||
"loop_count": 1,
|
||||
"start_node_id": "loop-start",
|
||||
"break_conditions": [],
|
||||
"logical_operator": "and",
|
||||
},
|
||||
}
|
||||
],
|
||||
"edges": [],
|
||||
}
|
||||
],
|
||||
"edges": [],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
_, _, graph_runtime_state = _make_graph_state()
|
||||
seen_configs: list[object] = []
|
||||
@@ -188,15 +193,16 @@ def test_run_adds_inputs_with_snippet_compatible_start_aliases() -> None:
|
||||
app_generate_entity.single_iteration_run = None
|
||||
app_generate_entity.single_loop_run = None
|
||||
|
||||
workflow = MagicMock(spec=Workflow)
|
||||
workflow.tenant_id = "tenant"
|
||||
workflow.app_id = "app"
|
||||
workflow.id = "workflow"
|
||||
workflow.type = "workflow"
|
||||
workflow.version = "v1"
|
||||
workflow.graph_dict = {"nodes": [], "edges": []}
|
||||
workflow = Workflow(
|
||||
tenant_id="tenant",
|
||||
app_id="app",
|
||||
id="workflow",
|
||||
type="workflow",
|
||||
version="v1",
|
||||
graph=json.dumps({"nodes": [], "edges": []}),
|
||||
kind=WorkflowKind.SNIPPET,
|
||||
)
|
||||
workflow.environment_variables = []
|
||||
workflow.kind_or_standard = "snippet"
|
||||
|
||||
runner = WorkflowAppRunner(
|
||||
application_generate_entity=app_generate_entity,
|
||||
|
||||
@@ -162,10 +162,8 @@ def _build_converter(*, invoke_from: InvokeFrom = InvokeFrom.SERVICE_API):
|
||||
workflow_id="workflow-id",
|
||||
workflow_execution_id="run-id",
|
||||
)
|
||||
user = MagicMock(spec=Account)
|
||||
user = Account(name="Tester", email="tester@example.com")
|
||||
user.id = "account-id"
|
||||
user.name = "Tester"
|
||||
user.email = "tester@example.com"
|
||||
return WorkflowResponseConverter(
|
||||
application_generate_entity=application_generate_entity,
|
||||
user=user,
|
||||
|
||||
@@ -29,15 +29,22 @@ class TestHandleMCPRequest:
|
||||
|
||||
def setup_method(self):
|
||||
"""Setup test fixtures"""
|
||||
self.app = Mock(spec=App)
|
||||
self.app = App()
|
||||
self.app.name = "test_app"
|
||||
self.app.mode = AppMode.CHAT
|
||||
|
||||
self.mcp_server = Mock(spec=AppMCPServer)
|
||||
self.mcp_server = AppMCPServer(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
name="Test Server",
|
||||
description="",
|
||||
server_code="test-server",
|
||||
status="active",
|
||||
parameters="{}",
|
||||
)
|
||||
self.mcp_server.description = "Test server"
|
||||
self.mcp_server.parameters_dict = {}
|
||||
|
||||
self.end_user = Mock(spec=EndUser)
|
||||
self.end_user = EndUser()
|
||||
self.user_input_form = []
|
||||
|
||||
# Create mock request
|
||||
@@ -336,8 +343,9 @@ class TestIndividualHandlers:
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool(self, mock_app_generate):
|
||||
"""Test call tool handler"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
# Create mock request
|
||||
mock_request = Mock()
|
||||
@@ -347,7 +355,7 @@ class TestIndividualHandlers:
|
||||
mock_request.root = mock_call_request
|
||||
|
||||
user_input_form: list[VariableEntity] = []
|
||||
end_user = Mock(spec=EndUser)
|
||||
end_user = EndUser()
|
||||
|
||||
# Mock app generate service response
|
||||
mock_response = {"answer": "test answer"}
|
||||
@@ -365,8 +373,9 @@ class TestIndividualHandlers:
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_structured_output_modern_client(self, mock_app_generate):
|
||||
"""structuredContent is attached alongside TextContent for >= 2025-06-18."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
mock_request = Mock()
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
@@ -376,7 +385,7 @@ class TestIndividualHandlers:
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2025-06-18")
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], EndUser(), "2025-06-18")
|
||||
|
||||
assert result.structuredContent == {"answer": "test answer"}
|
||||
assert result.content[0].text == "test answer"
|
||||
@@ -384,8 +393,9 @@ class TestIndividualHandlers:
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_no_structured_output_legacy_client(self, mock_app_generate):
|
||||
"""structuredContent is omitted for 2024-11-05 clients."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
mock_request = Mock()
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
@@ -395,14 +405,14 @@ class TestIndividualHandlers:
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2024-11-05")
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], EndUser(), "2024-11-05")
|
||||
|
||||
assert result.structuredContent is None
|
||||
assert result.content[0].text == "test answer"
|
||||
|
||||
def test_handle_call_tool_no_end_user(self):
|
||||
"""Test call tool handler without end user"""
|
||||
app = Mock(spec=App)
|
||||
app = App()
|
||||
mock_request = Mock()
|
||||
user_input_form: list[VariableEntity] = []
|
||||
|
||||
@@ -460,8 +470,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_prepare_tool_arguments_chat_mode(self):
|
||||
"""Test preparing tool arguments for chat mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
arguments = {"query": "test question", "name": "John"}
|
||||
|
||||
@@ -474,8 +485,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_prepare_tool_arguments_workflow_mode(self):
|
||||
"""Test preparing tool arguments for workflow mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
arguments = {"input_text": "test input"}
|
||||
|
||||
@@ -486,8 +498,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_prepare_tool_arguments_completion_mode(self):
|
||||
"""Test preparing tool arguments for completion mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.COMPLETION
|
||||
app = App(
|
||||
mode=AppMode.COMPLETION,
|
||||
)
|
||||
|
||||
arguments = {"name": "John"}
|
||||
|
||||
@@ -498,8 +511,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_extract_answer_from_mapping_response_chat(self):
|
||||
"""Test extracting answer from mapping response for chat mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
response = {"answer": "test answer", "other": "data"}
|
||||
|
||||
@@ -509,8 +523,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_extract_answer_from_mapping_response_workflow(self):
|
||||
"""Test extracting answer from mapping response for workflow mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
response = {"data": {"outputs": {"result": "test result"}}}
|
||||
|
||||
@@ -521,7 +536,7 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_extract_answer_from_streaming_response(self):
|
||||
"""Test extracting answer from streaming response"""
|
||||
app = Mock(spec=App)
|
||||
app = App()
|
||||
|
||||
# Mock RateLimitGenerator
|
||||
mock_generator = Mock(spec=RateLimitGenerator)
|
||||
@@ -538,8 +553,9 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_extract_structured_output_workflow(self):
|
||||
"""Workflow mode exposes the raw outputs mapping as structured content."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
response = {"data": {"outputs": {"result": "test result"}}}
|
||||
|
||||
@@ -547,58 +563,66 @@ class TestUtilityFunctions:
|
||||
|
||||
def test_extract_structured_output_chat(self):
|
||||
"""Chat mode wraps the answer string under an 'answer' key."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"}
|
||||
|
||||
def test_extract_structured_output_workflow_missing_outputs(self):
|
||||
"""Missing or malformed outputs fall back to None."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"data": {}}, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_response(self):
|
||||
"""A non-mapping workflow response yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, None, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_data(self):
|
||||
"""A non-mapping 'data' entry yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"data": "not a mapping"}, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_outputs(self):
|
||||
"""A non-mapping 'outputs' entry yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
app = App(
|
||||
mode=AppMode.WORKFLOW,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"data": {"outputs": ["not", "a", "mapping"]}}, "ignored") is None
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.AGENT_CHAT, AppMode.COMPLETION])
|
||||
def test_extract_structured_output_other_answer_modes(self, mode):
|
||||
"""Every chat-style mode wraps the answer string under an 'answer' key."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = mode
|
||||
app = App(
|
||||
mode=mode,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"}
|
||||
|
||||
def test_extract_structured_output_unknown_mode(self):
|
||||
"""Modes outside the MCP surface produce no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHANNEL
|
||||
app = App(
|
||||
mode=AppMode.CHANNEL,
|
||||
)
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") is None
|
||||
|
||||
def test_process_mapping_response_invalid_mode(self):
|
||||
"""Test processing mapping response with invalid app mode"""
|
||||
app = Mock(spec=App)
|
||||
app.mode = "invalid_mode"
|
||||
app = App(
|
||||
mode="invalid_mode",
|
||||
)
|
||||
|
||||
response = {"answer": "test"}
|
||||
|
||||
|
||||
+12
-12
@@ -5,7 +5,7 @@ These tests verify the Celery-based asynchronous storage functionality
|
||||
for workflow execution data.
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
@@ -14,7 +14,7 @@ from core.repositories.celery_workflow_execution_repository import CeleryWorkflo
|
||||
from graphon.entities import WorkflowExecution
|
||||
from graphon.enums import WorkflowType
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models import Account, EndUser
|
||||
from models import Account, EndUser, Tenant
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
|
||||
RESOURCE_TENANT_ID = "resource-tenant-id"
|
||||
@@ -34,18 +34,20 @@ def mock_session_factory():
|
||||
@pytest.fixture
|
||||
def mock_account():
|
||||
"""Mock Account user."""
|
||||
account = Mock(spec=Account)
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
account.id = str(uuid4())
|
||||
account.current_tenant_id = str(uuid4())
|
||||
account._current_tenant = Tenant(name="Test Tenant")
|
||||
account._current_tenant.id = str(uuid4())
|
||||
return account
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_end_user():
|
||||
"""Mock EndUser."""
|
||||
user = Mock(spec=EndUser)
|
||||
user.id = str(uuid4())
|
||||
user.tenant_id = str(uuid4())
|
||||
user = EndUser(
|
||||
id=str(uuid4()),
|
||||
tenant_id=str(uuid4()),
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
@@ -114,9 +116,8 @@ class TestCeleryWorkflowExecutionRepository:
|
||||
|
||||
def test_init_without_tenant_id_raises_error(self, mock_session_factory):
|
||||
"""Test that initialization fails without tenant_id."""
|
||||
# Create a mock Account with no tenant_id
|
||||
user = Mock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
# Create an Account with no tenant_id.
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = str(uuid4())
|
||||
|
||||
with pytest.raises(ValueError, match="tenant_id is required"):
|
||||
@@ -129,8 +130,7 @@ class TestCeleryWorkflowExecutionRepository:
|
||||
)
|
||||
|
||||
def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory):
|
||||
user = Mock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = str(uuid4())
|
||||
|
||||
repo = CeleryWorkflowExecutionRepository(
|
||||
|
||||
+11
-11
@@ -18,7 +18,7 @@ from graphon.entities.workflow_node_execution import (
|
||||
)
|
||||
from graphon.enums import BuiltinNodeTypes
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models import Account, EndUser
|
||||
from models import Account, EndUser, Tenant
|
||||
from models.workflow import WorkflowNodeExecutionTriggeredFrom
|
||||
|
||||
RESOURCE_TENANT_ID = "resource-tenant-id"
|
||||
@@ -38,18 +38,20 @@ def mock_session_factory():
|
||||
@pytest.fixture
|
||||
def mock_account():
|
||||
"""Mock Account user."""
|
||||
account = Mock(spec=Account)
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
account.id = str(uuid4())
|
||||
account.current_tenant_id = str(uuid4())
|
||||
account._current_tenant = Tenant(name="Test Tenant")
|
||||
account._current_tenant.id = str(uuid4())
|
||||
return account
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_end_user():
|
||||
"""Mock EndUser."""
|
||||
user = Mock(spec=EndUser)
|
||||
user.id = str(uuid4())
|
||||
user.tenant_id = str(uuid4())
|
||||
user = EndUser(
|
||||
id=str(uuid4()),
|
||||
tenant_id=str(uuid4()),
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
@@ -120,9 +122,8 @@ class TestCeleryWorkflowNodeExecutionRepository:
|
||||
|
||||
def test_init_without_tenant_id_raises_error(self, mock_session_factory):
|
||||
"""Test that initialization fails without tenant_id."""
|
||||
# Create a mock Account with no tenant_id
|
||||
user = Mock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
# Create an Account with no tenant_id.
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = str(uuid4())
|
||||
|
||||
with pytest.raises(ValueError, match="tenant_id is required"):
|
||||
@@ -135,8 +136,7 @@ class TestCeleryWorkflowNodeExecutionRepository:
|
||||
)
|
||||
|
||||
def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory):
|
||||
user = Mock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = str(uuid4())
|
||||
|
||||
repo = CeleryWorkflowNodeExecutionRepository(
|
||||
|
||||
@@ -69,7 +69,7 @@ class TestRepositoryFactory:
|
||||
mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock"
|
||||
|
||||
# Create non-database dependencies
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
app_id = "test-app-id"
|
||||
triggered_from = WorkflowRunTriggeredFrom.APP_RUN
|
||||
|
||||
@@ -104,7 +104,7 @@ class TestRepositoryFactory:
|
||||
# Setup mock configuration with invalid class path
|
||||
mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "invalid.module.InvalidClass"
|
||||
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
|
||||
with pytest.raises(RepositoryImportError) as exc_info:
|
||||
DifyCoreRepositoryFactory.create_workflow_execution_repository(
|
||||
@@ -122,7 +122,7 @@ class TestRepositoryFactory:
|
||||
# Setup mock configuration
|
||||
mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock"
|
||||
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
|
||||
# Create a mock repository class that raises exception on instantiation
|
||||
mock_repository_class = MagicMock()
|
||||
@@ -147,7 +147,7 @@ class TestRepositoryFactory:
|
||||
mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "unittest.mock.MagicMock"
|
||||
|
||||
# Create non-database dependencies
|
||||
mock_user = MagicMock(spec=EndUser)
|
||||
mock_user = EndUser()
|
||||
app_id = "test-app-id"
|
||||
triggered_from = WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP
|
||||
|
||||
@@ -182,7 +182,7 @@ class TestRepositoryFactory:
|
||||
# Setup mock configuration with invalid class path
|
||||
mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "invalid.module.InvalidClass"
|
||||
|
||||
mock_user = MagicMock(spec=EndUser)
|
||||
mock_user = EndUser()
|
||||
|
||||
with pytest.raises(RepositoryImportError) as exc_info:
|
||||
DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
|
||||
@@ -200,7 +200,7 @@ class TestRepositoryFactory:
|
||||
# Setup mock configuration
|
||||
mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "unittest.mock.MagicMock"
|
||||
|
||||
mock_user = MagicMock(spec=EndUser)
|
||||
mock_user = EndUser()
|
||||
|
||||
# Create a mock repository class that raises exception on instantiation
|
||||
mock_repository_class = MagicMock()
|
||||
@@ -231,7 +231,7 @@ class TestRepositoryFactory:
|
||||
mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock"
|
||||
|
||||
# Pass the real Engine directly instead of wrapping it in sessionmaker
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
app_id = "test-app-id"
|
||||
triggered_from = WorkflowRunTriggeredFrom.APP_RUN
|
||||
|
||||
|
||||
+58
-52
@@ -29,7 +29,7 @@ from graphon.enums import (
|
||||
WorkflowNodeExecutionMetadataKey,
|
||||
WorkflowNodeExecutionStatus,
|
||||
)
|
||||
from models import Account, EndUser
|
||||
from models import Account, EndUser, Tenant
|
||||
from models.enums import ExecutionOffLoadType
|
||||
from models.workflow import WorkflowNodeExecutionModel, WorkflowNodeExecutionOffload, WorkflowNodeExecutionTriggeredFrom
|
||||
|
||||
@@ -37,16 +37,18 @@ RESOURCE_TENANT_ID = "tenant"
|
||||
|
||||
|
||||
def _mock_account(*, tenant_id: str = "tenant", user_id: str = "user") -> Account:
|
||||
user = Mock(spec=Account)
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = user_id
|
||||
user.current_tenant_id = tenant_id
|
||||
user._current_tenant = Tenant(name="Test Tenant")
|
||||
user._current_tenant.id = tenant_id
|
||||
return user
|
||||
|
||||
|
||||
def _mock_end_user(*, tenant_id: str = "tenant", user_id: str = "user") -> EndUser:
|
||||
user = Mock(spec=EndUser)
|
||||
user.id = user_id
|
||||
user.tenant_id = tenant_id
|
||||
user = EndUser(
|
||||
id=user_id,
|
||||
tenant_id=tenant_id,
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
@@ -148,7 +150,7 @@ def test_init_requires_tenant_id(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
lambda *_: SimpleNamespace(upload_file=Mock()),
|
||||
)
|
||||
user = _mock_account()
|
||||
user.current_tenant_id = None
|
||||
user._current_tenant = None
|
||||
with pytest.raises(ValueError, match="tenant_id is required"):
|
||||
SQLAlchemyWorkflowNodeExecutionRepository(
|
||||
session_factory=Mock(spec=sessionmaker),
|
||||
@@ -165,7 +167,7 @@ def test_init_uses_resource_tenant_when_account_has_no_current_tenant(monkeypatc
|
||||
lambda *_: SimpleNamespace(upload_file=Mock()),
|
||||
)
|
||||
user = _mock_account()
|
||||
user.current_tenant_id = None
|
||||
user._current_tenant = None
|
||||
|
||||
repo = SQLAlchemyWorkflowNodeExecutionRepository(
|
||||
session_factory=Mock(spec=sessionmaker),
|
||||
@@ -305,8 +307,9 @@ def test_is_duplicate_key_error_and_regenerate_id(
|
||||
assert repo._is_duplicate_key_error(IntegrityError("other", params=None, orig=None)) is False
|
||||
|
||||
execution = _execution(execution_id="old-id")
|
||||
db_model = WorkflowNodeExecutionModel()
|
||||
db_model.id = "old-id"
|
||||
db_model = WorkflowNodeExecutionModel(
|
||||
id="old-id",
|
||||
)
|
||||
monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "new-id")
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo._regenerate_id_on_duplicate(execution, db_model)
|
||||
@@ -329,9 +332,10 @@ def test_persist_to_database_updates_existing_and_inserts_new(monkeypatch: pytes
|
||||
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
|
||||
)
|
||||
|
||||
db_model = WorkflowNodeExecutionModel()
|
||||
db_model.id = "id1"
|
||||
db_model.node_execution_id = "node1"
|
||||
db_model = WorkflowNodeExecutionModel(
|
||||
id="id1",
|
||||
node_execution_id="node1",
|
||||
)
|
||||
db_model.foo = "bar" # type: ignore[attr-defined]
|
||||
db_model.__dict__["_private"] = "x"
|
||||
|
||||
@@ -424,25 +428,26 @@ def test_to_domain_model_loads_offloaded_files(monkeypatch: pytest.MonkeyPatch)
|
||||
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
|
||||
)
|
||||
|
||||
db_model = WorkflowNodeExecutionModel()
|
||||
db_model.id = "id"
|
||||
db_model.node_execution_id = "node-exec"
|
||||
db_model.workflow_id = "wf"
|
||||
db_model.workflow_run_id = "run"
|
||||
db_model.index = 1
|
||||
db_model.predecessor_node_id = None
|
||||
db_model.node_id = "node"
|
||||
db_model.node_type = BuiltinNodeTypes.LLM
|
||||
db_model.title = "t"
|
||||
db_model.inputs = json.dumps({"trunc": "i"})
|
||||
db_model.process_data = json.dumps({"trunc": "p"})
|
||||
db_model.outputs = json.dumps({"trunc": "o"})
|
||||
db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED
|
||||
db_model.error = None
|
||||
db_model.elapsed_time = 0.1
|
||||
db_model.execution_metadata = json.dumps({"total_tokens": 3})
|
||||
db_model.created_at = datetime.now(UTC)
|
||||
db_model.finished_at = None
|
||||
db_model = WorkflowNodeExecutionModel(
|
||||
id="id",
|
||||
node_execution_id="node-exec",
|
||||
workflow_id="wf",
|
||||
workflow_run_id="run",
|
||||
index=1,
|
||||
predecessor_node_id=None,
|
||||
node_id="node",
|
||||
node_type=BuiltinNodeTypes.LLM,
|
||||
title="t",
|
||||
inputs=json.dumps({"trunc": "i"}),
|
||||
process_data=json.dumps({"trunc": "p"}),
|
||||
outputs=json.dumps({"trunc": "o"}),
|
||||
status=WorkflowNodeExecutionStatus.SUCCEEDED,
|
||||
error=None,
|
||||
elapsed_time=0.1,
|
||||
execution_metadata=json.dumps({"total_tokens": 3}),
|
||||
created_at=datetime.now(UTC),
|
||||
finished_at=None,
|
||||
)
|
||||
|
||||
off_in = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS)
|
||||
off_out = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.OUTPUTS)
|
||||
@@ -479,26 +484,27 @@ def test_to_domain_model_returns_early_when_no_offload_data(monkeypatch: pytest.
|
||||
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
|
||||
)
|
||||
|
||||
db_model = WorkflowNodeExecutionModel()
|
||||
db_model.id = "id"
|
||||
db_model.node_execution_id = "node-exec"
|
||||
db_model.workflow_id = "wf"
|
||||
db_model.workflow_run_id = "run"
|
||||
db_model.index = 1
|
||||
db_model.predecessor_node_id = None
|
||||
db_model.node_id = "node"
|
||||
db_model.node_type = BuiltinNodeTypes.LLM
|
||||
db_model.title = "t"
|
||||
db_model.inputs = json.dumps({"i": 1})
|
||||
db_model.process_data = json.dumps({"p": 2})
|
||||
db_model.outputs = json.dumps({"o": 3})
|
||||
db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED
|
||||
db_model.error = None
|
||||
db_model.elapsed_time = 0.1
|
||||
db_model.execution_metadata = "{}"
|
||||
db_model.created_at = datetime.now(UTC)
|
||||
db_model.finished_at = None
|
||||
db_model.offload_data = []
|
||||
db_model = WorkflowNodeExecutionModel(
|
||||
id="id",
|
||||
node_execution_id="node-exec",
|
||||
workflow_id="wf",
|
||||
workflow_run_id="run",
|
||||
index=1,
|
||||
predecessor_node_id=None,
|
||||
node_id="node",
|
||||
node_type=BuiltinNodeTypes.LLM,
|
||||
title="t",
|
||||
inputs=json.dumps({"i": 1}),
|
||||
process_data=json.dumps({"p": 2}),
|
||||
outputs=json.dumps({"o": 3}),
|
||||
status=WorkflowNodeExecutionStatus.SUCCEEDED,
|
||||
error=None,
|
||||
elapsed_time=0.1,
|
||||
execution_metadata="{}",
|
||||
created_at=datetime.now(UTC),
|
||||
finished_at=None,
|
||||
offload_data=[],
|
||||
)
|
||||
|
||||
domain = repo._to_domain_model(db_model)
|
||||
assert domain.inputs == {"i": 1}
|
||||
|
||||
@@ -9,7 +9,6 @@ import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -23,7 +22,7 @@ from graphon.entities.workflow_node_execution import (
|
||||
WorkflowNodeExecutionStatus,
|
||||
)
|
||||
from graphon.enums import BuiltinNodeTypes
|
||||
from models import Account, WorkflowNodeExecutionTriggeredFrom
|
||||
from models import Account, Tenant, WorkflowNodeExecutionTriggeredFrom
|
||||
from models.enums import ExecutionOffLoadType
|
||||
from models.workflow import WorkflowNodeExecutionModel, WorkflowNodeExecutionOffload
|
||||
|
||||
@@ -113,11 +112,12 @@ def create_workflow_node_execution(
|
||||
|
||||
|
||||
def mock_user() -> Account:
|
||||
"""Create a mock Account user for testing."""
|
||||
"""Create an Account user for testing."""
|
||||
|
||||
user = MagicMock(spec=Account)
|
||||
user = Account(name="Test Account", email="test@example.com")
|
||||
user.id = "test-user-id"
|
||||
user.current_tenant_id = "test-tenant-id"
|
||||
user._current_tenant = Tenant(name="Test Tenant")
|
||||
user._current_tenant.id = "test-tenant-id"
|
||||
return user
|
||||
|
||||
|
||||
@@ -143,26 +143,27 @@ class TestSQLAlchemyWorkflowNodeExecutionRepositoryTruncation:
|
||||
repo = self.create_repository(sqlite_engine)
|
||||
|
||||
# Create a database model without offload data
|
||||
db_model = WorkflowNodeExecutionModel()
|
||||
db_model.id = "test-id"
|
||||
db_model.node_execution_id = "node-exec-id"
|
||||
db_model.workflow_id = "workflow-id"
|
||||
db_model.workflow_run_id = "run-id"
|
||||
db_model.index = 1
|
||||
db_model.predecessor_node_id = None
|
||||
db_model.node_id = "node-id"
|
||||
db_model.node_type = BuiltinNodeTypes.LLM
|
||||
db_model.title = "Test Node"
|
||||
db_model.inputs = json.dumps({"value": "inputs"})
|
||||
db_model.process_data = json.dumps({"value": "process_data"})
|
||||
db_model.outputs = json.dumps({"value": "outputs"})
|
||||
db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED
|
||||
db_model.error = None
|
||||
db_model.elapsed_time = 1.0
|
||||
db_model.execution_metadata = "{}"
|
||||
db_model.created_at = datetime.now(UTC)
|
||||
db_model.finished_at = None
|
||||
db_model.offload_data = []
|
||||
db_model = WorkflowNodeExecutionModel(
|
||||
id="test-id",
|
||||
node_execution_id="node-exec-id",
|
||||
workflow_id="workflow-id",
|
||||
workflow_run_id="run-id",
|
||||
index=1,
|
||||
predecessor_node_id=None,
|
||||
node_id="node-id",
|
||||
node_type=BuiltinNodeTypes.LLM,
|
||||
title="Test Node",
|
||||
inputs=json.dumps({"value": "inputs"}),
|
||||
process_data=json.dumps({"value": "process_data"}),
|
||||
outputs=json.dumps({"value": "outputs"}),
|
||||
status=WorkflowNodeExecutionStatus.SUCCEEDED,
|
||||
error=None,
|
||||
elapsed_time=1.0,
|
||||
execution_metadata="{}",
|
||||
created_at=datetime.now(UTC),
|
||||
finished_at=None,
|
||||
offload_data=[],
|
||||
)
|
||||
|
||||
domain_model = repo._to_domain_model(db_model)
|
||||
|
||||
@@ -206,8 +207,9 @@ class TestWorkflowNodeExecutionModelTruncatedProperties:
|
||||
|
||||
def test_truncated_properties_without_offload_data(self):
|
||||
"""Test truncated properties when no offload data exists."""
|
||||
model = WorkflowNodeExecutionModel()
|
||||
model.offload_data = []
|
||||
model = WorkflowNodeExecutionModel(
|
||||
offload_data=[],
|
||||
)
|
||||
|
||||
assert model.inputs_truncated is False
|
||||
assert model.outputs_truncated is False
|
||||
|
||||
@@ -1,45 +1,56 @@
|
||||
"""Tests for _PrivateWorkflowPauseEntity implementation."""
|
||||
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
from models.workflow import WorkflowPause as WorkflowPauseModel
|
||||
from repositories.sqlalchemy_api_workflow_run_repository import _PrivateWorkflowPauseEntity
|
||||
|
||||
|
||||
def _make_workflow_pause(
|
||||
*,
|
||||
pause_id: str = "pause-123",
|
||||
workflow_run_id: str = "execution-456",
|
||||
state_object_key: str = "test-state-key",
|
||||
resumed_at: datetime | None = None,
|
||||
) -> WorkflowPauseModel:
|
||||
pause = WorkflowPauseModel(
|
||||
workflow_id="workflow-789",
|
||||
workflow_run_id=workflow_run_id,
|
||||
state_object_key=state_object_key,
|
||||
resumed_at=resumed_at,
|
||||
)
|
||||
pause.id = pause_id
|
||||
return pause
|
||||
|
||||
|
||||
class TestPrivateWorkflowPauseEntity:
|
||||
"""Test _PrivateWorkflowPauseEntity implementation."""
|
||||
|
||||
def test_entity_initialization(self):
|
||||
"""Test entity initialization with required parameters."""
|
||||
# Create mock models
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.id = "pause-123"
|
||||
mock_pause_model.workflow_run_id = "execution-456"
|
||||
mock_pause_model.resumed_at = None
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
# Create entity
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
# Verify initialization
|
||||
assert entity._pause_model is mock_pause_model
|
||||
assert entity._pause_model is pause_model
|
||||
assert entity._cached_state is None
|
||||
|
||||
def test_id_property(self):
|
||||
"""Test id property returns pause model ID."""
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.id = "pause-123"
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
assert entity.id == "pause-123"
|
||||
|
||||
def test_workflow_execution_id_property(self):
|
||||
"""Test workflow_execution_id property returns workflow run ID."""
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.workflow_run_id = "execution-456"
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
assert entity.workflow_execution_id == "execution-456"
|
||||
|
||||
@@ -47,19 +58,17 @@ class TestPrivateWorkflowPauseEntity:
|
||||
"""Test resumed_at property returns pause model resumed_at."""
|
||||
resumed_at = datetime(2023, 12, 25, 15, 30, 45)
|
||||
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.resumed_at = resumed_at
|
||||
pause_model = _make_workflow_pause(resumed_at=resumed_at)
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
assert entity.resumed_at == resumed_at
|
||||
|
||||
def test_resumed_at_property_none(self):
|
||||
"""Test resumed_at property returns None when not set."""
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.resumed_at = None
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
assert entity.resumed_at is None
|
||||
|
||||
@@ -69,10 +78,9 @@ class TestPrivateWorkflowPauseEntity:
|
||||
state_data = b'{"test": "data", "step": 5}'
|
||||
mock_storage.load.return_value = state_data
|
||||
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.state_object_key = "test-state-key"
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
# First call should load from storage
|
||||
result = entity.get_state()
|
||||
@@ -87,10 +95,9 @@ class TestPrivateWorkflowPauseEntity:
|
||||
state_data = b'{"test": "data", "step": 5}'
|
||||
mock_storage.load.return_value = state_data
|
||||
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
mock_pause_model.state_object_key = "test-state-key"
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
# First call
|
||||
result1 = entity.get_state()
|
||||
@@ -107,9 +114,9 @@ class TestPrivateWorkflowPauseEntity:
|
||||
"""Test get_state returns pre-cached data."""
|
||||
state_data = b'{"test": "data", "step": 5}'
|
||||
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
# Pre-cache data
|
||||
entity._cached_state = state_data
|
||||
@@ -128,9 +135,9 @@ class TestPrivateWorkflowPauseEntity:
|
||||
with patch("repositories.sqlalchemy_api_workflow_run_repository.storage", autospec=True) as mock_storage:
|
||||
mock_storage.load.return_value = binary_data
|
||||
|
||||
mock_pause_model = MagicMock(spec=WorkflowPauseModel)
|
||||
pause_model = _make_workflow_pause()
|
||||
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[])
|
||||
entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
|
||||
|
||||
result = entity.get_state()
|
||||
|
||||
|
||||
@@ -173,10 +173,10 @@ def test_workflow_account_getters_use_caller_session():
|
||||
created_by="created-account-id",
|
||||
environment_variables=[],
|
||||
conversation_variables=[],
|
||||
updated_by="updated-account-id",
|
||||
)
|
||||
workflow.updated_by = "updated-account-id"
|
||||
created_account = mock.Mock(spec=Account)
|
||||
updated_account = mock.Mock(spec=Account)
|
||||
created_account = Account(name="Test Account", email="test@example.com")
|
||||
updated_account = Account(name="Test Account", email="test@example.com")
|
||||
session = mock.Mock()
|
||||
session.get.side_effect = [created_account, updated_account]
|
||||
|
||||
@@ -245,8 +245,9 @@ def test_normalize_environment_variable_mappings_keeps_hidden_value():
|
||||
|
||||
class TestWorkflowNodeExecution:
|
||||
def test_execution_metadata_dict(self):
|
||||
node_exec = WorkflowNodeExecutionModel()
|
||||
node_exec.execution_metadata = None
|
||||
node_exec = WorkflowNodeExecutionModel(
|
||||
execution_metadata=None,
|
||||
)
|
||||
assert node_exec.execution_metadata_dict == {}
|
||||
|
||||
original = {"a": 1, "b": ["2"]}
|
||||
@@ -409,8 +410,9 @@ class TestWorkflowDraftVariableGetValue:
|
||||
size=12,
|
||||
storage_key="canonical-storage-key",
|
||||
)
|
||||
draft_var = WorkflowDraftVariable()
|
||||
draft_var.app_id = "app-1"
|
||||
draft_var = WorkflowDraftVariable(
|
||||
app_id="app-1",
|
||||
)
|
||||
draft_var.set_value(build_segment(persisted_file))
|
||||
draft_var._WorkflowDraftVariable__value = None
|
||||
|
||||
|
||||
@@ -79,10 +79,14 @@ def _make_pipeline(
|
||||
workflow_id: str | None = None,
|
||||
is_published: bool = False,
|
||||
) -> Pipeline:
|
||||
pipeline = Pipeline(tenant_id=tenant_id, name="Test Pipeline", description="test")
|
||||
pipeline = Pipeline(
|
||||
tenant_id=tenant_id,
|
||||
name="Test Pipeline",
|
||||
description="test",
|
||||
workflow_id=workflow_id,
|
||||
is_published=is_published,
|
||||
)
|
||||
pipeline.id = pipeline_id
|
||||
pipeline.workflow_id = workflow_id
|
||||
pipeline.is_published = is_published
|
||||
return pipeline
|
||||
|
||||
|
||||
@@ -122,8 +126,8 @@ def _make_dataset(*, dataset_id: str = "d1", pipeline_id: str = "p1", tenant_id:
|
||||
tenant_id=tenant_id,
|
||||
name="Test Dataset",
|
||||
created_by="u1",
|
||||
pipeline_id=pipeline_id,
|
||||
)
|
||||
dataset.pipeline_id = pipeline_id
|
||||
return dataset
|
||||
|
||||
|
||||
@@ -982,14 +986,22 @@ def test_retry_error_document_success(
|
||||
|
||||
# 1. Setup mocks
|
||||
dataset = mocker.Mock()
|
||||
document = mocker.Mock(spec=Document)
|
||||
document.id = "doc-1"
|
||||
document = Document(
|
||||
id="doc-1",
|
||||
)
|
||||
|
||||
log = mocker.Mock(spec=DocumentPipelineExecutionLog)
|
||||
log.pipeline_id = "p-1"
|
||||
log.datasource_info = "{}" # Ensure it's a string if it's used as JSON later
|
||||
log = DocumentPipelineExecutionLog(
|
||||
pipeline_id="p-1",
|
||||
document_id="document-id",
|
||||
datasource_type="upload_file",
|
||||
datasource_info="{}",
|
||||
datasource_node_id="node-id",
|
||||
input_data="{}",
|
||||
created_by="account-id",
|
||||
)
|
||||
# Ensure it's a string if it's used as JSON later
|
||||
|
||||
pipeline = mocker.Mock(spec=Pipeline)
|
||||
pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline", workflow_id="wf-1")
|
||||
pipeline.id = "p-1"
|
||||
|
||||
workflow = mocker.Mock()
|
||||
@@ -1019,9 +1031,11 @@ def test_set_datasource_variables_success(
|
||||
from models.dataset import Pipeline
|
||||
|
||||
# 1. Setup mocks
|
||||
pipeline = mocker.Mock(spec=Pipeline)
|
||||
pipeline = Pipeline(
|
||||
tenant_id="t1",
|
||||
name="Test Pipeline",
|
||||
)
|
||||
pipeline.id = "p-1"
|
||||
pipeline.tenant_id = "t1"
|
||||
|
||||
draft_wf = mocker.Mock()
|
||||
draft_wf.id = "wf-1"
|
||||
|
||||
@@ -2,7 +2,7 @@ import json
|
||||
import unittest
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, Mock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -325,20 +325,23 @@ class TestExtractScheduleConfig(unittest.TestCase):
|
||||
|
||||
def test_extract_schedule_config_with_cron_mode(self):
|
||||
"""Test extracting schedule config in cron mode."""
|
||||
workflow = Mock(spec=Workflow)
|
||||
workflow.graph_dict = {
|
||||
"nodes": [
|
||||
workflow = Workflow(
|
||||
graph=json.dumps(
|
||||
{
|
||||
"id": "schedule-node",
|
||||
"data": {
|
||||
"type": "trigger-schedule",
|
||||
"mode": "cron",
|
||||
"cron_expression": "0 10 * * *",
|
||||
"timezone": "America/New_York",
|
||||
},
|
||||
"nodes": [
|
||||
{
|
||||
"id": "schedule-node",
|
||||
"data": {
|
||||
"type": "trigger-schedule",
|
||||
"mode": "cron",
|
||||
"cron_expression": "0 10 * * *",
|
||||
"timezone": "America/New_York",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
config = ScheduleService.extract_schedule_config(workflow)
|
||||
|
||||
@@ -349,21 +352,24 @@ class TestExtractScheduleConfig(unittest.TestCase):
|
||||
|
||||
def test_extract_schedule_config_with_visual_mode(self):
|
||||
"""Test extracting schedule config in visual mode."""
|
||||
workflow = Mock(spec=Workflow)
|
||||
workflow.graph_dict = {
|
||||
"nodes": [
|
||||
workflow = Workflow(
|
||||
graph=json.dumps(
|
||||
{
|
||||
"id": "schedule-node",
|
||||
"data": {
|
||||
"type": "trigger-schedule",
|
||||
"mode": "visual",
|
||||
"frequency": "daily",
|
||||
"visual_config": {"time": "10:30 AM"},
|
||||
"timezone": "UTC",
|
||||
},
|
||||
"nodes": [
|
||||
{
|
||||
"id": "schedule-node",
|
||||
"data": {
|
||||
"type": "trigger-schedule",
|
||||
"mode": "visual",
|
||||
"frequency": "daily",
|
||||
"visual_config": {"time": "10:30 AM"},
|
||||
"timezone": "UTC",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
config = ScheduleService.extract_schedule_config(workflow)
|
||||
|
||||
@@ -374,23 +380,27 @@ class TestExtractScheduleConfig(unittest.TestCase):
|
||||
|
||||
def test_extract_schedule_config_no_schedule_node(self):
|
||||
"""Test extracting config when no schedule node exists."""
|
||||
workflow = Mock(spec=Workflow)
|
||||
workflow.graph_dict = {
|
||||
"nodes": [
|
||||
workflow = Workflow(
|
||||
graph=json.dumps(
|
||||
{
|
||||
"id": "other-node",
|
||||
"data": {"type": "llm"},
|
||||
"nodes": [
|
||||
{
|
||||
"id": "other-node",
|
||||
"data": {"type": "llm"},
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
config = ScheduleService.extract_schedule_config(workflow)
|
||||
assert config is None
|
||||
|
||||
def test_extract_schedule_config_invalid_graph(self):
|
||||
"""Test extracting config with invalid graph data."""
|
||||
workflow = Mock(spec=Workflow)
|
||||
workflow.graph_dict = None
|
||||
workflow = Workflow(
|
||||
graph="",
|
||||
)
|
||||
|
||||
with pytest.raises(ScheduleConfigError, match="Workflow graph is empty"):
|
||||
ScheduleService.extract_schedule_config(workflow)
|
||||
@@ -496,9 +506,8 @@ class TestScheduleWithTimezone(unittest.TestCase):
|
||||
assert summer_next.hour == 14
|
||||
|
||||
|
||||
def _workflow(**kwargs: Any) -> Workflow:
|
||||
graph_dict = kwargs.pop("graph_dict", {})
|
||||
workflow = Workflow.new(
|
||||
def _workflow(*, graph_dict: dict[str, Any]) -> Workflow:
|
||||
return Workflow.new(
|
||||
tenant_id="tenant-1",
|
||||
app_id="app-1",
|
||||
type=WorkflowType.WORKFLOW,
|
||||
@@ -510,9 +519,6 @@ def _workflow(**kwargs: Any) -> Workflow:
|
||||
conversation_variables=[],
|
||||
rag_pipeline_variables=[],
|
||||
)
|
||||
for key, value in kwargs.items():
|
||||
setattr(workflow, key, value)
|
||||
return workflow
|
||||
|
||||
|
||||
def test_to_schedule_config_should_build_from_cron_mode() -> None:
|
||||
|
||||
@@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
|
||||
from core.workflow.file_reference import build_file_reference
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
from graphon.variables.segments import ObjectSegment, StringSegment
|
||||
from graphon.variables.segments import StringSegment
|
||||
from graphon.variables.types import SegmentType
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import UploadFile
|
||||
@@ -77,28 +77,45 @@ class TestDraftVarLoaderSimple:
|
||||
def test_load_offloaded_variable_object_type_unit(self, draft_var_loader):
|
||||
"""Test _load_offloaded_variable with object type - isolated unit test."""
|
||||
# Create mock objects
|
||||
upload_file = Mock(spec=UploadFile)
|
||||
upload_file.key = "storage/key/test.json"
|
||||
upload_file = UploadFile(
|
||||
tenant_id="tenant-id",
|
||||
storage_type="opendal",
|
||||
key="storage/key/test.json",
|
||||
name="test.txt",
|
||||
size=0,
|
||||
extension="txt",
|
||||
mime_type="text/plain",
|
||||
created_by_role="account",
|
||||
created_by="account-id",
|
||||
created_at=datetime.now(),
|
||||
used=False,
|
||||
)
|
||||
|
||||
variable_file = Mock(spec=WorkflowDraftVariableFile)
|
||||
variable_file.value_type = SegmentType.OBJECT
|
||||
variable_file = WorkflowDraftVariableFile(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
upload_file_id="upload-file-id",
|
||||
size=0,
|
||||
length=0,
|
||||
value_type=SegmentType.OBJECT,
|
||||
)
|
||||
variable_file.upload_file = upload_file
|
||||
|
||||
draft_var = Mock(spec=WorkflowDraftVariable)
|
||||
draft_var.id = "draft-var-id"
|
||||
draft_var.node_id = "test-node-id"
|
||||
draft_var.name = "test_object"
|
||||
draft_var.description = "test description"
|
||||
draft_var.get_selector.return_value = ["test-node-id", "test_object"]
|
||||
draft_var.variable_file = variable_file
|
||||
draft_var = WorkflowDraftVariable(
|
||||
id="draft-var-id",
|
||||
node_id="test-node-id",
|
||||
name="test_object",
|
||||
description="test description",
|
||||
selector=json.dumps(["test-node-id", "test_object"]),
|
||||
variable_file=variable_file,
|
||||
)
|
||||
|
||||
test_object = {"key1": "value1", "key2": 42}
|
||||
test_json_content = json.dumps(test_object, ensure_ascii=False, separators=(",", ":"))
|
||||
|
||||
with patch("services.workflow_draft_variable_service.storage") as mock_storage:
|
||||
mock_storage.load.return_value = test_json_content.encode()
|
||||
mock_segment = ObjectSegment(value=test_object)
|
||||
draft_var.build_segment_from_serialized_value.return_value = mock_segment
|
||||
|
||||
# Execute the method
|
||||
selector_tuple, variable = draft_var_loader._load_offloaded_variable(draft_var)
|
||||
@@ -112,23 +129,32 @@ class TestDraftVarLoaderSimple:
|
||||
|
||||
# Verify method calls
|
||||
mock_storage.load.assert_called_once_with("storage/key/test.json")
|
||||
draft_var.build_segment_from_serialized_value.assert_called_once_with(SegmentType.OBJECT, test_object)
|
||||
|
||||
def test_load_offloaded_variable_missing_variable_file_unit(self, draft_var_loader):
|
||||
"""Test that assertion error is raised when variable_file is None."""
|
||||
draft_var = Mock(spec=WorkflowDraftVariable)
|
||||
draft_var.variable_file = None
|
||||
draft_var = WorkflowDraftVariable(
|
||||
variable_file=None,
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
draft_var_loader._load_offloaded_variable(draft_var)
|
||||
|
||||
def test_load_offloaded_variable_missing_upload_file_unit(self, draft_var_loader):
|
||||
"""Test that assertion error is raised when upload_file is None."""
|
||||
variable_file = Mock(spec=WorkflowDraftVariableFile)
|
||||
variable_file = WorkflowDraftVariableFile(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
upload_file_id="upload-file-id",
|
||||
size=0,
|
||||
length=0,
|
||||
value_type="file",
|
||||
)
|
||||
variable_file.upload_file = None
|
||||
|
||||
draft_var = Mock(spec=WorkflowDraftVariable)
|
||||
draft_var.variable_file = variable_file
|
||||
draft_var = WorkflowDraftVariable(
|
||||
variable_file=variable_file,
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
draft_var_loader._load_offloaded_variable(draft_var)
|
||||
@@ -147,31 +173,45 @@ class TestDraftVarLoaderSimple:
|
||||
def test_load_offloaded_variable_array_type_unit(self, draft_var_loader):
|
||||
"""Test _load_offloaded_variable with array type - isolated unit test."""
|
||||
# Create mock objects
|
||||
upload_file = Mock(spec=UploadFile)
|
||||
upload_file.key = "storage/key/test_array.json"
|
||||
upload_file = UploadFile(
|
||||
tenant_id="tenant-id",
|
||||
storage_type="opendal",
|
||||
key="storage/key/test_array.json",
|
||||
name="test.txt",
|
||||
size=0,
|
||||
extension="txt",
|
||||
mime_type="text/plain",
|
||||
created_by_role="account",
|
||||
created_by="account-id",
|
||||
created_at=datetime.now(),
|
||||
used=False,
|
||||
)
|
||||
|
||||
variable_file = Mock(spec=WorkflowDraftVariableFile)
|
||||
variable_file.value_type = SegmentType.ARRAY_ANY
|
||||
variable_file = WorkflowDraftVariableFile(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
upload_file_id="upload-file-id",
|
||||
size=0,
|
||||
length=0,
|
||||
value_type=SegmentType.ARRAY_ANY,
|
||||
)
|
||||
variable_file.upload_file = upload_file
|
||||
|
||||
draft_var = Mock(spec=WorkflowDraftVariable)
|
||||
draft_var.id = "draft-var-id"
|
||||
draft_var.node_id = "test-node-id"
|
||||
draft_var.name = "test_array"
|
||||
draft_var.description = "test array description"
|
||||
draft_var.get_selector.return_value = ["test-node-id", "test_array"]
|
||||
draft_var.variable_file = variable_file
|
||||
draft_var = WorkflowDraftVariable(
|
||||
id="draft-var-id",
|
||||
node_id="test-node-id",
|
||||
name="test_array",
|
||||
description="test array description",
|
||||
selector=json.dumps(["test-node-id", "test_array"]),
|
||||
variable_file=variable_file,
|
||||
)
|
||||
|
||||
test_array = ["item1", "item2", "item3"]
|
||||
test_array = ["item1", 2, True]
|
||||
test_json_content = json.dumps(test_array)
|
||||
|
||||
with patch("services.workflow_draft_variable_service.storage") as mock_storage:
|
||||
mock_storage.load.return_value = test_json_content.encode()
|
||||
from graphon.variables.segments import ArrayAnySegment
|
||||
|
||||
mock_segment = ArrayAnySegment(value=test_array)
|
||||
draft_var.build_segment_from_serialized_value.return_value = mock_segment
|
||||
|
||||
# Execute the method
|
||||
selector_tuple, variable = draft_var_loader._load_offloaded_variable(draft_var)
|
||||
|
||||
@@ -183,14 +223,31 @@ class TestDraftVarLoaderSimple:
|
||||
|
||||
# Verify method calls
|
||||
mock_storage.load.assert_called_once_with("storage/key/test_array.json")
|
||||
draft_var.build_segment_from_serialized_value.assert_called_once_with(SegmentType.ARRAY_ANY, test_array)
|
||||
|
||||
def test_load_offloaded_variable_file_type_rebuilds_storage_backed_payload(self, draft_var_loader):
|
||||
upload_file = Mock(spec=UploadFile)
|
||||
upload_file.key = "storage/key/test_file.json"
|
||||
upload_file = UploadFile(
|
||||
tenant_id="tenant-id",
|
||||
storage_type="opendal",
|
||||
key="storage/key/test_file.json",
|
||||
name="test.txt",
|
||||
size=0,
|
||||
extension="txt",
|
||||
mime_type="text/plain",
|
||||
created_by_role="account",
|
||||
created_by="account-id",
|
||||
created_at=datetime.now(),
|
||||
used=False,
|
||||
)
|
||||
|
||||
variable_file = Mock(spec=WorkflowDraftVariableFile)
|
||||
variable_file.value_type = SegmentType.FILE
|
||||
variable_file = WorkflowDraftVariableFile(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
upload_file_id="upload-file-id",
|
||||
size=0,
|
||||
length=0,
|
||||
value_type=SegmentType.FILE,
|
||||
)
|
||||
variable_file.upload_file = upload_file
|
||||
|
||||
draft_var = WorkflowDraftVariable(
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import dataclasses
|
||||
import json
|
||||
import secrets
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
@@ -128,7 +129,7 @@ class TestDraftVariableSaver:
|
||||
assert name == c.expected_name, fail_msg
|
||||
|
||||
def test_build_variables_from_start_mapping_rebuilds_system_files(self, sqlite_session: Session):
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
mock_user.id = str(uuid.uuid4())
|
||||
saver = DraftVariableSaver(
|
||||
session=sqlite_session,
|
||||
@@ -169,9 +170,8 @@ class TestDraftVariableSaver:
|
||||
def draft_saver(self, sqlite_session: Session):
|
||||
"""Create DraftVariableSaver instance with user context."""
|
||||
# Create a mock user
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
mock_user.id = "test-user-id"
|
||||
mock_user.tenant_id = "test-tenant-id"
|
||||
|
||||
return DraftVariableSaver(
|
||||
session=sqlite_session,
|
||||
@@ -218,9 +218,8 @@ class TestDraftVariableSaver:
|
||||
assert draft_var.file_id == mock_draft_var_file.id
|
||||
|
||||
def test_try_offload_large_variable_uses_resource_tenant(self, sqlite_session: Session):
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
mock_user.id = "test-user-id"
|
||||
mock_user.current_tenant_id = ""
|
||||
saver = DraftVariableSaver(
|
||||
session=sqlite_session,
|
||||
tenant_id="app-tenant-id",
|
||||
@@ -268,7 +267,7 @@ class TestDraftVariableSaver:
|
||||
self, mock_batch_upsert, sqlite_session: Session
|
||||
):
|
||||
"""Start node should persist common `sys.*` variables, not only `sys.files`."""
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
mock_user.id = "test-user-id"
|
||||
mock_user.tenant_id = "test-tenant-id"
|
||||
|
||||
@@ -303,7 +302,7 @@ class TestDraftVariableSaver:
|
||||
|
||||
@patch("services.workflow_draft_variable_service._batch_upsert_draft_variable", autospec=True)
|
||||
def test_start_node_save_normalizes_reserved_prefix_outputs(self, mock_batch_upsert, sqlite_session: Session):
|
||||
mock_user = MagicMock(spec=Account)
|
||||
mock_user = Account(name="Test Account", email="test@example.com")
|
||||
mock_user.id = "test-user-id"
|
||||
mock_user.tenant_id = "test-tenant-id"
|
||||
|
||||
@@ -474,13 +473,11 @@ class TestWorkflowDraftVariableService:
|
||||
"""Reset a node variable from its execution output and flush the restored value."""
|
||||
service = WorkflowDraftVariableService(sqlite_session)
|
||||
|
||||
# Create mock execution record
|
||||
mock_execution = Mock(spec=WorkflowNodeExecutionModel)
|
||||
mock_execution.load_full_outputs.return_value = {"test_var": "output_value"}
|
||||
execution = WorkflowNodeExecutionModel(outputs=json.dumps({"test_var": "output_value"}))
|
||||
|
||||
# Mock the repository to return the execution record
|
||||
service._api_node_execution_repo = Mock()
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = execution
|
||||
|
||||
test_app_id = self._get_test_app_id()
|
||||
workflow = self._create_test_workflow(test_app_id)
|
||||
@@ -549,13 +546,11 @@ class TestWorkflowDraftVariableService:
|
||||
sqlite_session.add(variable)
|
||||
sqlite_session.commit()
|
||||
|
||||
# Create mock execution record
|
||||
mock_execution = Mock(spec=WorkflowNodeExecutionModel)
|
||||
mock_execution.load_full_outputs.return_value = {"sys.files": "[]"}
|
||||
execution = WorkflowNodeExecutionModel(outputs=json.dumps({"sys.files": "[]"}))
|
||||
|
||||
# Mock the repository to return the execution record
|
||||
service._api_node_execution_repo = Mock()
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = execution
|
||||
|
||||
with patch.object(sqlite_session, "flush", wraps=sqlite_session.flush) as flush:
|
||||
result = service._reset_node_var_or_sys_var(workflow, variable)
|
||||
@@ -584,13 +579,11 @@ class TestWorkflowDraftVariableService:
|
||||
sqlite_session.add(variable)
|
||||
sqlite_session.commit()
|
||||
|
||||
# Create mock execution record
|
||||
mock_execution = Mock(spec=WorkflowNodeExecutionModel)
|
||||
mock_execution.load_full_outputs.return_value = {"sys.query": "reset query"}
|
||||
execution = WorkflowNodeExecutionModel(outputs=json.dumps({"sys.query": "reset query"}))
|
||||
|
||||
# Mock the repository to return the execution record
|
||||
service._api_node_execution_repo = Mock()
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution
|
||||
service._api_node_execution_repo.get_execution_by_id.return_value = execution
|
||||
|
||||
with patch.object(sqlite_session, "flush", wraps=sqlite_session.flush) as flush:
|
||||
result = service._reset_node_var_or_sys_var(workflow, variable)
|
||||
|
||||
Reference in New Issue
Block a user