mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-25 14:06:18 +00:00
fix(runtime): drain provider close across cancellation (#5622)
* fix(runtime): drain provider close across cancellation * test(runtime): cover redis provider teardown cancellation
This commit is contained in:
parent
6fae79047c
commit
656db1223d
@ -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`)
|
||||
|
||||
|
||||
@ -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}")
|
||||
|
||||
@ -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}")
|
||||
|
||||
140
backend/tests/test_runtime_provider_close_cancellation.py
Normal file
140
backend/tests/test_runtime_provider_close_cancellation.py
Normal file
@ -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,
|
||||
)
|
||||
Loading…
x
Reference in New Issue
Block a user