Files
hermes-agent/tests/agent/test_project_memory_switch_order.py

858 lines
34 KiB
Python

"""Manager lifecycle ordering, independent of any project-memory provider."""
import contextvars
import logging
import threading
from concurrent.futures import ThreadPoolExecutor, wait
import pytest
from agent.memory_manager import MemoryManager
from agent.memory_provider import MemoryProvider
class _StatefulProvider(MemoryProvider):
@property
def name(self):
return "builtin"
def __init__(self):
self.session_id = "A"
self.buffer = ["A buffer"]
self.extractions = []
self.switches = []
def is_available(self):
return True
def initialize(self, session_id, **kwargs):
self.session_id = session_id
def get_tool_schemas(self):
return []
def on_session_end(self, messages):
# Deliberately consult mutable binding at the write, like legacy providers.
self.extractions.append((self.session_id, list(messages)))
self.buffer.clear()
def on_session_switch(self, new_session_id, **kwargs):
self.session_id = new_session_id
self.buffer = [new_session_id + " buffer"]
self.switches.append((new_session_id, kwargs))
@pytest.mark.parametrize("targets,reason", [(('C',), 'resume'), (('C', 'D'), 'new_session')])
def test_queued_boundary_cannot_extract_or_rebind_after_newer_sync_switch(targets, reason, caplog):
# Approved FIFO contract: busy switches wait their turn, without dropping end.
provider = _StatefulProvider()
manager = MemoryManager()
manager.add_provider(provider)
entered = threading.Event()
release = threading.Event()
def occupy_worker():
entered.set()
assert release.wait(10)
with ThreadPoolExecutor(max_workers=1) as executor:
manager._sync_executor = executor
blocker = executor.submit(occupy_worker)
try:
assert entered.wait(5)
manager.commit_session_boundary_async(
[{"role": "user", "content": "A transcript"}],
new_session_id="B", parent_session_id="A",
)
for target in targets:
# The no-history /new path is synchronous, as is /resume.
manager.on_session_switch(target, reset=reason == "new_session", reason=reason)
assert provider.session_id == "A"
with caplog.at_level(logging.WARNING, logger="agent.memory_manager"):
release.set()
blocker.result(timeout=5)
assert manager.flush_pending(timeout=5)
assert provider.session_id == targets[-1]
assert provider.extractions == [("A", [{"role": "user", "content": "A transcript"}])]
assert provider.buffer == [targets[-1] + " buffer"]
assert [sid for sid, _ in provider.switches] == ["B", *targets]
finally:
release.set()
@pytest.mark.parametrize("end_raises", [False, True])
def test_active_extraction_owns_old_state_until_latest_switch_without_blocking(end_raises):
entered = threading.Event()
release = threading.Event()
profile = contextvars.ContextVar("test_boundary_profile", default="A profile")
switch_profiles = []
class _BlockingProvider(_StatefulProvider):
def on_session_switch(self, new_session_id, **kwargs):
switch_profiles.append(profile.get())
super().on_session_switch(new_session_id, **kwargs)
def on_session_end(self, messages):
entered.set()
assert release.wait(10)
super().on_session_end(messages)
if end_raises:
raise RuntimeError("extraction failed after touching old state")
provider = _BlockingProvider()
manager = MemoryManager()
manager.add_provider(provider)
messages = [{"role": "user", "content": "A transcript"}]
with ThreadPoolExecutor(max_workers=1) as executor, ThreadPoolExecutor(max_workers=1) as caller:
manager._sync_executor = executor
try:
caller.submit(manager.commit_session_boundary_async, messages,
new_session_id="B", parent_session_id="A").result(timeout=5)
assert entered.wait(5)
# Completion while extraction is still gated proves no caller waits on the LLM.
caller.submit(manager.on_session_switch, "C", reason="resume").result(timeout=5)
def switch_latest():
token = profile.set("D profile")
try:
manager.on_session_switch("D", reset=True, reason="new_session")
finally:
profile.reset(token)
caller.submit(switch_latest).result(timeout=5)
assert provider.session_id == "A"
assert provider.buffer == ["A buffer"]
release.set()
assert manager.flush_pending(timeout=5)
assert provider.extractions == [("A", messages)]
assert provider.session_id == "D"
assert provider.buffer == ["D buffer"]
assert [sid for sid, _ in provider.switches] == ["B", "C", "D"]
assert switch_profiles == ["A profile", "A profile", "D profile"]
finally:
release.set()
@pytest.mark.parametrize("sync_between", [False, True])
def test_async_boundaries_keep_fifo_transcript_ownership(sync_between, caplog):
provider = _StatefulProvider()
manager = MemoryManager()
manager.add_provider(provider)
entered = threading.Event()
release = threading.Event()
first = [{"role": "user", "content": "A transcript"}]
second = [{"role": "user", "content": "C transcript" if sync_between else "B transcript"}]
def occupy_worker():
entered.set()
assert release.wait(10)
with ThreadPoolExecutor(max_workers=1) as executor:
manager._sync_executor = executor
blocker = executor.submit(occupy_worker)
try:
assert entered.wait(5)
manager.commit_session_boundary_async(first, new_session_id="B", parent_session_id="A")
if sync_between:
manager.on_session_switch("C", parent_session_id="B", rewound=True)
manager.commit_session_boundary_async(
second, new_session_id="D", parent_session_id="C" if sync_between else "B",
)
release.set()
blocker.result(timeout=5)
assert manager.flush_pending(timeout=5)
expected = [("A", first), ("C" if sync_between else "B", second)]
assert provider.extractions == expected
assert provider.session_id == "D"
assert provider.buffer == ["D buffer"]
assert [sid for sid, _ in provider.switches] == (["B", "C", "D"] if sync_between else ["B", "D"])
if sync_between:
assert provider.switches[1][1]["rewound"] is True
finally:
release.set()
def test_reversed_worker_submission_keeps_intent_order_and_caller_context(monkeypatch):
first_submit = threading.Event()
release_submit = threading.Event()
profile = contextvars.ContextVar("submission_profile", default="default")
observed_profiles = []
class _Provider(_StatefulProvider):
def on_session_end(self, messages):
observed_profiles.append(profile.get())
super().on_session_end(messages)
provider = _Provider()
manager = MemoryManager()
manager.add_provider(provider)
wake_name = "_wake_background" if hasattr(manager, "_wake_background") else "_submit_background"
submit = getattr(manager, wake_name)
submissions = []
def paused_submit(*args, **kwargs):
submissions.append(threading.get_ident())
if len(submissions) == 1:
first_submit.set()
assert release_submit.wait(10)
return submit(*args, **kwargs)
monkeypatch.setattr(manager, wake_name, paused_submit)
first = [{"role": "user", "content": "A transcript"}]
second = [{"role": "user", "content": "B transcript"}]
def commit(messages, target, context):
token = profile.set(context)
try:
manager.commit_session_boundary_async(messages, new_session_id=target)
finally:
profile.reset(token)
with ThreadPoolExecutor(max_workers=1) as executor, ThreadPoolExecutor(max_workers=2) as callers:
manager._sync_executor = executor
first_call = callers.submit(commit, first, "B", "A profile")
try:
assert first_submit.wait(5)
# No executor task exists yet, but the accepted intent is not drained.
pending_before_submit = manager.flush_pending(timeout=0)
callers.submit(commit, second, "D", "B profile").result(timeout=5)
executor.submit(lambda: None).result(timeout=5)
finally:
release_submit.set()
first_call.result(timeout=5)
assert manager.flush_pending(timeout=5)
assert provider.extractions == [("A", first), ("B", second)]
assert not pending_before_submit
assert observed_profiles == ["A profile", "B profile"]
assert [sid for sid, _ in provider.switches] == ["B", "D"]
@pytest.mark.parametrize("finish", ["flush", "shutdown"])
def test_inline_fallback_reentrant_boundary_defers_without_self_wait(finish):
reentered = threading.Event()
release_end = threading.Event()
errors = []
flush_inside = []
first = [{"role": "user", "content": "A transcript"}]
second = [{"role": "user", "content": "B transcript"}]
class _Provider(_StatefulProvider):
def on_session_end(self, messages):
if not reentered.is_set():
reentered.set()
manager.commit_session_boundary_async(second, new_session_id="D")
flush_inside.append(manager.flush_pending(timeout=0))
assert release_end.wait(10)
super().on_session_end(messages)
provider = _Provider()
manager = MemoryManager()
manager.add_provider(provider)
# A real closed executor exercises the existing RuntimeError inline fallback.
executor = ThreadPoolExecutor(max_workers=1)
executor.shutdown()
manager._sync_executor = executor
def commit():
try:
manager.commit_session_boundary_async(first, new_session_id="B")
except BaseException as exc:
errors.append(exc)
caller = threading.Thread(target=commit, daemon=True)
caller.start()
try:
assert reentered.wait(5)
pending_during_end = manager.flush_pending(timeout=0)
release_end.set()
caller.join(timeout=2)
assert not caller.is_alive(), "inline boundary callback waited on its own owner"
assert not pending_during_end
assert not errors
assert flush_inside == [False]
if finish == "shutdown":
manager.shutdown_all()
assert manager.shutdown_drain_state["status"] == "drained"
else:
assert manager.flush_pending(timeout=5)
assert provider.extractions == [("A", first), ("B", second)]
assert [sid for sid, _ in provider.switches] == ["B", "D"]
finally:
release_end.set()
# Release the old implementation's self-wait on RED, without a hung pytest.
if caller.is_alive():
with manager._boundary_condition:
manager._boundary_active = False
manager._boundary_condition.notify_all()
caller.join(timeout=5)
assert not caller.is_alive()
def test_shutdown_counts_and_cancels_queued_boundary_receipts(monkeypatch):
entered = threading.Event()
release = threading.Event()
first = [{"role": "user", "content": "A transcript"}]
second = [{"role": "user", "content": "B transcript"}]
class _Provider(_StatefulProvider):
def on_session_end(self, messages):
entered.set()
assert release.wait(10)
super().on_session_end(messages)
manager = MemoryManager()
provider = _Provider()
manager.add_provider(provider)
monkeypatch.setattr("agent.memory_manager._SYNC_DRAIN_TIMEOUT_S", 0)
with ThreadPoolExecutor(max_workers=1) as executor:
manager._sync_executor = executor
try:
manager.commit_session_boundary_async(first, new_session_id="B")
assert entered.wait(5)
manager.commit_session_boundary_async(second, new_session_id="D")
manager.shutdown_all()
assert manager.shutdown_drain_state == {
"status": "timed_out", "abandoned_writes": 1,
"abandoned_prefetches": 0, "active_tasks": 1,
}
assert not manager.flush_pending(timeout=0)
manager.commit_session_boundary_async(second, new_session_id="late")
finally:
release.set()
assert manager.flush_pending(timeout=5)
assert provider.extractions == [("A", first)]
assert [sid for sid, _ in provider.switches] == ["B"]
@pytest.mark.parametrize("finish", ["later_work", "shutdown"])
def test_accepted_boundary_wakeup_pause_cannot_be_overtaken_or_lost(monkeypatch, finish):
accepted = threading.Event()
release = threading.Event()
calls = []
class _Provider(_StatefulProvider):
def on_session_end(self, messages):
calls.append(("end", self.session_id))
def on_session_switch(self, new_session_id, **kwargs):
super().on_session_switch(new_session_id, **kwargs)
calls.append(("switch", self.session_id))
def sync_turn(self, user_content, assistant_content, **kwargs):
calls.append(("sync", self.session_id))
def queue_prefetch(self, query, **kwargs):
calls.append(("prefetch", self.session_id))
def shutdown(self):
calls.append(("close", self.session_id))
manager = MemoryManager()
provider = _Provider()
manager.add_provider(provider)
wake_name = "_wake_background" if hasattr(manager, "_wake_background") else "_submit_background"
wake = getattr(manager, wake_name)
def paused_wakeup(*args, **kwargs):
if not accepted.is_set():
accepted.set()
assert release.wait(10)
return wake(*args, **kwargs)
monkeypatch.setattr(manager, wake_name, paused_wakeup)
monkeypatch.setattr("agent.memory_manager._SYNC_DRAIN_TIMEOUT_S", 0.2)
with ThreadPoolExecutor(max_workers=1) as executor, ThreadPoolExecutor(max_workers=1) as callers:
manager._sync_executor = executor
pending = callers.submit(manager.commit_session_boundary_async, [], new_session_id="B")
try:
assert accepted.wait(5)
assert not manager.flush_pending(timeout=0)
if finish == "shutdown":
manager.shutdown_all()
assert manager.shutdown_drain_state["status"] == "drained"
assert calls == [("end", "A"), ("switch", "B"), ("close", "B")]
else:
manager.sync_all("B turn", "response", session_id="B")
manager.queue_prefetch_all("B query", session_id="B")
executor.submit(lambda: None).result(timeout=5)
assert calls == [("end", "A"), ("switch", "B"), ("sync", "B"), ("prefetch", "B")]
finally:
release.set()
pending.result(timeout=5)
assert manager.flush_pending(timeout=5)
@pytest.mark.parametrize("kind", ["sync", "switch"])
def test_all_active_callbacks_have_receipts_and_defer_provider_close(monkeypatch, kind):
entered = threading.Event()
release = threading.Event()
closed = threading.Event()
inside_flush = []
calls = []
class _Provider(_StatefulProvider):
def blocked(self):
inside_flush.append(manager.flush_pending(timeout=0))
entered.set()
assert release.wait(10)
calls.append("callback done")
def sync_turn(self, user_content, assistant_content, **kwargs):
self.blocked()
def on_session_switch(self, new_session_id, **kwargs):
self.blocked()
super().on_session_switch(new_session_id, **kwargs)
def shutdown(self):
calls.append("close")
closed.set()
manager = MemoryManager()
manager.add_provider(_Provider())
monkeypatch.setattr("agent.memory_manager._SYNC_DRAIN_TIMEOUT_S", 0)
with ThreadPoolExecutor(max_workers=1) as callers:
caller = callers.submit(manager.sync_all, "turn", "reply") if kind == "sync" else callers.submit(manager.on_session_switch, "B")
try:
assert entered.wait(5)
assert inside_flush == [False]
assert not manager.flush_pending(timeout=0)
manager.shutdown_all()
assert manager.shutdown_drain_state["status"] == "timed_out"
assert manager.shutdown_drain_state["active_tasks"] == 1
assert not closed.is_set()
manager.on_session_switch("late")
manager.sync_all("late", "late")
finally:
release.set()
caller.result(timeout=5)
assert closed.wait(5)
assert manager.flush_pending(timeout=5)
assert calls == ["callback done", "close"]
manager.shutdown_all()
assert calls == ["callback done", "close"]
def test_cancelled_receipt_notifies_future_waiters_without_worker(monkeypatch):
entered = threading.Event()
release = threading.Event()
waiter_started = threading.Event()
waiter_finished = threading.Event()
manager = MemoryManager()
manager.add_provider(_StatefulProvider())
monkeypatch.setattr("agent.memory_manager._SYNC_DRAIN_TIMEOUT_S", 0)
with ThreadPoolExecutor(max_workers=1) as executor:
manager._sync_executor = executor
blocker = executor.submit(lambda: (entered.set(), release.wait(10)))
assert entered.wait(5)
manager.commit_session_boundary_async([], new_session_id="B")
receipts = tuple(manager._background_futures)
def waiting():
waiter_started.set()
if not wait(receipts, timeout=5)[1]:
waiter_finished.set()
waiter = threading.Thread(target=waiting, daemon=True)
waiter.start()
try:
assert waiter_started.wait(5)
manager.shutdown_all()
assert waiter_finished.wait(2), "cancelled receipt never notified concurrent.futures.wait"
# Timeout cleanup is scheduled, never run inline on shutdown's caller.
# It must still finish without releasing the occupied executor.
assert manager.flush_pending(timeout=5)
finally:
release.set()
blocker.result(timeout=5)
waiter.join(timeout=5)
def test_busy_sync_switch_orders_new_turn_and_prefetch_in_its_scope():
entered = threading.Event()
release = threading.Event()
calls = []
class _Provider(_StatefulProvider):
def sync_turn(self, user_content, assistant_content, **kwargs):
if user_content == "A":
entered.set()
assert release.wait(10)
calls.append((user_content, self.session_id))
def queue_prefetch(self, query, **kwargs):
calls.append((query, self.session_id))
manager = MemoryManager()
provider = _Provider()
manager.add_provider(provider)
try:
manager.sync_all("A", "reply", session_id="A")
assert entered.wait(5)
manager.on_session_switch("B")
assert provider.session_id == "A"
manager.sync_all("B", "reply", session_id="B")
manager.queue_prefetch_all("prefetch B", session_id="B")
finally:
release.set()
assert manager.flush_pending(timeout=5)
assert calls == [("A", "A"), ("B", "B"), ("prefetch B", "B")]
manager.shutdown_all()
@pytest.mark.parametrize("close_raises", [False, True])
def test_drained_shutdown_waits_for_finalizer_after_last_receipt(close_raises, monkeypatch):
work_entered = threading.Event()
release_work = threading.Event()
receipt_done = threading.Event()
release_owner = threading.Event()
finish_entered = threading.Event()
close_entered = threading.Event()
release_close = threading.Event()
calls = []
inside_flush = []
class _Provider(_StatefulProvider):
def sync_turn(self, *args, **kwargs):
work_entered.set()
assert release_work.wait(10)
def shutdown(self):
calls.append("close")
# Both are unbounded calls: a finalizer must not wait on itself.
manager.shutdown_all()
inside_flush.append(manager.flush_pending())
close_entered.set()
assert release_close.wait(10)
if close_raises:
raise SystemExit("finalizer failed")
manager = MemoryManager()
manager.add_provider(_Provider())
finish = manager._work_queue.finish
def observed_finish(*args, **kwargs):
finish_entered.set()
return finish(*args, **kwargs)
monkeypatch.setattr(manager._work_queue, "finish", observed_finish)
def hold_owner(receipt):
assert receipt.done()
receipt_done.set()
assert release_owner.wait(10)
with ThreadPoolExecutor(max_workers=1) as executor, ThreadPoolExecutor(max_workers=3) as callers:
manager._sync_executor = executor
manager.sync_all("A", "reply")
try:
assert work_entered.wait(5)
receipt = next(iter(manager._background_futures))
receipt.add_done_callback(hold_owner)
release_work.set()
assert receipt_done.wait(5)
assert manager.flush_pending(timeout=0)
closing = callers.submit(manager.shutdown_all)
assert finish_entered.wait(5)
assert manager.shutdown_drain_state["status"] == "drained"
# Wait until finish has registered its callback (or wrongly returned).
with manager._work_queue.condition:
assert manager._work_queue.condition.wait_for(
lambda: manager._work_queue._finalizer_set, timeout=5,
)
waiting_for_close = manager.flush_pending(timeout=0)
finalizer_receipt = manager._work_queue.finish(lambda: calls.append("duplicate close"))
assert finalizer_receipt is not None
assert not finalizer_receipt.cancel()
repeated = callers.submit(manager.shutdown_all)
release_owner.set()
assert close_entered.wait(5), "finalizer reentrant shutdown waited on itself"
assert not closing.done(), "normal drained shutdown returned before provider close"
assert not waiting_for_close, "flush ignored the registered finalizer"
assert inside_flush == [False]
assert not manager.flush_pending(timeout=0)
# Observe an actual waiter rather than relying on caller scheduling.
external_waiting = threading.Event()
original_wait = manager._work_queue.condition.wait
def observed_wait(timeout=None):
external_waiting.set()
return original_wait(timeout)
monkeypatch.setattr(manager._work_queue.condition, "wait", observed_wait)
flushing = callers.submit(manager.flush_pending)
assert external_waiting.wait(5)
assert not repeated.done(), "repeated normal shutdown skipped in-flight close"
assert not flushing.done()
finally:
release_work.set()
release_owner.set()
release_close.set()
closing.result(timeout=5)
repeated.result(timeout=5)
assert flushing.result(timeout=5)
assert not wait((finalizer_receipt,), timeout=0)[1]
if close_raises:
assert isinstance(finalizer_receipt.exception(), SystemExit)
else:
assert finalizer_receipt.exception() is None
manager.shutdown_all()
assert calls == ["close"]
assert manager.flush_pending(timeout=0)
@pytest.mark.parametrize("initiator", ["external", "owner"])
def test_timed_out_shutdown_defers_finalizer_but_flush_tracks_close(initiator, monkeypatch):
work_entered = threading.Event()
release_work = threading.Event()
owner_shutdown_done = threading.Event()
close_entered = threading.Event()
release_close = threading.Event()
calls = []
inside_flush = []
class _Provider(_StatefulProvider):
def sync_turn(self, *args, **kwargs):
if initiator == "owner":
manager.shutdown_all()
inside_flush.append(manager.flush_pending())
owner_shutdown_done.set()
work_entered.set()
assert release_work.wait(10)
calls.append("work done")
def shutdown(self):
calls.append("close")
manager.shutdown_all()
inside_flush.append(manager.flush_pending())
close_entered.set()
assert release_close.wait(10)
raise RuntimeError("provider close failed")
manager = MemoryManager()
manager.add_provider(_Provider())
monkeypatch.setattr("agent.memory_manager._SYNC_DRAIN_TIMEOUT_S", 0)
with ThreadPoolExecutor(max_workers=1) as executor, ThreadPoolExecutor(max_workers=1) as callers:
manager._sync_executor = executor
manager.sync_all("A", "reply")
try:
assert work_entered.wait(5), "owner shutdown waited on its own work"
if initiator == "owner":
assert owner_shutdown_done.is_set()
else:
callers.submit(manager.shutdown_all).result(timeout=5)
assert manager.shutdown_drain_state["status"] == "timed_out"
callers.submit(manager.shutdown_all).result(timeout=5)
assert not close_entered.is_set()
assert not manager.flush_pending(timeout=0)
release_work.set()
assert close_entered.wait(5)
# Repeated timed-out shutdown stays nonblocking, even during close.
callers.submit(manager.shutdown_all).result(timeout=5)
assert not manager.flush_pending(timeout=0), "flush ignored deferred close"
assert inside_flush == ([False, False] if initiator == "owner" else [False])
finally:
release_work.set()
release_close.set()
assert manager.flush_pending(timeout=5)
manager.shutdown_all()
assert calls == ["work done", "close"]
def _shutdown_thread(manager, *, repeat=False):
done = threading.Event()
errors = []
def run():
try:
manager.shutdown_all()
except BaseException as exc:
errors.append(exc)
finally:
done.set()
thread = threading.Thread(target=run, name="shutdown-repeat" if repeat else "shutdown-first", daemon=True)
thread.start()
return thread, done, errors
def _observe_repeat_wait(monkeypatch):
waiting = threading.Event()
original = threading.Condition.wait
def observed(condition, timeout=None):
if threading.current_thread().name == "shutdown-repeat":
waiting.set()
return original(condition, timeout)
monkeypatch.setattr(threading.Condition, "wait", observed)
return waiting
def test_repeat_during_drain_observes_timeout_before_blocked_close(monkeypatch):
work_entered, release_work = threading.Event(), threading.Event()
drain_entered, release_drain = threading.Event(), threading.Event()
close_entered, release_close = threading.Event(), threading.Event()
repeat_waiting = _observe_repeat_wait(monkeypatch)
calls = []
class _Provider(_StatefulProvider):
def sync_turn(self, *args, **kwargs):
work_entered.set()
assert release_work.wait(10)
def shutdown(self):
calls.append("close")
close_entered.set()
assert release_close.wait(10)
manager = MemoryManager()
manager.add_provider(_Provider())
flush, finish = manager._work_queue.flush, manager._work_queue.finish
def gated_drain(timeout=None, *, snapshot=None):
if timeout is not None and snapshot is not None:
drain_entered.set()
assert release_drain.wait(10)
return flush(0, snapshot=snapshot)
return flush(timeout, snapshot=snapshot)
def owner_exits_before_finish(*args, **kwargs):
release_work.set()
with manager._work_queue.condition:
assert manager._work_queue.condition.wait_for(lambda: manager._work_queue._owner is None, 5)
return finish(*args, **kwargs)
monkeypatch.setattr(manager._work_queue, "flush", gated_drain)
monkeypatch.setattr(manager._work_queue, "finish", owner_exits_before_finish)
manager.sync_all("A", "reply")
assert work_entered.wait(5)
first, first_done, first_errors = _shutdown_thread(manager)
repeated = None
try:
assert drain_entered.wait(5)
repeated, repeat_done, repeat_errors = _shutdown_thread(manager, repeat=True)
assert repeat_waiting.wait(5)
release_drain.set()
assert close_entered.wait(5)
first_returned = first_done.wait(2)
repeat_returned = repeat_done.wait(2)
assert first_returned and repeat_returned, "timeout must release initiator and existing repeat without waiting for close"
assert not manager.flush_pending(timeout=0)
assert manager.shutdown_drain_state["status"] == "timed_out"
finally:
release_drain.set()
release_work.set()
release_close.set()
first.join(5)
if repeated is not None:
repeated.join(5)
assert not first_errors and not repeat_errors
assert manager.flush_pending(timeout=5)
assert calls == ["close"]
def test_interrupted_drain_notifies_repeat_and_defers_safe_close(monkeypatch):
work_entered, release_work = threading.Event(), threading.Event()
drain_entered, release_drain = threading.Event(), threading.Event()
close_entered, release_close = threading.Event(), threading.Event()
repeat_waiting = _observe_repeat_wait(monkeypatch)
interrupt = KeyboardInterrupt("interrupted shutdown drain")
calls = []
class _Provider(_StatefulProvider):
def sync_turn(self, *args, **kwargs):
work_entered.set()
assert release_work.wait(10)
calls.append("work done")
def shutdown(self):
calls.append("close")
close_entered.set()
assert release_close.wait(10)
manager = MemoryManager()
manager.add_provider(_Provider())
flush = manager._work_queue.flush
def interrupted_drain(timeout=None, *, snapshot=None):
if timeout is not None and snapshot is not None:
drain_entered.set()
assert release_drain.wait(10)
raise interrupt
return flush(timeout, snapshot=snapshot)
monkeypatch.setattr(manager._work_queue, "flush", interrupted_drain)
manager.sync_all("A", "reply")
assert work_entered.wait(5)
first, first_done, first_errors = _shutdown_thread(manager)
repeated = None
try:
assert drain_entered.wait(5)
repeated, repeat_done, repeat_errors = _shutdown_thread(manager, repeat=True)
assert repeat_waiting.wait(5)
release_drain.set()
assert first_done.wait(2), "drain exception was blocked by cleanup"
assert first_errors == [interrupt]
repeat_returned = repeat_done.wait(2)
assert not close_entered.is_set(), "close raced active provider work"
release_work.set()
cleanup_started = close_entered.wait(2)
assert repeat_returned and cleanup_started, "failed drain stranded repeat or omitted finalizer"
assert manager.shutdown_drain_state["status"] == "failed"
assert not manager.flush_pending(timeout=0)
finally:
release_drain.set()
release_work.set()
release_close.set()
# Let the broken Event-based implementation's RED waiter exit too.
if hasattr(manager, "_shutdown_complete"):
manager._shutdown_complete.set()
first.join(5)
if repeated is not None:
repeated.join(5)
assert not repeat_errors
assert manager.flush_pending(timeout=5)
assert calls == ["work done", "close"]
def test_last_receipt_owner_shutdown_does_not_publish_close_completion(monkeypatch):
work_entered, release_work = threading.Event(), threading.Event()
owner_returned, release_owner = threading.Event(), threading.Event()
close_entered, release_close = threading.Event(), threading.Event()
repeat_waiting = _observe_repeat_wait(monkeypatch)
calls = []
class _Provider(_StatefulProvider):
def sync_turn(self, *args, **kwargs):
work_entered.set()
assert release_work.wait(10)
def shutdown(self):
calls.append("close")
manager.shutdown_all()
close_entered.set()
assert release_close.wait(10)
manager = MemoryManager()
manager.add_provider(_Provider())
def owner_shutdown(receipt):
assert receipt.done()
manager.shutdown_all()
owner_returned.set()
assert release_owner.wait(10)
manager.sync_all("A", "reply")
assert work_entered.wait(5)
receipt = next(iter(manager._background_futures))
receipt.add_done_callback(owner_shutdown)
release_work.set()
repeated = None
try:
assert owner_returned.wait(5), "owner shutdown self-waited"
assert manager.shutdown_drain_state["status"] == "drained"
repeated, repeat_done, repeat_errors = _shutdown_thread(manager, repeat=True)
assert repeat_waiting.wait(2), "owner return incorrectly published close completion"
assert not repeat_done.is_set()
assert not manager.flush_pending(timeout=0)
release_owner.set()
assert close_entered.wait(5)
assert not repeat_done.is_set(), "repeat returned during provider close"
finally:
release_work.set()
release_owner.set()
release_close.set()
if repeated is not None:
repeated.join(5)
assert repeat_done.is_set() and not repeat_errors
assert manager.flush_pending(timeout=5)
assert calls == ["close"]