fix(supermemory): keep pending turns across session switch until written

A failed flush at on_session_switch() restored the pending buffer and
then cleared it on the next line, so an unavailable service at the switch
boundary still lost every pending turn. Pending turns now carry their own
session_id; _write_turns() batches per session, so a later retry (next
turn, session end, shutdown) writes old-session turns under the old
session's custom_id even after the switch.

Addresses review on #109359.
This commit is contained in:
Mahesh Sanikommu
2026-09-12 12:23:00 -07:00
committed by kshitij
parent 03627dbf95
commit 1b76cfff83
2 changed files with 56 additions and 26 deletions

View File

@@ -315,7 +315,7 @@ class SupermemoryMemoryProvider(MemoryProvider):
self._client: Optional[_SupermemoryClient] = None
self._container_tag, self._turn_count, self._write_enabled, self._active = _DEFAULT_CONTAINER_TAG, 0, True, False
self._prefetch_thread = self._sync_thread = self._write_thread = None # only _write_thread is ever started
self._pending_turns: List[Dict[str, str]] = [] # turns whose write failed; retried on next write/end/shutdown
self._pending_turns: List[Dict[str, str]] = [] # failed writes, each tagged with its session_id; retried on next write/end/switch/shutdown
self._apply_config(_load_supermemory_config())
self._base_url, self._allowed_containers = _DEFAULT_BASE_URL, [] # env var is only consulted in initialize()
@@ -411,43 +411,47 @@ class SupermemoryMemoryProvider(MemoryProvider):
profile["search_results"], self._max_recall_results)
return _quietly(_recall, "Supermemory prefetch failed", default="")
def _write_turns(self, turns: List[Dict[str, str]], session_id: str, mode: str) -> None:
"""Append turns to the session's current 4h document (shared custom_id, API merges deltas). Failures stay pending."""
now = datetime.now(timezone.utc)
content = "\n\n".join(_format_turn(t["user"], t["assistant"]) for t in turns)
metadata = {"type": "conversation", "session_id": session_id, "timestamp": now.isoformat()} # no sm_capture_mode: Hermes policy
try:
self._client.add_memory(content, metadata=metadata, entity_context=self._entity_context,
custom_id=_capture_custom_id(session_id, now))
self._pending_turns = []
except Exception:
logger.log(logging.WARNING if mode != "turn" else logging.DEBUG, "Supermemory capture failed (%s, %d turns pending)",
mode, len(turns), exc_info=True)
self._pending_turns = turns
def _write_turns(self, turns: List[Dict[str, str]], mode: str) -> None:
"""One documents.add per session id in ``turns`` (custom_id = session + 4h bucket, so the API appends deltas).
Failed batches stay pending under their own session id, so a switch never re-homes them."""
failed: List[Dict[str, str]] = []
for sid in dict.fromkeys(t["session_id"] for t in turns):
batch = [t for t in turns if t["session_id"] == sid]
now = datetime.now(timezone.utc)
content = "\n\n".join(_format_turn(t["user"], t["assistant"]) for t in batch)
metadata = {"type": "conversation", "session_id": sid, "timestamp": now.isoformat()} # no sm_capture_mode: Hermes policy
try:
self._client.add_memory(content, metadata=metadata, entity_context=self._entity_context,
custom_id=_capture_custom_id(sid, now))
except Exception:
logger.log(logging.WARNING if mode != "turn" else logging.DEBUG, "Supermemory capture failed (%s, session=%s, %d turns pending)",
mode, sid, len(batch), exc_info=True)
failed += batch
self._pending_turns = failed
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None:
# Host runs this on a worker thread, so the blocking write is fine here.
if not self._can_write() or not self._auto_capture:
return
turn = {"user": _clean_text_for_capture(user_content), "assistant": _clean_text_for_capture(assistant_content)}
if any(turn.values()):
self._write_turns(self._pending_turns + [turn], session_id or self._session_id, "turn")
turn = {"user": _clean_text_for_capture(user_content), "assistant": _clean_text_for_capture(assistant_content),
"session_id": session_id or self._session_id}
if turn["user"] or turn["assistant"]:
self._write_turns(self._pending_turns + [turn], "turn")
def _flush_pending(self, session_id: str, mode: str) -> None:
def _flush_pending(self, mode: str) -> None:
if self._can_write() and self._pending_turns:
self._write_turns(list(self._pending_turns), session_id, mode)
self._write_turns(list(self._pending_turns), mode)
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
# Turns were already written as they completed; only retry what failed.
self._flush_pending(self._session_id, "session_end")
self._flush_pending("session_end")
def on_session_switch(self, new_session_id: str, *, parent_session_id: str = "", reset: bool = False, **kwargs) -> None:
old_session_id = self._session_id
self._flush_pending(old_session_id, "session_switch")
# Pending turns survive the switch: they carry their own session_id, so a later retry still lands on the old session.
self._flush_pending("session_switch")
if self._can_write():
self._turn_count = 0
self._session_id = str(new_session_id or "").strip() or old_session_id
self._pending_turns = []
self._session_id = str(new_session_id or "").strip() or self._session_id
def on_memory_write(self, action: str, target: str, content: str) -> None:
if not self._can_write() or action != "add" or not (content or "").strip():
@@ -461,7 +465,7 @@ class SupermemoryMemoryProvider(MemoryProvider):
self._write_thread.start()
def shutdown(self) -> None:
self._flush_pending(self._session_id, "shutdown")
self._flush_pending("shutdown")
if self._write_thread and self._write_thread.is_alive():
self._write_thread.join(timeout=5.0)
self._prefetch_thread = self._sync_thread = self._write_thread = None

View File

@@ -148,7 +148,7 @@ def test_failed_turn_write_is_retried_at_session_end(provider):
provider._client.fail_add = True
provider.sync_turn("hello", "hi there", session_id="session-1")
assert provider._client.add_calls == []
assert provider._pending_turns == [{"user": "hello", "assistant": "hi there"}]
assert provider._pending_turns == [{"user": "hello", "assistant": "hi there", "session_id": "session-1"}]
provider._client.fail_add = False
provider.on_session_end([])
@@ -179,6 +179,32 @@ def test_session_switch_flushes_pending_to_old_session(provider):
assert provider._pending_turns == []
def test_failed_switch_flush_keeps_old_session_turns_for_later_retry(provider):
provider._client.fail_add = True
provider.sync_turn("old turn", "old reply", session_id="session-1")
provider.on_session_switch("session-2", reset=True) # flush fails: service unavailable at the boundary
assert provider._session_id == "session-2"
assert provider._pending_turns == [{"user": "old turn", "assistant": "old reply", "session_id": "session-1"}]
provider._client.fail_add = False
provider.sync_turn("new turn", "new reply", session_id="session-2")
calls = provider._client.add_calls
assert [c["custom_id"] for c in calls] == [_capture_custom_id("session-1"), _capture_custom_id("session-2")]
assert calls[0]["metadata"]["session_id"] == "session-1" and "old turn" in calls[0]["content"]
assert calls[1]["metadata"]["session_id"] == "session-2" and "new turn" in calls[1]["content"]
assert provider._pending_turns == []
def test_failed_switch_flush_is_retried_at_shutdown(provider):
provider._client.fail_add = True
provider.sync_turn("old turn", "old reply", session_id="session-1")
provider.on_session_switch("session-2", reset=True)
provider._client.fail_add = False
provider.shutdown()
assert provider._client.add_calls[0]["custom_id"] == _capture_custom_id("session-1")
assert provider._pending_turns == []
def test_sync_turn_drops_inline_image_payloads(provider):
blob = "A" * 4096
provider.sync_turn(f"describe this data:image/png;base64,{blob}", "a screenshot", session_id="session-1")