diff --git a/EvoScientist/middleware/memory.py b/EvoScientist/middleware/memory.py index 5b64fac..6f51b16 100644 --- a/EvoScientist/middleware/memory.py +++ b/EvoScientist/middleware/memory.py @@ -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, diff --git a/tests/test_profile_memory_middleware.py b/tests/test_profile_memory_middleware.py index a5b3c5d..49452ea 100644 --- a/tests/test_profile_memory_middleware.py +++ b/tests/test_profile_memory_middleware.py @@ -417,9 +417,9 @@ async def test_profile_memory_async_path_inlines_content_under_blockbuster( call_threads = [] original_read = middleware._read_profile_memory - def tracked_read_profile_memory(): + def tracked_read_profile_memory(memory_dir=None): call_threads.append(threading.get_ident()) - return original_read() + return original_read(memory_dir) monkeypatch.setattr(middleware, "_read_profile_memory", tracked_read_profile_memory) @@ -625,3 +625,23 @@ def test_profile_memory_skips_legacy_unknown_placeholders(tmp_path, monkeypatch) assert "(unknown)" not in migrated_profile_text assert "Imported from legacy MEMORY.md" not in migrated_profile_text assert not (memories / "MEMORY.md").exists() + + +def test_read_profile_memory_honors_memory_dir_override(tmp_path, monkeypatch): + from EvoScientist import paths + + global_mem = tmp_path / "global" + user_mem = tmp_path / "user" + workspace = tmp_path / "workspace" + workspace.mkdir() + monkeypatch.setattr(paths, "WORKSPACE_ROOT", workspace) + + middleware = memory_module.create_memory_middleware(str(global_mem)) + middleware.modify_request(_request()) # bootstrap global + (user_mem / "profile").mkdir(parents=True) + (user_mem / "profile" / "USER_PROFILE.md").write_text( + "# User profile\n\nper-user marker", encoding="utf-8" + ) + + content = middleware._read_profile_memory(memory_dir=user_mem) + assert "per-user marker" in content