From 1b76cfff836e117a971bfa084399eb867958a397 Mon Sep 17 00:00:00 2001 From: Mahesh Sanikommu Date: Sat, 12 Sep 2026 12:23:00 -0700 Subject: [PATCH] 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. --- plugins/memory/supermemory/__init__.py | 54 ++++++++++--------- .../memory/test_supermemory_provider.py | 28 +++++++++- 2 files changed, 56 insertions(+), 26 deletions(-) diff --git a/plugins/memory/supermemory/__init__.py b/plugins/memory/supermemory/__init__.py index ac168a6846..fc3c639ac5 100644 --- a/plugins/memory/supermemory/__init__.py +++ b/plugins/memory/supermemory/__init__.py @@ -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 diff --git a/tests/plugins/memory/test_supermemory_provider.py b/tests/plugins/memory/test_supermemory_provider.py index 5993d5fbd2..c4d00baf0d 100644 --- a/tests/plugins/memory/test_supermemory_provider.py +++ b/tests/plugins/memory/test_supermemory_provider.py @@ -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")