diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 32dd8e288841..fb4fe4cc59ee 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -29,6 +29,7 @@ import json import logging import re import inspect +import os import threading from concurrent.futures import ThreadPoolExecutor from typing import Any, Callable, Dict, List, Optional @@ -357,10 +358,14 @@ class MemoryManager: provider is allowed. Failures in one provider never block the other. """ - def __init__(self) -> None: + def __init__(self, *, external_prefetch_timeout: Optional[float] = None) -> None: self._providers: List[MemoryProvider] = [] self._tool_to_provider: Dict[str, MemoryProvider] = {} self._has_external: bool = False # True once a non-builtin provider is added + self._external_prefetch_timeout = self._resolve_external_prefetch_timeout( + external_prefetch_timeout + ) + self._stuck_prefetch_threads: Dict[str, threading.Thread] = {} # Background executor for end-of-turn sync/prefetch. Lazily created on # first use so the common builtin-only path spawns no extra threads. # A single worker serializes a provider's writes (turn N must land @@ -369,6 +374,27 @@ class MemoryManager: self._sync_executor: Optional[ThreadPoolExecutor] = None self._sync_executor_lock = threading.Lock() + @staticmethod + def _resolve_external_prefetch_timeout(timeout: Optional[float]) -> float: + raw = timeout + if raw is None: + raw = os.getenv("HERMES_EXTERNAL_MEMORY_PREFETCH_TIMEOUT", "5") + try: + value = float(raw) + except (TypeError, ValueError): + logger.warning( + "Invalid external memory prefetch timeout %r; using 5.0s", + raw, + ) + return 5.0 + if value <= 0: + logger.warning( + "Non-positive external memory prefetch timeout %r; using 5.0s", + raw, + ) + return 5.0 + return value + # -- Registration -------------------------------------------------------- def add_provider(self, provider: MemoryProvider) -> None: @@ -504,7 +530,7 @@ class MemoryManager: parts = [] for provider in self._providers: try: - result = provider.prefetch(clean_query, session_id=session_id) + result = self._prefetch_provider(provider, clean_query, session_id=session_id) if result and result.strip(): parts.append(result) except Exception as e: @@ -514,6 +540,53 @@ class MemoryManager: ) return "\n\n".join(parts) + def _prefetch_provider( + self, provider: MemoryProvider, query: str, *, session_id: str = "" + ) -> str: + if provider.name == "builtin": + return provider.prefetch(query, session_id=session_id) + + existing = self._stuck_prefetch_threads.get(provider.name) + if existing is not None: + if existing.is_alive(): + logger.debug( + "Memory provider '%s' prefetch is still stuck from an earlier timeout; " + "skipping this turn", + provider.name, + ) + return "" + self._stuck_prefetch_threads.pop(provider.name, None) + + result_box: Dict[str, str] = {} + error_box: Dict[str, Exception] = {} + + def _run() -> None: + try: + result_box["value"] = provider.prefetch(query, session_id=session_id) or "" + except Exception as exc: # pragma: no cover - re-raised by caller + error_box["value"] = exc + + thread = threading.Thread( + target=_run, + daemon=True, + name=f"memory-prefetch-{provider.name}", + ) + thread.start() + thread.join(self._external_prefetch_timeout) + if thread.is_alive(): + self._stuck_prefetch_threads[provider.name] = thread + logger.warning( + "Memory provider '%s' prefetch timed out after %.1fs; skipping it until " + "the stuck call returns", + provider.name, + self._external_prefetch_timeout, + ) + return "" + + if error_box: + raise error_box["value"] + return result_box.get("value", "") + def queue_prefetch_all(self, query: str, *, session_id: str = "") -> None: """Queue background prefetch on all providers for the next turn. diff --git a/tests/agent/test_memory_provider.py b/tests/agent/test_memory_provider.py index 1a99977deb48..2bfd4e7a5e40 100644 --- a/tests/agent/test_memory_provider.py +++ b/tests/agent/test_memory_provider.py @@ -1,6 +1,8 @@ """Tests for the memory provider interface, manager, and builtin provider.""" import json +import threading +import time import pytest from types import SimpleNamespace from unittest.mock import MagicMock @@ -92,6 +94,21 @@ class MessagesMemoryProvider(FakeMemoryProvider): self.synced_turns.append((user_content, assistant_content, session_id, messages)) +class BlockingPrefetchProvider(FakeMemoryProvider): + """External provider whose prefetch call blocks until released.""" + + def __init__(self, name="external"): + super().__init__(name=name) + self.started = threading.Event() + self.release = threading.Event() + + def prefetch(self, query, *, session_id=""): + self.prefetch_queries.append(query) + self.started.set() + self.release.wait(timeout=5.0) + return self._prefetch_result + + # --------------------------------------------------------------------------- # MemoryProvider ABC tests # --------------------------------------------------------------------------- @@ -399,6 +416,34 @@ class TestMemoryManager: result = mgr.prefetch_all("query") assert "external memory" in result + def test_external_prefetch_timeout_skips_stuck_provider(self): + mgr = MemoryManager(external_prefetch_timeout=0.01) + builtin = FakeMemoryProvider("builtin") + builtin._prefetch_result = "builtin memory" + external = BlockingPrefetchProvider("hy-memory") + external._prefetch_result = "late external memory" + mgr.add_provider(builtin) + mgr.add_provider(external) + + started = time.monotonic() + result = mgr.prefetch_all("query") + elapsed = time.monotonic() - started + + assert result == "builtin memory" + assert elapsed < 0.5 + assert external.started.wait(timeout=1.0) + assert external.prefetch_queries == ["query"] + + started = time.monotonic() + result = mgr.prefetch_all("query 2") + elapsed = time.monotonic() - started + + assert result == "builtin memory" + assert elapsed < 0.2 + assert external.prefetch_queries == ["query"] + + external.release.set() + def test_system_prompt_failure_doesnt_block(self): mgr = MemoryManager() p1 = FakeMemoryProvider("builtin")