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
+29 -25
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