deerflow-code/offline-backend-20260512/backend/tests/test_workflow_schema_rebuild.py
2026-09-07 18:24:55 +08:00

91 lines
3.0 KiB
Python

"""Safety tests for the one-time MySQL Workflow Studio schema rebuild."""
from __future__ import annotations
from contextlib import asynccontextmanager
from types import SimpleNamespace
import pytest
import sqlalchemy as sa
import deerflow.persistence.engine as persistence_engine
from deerflow.persistence.base import Base
class _ScalarResult:
def __init__(self, value: str | None) -> None:
self._value = value
def scalar_one_or_none(self) -> str | None:
return self._value
class _SyncConnection:
def __init__(self, database: str | None) -> None:
self._database = database
def exec_driver_sql(self, statement: str) -> _ScalarResult:
assert statement == "SELECT DATABASE()"
return _ScalarResult(self._database)
class _AsyncConnection:
def __init__(self, sync_connection: _SyncConnection) -> None:
self._sync_connection = sync_connection
async def run_sync(self, callback):
return callback(self._sync_connection)
class _AsyncEngine:
def __init__(self, database: str, connected_database: str | None) -> None:
self.dialect = SimpleNamespace(name="mysql")
self.url = SimpleNamespace(database=database)
self._connection = _AsyncConnection(_SyncConnection(connected_database))
@asynccontextmanager
async def begin(self):
yield self._connection
@pytest.mark.asyncio
async def test_rebuild_workflow_schema_only_targets_known_tables(monkeypatch: pytest.MonkeyPatch) -> None:
engine = _AsyncEngine("deerflow", "deerflow")
monkeypatch.setattr(persistence_engine, "get_engine", lambda: engine)
monkeypatch.setattr(
sa,
"inspect",
lambda _connection: SimpleNamespace(
get_table_names=lambda: {"workflow_definitions", "unrelated_workflow_cache"}
),
)
calls: dict[str, list[str]] = {}
monkeypatch.setattr(
Base.metadata,
"drop_all",
lambda _connection, *, tables, checkfirst: calls.setdefault("drop", [table.name for table in tables]),
)
monkeypatch.setattr(
Base.metadata,
"create_all",
lambda _connection, *, tables, checkfirst: calls.setdefault("create", [table.name for table in tables]),
)
rebuilt = await persistence_engine.rebuild_workflow_schema_for_mysql()
expected = list(persistence_engine._WORKFLOW_SCHEMA_TABLE_NAMES)
assert rebuilt == ["workflow_definitions"]
assert calls == {"drop": expected, "create": expected}
assert "unrelated_workflow_cache" not in calls["drop"]
@pytest.mark.asyncio
async def test_rebuild_workflow_schema_refuses_unexpected_database(monkeypatch: pytest.MonkeyPatch) -> None:
engine = _AsyncEngine("deerflow", "another_database")
monkeypatch.setattr(persistence_engine, "get_engine", lambda: engine)
monkeypatch.setattr(sa, "inspect", lambda _connection: pytest.fail("must not inspect tables"))
with pytest.raises(RuntimeError, match="does not match configured database"):
await persistence_engine.rebuild_workflow_schema_for_mysql()