diff --git a/plugins/memory/honcho/session.py b/plugins/memory/honcho/session.py index e099da30fe..2ab6bebc56 100644 --- a/plugins/memory/honcho/session.py +++ b/plugins/memory/honcho/session.py @@ -402,6 +402,21 @@ class HonchoSessionManager(SessionAuthMixin, SessionPeersMixin, SessionContextMi with session._flush_lock: return self._flush_session_locked(session) + def _flush_session_before(self, session: HonchoSession, deadline: float | None) -> bool: + """_flush_session bounded by ``deadline``. False, with nothing sent, when the budget is spent or another + flush of this session holds the lock past it.""" + if deadline is None: + self._flush_session(session) + return True + remaining = deadline - time.monotonic() + if remaining <= 0 or not session._flush_lock.acquire(timeout=remaining): + return False + try: + self._flush_session_locked(session) + finally: + session._flush_lock.release() + return True + def _flush_session_locked(self, session: HonchoSession) -> bool: new_messages = [m for m in session.messages if not m.get("_synced")] if not new_messages: @@ -499,27 +514,39 @@ class HonchoSessionManager(SessionAuthMixin, SessionPeersMixin, SessionContextMi else: self._keep_until_flushed(session) - def flush_all(self) -> None: - """Flush unsynced messages for all cached sessions, then drain the async queue inline.""" + def flush_all(self, timeout: float | None = None) -> None: + """Flush unsynced messages for all cached sessions, then drain the async queue inline. ``timeout`` bounds + the whole pass: a session it cannot reach in time keeps its messages and is counted in one warning.""" + deadline = None if timeout is None else time.monotonic() + timeout with self._cache_lock: sessions = list(self._cache.values()) + skipped: list[HonchoSession] = [] for session in sessions: try: - self._flush_session(session) + if not self._flush_session_before(session, deadline): + skipped.append(session) except Exception as e: logger.error("Honcho flush_all error for %s: %s", session.key, e) - self._drain_async_queue() + skipped.extend(self._drain_async_queue(deadline)) + left = [s for s in {id(s): s for s in skipped}.values() if self._has_unsynced(s)] + if left: + unsynced = sum(1 for s in left for m in list(s.messages) if not m.get("_synced")) + logger.warning("Honcho flush ran out of time after %.1fs with %d message(s) in %d session(s) still unsynced", + timeout, unsynced, len(left)) - def _drain_async_queue(self) -> None: + def _drain_async_queue(self, deadline: float | None = None) -> list[HonchoSession]: + """Flush every queued session inline. Returns the sessions ``deadline`` left unflushed.""" + skipped: list[HonchoSession] = [] if self._async_queue is None: - return + return skipped while not self._async_queue.empty(): try: item = self._async_queue.get_nowait() - if item is not _ASYNC_SHUTDOWN: - self._flush_session(item) except queue.Empty: break + if item is not _ASYNC_SHUTDOWN and not self._flush_session_before(item, deadline): + skipped.append(item) + return skipped def _ensure_async_writer(self) -> None: """Start the async writer on first enqueue (idempotent, thread-safe).""" @@ -544,11 +571,12 @@ class HonchoSessionManager(SessionAuthMixin, SessionPeersMixin, SessionContextMi self._drain_async_queue() def shutdown(self, timeout: float = 10.0) -> None: - """Flush everything, then stop the async writer thread.""" + """Flush everything, then stop the async writer thread, both within ``timeout``.""" self._shutting_down = True if self._async_queue is not None: - self.flush_all() - self.stop_async_writer(timeout=timeout) + deadline = time.monotonic() + timeout + self.flush_all(timeout=timeout) + self.stop_async_writer(timeout=max(0.0, deadline - time.monotonic())) # ----- Prefetch cache ----- diff --git a/tests/honcho_plugin/test_async_memory.py b/tests/honcho_plugin/test_async_memory.py index e3d6d9d4c3..aaddf249e3 100644 --- a/tests/honcho_plugin/test_async_memory.py +++ b/tests/honcho_plugin/test_async_memory.py @@ -377,14 +377,14 @@ class TestStopAsyncWriterDrain: assert flushed == ["after"] assert mgr._async_queue.empty() - def test_shutdown_passes_the_join_timeout_through(self, make_manager, monkeypatch): + def test_shutdown_gives_the_writer_join_what_the_flush_left_of_the_timeout(self, make_manager, monkeypatch): mgr = make_manager("async") seen = {} monkeypatch.setattr(mgr, "stop_async_writer", lambda timeout=10.0: seen.setdefault("timeout", timeout)) mgr.shutdown(timeout=2.5) - assert seen["timeout"] == 2.5 + assert 2.0 < seen["timeout"] <= 2.5 class TestAsyncWriterRetry: diff --git a/tests/honcho_plugin/test_shutdown.py b/tests/honcho_plugin/test_shutdown.py index d6ae2a61af..1d99bdd8b2 100644 --- a/tests/honcho_plugin/test_shutdown.py +++ b/tests/honcho_plugin/test_shutdown.py @@ -17,7 +17,7 @@ from plugins.memory.honcho.client import ( spawn_context_thread, ) from plugins.memory.honcho.client_cache import _client_slots, _client_slots_lock -from plugins.memory.honcho.session import HonchoSessionManager +from plugins.memory.honcho.session import HonchoSession, HonchoSessionManager from plugins.plugin_utils import SingletonSlot @@ -157,6 +157,31 @@ class TestProviderShutdown: release.set() provider._recall_sync_thread.join(timeout=2) + def test_shutdown_skips_a_flush_whose_lock_an_upload_holds_past_the_budget(self, async_manager, monkeypatch, caplog): + provider = self._provider(async_manager) + monkeypatch.setattr(provider, "_shutdown_join_budget", lambda: 0.2) + session = HonchoSession("test-session", "user", "assistant", "sid") + session.add_message("user", "pending") + async_manager._cache["test-session"] = session + entered, release = threading.Event(), threading.Event() + + def upload(messages): + entered.set() + release.wait(timeout=5) + + async_manager._sessions_cache["sid"] = SimpleNamespace(add_messages=upload) + async_manager.save(session) + try: + assert entered.wait(timeout=1) + started = time.monotonic() + with caplog.at_level(logging.WARNING, logger="plugins.memory.honcho"): + provider.shutdown() + assert time.monotonic() - started < 1.0 + assert "1 message(s) in 1 session(s) still unsynced" in caplog.text + assert session.messages[0].get("_synced") is None + finally: + release.set() + def test_shutdown_does_not_close_the_shared_client(self, async_manager): provider = self._provider(async_manager) provider.shutdown()