Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 91 additions & 22 deletions hindsight-integrations/hermes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -391,6 +391,12 @@ def backup_paths(self) -> List[str]:

def __init__(self):
self._config = self._api_key = self._client = None
# LEAF lock guarding _client construction/retirement: 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()
self._embedded_url = None
self._client_lock = threading.Lock()
self._api_url, self._mode = _DEFAULT_API_URL, "cloud"
Expand Down Expand Up @@ -435,6 +441,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({})

Expand Down Expand Up @@ -783,16 +793,31 @@ def _announce_slow_first_start(self, profile: str) -> None:
platform=self._platform,
)

def _get_client(self):
"""Return the cached Hindsight client, creating it at most once.
def _get_client(self, *, retire=None):
"""Return the cached Hindsight client (created once, reused).

*retire* is the exact client instance the caller's failed operation ran
with — identity passed as an argument, never shared broken-client state a
concurrent retry could clobber. It is dropped and the client rebuilt
exactly once, under the lock; if a sibling already rebuilt after our
failure, that fresh client is returned as-is instead of being orphaned.

Locked because in local_embedded the creation *starts the daemon*: the background start
worker and the first memory operation both land here, and an unguarded check-then-set let
worker and the first memory operation both land here, and an unguarded check-and-set let
them each start one and build one client, of which the loser was dropped without being
closed (an unclosed aiohttp session). threading.Lock, not asyncio: these are threads, and
the guarded section never awaits.
"""
if retire is None and (client := self._client) is not None:
return client
with self._client_lock:
if retire is not None:
# Identity CAS: retire only the client the operation actually ran
# with. A sibling thread may have already replaced it — its fresh
# replacement must survive our retry.
if self._client is not None and self._client is not retire:
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()
Expand All @@ -807,14 +832,17 @@ def _run_hindsight_operation(self, operation):
"""Run an async client operation; for local_embedded, a stale-daemon
connection failure recreates the client and retries once."""
try:
return self._run_sync(operation(self._get_client()))
client = self._get_client()
return self._run_sync(operation(client))
except Exception as exc:
text = f"{type(exc).__name__}: {exc}".lower()
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()
# Retire the exact client this operation ran with (identity as an
# argument); the loser of a concurrent retry is not closed inline
# (closing against a dead daemon can hang) — only shutdown() closes.
client = self._get_client(retire=client)
return self._run_sync(operation(client))

# -- retain writer thread + server-side visibility -------------------------
Expand Down Expand Up @@ -1324,15 +1352,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()
Expand Down Expand Up @@ -1402,26 +1451,31 @@ def _retain_batch(
def _make_turn_retain_job(
self, turns: list[str], *, document_id: str, update_mode: str | None, label: str, track_ops: bool = True
) -> Callable[[], None]:
"""Writer job shipping *turns* as one document. Inputs are snapshotted NOW: the
writer runs after later sync_turn() calls mutate _session_turns/_turn_index/_session_id."""
content = "[" + ",".join(turns) + "]"
metadata = self._build_metadata(message_count=len(turns) * 2, turn_index=self._turn_index)
"""Writer job shipping *turns* as one document. Everything is snapshotted NOW —
the item included — because the writer runs after later sync_turn() calls or an
on_session_switch() mutate the provider's mutable item-config attributes
(_retain_tags, _observation_scopes, _retain_source). Building the item at run
time would stamp an OLD-session retain with NEW-session tags/scopes."""
lineage = (("session", self._session_id), ("parent", self._parent_session_id))
tags = [f"{kind}:{sid}" for kind, sid in lineage if sid] or None
bank_id, retain_async, retain_context = self._bank_id, self._retain_async, self._retain_context
item = self._build_retain_kwargs(
"[" + ",".join(turns) + "]",
context=self._retain_context,
metadata=self._build_metadata(message_count=len(turns) * 2, turn_index=self._turn_index),
tags=tags,
update_mode=update_mode,
)
bank_id, retain_async = self._bank_id, self._retain_async

def _job() -> None:
item = self._build_retain_kwargs(
content, context=retain_context, metadata=metadata, tags=tags, update_mode=update_mode
)
logger.debug(
"Hindsight %s: bank=%s, doc=%s, mode=%s, async=%s, content_len=%d, num_turns=%d",
label,
bank_id,
document_id,
update_mode,
retain_async,
len(content),
len(item["content"]),
len(turns),
)
resp = self._retain_batch(item, bank_id=bank_id, document_id=document_id, retain_async=retain_async)
Expand Down Expand Up @@ -1517,7 +1571,11 @@ def _tool_retain(self, args: dict) -> str:
content, context=context, tags=args.get("tags"), occurred_at=args.get("occurred_at")
)
logger.debug("Tool hindsight_retain: bank=%s, content_len=%d, context=%s", self._bank_id, len(content), context)
self._retain_batch(item, bank_id=self._bank_id)
# Forward the configured retain_async mode: it is a call-level arg
# (never an item key), so omitting it here drops the async/sync choice
# and aretain_batch falls back to its own server default. Matches the
# auto-retain path.
self._retain_batch(item, bank_id=self._bank_id, retain_async=self._retain_async)
logger.debug("Tool hindsight_retain: success")
return "Memory stored successfully."

Expand Down Expand Up @@ -1607,6 +1665,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.
Expand All @@ -1623,20 +1684,24 @@ def _flush():
self._document_id,
)

def _close_client(self) -> None:
def _close_client_of(self, client) -> None:
"""Both modes now hold a plain ``hindsight_client.Hindsight``, so one aclose on the
shared loop is the whole story. The old local_embedded branch existed because
``HindsightEmbedded.close()`` closed its inner sync client from the calling thread
("attached to a different loop" before aiohttp released the session); there is no
wrapper to unwind any more. The daemon deliberately outlives us — it is shared with
the hindsight-embed CLI and other profiles' clients.
"""
self._run_sync(self._client.aclose())
self._run_sync(client.aclose())

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():
Expand All @@ -1648,10 +1713,14 @@ def shutdown(self) -> None:
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
Expand Down
159 changes: 159 additions & 0 deletions hindsight-integrations/hermes/tests/test_client_lifecycle.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
"""Client lifecycle lock (#11923 in the original repo) — ``_client`` is read and
written from the retain writer thread, the background prefetch worker and the
turn/tool thread. ``_get_client()`` was check-then-act and the embedded
constructor takes seconds (runtime check, daemon spawn): two threads hitting a
cold or just-nulled client both construct, one wins, and the loser's client is
orphaned with an aiohttp session nothing ever closes. The stale-daemon retry
path (``self._client = None`` then ``_get_client()``) widened the window from
microseconds to seconds, and concurrent retries clobbered each other's
replacement. NousResearch/hermes-agent#117236."""

import threading
import time
import types


def test_client_created_once_under_concurrent_first_access(provider, monkeypatch):
"""8 threads against a cold cache and a slow factory: exactly one
construction, one shared object."""
instance, _ = provider({})
instance._client = None
built = []

def _slow_build():
time.sleep(0.2) # force the check-then-act window open
client = types.SimpleNamespace(name=f"client-{len(built)}")
built.append(client)
return client

monkeypatch.setattr(instance, "_new_cloud_client", _slow_build)

results = []

def _fetch():
results.append(instance._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 _op_for(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()

return op


def test_retry_does_not_orphan_a_sibling_client(provider, monkeypatch):
"""The stale-daemon retry retires the client while sibling threads 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."""
instance, _ = provider({})
instance._mode = "local_embedded"
built = []

def _build():
time.sleep(0.2) # widen the retire-and-rebuild window
client = types.SimpleNamespace()
built.append(client)
return client

monkeypatch.setattr(instance, "_new_embedded_client", _build)
broken = _build()
instance._client = broken

stop = threading.Event()

def _reader():
while not stop.is_set():
instance._get_client()

readers = [threading.Thread(target=_reader, daemon=True) for _ in range(4)]
for t in readers:
t.start()
try:
assert instance._run_hindsight_operation(_op_for(broken)) == "ok"
finally:
stop.set()
for t in readers:
t.join(timeout=5.0)

assert instance._client is not broken
orphans = [c for c in built if c is not broken and c is not instance._client]
assert orphans == []
assert sum(c is not broken for c in built) == 1


def test_concurrent_retries_retire_the_broken_client_exactly_once(provider, monkeypatch):
"""Concurrent stale-daemon retries each run with the same broken client; the
identity-argument retire (``_get_client(retire=client)``) must let exactly one
of them rebuild it — never a shared broken-client flag a sibling could clobber,
never two replacements orphaning each other."""
instance, _ = provider({})
instance._mode = "local_embedded"
built = []

def _build():
time.sleep(0.2)
client = types.SimpleNamespace()
built.append(client)
return client

monkeypatch.setattr(instance, "_new_embedded_client", _build)
broken = _build()
instance._client = broken

results = []
start = threading.Barrier(5)

def _worker():
start.wait(timeout=5.0)
results.append(instance._run_hindsight_operation(_op_for(broken)))

threads = [threading.Thread(target=_worker, daemon=True) for _ in range(4)]
for t in threads:
t.start()
start.wait(timeout=5.0) # release all four retries at once
for t in threads:
t.join(timeout=10.0)

assert sorted(results) == ["ok"] * 4
assert instance._client is not broken
# The broken client is replaced exactly once, and the lone replacement survives.
assert sum(c is not broken for c in built) == 1
assert [c for c in built if c is not broken and c is not instance._client] == []


def test_shutdown_closes_retired_client_and_allows_rebuild(provider, monkeypatch):
"""shutdown() retires the client under the lock BEFORE closing it, so a
concurrent ``_get_client()`` rebuilds instead of racing the close, and a later
``_get_client()`` must never hand back the closed object."""
instance, _ = provider({})
closed = []

class _CountingClient:
async def aclose(self):
closed.append(1)

retired = _CountingClient()
instance._client = retired
rebuilt = types.SimpleNamespace(name="rebuilt")
monkeypatch.setattr(instance, "_new_cloud_client", lambda: rebuilt)

instance.shutdown()

assert closed == [1]
assert instance._client is None
assert instance._get_client() is rebuilt
Loading
Loading