mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-25 14:06:18 +00:00
fix(persistence): drain engine close across cancellation (#5576)
* fix(persistence): drain engine close across cancellation * test(persistence): cover engine cleanup ownership
This commit is contained in:
parent
906c3d4554
commit
1ee0fa318b
@ -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())
|
||||
|
||||
92
backend/tests/test_persistence_engine_close_cancellation.py
Normal file
92
backend/tests/test_persistence_engine_close_cancellation.py
Normal file
@ -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
|
||||
Loading…
x
Reference in New Issue
Block a user