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")