mirror of
https://github.com/langgenius/dify.git
synced 2026-08-01 09:50:50 +08:00
test: use sqlite3 session in test_sqlalchemy_workflow_execution_repository (#38761)
This commit is contained in:
parent
133e974dff
commit
103c864b35
@ -1,56 +1,59 @@
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
from unittest.mock import MagicMock
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.repositories.sqlalchemy_workflow_execution_repository import SQLAlchemyWorkflowExecutionRepository
|
||||
from graphon.entities import WorkflowExecution
|
||||
from graphon.enums import WorkflowExecutionStatus, WorkflowType
|
||||
from models import Account, CreatorUserRole, EndUser, WorkflowRun
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
from models import Account, CreatorUserRole, EndUser, Tenant, WorkflowRun
|
||||
from models.enums import EndUserType, WorkflowRunTriggeredFrom
|
||||
from models.workflow import WorkflowType as ModelWorkflowType
|
||||
|
||||
TABLES = (WorkflowRun,)
|
||||
|
||||
RESOURCE_TENANT_ID = "resource-tenant-id"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session_factory():
|
||||
"""Mock SQLAlchemy session factory."""
|
||||
session_factory = MagicMock(spec=sessionmaker)
|
||||
session = MagicMock()
|
||||
session.get.return_value = None
|
||||
session_factory.return_value.__enter__.return_value = session
|
||||
return session_factory
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_engine():
|
||||
"""Mock SQLAlchemy Engine."""
|
||||
return MagicMock(spec=Engine)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_account():
|
||||
"""Mock Account user."""
|
||||
account = MagicMock(spec=Account)
|
||||
def _make_account(*, tenant_id: str | None = None) -> Account:
|
||||
account = Account(name="Repository User", email=f"{uuid4()}@example.com")
|
||||
account.id = str(uuid4())
|
||||
account.current_tenant_id = str(uuid4())
|
||||
if tenant_id is not None:
|
||||
tenant = Tenant(name="Repository Tenant")
|
||||
tenant.id = tenant_id
|
||||
account._current_tenant = tenant
|
||||
return account
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_end_user():
|
||||
"""Mock EndUser."""
|
||||
user = MagicMock(spec=EndUser)
|
||||
user.id = str(uuid4())
|
||||
user.tenant_id = str(uuid4())
|
||||
return user
|
||||
def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]:
|
||||
"""Create repository-owned sessions bound to the isolated SQLite engine."""
|
||||
return sessionmaker(bind=sqlite_engine, expire_on_commit=False)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_workflow_execution():
|
||||
def account() -> Account:
|
||||
return _make_account(tenant_id=str(uuid4()))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def end_user() -> EndUser:
|
||||
return EndUser(
|
||||
id=str(uuid4()),
|
||||
tenant_id=str(uuid4()),
|
||||
app_id=None,
|
||||
type=EndUserType.SERVICE_API,
|
||||
external_user_id=None,
|
||||
name="Repository End User",
|
||||
session_id=str(uuid4()),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_workflow_execution() -> WorkflowExecution:
|
||||
"""Sample WorkflowExecution for testing."""
|
||||
return WorkflowExecution(
|
||||
id_=str(uuid4()),
|
||||
@ -71,125 +74,147 @@ def sample_workflow_execution():
|
||||
|
||||
|
||||
class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
def test_init_with_sessionmaker(self, mock_session_factory, mock_account):
|
||||
def test_init_with_sessionmaker(self, sqlite_session_factory: sessionmaker[Session], account: Account):
|
||||
app_id = "test_app_id"
|
||||
triggered_from = WorkflowRunTriggeredFrom.APP_RUN
|
||||
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id=app_id,
|
||||
triggered_from=triggered_from,
|
||||
)
|
||||
|
||||
assert repo._session_factory == mock_session_factory
|
||||
assert repo._session_factory is sqlite_session_factory
|
||||
assert repo._tenant_id == RESOURCE_TENANT_ID
|
||||
assert repo._app_id == app_id
|
||||
assert repo._triggered_from == triggered_from
|
||||
assert repo._creator_user_id == mock_account.id
|
||||
assert repo._creator_user_id == account.id
|
||||
assert repo._creator_user_role == CreatorUserRole.ACCOUNT
|
||||
|
||||
def test_init_with_engine(self, mock_engine, mock_account):
|
||||
def test_init_with_engine(self, sqlite_engine: Engine, account: Account):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_engine,
|
||||
session_factory=sqlite_engine,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app_id",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
|
||||
assert isinstance(repo._session_factory, sessionmaker)
|
||||
assert repo._session_factory.kw["bind"] == mock_engine
|
||||
assert repo._session_factory.kw["bind"] is sqlite_engine
|
||||
|
||||
def test_init_invalid_session_factory(self, mock_account):
|
||||
def test_init_invalid_session_factory(self, account: Account):
|
||||
with pytest.raises(ValueError, match="Invalid session_factory type"):
|
||||
SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory="invalid",
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id=None,
|
||||
triggered_from=None,
|
||||
)
|
||||
|
||||
def test_init_no_tenant_id(self, mock_session_factory):
|
||||
user = MagicMock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
def test_init_no_tenant_id(self, sqlite_session_factory: sessionmaker[Session]):
|
||||
user = _make_account()
|
||||
|
||||
with pytest.raises(ValueError, match="tenant_id is required"):
|
||||
SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id="",
|
||||
user=user,
|
||||
app_id=None,
|
||||
triggered_from=None,
|
||||
)
|
||||
|
||||
def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory):
|
||||
user = MagicMock(spec=Account)
|
||||
user.current_tenant_id = None
|
||||
user.id = str(uuid4())
|
||||
def test_init_uses_resource_tenant_when_account_has_no_current_tenant(
|
||||
self, sqlite_session_factory: sessionmaker[Session]
|
||||
):
|
||||
user = _make_account()
|
||||
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
tenant_id="resource-tenant-id",
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=user,
|
||||
app_id="test-app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
|
||||
assert repo._tenant_id == "resource-tenant-id"
|
||||
assert repo._tenant_id == RESOURCE_TENANT_ID
|
||||
assert repo._creator_user_id == user.id
|
||||
|
||||
def test_init_with_end_user(self, mock_session_factory, mock_end_user):
|
||||
def test_init_with_end_user(self, sqlite_session_factory: sessionmaker[Session], end_user: EndUser):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_end_user,
|
||||
user=end_user,
|
||||
app_id=None,
|
||||
triggered_from=None,
|
||||
)
|
||||
assert repo._tenant_id == RESOURCE_TENANT_ID
|
||||
assert repo._creator_user_role == CreatorUserRole.END_USER
|
||||
|
||||
def test_to_domain_model(self, mock_session_factory, mock_account):
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
def test_to_domain_model(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
account: Account,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id=None,
|
||||
triggered_from=None,
|
||||
)
|
||||
|
||||
db_model = MagicMock(spec=WorkflowRun)
|
||||
db_model.id = str(uuid4())
|
||||
db_model.workflow_id = str(uuid4())
|
||||
db_model.type = "workflow"
|
||||
db_model.version = "1.0"
|
||||
db_model.inputs_dict = {"in": "val"}
|
||||
db_model.outputs_dict = {"out": "val"}
|
||||
db_model.graph_dict = {"nodes": []}
|
||||
db_model.status = "succeeded"
|
||||
db_model.error = "some error"
|
||||
db_model.total_tokens = 50
|
||||
db_model.total_steps = 3
|
||||
db_model.exceptions_count = 1
|
||||
db_model.created_at = datetime.now(UTC)
|
||||
db_model.finished_at = datetime.now(UTC)
|
||||
db_model = WorkflowRun(
|
||||
id=str(uuid4()),
|
||||
tenant_id=account.current_tenant_id,
|
||||
app_id=str(uuid4()),
|
||||
workflow_id=str(uuid4()),
|
||||
type=ModelWorkflowType.WORKFLOW,
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
version="1.0",
|
||||
inputs=json.dumps({"in": "val"}),
|
||||
outputs=json.dumps({"out": "val"}),
|
||||
graph=json.dumps({"nodes": []}),
|
||||
status=WorkflowExecutionStatus.SUCCEEDED,
|
||||
error="some error",
|
||||
elapsed_time=1.0,
|
||||
total_tokens=50,
|
||||
total_steps=3,
|
||||
exceptions_count=1,
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by=account.id,
|
||||
created_at=datetime.now(UTC),
|
||||
finished_at=datetime.now(UTC),
|
||||
)
|
||||
sqlite_session.add(db_model)
|
||||
sqlite_session.commit()
|
||||
sqlite_session.expunge_all()
|
||||
persisted_model = sqlite_session.get(WorkflowRun, db_model.id)
|
||||
assert persisted_model is not None
|
||||
|
||||
domain_model = repo._to_domain_model(db_model)
|
||||
domain_model = repo._to_domain_model(persisted_model)
|
||||
|
||||
assert domain_model.id_ == db_model.id
|
||||
assert domain_model.workflow_id == db_model.workflow_id
|
||||
assert domain_model.id_ == persisted_model.id
|
||||
assert domain_model.workflow_id == persisted_model.workflow_id
|
||||
assert domain_model.status == WorkflowExecutionStatus.SUCCEEDED
|
||||
assert domain_model.inputs == db_model.inputs_dict
|
||||
assert domain_model.inputs == {"in": "val"}
|
||||
assert domain_model.error_message == "some error"
|
||||
|
||||
def test_to_db_model(self, mock_session_factory, mock_account, sample_workflow_execution):
|
||||
def test_to_db_model(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
|
||||
)
|
||||
@ -208,11 +233,16 @@ class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
assert db_model.total_tokens == sample_workflow_execution.total_tokens
|
||||
assert db_model.elapsed_time == 10.0
|
||||
|
||||
def test_to_db_model_edge_cases(self, mock_session_factory, mock_account, sample_workflow_execution):
|
||||
def test_to_db_model_edge_cases(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.DEBUGGING,
|
||||
)
|
||||
@ -231,11 +261,16 @@ class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
assert db_model.error is None
|
||||
assert db_model.elapsed_time == 0
|
||||
|
||||
def test_to_db_model_app_id_none(self, mock_session_factory, mock_account, sample_workflow_execution):
|
||||
def test_to_db_model_app_id_none(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id=None,
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
@ -244,11 +279,16 @@ class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
assert not hasattr(db_model, "app_id") or db_model.app_id is None
|
||||
assert db_model.tenant_id == repo._tenant_id
|
||||
|
||||
def test_to_db_model_missing_context(self, mock_session_factory, mock_account, sample_workflow_execution):
|
||||
def test_to_db_model_missing_context(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id=None,
|
||||
triggered_from=None,
|
||||
)
|
||||
@ -267,33 +307,47 @@ class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
with pytest.raises(ValueError, match="created_by_role is required"):
|
||||
repo._to_db_model(sample_workflow_execution)
|
||||
|
||||
def test_save(self, mock_session_factory, mock_account, sample_workflow_execution):
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
def test_save(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
|
||||
repo.save(sample_workflow_execution)
|
||||
|
||||
session = mock_session_factory.return_value.__enter__.return_value
|
||||
session.merge.assert_called_once()
|
||||
session.commit.assert_called_once()
|
||||
persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_)
|
||||
assert persisted_model is not None
|
||||
assert persisted_model.tenant_id == RESOURCE_TENANT_ID
|
||||
assert persisted_model.inputs_dict == sample_workflow_execution.inputs
|
||||
assert persisted_model.outputs_dict == sample_workflow_execution.outputs
|
||||
|
||||
# Check cache
|
||||
assert sample_workflow_execution.id_ in repo._execution_cache
|
||||
cached_model = repo._execution_cache[sample_workflow_execution.id_]
|
||||
assert cached_model.id == sample_workflow_execution.id_
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
def test_save_uses_execution_started_at_when_record_does_not_exist(
|
||||
self, mock_session_factory, mock_account, sample_workflow_execution
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
@ -301,41 +355,71 @@ class TestSQLAlchemyWorkflowExecutionRepository:
|
||||
started_at = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC)
|
||||
sample_workflow_execution.started_at = started_at
|
||||
|
||||
session = mock_session_factory.return_value.__enter__.return_value
|
||||
session.get.return_value = None
|
||||
|
||||
repo.save(sample_workflow_execution)
|
||||
|
||||
saved_model = session.merge.call_args.args[0]
|
||||
assert saved_model.created_at == started_at
|
||||
session.commit.assert_called_once()
|
||||
persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_)
|
||||
assert persisted_model is not None
|
||||
assert persisted_model.created_at == started_at.replace(tzinfo=None)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
def test_save_preserves_existing_created_at_when_record_already_exists(
|
||||
self, mock_session_factory, mock_account, sample_workflow_execution
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=mock_session_factory,
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=mock_account,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
|
||||
execution_id = sample_workflow_execution.id_
|
||||
existing_created_at = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC)
|
||||
|
||||
existing_run = WorkflowRun()
|
||||
existing_run.id = execution_id
|
||||
existing_run.tenant_id = repo._tenant_id
|
||||
existing_run.created_at = existing_created_at
|
||||
|
||||
session = mock_session_factory.return_value.__enter__.return_value
|
||||
session.get.return_value = existing_run
|
||||
sample_workflow_execution.started_at = existing_created_at
|
||||
repo.save(sample_workflow_execution)
|
||||
|
||||
sample_workflow_execution.started_at = datetime(2026, 1, 1, 12, 30, 0, tzinfo=UTC)
|
||||
|
||||
repo.save(sample_workflow_execution)
|
||||
|
||||
saved_model = session.merge.call_args.args[0]
|
||||
assert saved_model.created_at == existing_created_at
|
||||
session.commit.assert_called_once()
|
||||
persisted_model = sqlite_session.get(WorkflowRun, execution_id)
|
||||
assert persisted_model is not None
|
||||
assert persisted_model.created_at == existing_created_at.replace(tzinfo=None)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
def test_save_rejects_execution_owned_by_another_tenant(
|
||||
self,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
account: Account,
|
||||
sample_workflow_execution: WorkflowExecution,
|
||||
):
|
||||
other_tenant_id = str(uuid4())
|
||||
other_account = _make_account(tenant_id=str(uuid4()))
|
||||
other_repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=other_tenant_id,
|
||||
user=other_account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
other_repo.save(sample_workflow_execution)
|
||||
|
||||
repo = SQLAlchemyWorkflowExecutionRepository(
|
||||
session_factory=sqlite_session_factory,
|
||||
tenant_id=RESOURCE_TENANT_ID,
|
||||
user=account,
|
||||
app_id="test_app",
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
)
|
||||
with pytest.raises(ValueError, match="Unauthorized access to workflow run"):
|
||||
repo.save(sample_workflow_execution)
|
||||
|
||||
sqlite_session.expire_all()
|
||||
persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_)
|
||||
assert persisted_model is not None
|
||||
assert persisted_model.tenant_id == other_tenant_id
|
||||
|
||||
Loading…
Reference in New Issue
Block a user