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
117 changes: 66 additions & 51 deletions plugins/memory/hindsight/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1653,65 +1653,80 @@ def on_session_switch(
# everything before mutating self._* so metadata + tags + doc_id
# all reference the old session consistently.
if self._session_turns:
old_turns = list(self._session_turns)
old_session_id = self._session_id
old_parent_session_id = self._parent_session_id
old_turn_index = self._turn_index
old_metadata = self._build_metadata(
message_count=len(old_turns) * 2,
turn_index=old_turn_index,
)
old_lineage_tags: list[str] = []
if old_session_id:
old_lineage_tags.append(f"session:{old_session_id}")
if old_parent_session_id:
old_lineage_tags.append(f"parent:{old_parent_session_id}")
old_content = "[" + ",".join(old_turns) + "]"
# Resolve doc_id + update_mode against the OLD session BEFORE
# we rotate _session_id, so the flush lands in the old
# session's document either way (legacy: per-process unique;
# β‰₯0.5.0: stable session-scoped + append).
old_document_id, old_update_mode = self._resolve_retain_target(
self._document_id
)

def _flush():
try:
item = self._build_retain_kwargs(
old_content,
context=self._retain_context,
metadata=old_metadata,
tags=old_lineage_tags or None,
)
item.pop("bank_id", None)
item.pop("retain_async", None)
if old_update_mode is not None:
item["update_mode"] = old_update_mode
logger.debug(
"Hindsight flush-on-switch: bank=%s, doc=%s, mode=%s, num_turns=%d",
self._bank_id, old_document_id, old_update_mode, len(old_turns),
)
self._run_hindsight_operation(
lambda client: client.aretain_batch(
bank_id=self._bank_id,
items=[item],
document_id=old_document_id,
retain_async=self._retain_async,
# In append mode each sync_turn already shipped its delta and
# advanced _last_retained_turn_count, so the flush must send only
# the turns past that watermark β€” re-sending the whole buffer would
# append duplicate copies of already-retained turns to the document
# (the delta invariant sync_turn enforces). On legacy/overwrite the
# document is replaced each retain, so the flush must carry the
# entire session.
if old_update_mode == "append":
old_turns = self._session_turns[self._last_retained_turn_count:]
else:
old_turns = list(self._session_turns)

# In append mode the whole buffer may already be retained (the
# watermark caught up), leaving nothing to flush. Skip the retain
# entirely rather than ship an empty payload.
if old_turns:
old_session_id = self._session_id
old_parent_session_id = self._parent_session_id
old_turn_index = self._turn_index
old_metadata = self._build_metadata(
message_count=len(old_turns) * 2,
turn_index=old_turn_index,
)
old_lineage_tags: list[str] = []
if old_session_id:
old_lineage_tags.append(f"session:{old_session_id}")
if old_parent_session_id:
old_lineage_tags.append(f"parent:{old_parent_session_id}")
old_content = "[" + ",".join(old_turns) + "]"

def _flush():
try:
item = self._build_retain_kwargs(
old_content,
context=self._retain_context,
metadata=old_metadata,
tags=old_lineage_tags or None,
)
)
except Exception as e:
logger.warning("Hindsight flush-on-switch failed: %s", e, exc_info=True)

# Route the flush through the same writer queue sync_turn
# uses. That serializes it behind any still-queued retains
# from the old session (FIFO by document_id), avoids racing
# two threads on aretain_batch against the same document, and
# keeps shutdown's drain semantics intact. Skip enqueue if
# shutdown has already fired β€” the writer is draining/gone.
if not self._shutting_down.is_set():
self._ensure_writer()
self._register_atexit()
self._retain_queue.put(_flush)
item.pop("bank_id", None)
item.pop("retain_async", None)
if old_update_mode is not None:
item["update_mode"] = old_update_mode
logger.debug(
"Hindsight flush-on-switch: bank=%s, doc=%s, mode=%s, num_turns=%d",
self._bank_id, old_document_id, old_update_mode, len(old_turns),
)
self._run_hindsight_operation(
lambda client: client.aretain_batch(
bank_id=self._bank_id,
items=[item],
document_id=old_document_id,
retain_async=self._retain_async,
)
)
except Exception as e:
logger.warning("Hindsight flush-on-switch failed: %s", e, exc_info=True)

# Route the flush through the same writer queue sync_turn
# uses. That serializes it behind any still-queued retains
# from the old session (FIFO by document_id), avoids racing
# two threads on aretain_batch against the same document, and
# keeps shutdown's drain semantics intact. Skip enqueue if
# shutdown has already fired β€” the writer is draining/gone.
if not self._shutting_down.is_set():
self._ensure_writer()
self._register_atexit()
self._retain_queue.put(_flush)

# 2. Drain any in-flight prefetch from the old session and drop
# its cached result so the new session doesn't see stale recall.
Expand Down
74 changes: 74 additions & 0 deletions tests/plugins/memory/test_hindsight_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -1245,6 +1245,80 @@ def test_session_switch_flush_picks_capability_against_old_session(
assert kw["document_id"] == "test-session"
assert kw["items"][0]["update_mode"] == "append"

def test_session_switch_does_not_reflush_already_appended_turns(
self, provider_with_config, monkeypatch
):
"""Regression: in append mode sync_turn already ships each turn as a
delta, so on_session_switch must NOT re-append the whole session.

With the default retain_every_n_turns=1 every turn is retained the
moment it lands and the delta watermark catches up to the buffer.
The flush-on-switch then has nothing past the watermark to send, so
no retain should fire at all. Before the fix it snapshotted the full
_session_turns and re-appended every already-retained turn, silently
duplicating them in the stored document on every /resume, /branch,
/new, /reset and context compression.
"""
self._clear_capability_cache()
monkeypatch.setattr(
"plugins.memory.hindsight._fetch_hindsight_api_version",
lambda *a, **kw: "0.5.6",
)
p = provider_with_config(retain_every_n_turns=1, retain_async=False)
p.sync_turn("turn1-user", "turn1-asst")
p.sync_turn("turn2-user", "turn2-asst")
p._retain_queue.join()

# Every buffered turn is already appended β€” watermark == buffer len.
assert p._last_retained_turn_count == len(p._session_turns) == 2

p._client.aretain_batch.reset_mock()

p.on_session_switch("new-sid", parent_session_id="test-session", reset=True)
p._retain_queue.join()

# No turns past the watermark β†’ no flush retain. Without the fix this
# would re-ship [turn1, turn2] under update_mode='append'.
p._client.aretain_batch.assert_not_called()
self._clear_capability_cache()

def test_session_switch_flush_ships_only_unretained_delta_in_append_mode(
self, provider_with_config, monkeypatch
):
"""When retain_every_n_turns>1 some turns are retained and some are
still buffered at switch time. The flush must carry only the buffered
delta past the watermark, never the already-appended turns."""
self._clear_capability_cache()
monkeypatch.setattr(
"plugins.memory.hindsight._fetch_hindsight_api_version",
lambda *a, **kw: "0.5.6",
)
p = provider_with_config(retain_every_n_turns=2, retain_async=False)
# Turns 1+2 hit the boundary β†’ retained as a delta; watermark = 2.
p.sync_turn("turn1-user", "turn1-asst")
p.sync_turn("turn2-user", "turn2-asst")
# Turn 3 is buffered, below the next boundary, so not yet retained.
p.sync_turn("turn3-user", "turn3-asst")
p._retain_queue.join()
assert p._last_retained_turn_count == 2
assert len(p._session_turns) == 3

p._client.aretain_batch.reset_mock()

p.on_session_switch("new-sid", parent_session_id="test-session", reset=True)
p._retain_queue.join()

p._client.aretain_batch.assert_called_once()
item = p._client.aretain_batch.call_args.kwargs["items"][0]
assert item["update_mode"] == "append"
# Only the un-retained turn 3 β€” turns 1+2 were already appended.
assert "turn3-user" in item["content"]
assert "turn1-user" not in item["content"]
assert "turn2-user" not in item["content"]
# message_count reflects only the flushed delta (1 turn -> 2 messages).
assert item["metadata"]["message_count"] == "2"
self._clear_capability_cache()


# ---------------------------------------------------------------------------
# System prompt tests
Expand Down