diff --git a/backend/packages/harness/deerflow/runtime/AGENTS.md b/backend/packages/harness/deerflow/runtime/AGENTS.md index 1f289c0a4..70725fe42 100644 --- a/backend/packages/harness/deerflow/runtime/AGENTS.md +++ b/backend/packages/harness/deerflow/runtime/AGENTS.md @@ -1,6 +1,6 @@ ### Stream Bridge Heartbeats -Memory and Redis bridges take their default idle heartbeat cadence from the startup-only `stream_bridge.heartbeat_interval_seconds` setting. Keep the default on the bridge instance so SSE, `/wait`, and internal subscribers stay aligned; an explicit `subscribe(..., heartbeat_interval=...)` remains a per-subscription override. +Memory and Redis bridges keep the startup-only `stream_bridge.heartbeat_interval_seconds` default on the instance; explicit `subscribe(..., heartbeat_interval=...)` overrides it per subscription. Provider context managers own cache/bridge backends through exit: drain `aclose()` / `close()` across caller cancellation before propagating cancellation. ### Checkpoint Channel Modes (`full` / `delta`) diff --git a/backend/packages/harness/deerflow/runtime/checkpoint_cache/provider.py b/backend/packages/harness/deerflow/runtime/checkpoint_cache/provider.py index 52dcf1a3b..28e7c3bd9 100644 --- a/backend/packages/harness/deerflow/runtime/checkpoint_cache/provider.py +++ b/backend/packages/harness/deerflow/runtime/checkpoint_cache/provider.py @@ -12,6 +12,7 @@ from typing import Any from deerflow.config.app_config import AppConfig from deerflow.runtime.checkpoint_cache.base import CACHE_FORMAT_VERSION, CheckpointHistoryCache from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache +from deerflow.utils.file_io import await_drained logger = logging.getLogger(__name__) @@ -80,7 +81,7 @@ async def make_checkpoint_cache( try: yield cache finally: - await cache.aclose() + await await_drained(cache.aclose()) return if config.type == "redis": @@ -95,7 +96,7 @@ async def make_checkpoint_cache( try: yield cache finally: - await cache.aclose() + await await_drained(cache.aclose()) return raise ValueError(f"Unknown checkpoint cache type: {config.type!r}") diff --git a/backend/packages/harness/deerflow/runtime/stream_bridge/async_provider.py b/backend/packages/harness/deerflow/runtime/stream_bridge/async_provider.py index 5376f3f30..1f321573e 100644 --- a/backend/packages/harness/deerflow/runtime/stream_bridge/async_provider.py +++ b/backend/packages/harness/deerflow/runtime/stream_bridge/async_provider.py @@ -20,6 +20,7 @@ from collections.abc import AsyncIterator from deerflow.config.app_config import AppConfig from deerflow.config.stream_bridge_config import StreamBridgeConfig, get_stream_bridge_config +from deerflow.utils.file_io import await_drained from .base import DEFAULT_HEARTBEAT_INTERVAL_SECONDS, StreamBridge @@ -71,7 +72,7 @@ async def make_stream_bridge(app_config: AppConfig | None = None) -> AsyncIterat try: yield bridge finally: - await bridge.close() + await await_drained(bridge.close()) return if config.type == "redis": @@ -95,7 +96,7 @@ async def make_stream_bridge(app_config: AppConfig | None = None) -> AsyncIterat try: yield bridge finally: - await bridge.close() + await await_drained(bridge.close()) return raise ValueError(f"Unknown stream bridge type: {config.type!r}") diff --git a/backend/tests/test_runtime_provider_close_cancellation.py b/backend/tests/test_runtime_provider_close_cancellation.py new file mode 100644 index 000000000..c93b10ad6 --- /dev/null +++ b/backend/tests/test_runtime_provider_close_cancellation.py @@ -0,0 +1,140 @@ +import asyncio +from types import SimpleNamespace + +import pytest + +from deerflow.runtime.checkpoint_cache import provider as checkpoint_provider +from deerflow.runtime.checkpoint_cache import redis as checkpoint_redis +from deerflow.runtime.stream_bridge import async_provider as stream_provider +from deerflow.runtime.stream_bridge import memory as stream_memory +from deerflow.runtime.stream_bridge import redis as stream_redis + + +class _BlockingCheckpointCache: + def __init__(self) -> None: + self.close_started = asyncio.Event() + self.allow_close = asyncio.Event() + + async def aclose(self) -> None: + self.close_started.set() + await self.allow_close.wait() + + +class _BlockingStreamBridge: + def __init__(self) -> None: + self.close_started = asyncio.Event() + self.allow_close = asyncio.Event() + + async def close(self) -> None: + self.close_started.set() + await self.allow_close.wait() + + +async def _assert_close_is_drained(cm, close_started: asyncio.Event, allow_close: asyncio.Event) -> None: + entered = asyncio.Event() + leave = asyncio.Event() + + async def owner() -> None: + async with cm: + entered.set() + await leave.wait() + + task: asyncio.Task[None] | None = None + try: + task = asyncio.create_task(owner()) + await asyncio.wait_for(entered.wait(), timeout=1) + leave.set() + await asyncio.wait_for(close_started.wait(), timeout=1) + + task.cancel() + for _ in range(5): + await asyncio.sleep(0) + assert not task.done(), "provider teardown returned before close finished" + + task.cancel() + for _ in range(5): + await asyncio.sleep(0) + assert not task.done(), "repeated cancellation interrupted provider close" + + allow_close.set() + with pytest.raises(asyncio.CancelledError): + await task + finally: + allow_close.set() + if task is not None and not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +@pytest.mark.asyncio +async def test_checkpoint_cache_context_drains_close_across_repeated_cancellation(monkeypatch: pytest.MonkeyPatch) -> None: + cache = _BlockingCheckpointCache() + monkeypatch.setattr(checkpoint_provider, "MemoryCheckpointHistoryCache", lambda **_kwargs: cache) + config = SimpleNamespace(database=SimpleNamespace(checkpoint_cache=SimpleNamespace(type="memory", max_entries=128))) + + await _assert_close_is_drained( + checkpoint_provider.make_checkpoint_cache(config, serde=object()), + cache.close_started, + cache.allow_close, + ) + + +@pytest.mark.asyncio +async def test_checkpoint_cache_redis_context_drains_close_across_repeated_cancellation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + cache = _BlockingCheckpointCache() + monkeypatch.setattr(checkpoint_redis, "RedisCheckpointHistoryCache", lambda *_args, **_kwargs: cache) + config = SimpleNamespace( + database=SimpleNamespace( + checkpoint_cache=SimpleNamespace( + type="redis", + max_entries=128, + redis_url="redis://localhost:6379/0", + ttl_seconds=60, + ) + ) + ) + + await _assert_close_is_drained( + checkpoint_provider.make_checkpoint_cache(config, serde=object()), + cache.close_started, + cache.allow_close, + ) + + +@pytest.mark.asyncio +async def test_stream_bridge_context_drains_close_across_repeated_cancellation(monkeypatch: pytest.MonkeyPatch) -> None: + bridge = _BlockingStreamBridge() + monkeypatch.setattr(stream_memory, "MemoryStreamBridge", lambda **_kwargs: bridge) + config = SimpleNamespace(stream_bridge=SimpleNamespace(type="memory", queue_maxsize=8, heartbeat_interval_seconds=1.0)) + + await _assert_close_is_drained( + stream_provider.make_stream_bridge(config), + bridge.close_started, + bridge.allow_close, + ) + + +@pytest.mark.asyncio +async def test_stream_bridge_redis_context_drains_close_across_repeated_cancellation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + bridge = _BlockingStreamBridge() + monkeypatch.setattr(stream_redis, "RedisStreamBridge", lambda *_args, **_kwargs: bridge) + config = SimpleNamespace( + stream_bridge=SimpleNamespace( + type="redis", + redis_url="redis://localhost:6379/0", + queue_maxsize=8, + heartbeat_interval_seconds=1.0, + max_connections=4, + stream_ttl_seconds=60, + ) + ) + + await _assert_close_is_drained( + stream_provider.make_stream_bridge(config), + bridge.close_started, + bridge.allow_close, + )