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:
committed by
kshitij
parent
03627dbf95
commit
1b76cfff83
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user