diff --git a/backend/packages/harness/deerflow/persistence/engine.py b/backend/packages/harness/deerflow/persistence/engine.py index 1503e82e6..565ced5b9 100644 --- a/backend/packages/harness/deerflow/persistence/engine.py +++ b/backend/packages/harness/deerflow/persistence/engine.py @@ -16,6 +16,8 @@ import logging from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine +from deerflow.utils.file_io import await_drained + # Recycle pooled Postgres connections before stale idle sockets can hang # pool_pre_ping. The command timeout bounds stalled ORM queries independently. POSTGRES_POOL_RECYCLE_SECONDS = 300 @@ -266,10 +268,21 @@ def get_engine() -> AsyncEngine | None: async def close_engine() -> None: - """Dispose the engine, release all connections.""" + """Dispose the engine before releasing the process-global ownership.""" global _engine, _session_factory - if _engine is not None: - await _engine.dispose() + + engine = _engine + if engine is None: + _session_factory = None + return + + async def dispose_and_clear() -> None: + global _engine, _session_factory + + await engine.dispose() logger.info("Persistence engine closed") - _engine = None - _session_factory = None + if _engine is engine: + _engine = None + _session_factory = None + + await await_drained(dispose_and_clear()) diff --git a/backend/tests/test_persistence_engine_close_cancellation.py b/backend/tests/test_persistence_engine_close_cancellation.py new file mode 100644 index 000000000..f7994241e --- /dev/null +++ b/backend/tests/test_persistence_engine_close_cancellation.py @@ -0,0 +1,92 @@ +import asyncio + +import pytest + +from deerflow.persistence import engine as engine_mod + + +class _BlockingEngine: + def __init__(self) -> None: + self.dispose_started = asyncio.Event() + self.allow_dispose = asyncio.Event() + self.dispose_finished = asyncio.Event() + + async def dispose(self) -> None: + self.dispose_started.set() + await self.allow_dispose.wait() + self.dispose_finished.set() + + +@pytest.mark.asyncio +async def test_close_engine_drains_dispose_across_repeated_cancellation() -> None: + fake_engine = _BlockingEngine() + fake_factory = object() + previous_engine = engine_mod._engine + previous_factory = engine_mod._session_factory + task: asyncio.Task[None] | None = None + + try: + engine_mod._engine = fake_engine + engine_mod._session_factory = fake_factory + task = asyncio.create_task(engine_mod.close_engine()) + await asyncio.wait_for(fake_engine.dispose_started.wait(), timeout=1) + + task.cancel() + for _ in range(5): + await asyncio.sleep(0) + assert not task.done(), "engine teardown returned before dispose finished" + assert engine_mod._engine is fake_engine + assert engine_mod._session_factory is fake_factory + + task.cancel() + for _ in range(5): + await asyncio.sleep(0) + assert not task.done(), "repeated cancellation interrupted engine dispose" + + fake_engine.allow_dispose.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert fake_engine.dispose_finished.is_set() + assert engine_mod._engine is None + assert engine_mod._session_factory is None + finally: + fake_engine.allow_dispose.set() + if task is not None and not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + engine_mod._engine = previous_engine + engine_mod._session_factory = previous_factory + + +@pytest.mark.asyncio +async def test_close_engine_preserves_replacement_globals() -> None: + closing_engine = _BlockingEngine() + closing_factory = object() + replacement_engine = _BlockingEngine() + replacement_factory = object() + previous_engine = engine_mod._engine + previous_factory = engine_mod._session_factory + task: asyncio.Task[None] | None = None + + try: + engine_mod._engine = closing_engine + engine_mod._session_factory = closing_factory + task = asyncio.create_task(engine_mod.close_engine()) + await asyncio.wait_for(closing_engine.dispose_started.wait(), timeout=1) + + engine_mod._engine = replacement_engine + engine_mod._session_factory = replacement_factory + closing_engine.allow_dispose.set() + await task + + assert closing_engine.dispose_finished.is_set() + assert engine_mod._engine is replacement_engine + assert engine_mod._session_factory is replacement_factory + finally: + closing_engine.allow_dispose.set() + if task is not None and not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + engine_mod._engine = previous_engine + engine_mod._session_factory = previous_factory