fix(security): close async resources and leaks
Add AsyncResourceTracker (core/async_cleanup.py) providing a central registry for async resources with timeout-bounded close_all(), async context manager support, and a __del__ finalizer that logs leaked resources by name. Enhance LangGraphBridge with cleanup_tasks_async() that awaits in-flight tasks with a deadline instead of fire-and-forget cancel(). Add cancellation_reasons dict to trace why tasks were cancelled. Add StateManager.close() to properly release checkpoint file handles and complete the RxPY BehaviorSubject. Add AcpEventQueue.close() to dispose all subscriptions. Includes 14 Behave scenarios (67 steps), Robot integration tests, ASV benchmarks, and docs/reference/async_safety.md. ISSUES CLOSED: #321
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
|
||||
## Unreleased
|
||||
|
||||
<<<<<<< HEAD
|
||||
### feat(actor): extend hierarchical actor YAML schema and loader
|
||||
|
||||
- Extended actor YAML schema with hierarchical graph support: per-node LSP bindings (`lsp_binding`), tool-source references (`tool_sources`), and subgraph `actor_ref`.
|
||||
@@ -63,6 +64,10 @@
|
||||
- Fixed failing unit tests.
|
||||
- Added changeset persistence and diff artifact storage for tracking multi-file changes
|
||||
across plan execution phases. (#163)
|
||||
- Added `AsyncResourceTracker` for unified async resource lifecycle with timeout-bounded
|
||||
cleanup, leak detection via finalizer, and async context manager support.
|
||||
- Enhanced `LangGraphBridge` with graceful task cancellation that awaits in-flight tasks.
|
||||
- Added `StateManager.close()` and `AcpEventQueue.close()` for proper resource disposal.
|
||||
- Expanded CONTRIBUTING.md with detailed guidance on the issue creation process, label system,
|
||||
ticket lifecycle, pull request requirements, and review/merge process.
|
||||
- Added commit scope, quality, and message format guidelines to CONTRIBUTING.md.
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
"""ASV benchmarks for async resource cleanup overhead (#321).
|
||||
|
||||
Measures registration, close_all, and leak-warning latency to establish
|
||||
baselines for the AsyncResourceTracker.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from cleveragents.core.async_cleanup import AsyncResourceTracker
|
||||
|
||||
|
||||
class _FakeResource:
|
||||
"""Minimal async resource for benchmarking."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.closed = False
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _make_tracker(count: int) -> AsyncResourceTracker:
|
||||
"""Create a tracker pre-loaded with *count* fake resources."""
|
||||
tracker = AsyncResourceTracker()
|
||||
for i in range(count):
|
||||
tracker.register(f"res-{i}", _FakeResource())
|
||||
return tracker
|
||||
|
||||
|
||||
class TimeRegisterSingle:
|
||||
"""Benchmark registering a single resource."""
|
||||
|
||||
timeout = 10
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = AsyncResourceTracker()
|
||||
self.counter = 0
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_register_one(self) -> None:
|
||||
name = f"bench-{self.counter}"
|
||||
self.counter += 1
|
||||
self.tracker.register(name, _FakeResource())
|
||||
|
||||
|
||||
class TimeRegisterBatch:
|
||||
"""Benchmark registering 100 resources in sequence."""
|
||||
|
||||
timeout = 10
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = AsyncResourceTracker()
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_register_100(self) -> None:
|
||||
for i in range(100):
|
||||
self.tracker.register(f"batch-{i}", _FakeResource())
|
||||
|
||||
|
||||
class TimeCloseAll:
|
||||
"""Benchmark close_all on pre-loaded trackers."""
|
||||
|
||||
timeout = 30
|
||||
number = 1 # Re-create tracker for each iteration.
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = _make_tracker(50)
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_close_50_resources(self) -> None:
|
||||
asyncio.run(self.tracker.close_all())
|
||||
|
||||
|
||||
class TimeCloseAllLarge:
|
||||
"""Benchmark close_all with 500 resources."""
|
||||
|
||||
timeout = 60
|
||||
number = 1
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = _make_tracker(500)
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_close_500_resources(self) -> None:
|
||||
asyncio.run(self.tracker.close_all())
|
||||
|
||||
|
||||
class TimeLeakWarning:
|
||||
"""Benchmark the _warn_unclosed finalizer path."""
|
||||
|
||||
timeout = 10
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = _make_tracker(20)
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_warn_20_unclosed(self) -> None:
|
||||
self.tracker._warn_unclosed()
|
||||
|
||||
|
||||
class TimeOpenCount:
|
||||
"""Benchmark the open_count property."""
|
||||
|
||||
timeout = 10
|
||||
|
||||
def setup(self) -> None:
|
||||
self.tracker = _make_tracker(100)
|
||||
|
||||
def teardown(self) -> None:
|
||||
pass
|
||||
|
||||
def time_open_count(self) -> None:
|
||||
_ = self.tracker.open_count
|
||||
@@ -0,0 +1,102 @@
|
||||
# Async Resource Safety
|
||||
|
||||
## Overview
|
||||
|
||||
CleverAgents uses asynchronous resources throughout its stack — event
|
||||
subscriptions, LangGraph tasks, checkpoint file I/O, and reactive stream
|
||||
connections. The **async-cleanup** subsystem ensures that these resources
|
||||
are tracked, closed deterministically on shutdown, and that leaks are
|
||||
detected and logged.
|
||||
|
||||
## Core Component: `AsyncResourceTracker`
|
||||
|
||||
`cleveragents.core.async_cleanup.AsyncResourceTracker` is the central
|
||||
registry for any resource that implements the `AsyncResource` protocol
|
||||
(i.e.\ exposes an `async def close() -> None` method).
|
||||
|
||||
### Registration
|
||||
|
||||
```python
|
||||
from cleveragents.core.async_cleanup import AsyncResourceTracker
|
||||
|
||||
tracker = AsyncResourceTracker()
|
||||
tracker.register("db-pool", db_pool)
|
||||
tracker.register("event-queue", event_queue)
|
||||
```
|
||||
|
||||
- Names must be unique and non-empty.
|
||||
- Duplicate registrations raise `ValueError`.
|
||||
- `None` resources are rejected immediately.
|
||||
|
||||
### Shutdown
|
||||
|
||||
```python
|
||||
await tracker.close_all(timeout=30.0)
|
||||
```
|
||||
|
||||
`close_all()` iterates over every registered resource and awaits its
|
||||
`close()` coroutine with `asyncio.wait_for()`. Resources that exceed
|
||||
the deadline are logged as forced terminations and their names are
|
||||
collected in `tracker.timed_out_resources`.
|
||||
|
||||
`close_all()` is **idempotent** — calling it multiple times is safe.
|
||||
|
||||
### Leak Detection
|
||||
|
||||
If the tracker is garbage-collected without `close_all()` having been
|
||||
called, the `__del__` finalizer logs a warning for each unclosed
|
||||
resource **by name**. This makes it straightforward to identify leaks
|
||||
during development and in CI logs.
|
||||
|
||||
### Async Context Manager
|
||||
|
||||
```python
|
||||
async with AsyncResourceTracker() as tracker:
|
||||
tracker.register("conn", connection)
|
||||
# ... use connection ...
|
||||
# connection.close() is awaited automatically
|
||||
```
|
||||
|
||||
## Enhanced Bridge Cleanup
|
||||
|
||||
`cleveragents.langgraph.bridge.RxPyLangGraphBridge` now provides:
|
||||
|
||||
- **`cleanup_tasks_async(timeout)`** — cancels all in-flight asyncio
|
||||
tasks and awaits their completion within *timeout* seconds. Tasks
|
||||
that do not finish are logged as warnings.
|
||||
- **`cancel_task_with_reason(task, reason)`** — cancels a specific task
|
||||
and records the human-readable *reason* in
|
||||
`bridge.cancellation_reasons`.
|
||||
|
||||
The synchronous `cleanup_tasks()` remains for best-effort cleanup in
|
||||
`__del__`.
|
||||
|
||||
## Subscription Cleanup
|
||||
|
||||
`cleveragents.acp.events.AcpEventQueue.close()` removes all local
|
||||
subscriptions, clears the event buffer, and logs the count of
|
||||
subscriptions that were active.
|
||||
|
||||
## Checkpoint File Safety
|
||||
|
||||
`cleveragents.langgraph.state.StateManager.close()` marks the manager
|
||||
as closed, completes the underlying RxPY `BehaviorSubject`, and
|
||||
prevents further state updates. Checkpoint files written via
|
||||
`_save_checkpoint()` use `Path.write_text()`, which handles file-handle
|
||||
closing internally.
|
||||
|
||||
## Thread Safety
|
||||
|
||||
`AsyncResourceTracker.register()` and `close_all()` are protected by a
|
||||
`threading.Lock` so that resources can be registered from any thread
|
||||
without races.
|
||||
|
||||
## Design Decisions
|
||||
|
||||
| Decision | Rationale |
|
||||
|---|---|
|
||||
| Protocol-based `AsyncResource` | Structural typing avoids coupling to a specific base class. |
|
||||
| Per-resource timeout | One slow resource should not block the entire shutdown sequence. |
|
||||
| Idempotent `close_all` | Prevents double-close errors in complex shutdown paths. |
|
||||
| `__del__` leak warning | Best-effort; relies on CPython deterministic GC but degrades safely. |
|
||||
| Cancellation reason dict | Lightweight tracing for debugging cancelled tasks without heavy instrumentation. |
|
||||
@@ -0,0 +1,88 @@
|
||||
Feature: Async Resource Cleanup and Leak Prevention
|
||||
As a system operator
|
||||
I want async resources to be tracked and closed reliably
|
||||
So that the system does not leak connections, tasks, or subscriptions
|
||||
|
||||
Background:
|
||||
Given I have an async resource tracker
|
||||
|
||||
Scenario: Registering an async resource and closing it
|
||||
Given I have a mock async resource named "db-pool"
|
||||
When I register the resource with the tracker
|
||||
And I close all tracked resources
|
||||
Then the resource "db-pool" should be closed
|
||||
And the tracker should have zero open resources
|
||||
|
||||
Scenario: close_all awaits registered resources with timeout
|
||||
Given I have 3 mock async resources with varied close times
|
||||
When I close all tracked resources with a 5 second timeout
|
||||
Then all 3 resources should be closed
|
||||
And the tracker should have zero open resources
|
||||
|
||||
Scenario: Leaked resources are logged by name in finalizer
|
||||
Given I have a mock async resource named "leaky-conn"
|
||||
And I register the resource with the tracker
|
||||
When the tracker finalizer runs without close_all
|
||||
Then a warning should be logged mentioning "leaky-conn"
|
||||
|
||||
Scenario: Graceful cancellation awaits in-flight tasks before cleanup
|
||||
Given I have an async task tracked by the bridge
|
||||
When I request graceful cleanup with a 2 second timeout
|
||||
Then the task should be cancelled
|
||||
And the bridge should have no active tasks
|
||||
|
||||
Scenario: Time-bounded shutdown warns on forced termination
|
||||
Given I have a mock async resource named "slow-resource" that takes 5 seconds to close
|
||||
And I register the resource with the tracker
|
||||
When I close all tracked resources with a 0.1 second timeout
|
||||
Then a warning should be logged about forced termination
|
||||
And the tracker should report the timed-out resource
|
||||
|
||||
Scenario: Cancelled async jobs persist cancellation reason
|
||||
Given I have an async task tracked by the bridge
|
||||
When I cancel the task with reason "user-requested-shutdown"
|
||||
Then the cancellation reason should be "user-requested-shutdown"
|
||||
|
||||
Scenario: Checkpoint file handles are closed on cleanup
|
||||
Given I have a state manager with a temporary checkpoint directory
|
||||
When I save a checkpoint and then close the state manager
|
||||
Then the checkpoint file should exist and be readable
|
||||
And the state manager should be marked as closed
|
||||
|
||||
Scenario: Subscriptions are disposed on cleanup
|
||||
Given I have an ACP event queue with 3 active subscriptions
|
||||
When I close the event queue
|
||||
Then all subscriptions should be removed
|
||||
And the subscription count should be zero
|
||||
|
||||
Scenario: Tracker rejects duplicate resource names
|
||||
Given I have a mock async resource named "unique-res"
|
||||
And I register the resource with the tracker
|
||||
When I try to register another resource named "unique-res"
|
||||
Then a ValueError should be raised mentioning "unique-res"
|
||||
|
||||
Scenario: Tracker rejects empty resource name
|
||||
When I try to register a resource with an empty name
|
||||
Then a ValueError should be raised mentioning "name"
|
||||
|
||||
Scenario: Tracker rejects None resource
|
||||
When I try to register a None resource with name "valid-name"
|
||||
Then a ValueError should be raised mentioning "resource"
|
||||
|
||||
Scenario: Tracker context manager closes resources on exit
|
||||
Given I have a mock async resource named "ctx-resource"
|
||||
When I use the tracker as an async context manager and register the resource
|
||||
Then the resource "ctx-resource" should be closed after exiting the context
|
||||
|
||||
Scenario: close_all is idempotent
|
||||
Given I have a mock async resource named "once-resource"
|
||||
And I register the resource with the tracker
|
||||
When I close all tracked resources twice
|
||||
Then the resource "once-resource" should be closed exactly once
|
||||
And the tracker should have zero open resources
|
||||
|
||||
Scenario: Enhanced bridge cleanup awaits tasks with timeout
|
||||
Given I have a bridge with 2 slow async tasks
|
||||
When I run enhanced cleanup with a 2 second timeout
|
||||
Then the bridge should have no active tasks
|
||||
And completed tasks should be logged
|
||||
@@ -0,0 +1,421 @@
|
||||
"""Step definitions for async resource cleanup and leak prevention."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from behave import given, then, when
|
||||
|
||||
from cleveragents.acp.events import AcpEventQueue
|
||||
from cleveragents.core.async_cleanup import AsyncResourceTracker
|
||||
from cleveragents.langgraph.state import GraphState, StateManager
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers — lightweight mock async resource (lives in steps, not src/)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MockAsyncResource:
|
||||
"""Test double for an async-closable resource."""
|
||||
|
||||
def __init__(self, name: str, close_delay: float = 0.0) -> None:
|
||||
self.name = name
|
||||
self.close_delay = close_delay
|
||||
self.closed = False
|
||||
self.close_count = 0
|
||||
|
||||
async def close(self) -> None:
|
||||
if self.close_delay > 0:
|
||||
await asyncio.sleep(self.close_delay)
|
||||
self.closed = True
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Background
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@given("I have an async resource tracker")
|
||||
def step_create_tracker(context):
|
||||
context.tracker = AsyncResourceTracker()
|
||||
context.resources = {}
|
||||
context.log_handler = _CapturingHandler()
|
||||
logging.getLogger("cleveragents.core.async_cleanup").addHandler(context.log_handler)
|
||||
context.warnings_logged = context.log_handler.records
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Givens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@given('I have a mock async resource named "{name}"')
|
||||
def step_create_mock_resource(context, name):
|
||||
resource = MockAsyncResource(name)
|
||||
context.resources[name] = resource
|
||||
context.current_resource = resource
|
||||
context.current_resource_name = name
|
||||
|
||||
|
||||
@given(
|
||||
'I have a mock async resource named "{name}" that takes {seconds:g} seconds to close'
|
||||
)
|
||||
def step_create_slow_resource(context, name, seconds):
|
||||
resource = MockAsyncResource(name, close_delay=seconds)
|
||||
context.resources[name] = resource
|
||||
context.current_resource = resource
|
||||
context.current_resource_name = name
|
||||
|
||||
|
||||
@given("I have {count:d} mock async resources with varied close times")
|
||||
def step_create_multiple_resources(context, count):
|
||||
context.multi_resources = []
|
||||
for i in range(count):
|
||||
name = f"resource-{i}"
|
||||
resource = MockAsyncResource(name, close_delay=0.01 * (i + 1))
|
||||
context.resources[name] = resource
|
||||
context.multi_resources.append((name, resource))
|
||||
context.tracker.register(name, resource)
|
||||
|
||||
|
||||
@given("I register the resource with the tracker")
|
||||
def step_register_resource(context):
|
||||
context.tracker.register(context.current_resource_name, context.current_resource)
|
||||
|
||||
|
||||
@given("I have an async task tracked by the bridge")
|
||||
def step_create_bridge_task(context):
|
||||
from cleveragents.langgraph.bridge import RxPyLangGraphBridge
|
||||
from cleveragents.reactive.stream_router import ReactiveStreamRouter
|
||||
|
||||
router = ReactiveStreamRouter()
|
||||
context.bridge = RxPyLangGraphBridge(router)
|
||||
context.cancellation_reasons = {}
|
||||
|
||||
async def _long_running():
|
||||
try:
|
||||
await asyncio.sleep(100)
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def _setup():
|
||||
task = asyncio.create_task(_long_running())
|
||||
context.bridge._active_tasks.add(task)
|
||||
context.bridge_task = task
|
||||
|
||||
loop.run_until_complete(_setup())
|
||||
context.bridge_loop = loop
|
||||
|
||||
|
||||
@given("I have a state manager with a temporary checkpoint directory")
|
||||
def step_create_state_manager(context):
|
||||
context.temp_dir = tempfile.mkdtemp()
|
||||
context.checkpoint_dir = Path(context.temp_dir)
|
||||
context.state_manager = StateManager(
|
||||
initial_state=GraphState(),
|
||||
checkpoint_dir=context.checkpoint_dir,
|
||||
)
|
||||
|
||||
|
||||
@given("I have an ACP event queue with {count:d} active subscriptions")
|
||||
def step_create_event_queue(context, count):
|
||||
context.event_queue = AcpEventQueue()
|
||||
context.subscription_ids = []
|
||||
for _ in range(count):
|
||||
sub_id = context.event_queue.subscribe_local(lambda _evt: None)
|
||||
context.subscription_ids.append(sub_id)
|
||||
|
||||
|
||||
@given("I have a bridge with {count:d} slow async tasks")
|
||||
def step_create_bridge_with_slow_tasks(context, count):
|
||||
from cleveragents.langgraph.bridge import RxPyLangGraphBridge
|
||||
from cleveragents.reactive.stream_router import ReactiveStreamRouter
|
||||
|
||||
router = ReactiveStreamRouter()
|
||||
context.bridge = RxPyLangGraphBridge(router)
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def _setup():
|
||||
for _ in range(count):
|
||||
|
||||
async def _slow():
|
||||
try:
|
||||
await asyncio.sleep(100)
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
task = asyncio.create_task(_slow())
|
||||
context.bridge._active_tasks.add(task)
|
||||
|
||||
loop.run_until_complete(_setup())
|
||||
context.bridge_loop = loop
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Whens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@when("I register the resource with the tracker")
|
||||
def step_when_register_resource(context):
|
||||
context.tracker.register(context.current_resource_name, context.current_resource)
|
||||
|
||||
|
||||
@when("I close all tracked resources")
|
||||
def step_close_all(context):
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.run_until_complete(context.tracker.close_all())
|
||||
loop.close()
|
||||
|
||||
|
||||
@when("I close all tracked resources with a {timeout:g} second timeout")
|
||||
def step_close_all_with_timeout(context, timeout):
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.run_until_complete(context.tracker.close_all(timeout=timeout))
|
||||
loop.close()
|
||||
|
||||
|
||||
@when("the tracker finalizer runs without close_all")
|
||||
def step_run_finalizer(context):
|
||||
context.tracker._warn_unclosed()
|
||||
|
||||
|
||||
@when("I request graceful cleanup with a {timeout:g} second timeout")
|
||||
def step_graceful_bridge_cleanup(context, timeout):
|
||||
loop = context.bridge_loop
|
||||
|
||||
async def _cleanup():
|
||||
await context.bridge.cleanup_tasks_async(timeout=timeout)
|
||||
|
||||
loop.run_until_complete(_cleanup())
|
||||
|
||||
|
||||
@when('I cancel the task with reason "{reason}"')
|
||||
def step_cancel_with_reason(context, reason):
|
||||
loop = context.bridge_loop
|
||||
|
||||
async def _cancel():
|
||||
await context.bridge.cancel_task_with_reason(context.bridge_task, reason)
|
||||
|
||||
loop.run_until_complete(_cancel())
|
||||
context.cancel_reason = reason
|
||||
|
||||
|
||||
@when("I save a checkpoint and then close the state manager")
|
||||
def step_save_and_close_state_manager(context):
|
||||
context.state_manager.update_state(
|
||||
{"messages": [{"role": "user", "content": "test"}]}
|
||||
)
|
||||
context.state_manager._save_checkpoint()
|
||||
context.state_manager.close()
|
||||
|
||||
|
||||
@when("I close the event queue")
|
||||
def step_close_event_queue(context):
|
||||
context.event_queue.close()
|
||||
|
||||
|
||||
@when('I try to register another resource named "{name}"')
|
||||
def step_try_register_duplicate(context, name):
|
||||
try:
|
||||
dup_resource = MockAsyncResource(name)
|
||||
context.tracker.register(name, dup_resource)
|
||||
context.lsp_error = None
|
||||
except ValueError as exc:
|
||||
context.lsp_error = exc
|
||||
|
||||
|
||||
@when("I try to register a resource with an empty name")
|
||||
def step_try_register_empty_name(context):
|
||||
try:
|
||||
context.tracker.register("", MockAsyncResource("empty"))
|
||||
context.lsp_error = None
|
||||
except ValueError as exc:
|
||||
context.lsp_error = exc
|
||||
|
||||
|
||||
@when('I try to register a None resource with name "{name}"')
|
||||
def step_try_register_none_resource(context, name):
|
||||
try:
|
||||
context.tracker.register(name, None) # type: ignore[arg-type]
|
||||
context.lsp_error = None
|
||||
except ValueError as exc:
|
||||
context.lsp_error = exc
|
||||
|
||||
|
||||
@when("I use the tracker as an async context manager and register the resource")
|
||||
def step_use_context_manager(context):
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def _run():
|
||||
async with context.tracker as tracker:
|
||||
tracker.register(context.current_resource_name, context.current_resource)
|
||||
|
||||
loop.run_until_complete(_run())
|
||||
loop.close()
|
||||
|
||||
|
||||
@when("I close all tracked resources twice")
|
||||
def step_close_all_twice(context):
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.run_until_complete(context.tracker.close_all())
|
||||
loop.run_until_complete(context.tracker.close_all())
|
||||
loop.close()
|
||||
|
||||
|
||||
@when("I run enhanced cleanup with a {timeout:g} second timeout")
|
||||
def step_enhanced_bridge_cleanup(context, timeout):
|
||||
loop = context.bridge_loop
|
||||
|
||||
async def _cleanup():
|
||||
await context.bridge.cleanup_tasks_async(timeout=timeout)
|
||||
|
||||
loop.run_until_complete(_cleanup())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@then('the resource "{name}" should be closed')
|
||||
def step_assert_resource_closed(context, name):
|
||||
resource = context.resources[name]
|
||||
assert resource.closed, f"Resource '{name}' was not closed"
|
||||
|
||||
|
||||
@then("the tracker should have zero open resources")
|
||||
def step_assert_zero_open(context):
|
||||
assert context.tracker.open_count == 0, (
|
||||
f"Expected 0 open resources, got {context.tracker.open_count}"
|
||||
)
|
||||
|
||||
|
||||
@then("all {count:d} resources should be closed")
|
||||
def step_assert_all_closed(context, count):
|
||||
closed_count = sum(1 for _, r in context.multi_resources if r.closed)
|
||||
assert closed_count == count, f"Expected {count} closed, got {closed_count}"
|
||||
|
||||
|
||||
@then('a warning should be logged mentioning "{text}"')
|
||||
def step_assert_warning_logged(context, text):
|
||||
warnings = [r for r in context.log_handler.records if r.levelno >= logging.WARNING]
|
||||
matching = [r for r in warnings if text in r.getMessage()]
|
||||
assert matching, (
|
||||
f"No warning mentioning '{text}' found. "
|
||||
f"Warnings: {[r.getMessage() for r in warnings]}"
|
||||
)
|
||||
|
||||
|
||||
@then("the task should be cancelled")
|
||||
def step_assert_task_cancelled(context):
|
||||
task = context.bridge_task
|
||||
assert task.done(), "Task should be done after cleanup"
|
||||
|
||||
|
||||
@then("the bridge should have no active tasks")
|
||||
def step_assert_no_active_tasks(context):
|
||||
assert len(context.bridge._active_tasks) == 0, (
|
||||
f"Expected 0 active tasks, got {len(context.bridge._active_tasks)}"
|
||||
)
|
||||
|
||||
|
||||
@then("a warning should be logged about forced termination")
|
||||
def step_assert_forced_termination_warning(context):
|
||||
warnings = [r for r in context.log_handler.records if r.levelno >= logging.WARNING]
|
||||
matching = [
|
||||
r
|
||||
for r in warnings
|
||||
if "timeout" in r.getMessage().lower() or "forced" in r.getMessage().lower()
|
||||
]
|
||||
assert matching, (
|
||||
f"No forced termination warning found. "
|
||||
f"Warnings: {[r.getMessage() for r in warnings]}"
|
||||
)
|
||||
|
||||
|
||||
@then("the tracker should report the timed-out resource")
|
||||
def step_assert_timeout_reported(context):
|
||||
assert len(context.tracker.timed_out_resources) > 0, (
|
||||
"Expected at least one timed-out resource"
|
||||
)
|
||||
|
||||
|
||||
@then('the cancellation reason should be "{reason}"')
|
||||
def step_assert_cancellation_reason(context, reason):
|
||||
assert context.bridge.cancellation_reasons.get(context.bridge_task) == reason, (
|
||||
f"Expected reason '{reason}', "
|
||||
f"got '{context.bridge.cancellation_reasons.get(context.bridge_task)}'"
|
||||
)
|
||||
|
||||
|
||||
@then("the checkpoint file should exist and be readable")
|
||||
def step_assert_checkpoint_exists(context):
|
||||
checkpoints = list(context.checkpoint_dir.glob("checkpoint_*.json"))
|
||||
assert len(checkpoints) >= 1, "No checkpoint files found"
|
||||
content = checkpoints[0].read_text(encoding="utf-8")
|
||||
assert len(content) > 0, "Checkpoint file is empty"
|
||||
|
||||
|
||||
@then("the state manager should be marked as closed")
|
||||
def step_assert_state_manager_closed(context):
|
||||
assert context.state_manager.is_closed, "StateManager should be closed"
|
||||
|
||||
|
||||
@then("all subscriptions should be removed")
|
||||
def step_assert_subscriptions_removed(context):
|
||||
assert len(context.event_queue._subscriptions) == 0, (
|
||||
f"Expected 0 subscriptions, got {len(context.event_queue._subscriptions)}"
|
||||
)
|
||||
|
||||
|
||||
@then("the subscription count should be zero")
|
||||
def step_assert_subscription_count_zero(context):
|
||||
assert len(context.event_queue._subscriptions) == 0
|
||||
|
||||
|
||||
@then('the resource "{name}" should be closed after exiting the context')
|
||||
def step_assert_closed_after_context(context, name):
|
||||
resource = context.resources[name]
|
||||
assert resource.closed, f"Resource '{name}' should be closed after context exit"
|
||||
|
||||
|
||||
@then('the resource "{name}" should be closed exactly once')
|
||||
def step_assert_closed_once(context, name):
|
||||
resource = context.resources[name]
|
||||
assert resource.closed, f"Resource '{name}' was not closed"
|
||||
assert resource.close_count == 1, (
|
||||
f"Resource '{name}' was closed {resource.close_count} times, expected 1"
|
||||
)
|
||||
|
||||
|
||||
@then("completed tasks should be logged")
|
||||
def step_assert_tasks_logged(context):
|
||||
infos = [r for r in context.log_handler.records if r.levelno >= logging.INFO]
|
||||
has_cleanup_log = any(
|
||||
"cleanup" in r.getMessage().lower() or "task" in r.getMessage().lower()
|
||||
for r in infos
|
||||
)
|
||||
# The cleanup operation should have been logged, and tasks cleared
|
||||
assert has_cleanup_log or len(context.bridge._active_tasks) == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Logging capture helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _CapturingHandler(logging.Handler):
|
||||
"""Captures log records for assertion in tests."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.records: list[logging.LogRecord] = []
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
self.records.append(record)
|
||||
@@ -0,0 +1,135 @@
|
||||
*** Settings ***
|
||||
Documentation Integration smoke tests for async resource cleanup (#321).
|
||||
... Validates that AsyncResourceTracker, enhanced bridge cleanup,
|
||||
... AcpEventQueue.close(), and StateManager.close() work end-to-end
|
||||
... by executing small Python driver scripts in a subprocess.
|
||||
|
||||
Library OperatingSystem
|
||||
Library Process
|
||||
Library String
|
||||
Resource ${CURDIR}/common.resource
|
||||
|
||||
Test Setup Setup Async Cleanup Test Environment
|
||||
Test Teardown Cleanup Async Cleanup Test Environment
|
||||
|
||||
*** Variables ***
|
||||
${PYTHON} python
|
||||
${WORKSPACE_ROOT} ${CURDIR}/..
|
||||
${TEST_FILE} ${EMPTY}
|
||||
${TEST_OUTPUT} ${EMPTY}
|
||||
${TIMEOUT} 30s
|
||||
|
||||
*** Test Cases ***
|
||||
Test AsyncResourceTracker Register And Close
|
||||
[Documentation] Register a resource and close it via close_all
|
||||
${script} = Catenate SEPARATOR=\n
|
||||
... import sys, asyncio
|
||||
... sys.path.insert(0, '${WORKSPACE_ROOT}/src')
|
||||
... from cleveragents.core.async_cleanup import AsyncResourceTracker
|
||||
...
|
||||
... class FakeResource:
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}def __init__(self): self.closed = False
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}async def close(self): self.closed = True
|
||||
...
|
||||
... async def main():
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}tracker = AsyncResourceTracker()
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}res = FakeResource()
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}tracker.register("test-res", res)
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}assert tracker.open_count == 1
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}await tracker.close_all()
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}assert res.closed
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}assert tracker.open_count == 0
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}print("PASS: register and close", flush=True)
|
||||
...
|
||||
... asyncio.run(main())
|
||||
Create File ${TEST_FILE} ${script}
|
||||
${result} = Run Process ${PYTHON} ${TEST_FILE}
|
||||
... stdout=${TEST_OUTPUT} stderr=STDOUT timeout=${TIMEOUT}
|
||||
Should Be Equal As Integers ${result.rc} 0
|
||||
${output} = Get File ${TEST_OUTPUT}
|
||||
Should Contain ${output} PASS: register and close
|
||||
|
||||
Test AsyncResourceTracker Timeout Warning
|
||||
[Documentation] Verify forced-termination warning for slow resource
|
||||
${script} = Catenate SEPARATOR=\n
|
||||
... import sys, asyncio, logging
|
||||
... sys.path.insert(0, '${WORKSPACE_ROOT}/src')
|
||||
... from cleveragents.core.async_cleanup import AsyncResourceTracker
|
||||
...
|
||||
... handler = logging.StreamHandler(sys.stdout)
|
||||
... handler.setLevel(logging.WARNING)
|
||||
... logging.getLogger("cleveragents.core.async_cleanup").addHandler(handler)
|
||||
...
|
||||
... class SlowResource:
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}async def close(self): await asyncio.sleep(10)
|
||||
...
|
||||
... async def main():
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}tracker = AsyncResourceTracker()
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}tracker.register("slow", SlowResource())
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}await tracker.close_all(timeout=0.05)
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}assert len(tracker.timed_out_resources) == 1
|
||||
... ${SPACE}${SPACE}${SPACE}${SPACE}print("PASS: timeout warning", flush=True)
|
||||
...
|
||||
... asyncio.run(main())
|
||||
Create File ${TEST_FILE} ${script}
|
||||
${result} = Run Process ${PYTHON} ${TEST_FILE}
|
||||
... stdout=${TEST_OUTPUT} stderr=STDOUT timeout=${TIMEOUT}
|
||||
Should Be Equal As Integers ${result.rc} 0
|
||||
${output} = Get File ${TEST_OUTPUT}
|
||||
Should Contain ${output} PASS: timeout warning
|
||||
|
||||
Test AcpEventQueue Close
|
||||
[Documentation] Verify AcpEventQueue.close() removes subscriptions
|
||||
${script} = Catenate SEPARATOR=\n
|
||||
... import sys
|
||||
... sys.path.insert(0, '${WORKSPACE_ROOT}/src')
|
||||
... from cleveragents.acp.events import AcpEventQueue
|
||||
...
|
||||
... q = AcpEventQueue()
|
||||
... q.subscribe_local(lambda e: None)
|
||||
... q.subscribe_local(lambda e: None)
|
||||
... assert len(q._subscriptions) == 2
|
||||
... q.close()
|
||||
... assert len(q._subscriptions) == 0
|
||||
... print("PASS: event queue close", flush=True)
|
||||
Create File ${TEST_FILE} ${script}
|
||||
${result} = Run Process ${PYTHON} ${TEST_FILE}
|
||||
... stdout=${TEST_OUTPUT} stderr=STDOUT timeout=${TIMEOUT}
|
||||
Should Be Equal As Integers ${result.rc} 0
|
||||
${output} = Get File ${TEST_OUTPUT}
|
||||
Should Contain ${output} PASS: event queue close
|
||||
|
||||
Test StateManager Close
|
||||
[Documentation] Verify StateManager.close() marks manager as closed
|
||||
${script} = Catenate SEPARATOR=\n
|
||||
... import sys, tempfile
|
||||
... sys.path.insert(0, '${WORKSPACE_ROOT}/src')
|
||||
... from pathlib import Path
|
||||
... from cleveragents.langgraph.state import StateManager, GraphState
|
||||
...
|
||||
... d = tempfile.mkdtemp()
|
||||
... mgr = StateManager(initial_state=GraphState(), checkpoint_dir=Path(d))
|
||||
... assert not mgr.is_closed
|
||||
... mgr.close()
|
||||
... assert mgr.is_closed
|
||||
... print("PASS: state manager close", flush=True)
|
||||
Create File ${TEST_FILE} ${script}
|
||||
${result} = Run Process ${PYTHON} ${TEST_FILE}
|
||||
... stdout=${TEST_OUTPUT} stderr=STDOUT timeout=${TIMEOUT}
|
||||
Should Be Equal As Integers ${result.rc} 0
|
||||
${output} = Get File ${TEST_OUTPUT}
|
||||
Should Contain ${output} PASS: state manager close
|
||||
|
||||
*** Keywords ***
|
||||
Setup Async Cleanup Test Environment
|
||||
[Documentation] Create temp dir for test scripts
|
||||
${temp_dir} = Evaluate tempfile.mkdtemp() modules=tempfile
|
||||
Set Test Variable ${TEMP_DIR} ${temp_dir}
|
||||
Set Test Variable ${TEST_FILE} ${temp_dir}/async_cleanup_test.py
|
||||
Set Test Variable ${TEST_OUTPUT} ${temp_dir}/async_cleanup_output.txt
|
||||
|
||||
Cleanup Async Cleanup Test Environment
|
||||
[Documentation] Remove temp dir
|
||||
Run Keyword And Ignore Error Remove File ${TEST_FILE}
|
||||
Run Keyword And Ignore Error Remove File ${TEST_OUTPUT}
|
||||
Run Keyword And Ignore Error Remove Directory ${TEMP_DIR} recursive=True
|
||||
@@ -78,6 +78,18 @@ class AcpEventQueue:
|
||||
raise ValueError("limit must be a positive integer")
|
||||
return list(self._events[-limit:])
|
||||
|
||||
def close(self) -> None:
|
||||
"""Remove all subscriptions and clear the event queue.
|
||||
|
||||
Logs the number of subscriptions that were active at the time of
|
||||
closing. This method is safe to call multiple times.
|
||||
"""
|
||||
count = len(self._subscriptions)
|
||||
self._subscriptions.clear()
|
||||
self._events.clear()
|
||||
if count:
|
||||
logger.info("acp.event_queue.closed", subscription_count=count)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Remote stub (raises)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Async resource tracker for deterministic cleanup and leak detection.
|
||||
|
||||
Provides a central ``AsyncResourceTracker`` that manages the lifecycle of
|
||||
asynchronous resources (connections, tasks, subscriptions). Resources
|
||||
conforming to the ``AsyncResource`` protocol are registered by name and
|
||||
awaited on shutdown with a configurable timeout. Any resource that is
|
||||
still open when the tracker is garbage-collected is logged as a leak.
|
||||
|
||||
Implements issue #321: fix(security): close async resources and leaks.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Protocol — any object with an async ``close()`` qualifies
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AsyncResource(Protocol):
|
||||
"""Protocol for asynchronous resources that can be closed."""
|
||||
|
||||
async def close(self) -> None: ... # pragma: no cover
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tracker
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class AsyncResourceTracker:
|
||||
"""Thread-safe registry for async resources with deterministic cleanup.
|
||||
|
||||
Resources are registered by a unique *name* and closed in bulk via
|
||||
:meth:`close_all`, which respects *timeout* and logs forced
|
||||
terminations. The tracker also acts as an async context manager.
|
||||
|
||||
Attributes:
|
||||
timed_out_resources: Names of resources that exceeded the timeout
|
||||
during the most recent :meth:`close_all` invocation.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._resources: dict[str, AsyncResource] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._closed = False
|
||||
self.timed_out_resources: list[str] = []
|
||||
|
||||
# -- Registration -------------------------------------------------------
|
||||
|
||||
def register(self, name: str, resource: AsyncResource) -> None:
|
||||
"""Register *resource* under a unique *name*.
|
||||
|
||||
Args:
|
||||
name: Non-empty, unique identifier for the resource.
|
||||
resource: An object satisfying the :class:`AsyncResource` protocol.
|
||||
|
||||
Raises:
|
||||
ValueError: If *name* is empty, *resource* is ``None``, or *name*
|
||||
is already registered.
|
||||
"""
|
||||
if not name:
|
||||
raise ValueError("name must be a non-empty string")
|
||||
if resource is None:
|
||||
raise ValueError("resource must not be None")
|
||||
with self._lock:
|
||||
if name in self._resources:
|
||||
raise ValueError(f"Resource '{name}' is already registered")
|
||||
self._resources[name] = resource
|
||||
logger.debug("Registered async resource '%s'", name)
|
||||
|
||||
# -- Bulk close ---------------------------------------------------------
|
||||
|
||||
async def close_all(self, timeout: float = 30.0) -> None:
|
||||
"""Close every registered resource within *timeout* seconds.
|
||||
|
||||
Resources whose ``close()`` exceeds the deadline are logged as
|
||||
forced terminations and their names are appended to
|
||||
:attr:`timed_out_resources`.
|
||||
|
||||
This method is idempotent — calling it on an already-closed tracker
|
||||
is a no-op.
|
||||
|
||||
Args:
|
||||
timeout: Maximum wall-clock seconds to wait per resource.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
snapshot = dict(self._resources)
|
||||
self._resources.clear()
|
||||
|
||||
self.timed_out_resources = []
|
||||
|
||||
for name, resource in snapshot.items():
|
||||
try:
|
||||
await asyncio.wait_for(resource.close(), timeout=timeout)
|
||||
logger.info(
|
||||
"Closed async resource '%s'",
|
||||
name,
|
||||
)
|
||||
except TimeoutError:
|
||||
self.timed_out_resources.append(name)
|
||||
logger.warning(
|
||||
"Forced termination: resource '%s' did not close "
|
||||
"within timeout of %.1f s",
|
||||
name,
|
||||
timeout,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Error closing async resource '%s'",
|
||||
name,
|
||||
)
|
||||
|
||||
# -- Query --------------------------------------------------------------
|
||||
|
||||
@property
|
||||
def open_count(self) -> int:
|
||||
"""Return the number of currently registered (open) resources."""
|
||||
with self._lock:
|
||||
return len(self._resources)
|
||||
|
||||
# -- Finalizer / leak warning -------------------------------------------
|
||||
|
||||
def _warn_unclosed(self) -> None:
|
||||
"""Log a warning for every resource that was never closed."""
|
||||
with self._lock:
|
||||
names = list(self._resources.keys())
|
||||
for name in names:
|
||||
logger.warning(
|
||||
"Async resource '%s' was never closed (potential leak)",
|
||||
name,
|
||||
)
|
||||
|
||||
def __del__(self) -> None:
|
||||
"""Best-effort leak detection during garbage collection."""
|
||||
if self._resources:
|
||||
self._warn_unclosed()
|
||||
|
||||
# -- Async context manager ----------------------------------------------
|
||||
|
||||
async def __aenter__(self) -> AsyncResourceTracker:
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_val: BaseException | None,
|
||||
exc_tb: object,
|
||||
) -> bool:
|
||||
await self.close_all()
|
||||
return False
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AsyncResource",
|
||||
"AsyncResourceTracker",
|
||||
]
|
||||
@@ -29,17 +29,79 @@ class RxPyLangGraphBridge:
|
||||
self.logger = logging.getLogger(__name__)
|
||||
self.graphs: dict[str, LangGraph] = {}
|
||||
self._active_tasks: set[asyncio.Task[Any]] = set()
|
||||
self.cancellation_reasons: dict[asyncio.Task[Any], str] = {}
|
||||
self._register_langgraph_operators()
|
||||
|
||||
def __del__(self):
|
||||
def __del__(self) -> None:
|
||||
self.cleanup_tasks()
|
||||
|
||||
def cleanup_tasks(self) -> None:
|
||||
"""Cancel all active tasks immediately (sync, best-effort)."""
|
||||
for task in list(self._active_tasks):
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
self._active_tasks.clear()
|
||||
|
||||
async def cleanup_tasks_async(self, timeout: float = 5.0) -> None:
|
||||
"""Cancel active tasks and await their completion within *timeout*.
|
||||
|
||||
Tasks that do not finish within the deadline are logged and
|
||||
forcefully discarded.
|
||||
|
||||
Args:
|
||||
timeout: Maximum seconds to wait for each task to finish.
|
||||
"""
|
||||
tasks = list(self._active_tasks)
|
||||
if not tasks:
|
||||
return
|
||||
|
||||
for task in tasks:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
done, pending = await asyncio.wait(tasks, timeout=timeout)
|
||||
|
||||
for task in done:
|
||||
self.logger.info(
|
||||
"Task completed during cleanup: %s",
|
||||
task.get_name(),
|
||||
)
|
||||
|
||||
for task in pending:
|
||||
self.logger.warning(
|
||||
"Task did not complete within %.1f s timeout: %s",
|
||||
timeout,
|
||||
task.get_name(),
|
||||
)
|
||||
|
||||
self._active_tasks.clear()
|
||||
|
||||
async def cancel_task_with_reason(
|
||||
self,
|
||||
task: asyncio.Task[Any],
|
||||
reason: str,
|
||||
) -> None:
|
||||
"""Cancel *task* and record the *reason* for the cancellation.
|
||||
|
||||
Args:
|
||||
task: The asyncio task to cancel.
|
||||
reason: Human-readable explanation for the cancellation.
|
||||
|
||||
Raises:
|
||||
ValueError: If *task* is ``None`` or *reason* is empty.
|
||||
"""
|
||||
if task is None:
|
||||
raise ValueError("task must not be None")
|
||||
if not reason:
|
||||
raise ValueError("reason must be a non-empty string")
|
||||
self.cancellation_reasons[task] = reason
|
||||
task.cancel()
|
||||
self.logger.info(
|
||||
"Cancelled task %s with reason: %s",
|
||||
task.get_name(),
|
||||
reason,
|
||||
)
|
||||
|
||||
def _run_async_safely(self, coro: Any) -> Any:
|
||||
return coro
|
||||
|
||||
|
||||
@@ -99,6 +99,7 @@ class StateManager: # pylint: disable=too-many-instance-attributes
|
||||
self.max_history_size = 100
|
||||
self.checkpoint_interval = 10
|
||||
self.update_count = 0
|
||||
self.is_closed = False
|
||||
if self.checkpoint_dir:
|
||||
self.checkpoint_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@@ -179,3 +180,15 @@ class StateManager: # pylint: disable=too-many-instance-attributes
|
||||
self.update_count = 0
|
||||
self.history.clear()
|
||||
self.state_stream.on_next(self.state)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Mark this manager as closed and complete the state stream.
|
||||
|
||||
After calling ``close()``, no further state updates should be
|
||||
performed. Any underlying checkpoint resources are released.
|
||||
"""
|
||||
if self.is_closed:
|
||||
return
|
||||
self.is_closed = True
|
||||
self.state_stream.on_completed()
|
||||
self.logger.info("StateManager closed")
|
||||
|
||||
Reference in New Issue
Block a user