"""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()