feat: resolve memory middleware dir per-call from configurable
This commit is contained in:
@@ -306,9 +306,22 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
)
|
||||
)
|
||||
|
||||
def _file_path(self, memory_path: str) -> Path:
|
||||
def _runtime_memory_dir(self) -> str | None:
|
||||
try:
|
||||
from langgraph.config import get_config
|
||||
cfg = get_config()
|
||||
except Exception:
|
||||
return None
|
||||
configurable = cfg.get("configurable") if isinstance(cfg, dict) else None
|
||||
if not isinstance(configurable, dict):
|
||||
return None
|
||||
value = configurable.get("ai4sci_memory_dir")
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
def _file_path(self, memory_path: str, memory_dir: str | Path | None = None) -> Path:
|
||||
"""Resolve a memory-relative path against the memory directory."""
|
||||
return self._memory_dir / memory_path.lstrip("/")
|
||||
root = Path(memory_dir) if memory_dir is not None else self._memory_dir
|
||||
return root / memory_path.lstrip("/")
|
||||
|
||||
def _read_text(self, path: Path) -> str | None:
|
||||
"""Read UTF-8 text, returning None only when the file is absent."""
|
||||
@@ -341,18 +354,18 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
return False
|
||||
return True
|
||||
|
||||
def _ensure_observation_dirs(self) -> None:
|
||||
def _ensure_observation_dirs(self, memory_dir: str | Path | None = None) -> None:
|
||||
"""Create non-project observation directories agents are prompted to search."""
|
||||
try:
|
||||
self._file_path("/observations/global").mkdir(parents=True, exist_ok=True)
|
||||
self._file_path("/observations/global", memory_dir).mkdir(parents=True, exist_ok=True)
|
||||
except OSError as e:
|
||||
logger.warning("Failed to create observation memory dir: %s", e)
|
||||
|
||||
def _ensure_profile_files(self) -> list[tuple[str, str]]:
|
||||
def _ensure_profile_files(self, memory_dir: str | Path | None = None) -> list[tuple[str, str]]:
|
||||
"""Create the expected profile files if needed and return their contents."""
|
||||
records = []
|
||||
for memory_path, template in self._profile_specs:
|
||||
path = self._file_path(memory_path)
|
||||
path = self._file_path(memory_path, memory_dir)
|
||||
content = self._read_text(path)
|
||||
if content is None:
|
||||
if not self._write_text(path, template):
|
||||
@@ -361,13 +374,13 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
records.append((memory_path, content))
|
||||
return records
|
||||
|
||||
def _migrate_legacy_memory(self) -> bool:
|
||||
def _migrate_legacy_memory(self, memory_dir: str | Path | None = None) -> bool:
|
||||
"""Import recognized sections from legacy ``MEMORY.md`` into profiles.
|
||||
|
||||
The legacy file is removed only after real content is copied or the file
|
||||
is found to contain only old template placeholders.
|
||||
"""
|
||||
legacy_path = self._memory_dir / _LEGACY_MEMORY_FILENAME
|
||||
legacy_path = (Path(memory_dir) if memory_dir is not None else self._memory_dir) / _LEGACY_MEMORY_FILENAME
|
||||
legacy = self._read_text(legacy_path)
|
||||
if legacy is None:
|
||||
return True
|
||||
@@ -402,7 +415,7 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
for memory_path, bodies in imports.items():
|
||||
if not bodies:
|
||||
continue
|
||||
path = self._file_path(memory_path)
|
||||
path = self._file_path(memory_path, memory_dir)
|
||||
content = self._read_text(path)
|
||||
if content is None:
|
||||
logger.warning(
|
||||
@@ -419,22 +432,22 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
|
||||
return self._delete_legacy_memory(legacy_path)
|
||||
|
||||
def _read_bootstrapped_profile_records(self) -> list[tuple[str, str]]:
|
||||
records = self._ensure_profile_files()
|
||||
if self._migrate_legacy_memory():
|
||||
def _read_bootstrapped_profile_records(self, memory_dir=None) -> list[tuple[str, str]]:
|
||||
records = self._ensure_profile_files(memory_dir)
|
||||
if self._migrate_legacy_memory(memory_dir):
|
||||
records = [
|
||||
(memory_path, self._read_text(self._file_path(memory_path)) or "")
|
||||
(memory_path, self._read_text(self._file_path(memory_path, memory_dir)) or "")
|
||||
for memory_path, _ in records
|
||||
]
|
||||
return records
|
||||
|
||||
def _read_profile_records(self) -> list[tuple[str, str]]:
|
||||
def _read_profile_records(self, memory_dir=None) -> list[tuple[str, str]]:
|
||||
"""Load all profile files after bootstrapping and legacy migration."""
|
||||
if not self._enable_observation_memory:
|
||||
return self._read_bootstrapped_profile_records()
|
||||
return self._read_bootstrapped_profile_records(memory_dir)
|
||||
|
||||
self._ensure_observation_dirs()
|
||||
return self._read_bootstrapped_profile_records()
|
||||
self._ensure_observation_dirs(memory_dir)
|
||||
return self._read_bootstrapped_profile_records(memory_dir)
|
||||
|
||||
def _profile_context_from_records(self, records: list[tuple[str, str]]) -> str:
|
||||
"""Inline profile contents unless they exceed the prompt budget."""
|
||||
@@ -447,10 +460,10 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
return full
|
||||
return self._profile_pointer_context
|
||||
|
||||
def _read_profile_memory(self) -> str:
|
||||
def _read_profile_memory(self, memory_dir=None) -> str:
|
||||
"""Return profile context, falling back to file pointers."""
|
||||
try:
|
||||
records = self._read_profile_records()
|
||||
records = self._read_profile_records(memory_dir)
|
||||
return (
|
||||
self._profile_context_from_records(records)
|
||||
or self._profile_pointer_context
|
||||
@@ -459,14 +472,14 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
logger.debug("Failed to read profile memory: %s", e)
|
||||
return self._profile_pointer_context
|
||||
|
||||
def _refresh_observation_index_context(self) -> str:
|
||||
def _refresh_observation_index_context(self, memory_dir=None) -> str:
|
||||
"""Refresh the prompt observation index from current memory files."""
|
||||
if not self._enable_observation_memory:
|
||||
return ""
|
||||
try:
|
||||
self._ensure_observation_dirs()
|
||||
self._ensure_observation_dirs(memory_dir)
|
||||
context = build_observation_index_context(
|
||||
memory_dir=self._memory_dir,
|
||||
memory_dir=(Path(memory_dir) if memory_dir is not None else self._memory_dir),
|
||||
project_id=self._project_id,
|
||||
)
|
||||
except OSError as e:
|
||||
@@ -570,20 +583,21 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
|
||||
async def amodify_request(self, request: ModelRequest) -> ModelRequest:
|
||||
"""Apply memory injection for asynchronous model calls."""
|
||||
memory_dir = self._runtime_memory_dir()
|
||||
observation_index_context = ""
|
||||
profile_context = ""
|
||||
|
||||
if self._enable_observation_memory and self._enable_profile_memory:
|
||||
observation_index_context, profile_context = await asyncio.gather(
|
||||
asyncio.to_thread(self._refresh_observation_index_context),
|
||||
asyncio.to_thread(self._read_profile_memory),
|
||||
asyncio.to_thread(self._refresh_observation_index_context, memory_dir),
|
||||
asyncio.to_thread(self._read_profile_memory, memory_dir),
|
||||
)
|
||||
elif self._enable_observation_memory:
|
||||
observation_index_context = await asyncio.to_thread(
|
||||
self._refresh_observation_index_context
|
||||
self._refresh_observation_index_context, memory_dir
|
||||
)
|
||||
elif self._enable_profile_memory:
|
||||
profile_context = await asyncio.to_thread(self._read_profile_memory)
|
||||
profile_context = await asyncio.to_thread(self._read_profile_memory, memory_dir)
|
||||
|
||||
return self._inject_memory_context(
|
||||
request,
|
||||
|
||||
Reference in New Issue
Block a user