From 8d1c96fd2f60c8b47bd50ba8817a3fd1c23f834c Mon Sep 17 00:00:00 2001 From: Erosika Date: Thu, 16 Jul 2026 14:32:25 -0400 Subject: [PATCH] fix(memory): align external prefetch guard with fail-open contracts --- agent/memory_manager.py | 63 ++++++++++++----------------- tests/agent/test_memory_provider.py | 14 +++++++ 2 files changed, 39 insertions(+), 38 deletions(-) diff --git a/agent/memory_manager.py b/agent/memory_manager.py index fb4fe4cc59ee..37a0411ade88 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -29,7 +29,6 @@ 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 @@ -45,6 +44,7 @@ logger = logging.getLogger(__name__) # teardown indefinitely — the worker threads are daemon, so anything still # running past this window dies with the interpreter. _SYNC_DRAIN_TIMEOUT_S = 5.0 +_EXTERNAL_PREFETCH_TIMEOUT_S = 8.0 def normalize_tool_schema(schema: Any) -> Optional[Dict[str, Any]]: @@ -362,10 +362,15 @@ class MemoryManager: 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._external_prefetch_timeout = ( + _EXTERNAL_PREFETCH_TIMEOUT_S + if external_prefetch_timeout is None + else float(external_prefetch_timeout) ) - self._stuck_prefetch_threads: Dict[str, threading.Thread] = {} + if self._external_prefetch_timeout <= 0: + raise ValueError("external_prefetch_timeout must be positive") + self._external_prefetch_threads: Dict[str, threading.Thread] = {} + self._external_prefetch_lock = threading.Lock() # 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 @@ -374,27 +379,6 @@ 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: @@ -546,17 +530,6 @@ class MemoryManager: 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] = {} @@ -571,10 +544,21 @@ class MemoryManager: daemon=True, name=f"memory-prefetch-{provider.name}", ) - thread.start() + with self._external_prefetch_lock: + existing = self._external_prefetch_threads.get(provider.name) + if existing is not None: + if existing.is_alive(): + logger.debug( + "Memory provider '%s' prefetch is still running; skipping this turn", + provider.name, + ) + return "" + self._external_prefetch_threads.pop(provider.name, None) + self._external_prefetch_threads[provider.name] = thread + 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", @@ -583,6 +567,9 @@ class MemoryManager: ) return "" + with self._external_prefetch_lock: + if self._external_prefetch_threads.get(provider.name) is thread: + self._external_prefetch_threads.pop(provider.name, None) if error_box: raise error_box["value"] return result_box.get("value", "") diff --git a/tests/agent/test_memory_provider.py b/tests/agent/test_memory_provider.py index 2bfd4e7a5e40..01e68b279dd3 100644 --- a/tests/agent/test_memory_provider.py +++ b/tests/agent/test_memory_provider.py @@ -444,6 +444,20 @@ class TestMemoryManager: external.release.set() + deadline = time.monotonic() + 1.0 + while ( + external.name in mgr._external_prefetch_threads + and mgr._external_prefetch_threads[external.name].is_alive() + and time.monotonic() < deadline + ): + time.sleep(0.01) + + result = mgr.prefetch_all("query 3") + + assert result == "builtin memory\n\nlate external memory" + assert external.prefetch_queries == ["query", "query 3"] + assert external.name not in mgr._external_prefetch_threads + def test_system_prompt_failure_doesnt_block(self): mgr = MemoryManager() p1 = FakeMemoryProvider("builtin")