91 lines
3.0 KiB
Python
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()
|