858 lines
34 KiB
Python
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"]
|