From 44e607be61cc71686dfe78c9c689a24469232449 Mon Sep 17 00:00:00 2001 From: Yingliang Zhang Date: Sun, 20 Sep 2026 20:51:26 +0800 Subject: [PATCH 1/2] fix(hindsight): fence prefetch publication to the owning session generation (#64745) The background prefetch worker published its recall into the session slot unconditionally; a worker outliving on_session_switch's 3s join wrote the old session's memories into the new session's slot. queue_prefetch also spawned unbounded threads with the last finisher winning the slot. Workers now capture a slot generation at spawn, queue_prefetch bumps it and skips while a prior worker runs, on_session_switch/shutdown bump it to fence late publishers, and the publish + recall are gated on the current generation. --- plugins/memory/hindsight/__init__.py | 34 ++++- .../plugins/memory/test_hindsight_provider.py | 141 ++++++++++++++++++ 2 files changed, 174 insertions(+), 1 deletion(-) diff --git a/plugins/memory/hindsight/__init__.py b/plugins/memory/hindsight/__init__.py index 20236b0a4e69f..23d70820664f3 100644 --- a/plugins/memory/hindsight/__init__.py +++ b/plugins/memory/hindsight/__init__.py @@ -367,6 +367,10 @@ def __init__(self): self._prefetch_result, self._prefetch_count = "", 0 self._prefetch_lock = threading.Lock() self._prefetch_thread = None + # Prefetch-slot ownership (guarded by _prefetch_lock): bumped per spawned + # worker and on session switch/shutdown, so a late or superseded worker + # whose captured generation is stale can never publish into the slot. + self._prefetch_generation = 0 self._last_recall_returned, self._last_recall_count = False, 0 self._apply_recall_settings({}) @@ -963,15 +967,36 @@ def queue_prefetch(self, query: str, *, session_id: str = "") -> None: if self._recall_sync or self._recall_disabled(): return + # One worker at a time (mirrors _ensure_writer): rapid turns must not + # stack concurrent recalls against one embedded daemon — the LAST + # finisher would win the slot, not the newest warm. + if self._prefetch_thread is not None and self._prefetch_thread.is_alive(): + logger.debug("Prefetch: prior worker still running; skipping this warm") + return + + with self._prefetch_lock: + self._prefetch_generation += 1 + generation = self._prefetch_generation + def _run(): # Wait (bounded, off the reply path) for the just-completed turn's # retain to be recall-visible so the warmed context includes it. if self._prefetch_waits_for_retain: self._wait_for_retains_drained(self._prefetch_retain_drain_timeout) + # A superseded/shut-down worker must not even spend a recall + # against the daemon on a result nobody will use. + with self._prefetch_lock: + if generation != self._prefetch_generation or self._shutting_down.is_set(): + return text, count = self._do_recall(query) if text: + # Fenced publish: the slot can outlive the switching session's + # 3s join (drain up to 10s + recall up to 120s), so only the + # generation that owns the slot may write it — a stale worker + # must not inject the old session's memories into the new one. with self._prefetch_lock: - self._prefetch_result, self._prefetch_count = text, count + if generation == self._prefetch_generation and not self._shutting_down.is_set(): + self._prefetch_result, self._prefetch_count = text, count self._prefetch_thread = spawn_context_thread(_run, name="hindsight-prefetch") self._prefetch_thread.start() @@ -1193,6 +1218,9 @@ def _flush(): # 2. Drain the old session's in-flight prefetch and drop its result. self._join_prefetch(3.0) with self._prefetch_lock: + # Fence out any worker still running past the join: the slot now + # belongs to the new session. + self._prefetch_generation += 1 self._prefetch_result = "" # 3. Rotate to the new session. @@ -1223,6 +1251,10 @@ def shutdown(self) -> None: logger.debug("Hindsight shutdown: stopping writer + waiting for background threads") # Stop accepting retain jobs first so late sync_turn() calls are dropped. self._shutting_down.set() + # Fence any in-flight prefetch worker: a recall finishing during or + # after the joins below must never publish into a closing provider. + with self._prefetch_lock: + self._prefetch_generation += 1 # The writer finishes in-flight work then exits on the sentinel; the # bounded join keeps shutdown predictable even if the daemon is wedged. if (writer := self._writer_thread) is not None and writer.is_alive(): diff --git a/tests/plugins/memory/test_hindsight_provider.py b/tests/plugins/memory/test_hindsight_provider.py index 93a429464e55a..20342be7e5ff9 100644 --- a/tests/plugins/memory/test_hindsight_provider.py +++ b/tests/plugins/memory/test_hindsight_provider.py @@ -1256,6 +1256,147 @@ def _aretain_batch_tracking(**kw): assert call_order[1] == "3" +# --------------------------------------------------------------------------- +# Prefetch generation fence (#64745: stale/superseded worker must not publish) +# --------------------------------------------------------------------------- + + +class TestPrefetchGenerationFence: + def test_stale_worker_cannot_publish_after_switch_join_timeout(self, provider, monkeypatch): + """A prefetch worker can legitimately outlive on_session_switch's 3s + join (10s drain + 120s recall). Its late publish lands in the NEW + session's slot upstream; the generation fence must drop it, and the + new session's first prefetch() must inject nothing of the old one.""" + import threading + + gate = threading.Event() + entered = threading.Event() + + def _gated_recall(query): + entered.set() + gate.wait(timeout=10.0) + return "- old-session memory", 1 + + monkeypatch.setattr(provider, "_do_recall", _gated_recall) + provider.queue_prefetch("old session query") + assert entered.wait(timeout=5.0), "prefetch worker never reached the recall" + + # The worker is still gated when the switch's 3.0s join times out, so + # the switch returns with the worker alive — the leaked window upstream. + provider.on_session_switch("new-sid") + gate.set() + provider._prefetch_thread.join(timeout=5.0) + + assert provider._prefetch_result == "" + assert provider.prefetch("anything") == "" + + def test_superseded_worker_cannot_overwrite_newer_result(self, provider, monkeypatch): + """Two overlapping workers: the older finishing after the newer must + not clobber the newer result sitting in the slot.""" + import threading + + gate_first = threading.Event() + + def _routed_recall(query): + if query == "first query": + gate_first.wait(timeout=10.0) + return "- first (older) memory", 1 + return "- second (newer) memory", 1 + + monkeypatch.setattr(provider, "_do_recall", _routed_recall) + + provider.queue_prefetch("first query") + worker_first = provider._prefetch_thread + # Simulate the lost-thread-handle hazard: upstream overwrote + # _prefetch_thread freely, so a dead handle lets the next warm spawn. + provider._prefetch_thread = None + provider.queue_prefetch("second query") + worker_second = provider._prefetch_thread + + # The newer worker finishes first and wins the slot. + worker_second.join(timeout=5.0) + with provider._prefetch_lock: + assert provider._prefetch_result == "- second (newer) memory" + + # The older worker finishes LAST; its result must be fenced out. + gate_first.set() + worker_first.join(timeout=5.0) + with provider._prefetch_lock: + assert provider._prefetch_result == "- second (newer) memory" + + def test_shutdown_fences_inflight_prefetch_publish(self, provider): + """A recall resuming after shutdown() began must not publish, and the + fence must prevent a recall against a client shutdown already closed.""" + import threading + + gate = threading.Event() + entered = threading.Event() + + async def _gated_arecall(**kwargs): + entered.set() + gate.wait(timeout=10.0) + return SimpleNamespace(results=[SimpleNamespace(text="late memory")]) + + provider._client.arecall = AsyncMock(side_effect=_gated_arecall) + + provider.queue_prefetch("q") + assert entered.wait(timeout=5.0), "prefetch worker never reached the recall" + + # Release the recall the moment shutdown starts so shutdown's prefetch + # join returns promptly instead of burning its 5s budget. + def _release_on_shutdown(): + provider._shutting_down.wait(timeout=10.0) + gate.set() + + threading.Thread(target=_release_on_shutdown, daemon=True).start() + provider.shutdown() + provider._prefetch_thread.join(timeout=5.0) + + assert provider._prefetch_result == "" + assert provider._client is None + assert provider.prefetch("q") == "" + + def test_queue_prefetch_skips_while_prior_worker_running(self, provider): + """Rapid turns must warm serially: while one worker runs, further + queue_prefetch calls neither spawn nor bump — only ONE recall happens.""" + import threading + + gate = threading.Event() + entered = threading.Event() + + async def _gated_arecall(**kwargs): + entered.set() + gate.wait(timeout=10.0) + return SimpleNamespace(results=[SimpleNamespace(text="m")]) + + provider._client.arecall = AsyncMock(side_effect=_gated_arecall) + + provider.queue_prefetch("q1") + assert entered.wait(timeout=5.0), "prefetch worker never reached the recall" + first = provider._prefetch_thread + + provider.queue_prefetch("q2") + provider.queue_prefetch("q3") + + gate.set() + first.join(timeout=5.0) + + # The live worker was never replaced and nothing else ever recalled. + assert provider._prefetch_thread is first + assert provider._client.arecall.await_count == 1 + + def test_same_session_boundary_still_publishes(self, provider): + """Positive control: with no session switch the fence must be a no-op — + an uninterrupted warm → prefetch cycle delivers exactly as upstream.""" + provider.queue_prefetch("test") + if provider._prefetch_thread: + provider._prefetch_thread.join(timeout=5.0) + + result = provider.prefetch("test") + assert "Memory 1" in result + assert "Memory 2" in result + + # --------------------------------------------------------------------------- # update_mode='append' capability probe + retain dispatch # --------------------------------------------------------------------------- From 1000b62f0df1193610a36b6f08e99d868acccd8e Mon Sep 17 00:00:00 2001 From: Yingliang Zhang Date: Sun, 20 Sep 2026 21:21:01 +0800 Subject: [PATCH 2/2] fix(hindsight): guard client lifecycle with a leaf lock --- plugins/memory/hindsight/__init__.py | 58 +++++++--- .../plugins/memory/test_hindsight_provider.py | 101 +++++++++++++++++- 2 files changed, 143 insertions(+), 16 deletions(-) diff --git a/plugins/memory/hindsight/__init__.py b/plugins/memory/hindsight/__init__.py index 23d70820664f3..d5a95761f0241 100644 --- a/plugins/memory/hindsight/__init__.py +++ b/plugins/memory/hindsight/__init__.py @@ -329,6 +329,15 @@ def backup_paths(self) -> List[str]: def __init__(self): self._config = self._api_key = self._client = None + # LEAF lock guarding _client construction/retirement (#11923): two threads + # racing a cold or just-nulled client would both build one, and the loser + # owns an aiohttp session nothing ever closes. Never held across + # _run_sync/operation(client) or while taking _prefetch_lock / + # _pending_retain_ops_lock. + self._client_lock = threading.Lock() + # Client observed broken by the stale-daemon retry; only its observer may + # retire it (identity-checked against _client under _client_lock). + self._broken_client = None self._api_url, self._llm_base_url, self._mode = _DEFAULT_API_URL, "", "cloud" self._timeout, self._idle_timeout = _DEFAULT_TIMEOUT, _DEFAULT_IDLE_TIMEOUT self._bank_id, self._budget, self._bank_id_template = "hermes", "mid", "" @@ -501,11 +510,25 @@ def _new_cloud_client(self): self._api_url, bool(self._api_key), kwargs["timeout"]) return Hindsight(**kwargs) - def _get_client(self): - """Return the cached Hindsight client (created once, reused).""" - if self._client is None: - self._client = self._new_embedded_client() if self._mode == "local_embedded" else self._new_cloud_client() - return self._client + def _get_client(self, *, recreate: bool = False): + """Return the cached Hindsight client (created once, reused). + + *recreate* is used by the stale-daemon retry path: it drops the cached + client (if it is still the broken one we observed) and rebuilds it + exactly once, under the lock. + """ + if not recreate and (client := self._client) is not None: + return client + with self._client_lock: + if recreate: + # Only retire the client we actually observed as broken; a sibling + # thread may have already rebuilt it after our failure. + if self._client is not None and self._client is not getattr(self, "_broken_client", None): + return self._client + self._client = None + if self._client is None: + self._client = self._new_embedded_client() if self._mode == "local_embedded" else self._new_cloud_client() + return self._client def _run_sync(self, coro): """Schedule *coro* on the shared loop using the configured timeout.""" @@ -521,8 +544,9 @@ def _run_hindsight_operation(self, operation): if self._mode != "local_embedded" or not any(m in text for m in _RETRIABLE_CONNECTION_MARKERS): raise logger.info("Hindsight embedded daemon appears unreachable; recreating client and retrying once: %s", exc) - self._client = None - self._client = client = self._get_client() + self._broken_client = self._client + self._client = client = self._get_client(recreate=True) + self._broken_client = None return self._run_sync(operation(client)) # -- retain writer thread + server-side visibility ------------------------- @@ -1232,20 +1256,20 @@ def _flush(): logger.debug("Hindsight on_session_switch: new_session=%s parent=%s reset=%s doc=%s", self._session_id, self._parent_session_id, reset, self._document_id) - def _close_client(self) -> None: + def _close_client_of(self, client) -> None: if self._mode != "local_embedded": - self._run_sync(self._client.aclose()) + self._run_sync(client.aclose()) return # HindsightEmbedded.close() closes its sync client from this thread ("attached # to a different loop" before aiohttp releases the session): aclose the inner # client on the shared loop first, then let the wrapper clean up bookkeeping. - inner_client = getattr(self._client, "_client", None) + inner_client = getattr(client, "_client", None) if inner_client is not None and hasattr(inner_client, "aclose"): _run_sync(inner_client.aclose()) with contextlib.suppress(Exception): - self._client._client = None + client._client = None with contextlib.suppress(RuntimeError): - self._client.close() + client.close() def shutdown(self) -> None: logger.debug("Hindsight shutdown: stopping writer + waiting for background threads") @@ -1264,10 +1288,14 @@ def shutdown(self) -> None: logger.warning("Hindsight writer did not stop within 10s; abandoning %d pending retain(s)", self._retain_queue.qsize()) self._join_prefetch(5.0) - if self._client is not None: - with contextlib.suppress(Exception): - self._close_client() + # Retire the client under the lock BEFORE closing it, so a concurrent + # _get_client() sees None and rebuilds instead of racing the close. + with self._client_lock: + client = self._client self._client = None + if client is not None: + with contextlib.suppress(Exception): + self._close_client_of(client) # The module-global loop is intentionally NOT stopped: it's shared by every # provider in the process (one per gateway chat session); stopping it would # strand siblings' aiohttp sessions ("Unclosed client session"). Daemon diff --git a/tests/plugins/memory/test_hindsight_provider.py b/tests/plugins/memory/test_hindsight_provider.py index 20342be7e5ff9..701c7e8211489 100644 --- a/tests/plugins/memory/test_hindsight_provider.py +++ b/tests/plugins/memory/test_hindsight_provider.py @@ -561,7 +561,7 @@ def test_local_embedded_recall_reconnects_after_idle_shutdown(self, provider, mo provider._mode = "local_embedded" provider._client = first_client - monkeypatch.setattr(provider, "_get_client", lambda: next(clients)) + monkeypatch.setattr(provider, "_get_client", lambda **_: next(clients)) result = json.loads(provider.handle_tool_call( "hindsight_recall", {"query": "test"} @@ -1397,6 +1397,105 @@ def test_same_session_boundary_still_publishes(self, provider): assert "Memory 2" in result +# --------------------------------------------------------------------------- +# Client lifecycle lock (#11923: concurrent construction must not orphan a client) +# --------------------------------------------------------------------------- + + +class TestClientLifecycleLock: + def test_client_created_once_under_concurrent_first_access(self, provider, monkeypatch): + """N threads hitting a cold client cache must construct exactly one: the + slow embedded constructor (runtime check, daemon spawn) used to let N-1 + duplicates through a check-then-act window, each leaked client owning an + aiohttp session nothing ever closes.""" + provider._client = None + + built = [] + + def _slow_build(): + time.sleep(0.2) # force the check-then-act window open + client = SimpleNamespace(name=f"client-{len(built)}") + built.append(client) + return client + + monkeypatch.setattr(provider, "_new_cloud_client", _slow_build) + + results = [] + + def _fetch(): + results.append(provider._get_client()) + + threads = [threading.Thread(target=_fetch, daemon=True) for _ in range(8)] + for t in threads: + t.start() + for t in threads: + t.join(timeout=10.0) + + assert len(built) == 1 + assert len(results) == 8 + assert all(client is built[0] for client in results) + + def test_retry_does_not_orphan_a_sibling_client(self, provider, monkeypatch): + """The stale-daemon retry retires _client while sibling threads may sit in + _get_client(); without the lock both build and the overwritten client is an + orphan nobody closes. Exactly one replacement must ever be constructed.""" + provider._mode = "local_embedded" + + built = [] + + def _build(): + time.sleep(0.2) # widen the null-and-rebuild window + client = SimpleNamespace() + built.append(client) + return client + + monkeypatch.setattr(provider, "_new_embedded_client", _build) + + broken = _build() + provider._client = broken + + def op(client): + async def _attempt(): + if client is broken: + raise RuntimeError("Cannot connect to host 127.0.0.1:8888") + return "ok" + return _attempt() + + stop = threading.Event() + + def _reader(): + while not stop.is_set(): + provider._get_client() + + readers = [threading.Thread(target=_reader, daemon=True) for _ in range(4)] + for t in readers: + t.start() + try: + assert provider._run_hindsight_operation(op) == "ok" + finally: + stop.set() + for t in readers: + t.join(timeout=5.0) + + assert provider._client is not broken + orphans = [c for c in built if c is not broken and c is not provider._client] + assert orphans == [] + assert sum(c is not broken for c in built) == 1 + + def test_shutdown_closes_retired_client_and_allows_rebuild(self, provider, monkeypatch): + """shutdown() retires the client before closing it, so a later + _get_client() must rebuild rather than hand back the closed object.""" + retired = provider._client + rebuilt = _make_mock_client() + monkeypatch.setattr(provider, "_new_cloud_client", lambda: rebuilt) + + provider.shutdown() + + assert retired.aclose.await_count == 1 + assert provider._client is None + assert provider._get_client() is rebuilt + + # --------------------------------------------------------------------------- # update_mode='append' capability probe + retain dispatch # ---------------------------------------------------------------------------