fix(memory): fail fast on stuck external prefetch

This commit is contained in:
LeonSGP43 2026-07-11 22:18:01 -07:00 committed by Teknium
parent 2ad6ab17e3
commit d77c455d7d
2 changed files with 120 additions and 2 deletions

View file

@ -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.

View file

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