From 1cd957c14fa48677e4112a3739bed73bcf9f32ed Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 19 Feb 2026 23:58:54 +0800 Subject: [PATCH 1/3] feat(memory_store): add pure-python local memory store implementation --- reme/__init__.py | 2 +- reme/config/cli.yaml | 3 +- reme/config/fs.yaml | 3 +- reme/core/context/service_context.py | 16 +- reme/core/embedding/base_embedding_model.py | 112 +++- reme/core/embedding/openai_embedding_model.py | 1 + .../embedding/openai_embedding_model_sync.py | 1 + reme/core/llm/openai_llm.py | 1 + reme/core/llm/openai_llm_sync.py | 1 + reme/core/memory_store/__init__.py | 6 +- reme/core/memory_store/chroma_memory_store.py | 1 + reme/core/memory_store/local_memory_store.py | 502 ++++++++++++++++++ reme/core/memory_store/sqlite_memory_store.py | 1 + tests/test_embedding_cache.py | 483 +++++++++-------- tests/test_memory_store.py | 59 +- 15 files changed, 951 insertions(+), 241 deletions(-) create mode 100644 reme/core/memory_store/local_memory_store.py diff --git a/reme/__init__.py b/reme/__init__.py index c280ab02..100fa70f 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -20,7 +20,7 @@ __all__ = [ "ReMeFs", ] -__version__ = "0.3.0.0b1" +__version__ = "0.3.0.0b2" """ diff --git a/reme/config/cli.yaml b/reme/config/cli.yaml index a126ccfd..2f259869 100644 --- a/reme/config/cli.yaml +++ b/reme/config/cli.yaml @@ -22,7 +22,8 @@ embedding_models: memory_stores: default: - backend: chroma +# backend: chroma + backend: local db_name: reme.db store_name: reme embedding_model: default diff --git a/reme/config/fs.yaml b/reme/config/fs.yaml index dc07e41d..2d92c7bc 100644 --- a/reme/config/fs.yaml +++ b/reme/config/fs.yaml @@ -18,7 +18,8 @@ embedding_models: memory_stores: default: # backend: sqlite - backend: chroma +# backend: chroma + backend: local db_name: reme.db store_name: reme embedding_model: default diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index 3f5a6403..febf956d 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -193,6 +193,7 @@ class ServiceContext(BaseContext): logger.warning(f"Embedding model backend {config.backend} is not supported.") else: self.embedding_models[name] = R.embedding_models[config.backend]( + cache_dir=self.working_path / "embedding_cache", model_name=config.model_name, **config.model_extra, ) @@ -276,19 +277,24 @@ class ServiceContext(BaseContext): async def close(self): """Close all service components asynchronously.""" - for _, vector_store in self.vector_stores.items(): + for name, vector_store in self.vector_stores.items(): + logger.info(f"Closing vector store: {name}") await vector_store.close() - for _, memory_store in self.memory_stores.items(): + for name, memory_store in self.memory_stores.items(): + logger.info(f"Closing memory store: {name}") await memory_store.close() - for _, file_watcher in self.file_watchers.items(): + for name, file_watcher in self.file_watchers.items(): + logger.info(f"Closing file watcher: {name}") await file_watcher.close() - for _, llm in self.llms.items(): + for name, llm in self.llms.items(): + logger.info(f"Closing LLM: {name}") await llm.close() - for _, embedding_model in self.embedding_models.items(): + for name, embedding_model in self.embedding_models.items(): + logger.info(f"Closing embedding model: {name}") await embedding_model.close() self.shutdown_thread_pool() diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py index cbcdb048..7d467f17 100644 --- a/reme/core/embedding/base_embedding_model.py +++ b/reme/core/embedding/base_embedding_model.py @@ -5,10 +5,12 @@ Defines the abstract base class and standard API for all embedding model impleme import asyncio import hashlib +import json import os import time from abc import ABC from collections import OrderedDict +from pathlib import Path from loguru import logger @@ -33,7 +35,8 @@ class BaseEmbeddingModel(ABC): max_retries: int = 3, raise_exception: bool = True, max_input_length: int = 8192, - max_cache_size: int = 10000, + cache_dir: str | Path = ".reme", + max_cache_size: int = 2000, **kwargs, ): """Initialize model configuration and parameters. @@ -58,6 +61,7 @@ class BaseEmbeddingModel(ABC): self.max_retries = max_retries self.raise_exception = raise_exception self.max_input_length = max_input_length + self.cache_dir = cache_dir self.max_cache_size = max_cache_size self.kwargs = kwargs @@ -66,6 +70,12 @@ class BaseEmbeddingModel(ABC): self._cache_hits = 0 self._cache_misses = 0 + self.cache_path: Path = Path(self.cache_dir) + self.cache_path.mkdir(parents=True, exist_ok=True) + + # Load cache from disk if available + self._load_cache() + @property def api_key(self) -> str | None: """Get API key from environment variable.""" @@ -90,15 +100,103 @@ class BaseEmbeddingModel(ABC): return [self._truncate_text(text) for text in texts] def _get_cache_key(self, text: str) -> str: - """Generate a cache key by hashing the input text. + """Generate a cache key by hashing text + model_name + dimensions. + + This ensures that the same text produces different cache keys when + using different models or dimensions. Args: text: Input text to hash Returns: - SHA256 hash of the text as hexadecimal string + SHA256 hash combining text, model name, and dimensions """ - return hashlib.sha256(text.encode("utf-8")).hexdigest() + # Combine text, model_name, and dimensions to create unique cache key + cache_string = f"{text}|{self.model_name}|{self.dimensions}" + return hashlib.sha256(cache_string.encode("utf-8")).hexdigest() + + def _get_cache_file_path(self) -> Path: + """Get the path to the cache file. + + Returns: + Path to the embedding cache JSONL file + """ + return self.cache_path / "embedding_cache.jsonl" + + def _load_cache(self) -> None: + """Load embedding cache from disk (JSONL format). + + Each line in the JSONL file contains a JSON object with: + - key: the cache key (SHA256 hash) + - embedding: the embedding vector (list of floats) + + Loads in reverse order (newest first) to prioritize recent embeddings + when max_cache_size is smaller than the file content. + """ + cache_file = self._get_cache_file_path() + if not cache_file.exists(): + logger.info(f"No cache file found at {cache_file}, starting with empty cache") + return + + try: + # Read all lines first (to load in reverse order) + with open(cache_file, "r", encoding="utf-8") as f: + lines = f.readlines() + + loaded_count = 0 + # Load in reverse order (newest entries first) + for _, line in enumerate(reversed(lines), 1): + line = line.strip() + if not line: + continue + try: + data = json.loads(line) + cache_key = data.get("key") + embedding = data.get("embedding") + + if cache_key and embedding: + # Skip if already loaded (keep the newest) + if cache_key in self._embedding_cache: + continue + + # Respect max_cache_size during loading + if len(self._embedding_cache) >= self.max_cache_size: + logger.info( + f"Cache size limit reached ({self.max_cache_size}), " + f"loaded {loaded_count} newest entries", + ) + break + self._embedding_cache[cache_key] = embedding + loaded_count += 1 + except json.JSONDecodeError as e: + logger.warning(f"Failed to parse line in cache file: {e}") + continue + + logger.info(f"Loaded {loaded_count} embeddings from cache file: {cache_file}") + except Exception as e: + logger.error(f"Failed to load cache from {cache_file}: {e}") + + def _save_cache(self) -> None: + """Save embedding cache to disk (JSONL format). + + Each line contains a JSON object with the cache key and embedding vector. + Only saves if cache is non-empty. + """ + logger.info(f"Attempting to save cache, current size: {len(self._embedding_cache)}") + if not self._embedding_cache: + logger.info("Cache is empty, skipping save") + return + + cache_file = self._get_cache_file_path() + try: + with open(cache_file, "w", encoding="utf-8") as f: + for cache_key, embedding in self._embedding_cache.items(): + cache_entry = {"key": cache_key, "embedding": embedding} + f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n") + + logger.info(f"Saved {len(self._embedding_cache)} embeddings to cache file: {cache_file}") + except Exception as e: + logger.error(f"Failed to save cache to {cache_file}: {e}") def _get_from_cache(self, text: str) -> list[float] | None: """Retrieve embedding from cache if it exists. @@ -114,6 +212,10 @@ class BaseEmbeddingModel(ABC): # Move to end (most recently used) self._embedding_cache.move_to_end(cache_key) self._cache_hits += 1 + text_preview = text[:50] + "..." if len(text) > 50 else text + logger.info( + f"Cache hit for text: '{text_preview}' (hits: {self._cache_hits}, misses: {self._cache_misses})", + ) return self._embedding_cache[cache_key] self._cache_misses += 1 return None @@ -411,6 +513,8 @@ class BaseEmbeddingModel(ABC): def close_sync(self): """Synchronously release resources and close connections.""" + self._save_cache() async def close(self): """Asynchronously release resources and close connections.""" + self._save_cache() diff --git a/reme/core/embedding/openai_embedding_model.py b/reme/core/embedding/openai_embedding_model.py index 2fb0f8c1..a068b35d 100644 --- a/reme/core/embedding/openai_embedding_model.py +++ b/reme/core/embedding/openai_embedding_model.py @@ -50,3 +50,4 @@ class OpenAIEmbeddingModel(BaseEmbeddingModel): if self._client is not None: await self._client.close() self._client = None + await super().close() diff --git a/reme/core/embedding/openai_embedding_model_sync.py b/reme/core/embedding/openai_embedding_model_sync.py index 25242b06..bb4b7172 100644 --- a/reme/core/embedding/openai_embedding_model_sync.py +++ b/reme/core/embedding/openai_embedding_model_sync.py @@ -33,3 +33,4 @@ class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel): if self._client is not None: self._client.close() self._client = None + super().close_sync() diff --git a/reme/core/llm/openai_llm.py b/reme/core/llm/openai_llm.py index c757ed50..10030bb2 100644 --- a/reme/core/llm/openai_llm.py +++ b/reme/core/llm/openai_llm.py @@ -112,3 +112,4 @@ class OpenAILLM(BaseLLM): if self._client is not None: await self._client.close() self._client = None + await super().close() diff --git a/reme/core/llm/openai_llm_sync.py b/reme/core/llm/openai_llm_sync.py index 1cdce427..8f1f4517 100644 --- a/reme/core/llm/openai_llm_sync.py +++ b/reme/core/llm/openai_llm_sync.py @@ -55,3 +55,4 @@ class OpenAILLMSync(OpenAILLM): if self._client is not None: self._client.close() self._client = None + super().close_sync() diff --git a/reme/core/memory_store/__init__.py b/reme/core/memory_store/__init__.py index ab7ba78c..a3079db7 100644 --- a/reme/core/memory_store/__init__.py +++ b/reme/core/memory_store/__init__.py @@ -1,19 +1,23 @@ """Memory store module for persistent memory management. This module provides storage backends for memory chunks and file metadata, -including SQLite-based and ChromaDB-based implementations with vector and full-text search. +including SQLite-based, ChromaDB-based, and pure-Python local implementations +with vector and full-text search. """ from .base_memory_store import BaseMemoryStore from .chroma_memory_store import ChromaMemoryStore +from .local_memory_store import LocalMemoryStore from .sqlite_memory_store import SqliteMemoryStore from ..context import R __all__ = [ "BaseMemoryStore", "ChromaMemoryStore", + "LocalMemoryStore", "SqliteMemoryStore", ] R.memory_stores.register("sqlite")(SqliteMemoryStore) R.memory_stores.register("chroma")(ChromaMemoryStore) +R.memory_stores.register("local")(LocalMemoryStore) diff --git a/reme/core/memory_store/chroma_memory_store.py b/reme/core/memory_store/chroma_memory_store.py index 186cdfdf..abb81a42 100644 --- a/reme/core/memory_store/chroma_memory_store.py +++ b/reme/core/memory_store/chroma_memory_store.py @@ -609,3 +609,4 @@ class ChromaMemoryStore(BaseMemoryStore): # ChromaDB PersistentClient handles persistence automatically self.client = None self.chunks_collection = None + await super().close() diff --git a/reme/core/memory_store/local_memory_store.py b/reme/core/memory_store/local_memory_store.py new file mode 100644 index 00000000..12863f16 --- /dev/null +++ b/reme/core/memory_store/local_memory_store.py @@ -0,0 +1,502 @@ +"""Pure-Python in-memory storage backend for memory index, with JSON file persistence.""" + +import json +import time +from dataclasses import dataclass +from pathlib import Path + +from loguru import logger + +from .base_memory_store import BaseMemoryStore +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk, MemorySearchResult +from ..utils.common_utils import cosine_similarity + + +@dataclass +class _ChunkRecord: + """Internal in-memory representation of a stored chunk.""" + + id: str + path: str + source: str + start_line: int + end_line: int + text: str + hash: str + embedding: list[float] | None + updated_at: int + + +class LocalMemoryStore(BaseMemoryStore): + """Pure-Python in-memory memory storage with JSONL file persistence. + + No external dependencies required. All data lives in Python dicts; + writes are persisted to JSONL files on disk so state survives restarts. + + Inherits embedding methods from BaseMemoryStore: + - get_chunk_embedding / get_chunk_embeddings (async) + - get_embedding / get_embeddings (async) + + Provides: + - Vector similarity search (cosine similarity, pure Python) + - Full-text / keyword search (Python substring matching) + - Efficient chunk and file metadata management + """ + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._started: bool = False + # In-memory indexes + self._chunks: dict[str, _ChunkRecord] = {} + self._files: dict[str, dict[str, FileMetadata]] = {} # source -> path -> meta + # Persistence paths (mirror ChromaMemoryStore convention) + self._chunks_file: Path = self.db_path.parent / f"{self.store_name}_chunks.jsonl" + self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json" + + # ------------------------------------------------------------------ + # Persistence helpers + # ------------------------------------------------------------------ + + async def _load_chunks(self) -> None: + """Load chunks from JSONL file into memory.""" + if not self._chunks_file.exists(): + return + try: + data = await self._run_sync_in_executor( + self._chunks_file.read_text, + encoding="utf-8", + ) + self._chunks = {} + for line in data.strip().split("\n"): + if not line: + continue + rec = json.loads(line) + chunk_id = rec["id"] + self._chunks[chunk_id] = _ChunkRecord(**rec) + logger.debug(f"Loaded {len(self._chunks)} chunks from {self._chunks_file}") + except Exception as e: + logger.warning(f"Failed to load chunks from {self._chunks_file}: {e}") + + async def _save_chunks(self) -> None: + """Persist chunks to JSONL file.""" + try: + lines = [] + for rec in self._chunks.values(): + chunk_dict = { + "id": rec.id, + "path": rec.path, + "source": rec.source, + "start_line": rec.start_line, + "end_line": rec.end_line, + "text": rec.text, + "hash": rec.hash, + "embedding": rec.embedding, + "updated_at": rec.updated_at, + } + lines.append(json.dumps(chunk_dict, ensure_ascii=False)) + data = "\n".join(lines) + await self._run_sync_in_executor( + self._chunks_file.write_text, + data, + encoding="utf-8", + ) + logger.debug(f"Saved {len(self._chunks)} chunks to {self._chunks_file}") + except Exception as e: + logger.error(f"Failed to save chunks to {self._chunks_file}: {e}") + + async def _load_metadata(self) -> None: + """Load file metadata from JSON file into memory.""" + if not self._metadata_file.exists(): + return + try: + data = await self._run_sync_in_executor( + self._metadata_file.read_text, + encoding="utf-8", + ) + raw: dict = json.loads(data) + self._files = { + source: {path: FileMetadata(**meta) for path, meta in files.items()} for source, files in raw.items() + } + logger.debug(f"Loaded file metadata from {self._metadata_file}") + except Exception as e: + logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") + + async def _save_metadata(self) -> None: + """Persist file metadata to JSON file.""" + try: + raw: dict = {} + for source, files in self._files.items(): + raw[source] = { + path: { + "path": meta.path, + "hash": meta.hash, + "mtime_ms": meta.mtime_ms, + "size": meta.size, + "chunk_count": meta.chunk_count, + } + for path, meta in files.items() + } + data = json.dumps(raw, indent=2, ensure_ascii=False) + await self._run_sync_in_executor( + self._metadata_file.write_text, + data, + encoding="utf-8", + ) + logger.debug(f"Saved file metadata to {self._metadata_file}") + except Exception as e: + logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") + + # ------------------------------------------------------------------ + # Lifecycle + # ------------------------------------------------------------------ + + async def start(self) -> None: + """Load persisted data into memory.""" + if self._started: + return + self._started = True + self.db_path.mkdir(parents=True, exist_ok=True) + await self._load_metadata() + await self._load_chunks() + logger.info( + f"LocalMemoryStore '{self.store_name}' ready: " + f"{len(self._chunks)} chunks, metadata at {self._metadata_file}", + ) + + async def close(self) -> None: + """Flush state to disk and release memory.""" + await self._save_metadata() + await self._save_chunks() + self._chunks.clear() + self._files.clear() + self._started = False + + # ------------------------------------------------------------------ + # Write operations + # ------------------------------------------------------------------ + + async def upsert_file( + self, + file_meta: FileMetadata, + source: MemorySource, + chunks: list[MemoryChunk], + ) -> None: + """Insert or update file and its chunks.""" + if not chunks: + return + + # Remove existing chunks for this file/source first + await self.delete_file(file_meta.path, source) + + # Batch generate embeddings (base class returns mock embeddings when vector_enabled=False) + chunks = await self.get_chunk_embeddings(chunks) + + now = int(time.time() * 1000) + for chunk in chunks: + self._chunks[chunk.id] = _ChunkRecord( + id=chunk.id, + path=file_meta.path, + source=source.value, + start_line=chunk.start_line, + end_line=chunk.end_line, + text=chunk.text, + hash=chunk.hash, + embedding=chunk.embedding, + updated_at=now, + ) + + if source.value not in self._files: + self._files[source.value] = {} + self._files[source.value][file_meta.path] = FileMetadata( + hash=file_meta.hash, + mtime_ms=file_meta.mtime_ms, + size=file_meta.size, + path=file_meta.path, + chunk_count=len(chunks), + ) + + await self._save_chunks() + await self._save_metadata() + + async def delete_file(self, path: str, source: MemorySource) -> None: + """Delete file and all its chunks.""" + to_delete = [cid for cid, rec in self._chunks.items() if rec.path == path and rec.source == source.value] + for cid in to_delete: + del self._chunks[cid] + + if source.value in self._files: + self._files[source.value].pop(path, None) + + await self._save_chunks() + await self._save_metadata() + + async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: + """Delete specific chunks for a file.""" + if not chunk_ids: + return + + for cid in chunk_ids: + self._chunks.pop(cid, None) + + # Recalculate chunk_count in file metadata + for source_meta in self._files.values(): + if path in source_meta: + source_meta[path].chunk_count = sum(1 for rec in self._chunks.values() if rec.path == path) + + await self._save_chunks() + await self._save_metadata() + + async def upsert_chunks( + self, + chunks: list[MemoryChunk], + source: MemorySource, + ) -> None: + """Insert or update specific chunks without affecting other chunks.""" + if not chunks: + return + + chunks = await self.get_chunk_embeddings(chunks) + + now = int(time.time() * 1000) + for chunk in chunks: + self._chunks[chunk.id] = _ChunkRecord( + id=chunk.id, + path=chunk.path, + source=source.value, + start_line=chunk.start_line, + end_line=chunk.end_line, + text=chunk.text, + hash=chunk.hash, + embedding=chunk.embedding, + updated_at=now, + ) + + await self._save_chunks() + + # ------------------------------------------------------------------ + # Read operations + # ------------------------------------------------------------------ + + async def list_files(self, source: MemorySource) -> list[str]: + """List all indexed files for a source.""" + return list(self._files.get(source.value, {}).keys()) + + async def get_file_metadata( + self, + path: str, + source: MemorySource, + ) -> FileMetadata | None: + """Get file metadata.""" + return self._files.get(source.value, {}).get(path) + + async def get_file_chunks( + self, + path: str, + source: MemorySource, + ) -> list[MemoryChunk]: + """Get all chunks for a file, sorted by start_line.""" + records = [rec for rec in self._chunks.values() if rec.path == path and rec.source == source.value] + records.sort(key=lambda r: r.start_line) + return [ + MemoryChunk( + id=rec.id, + path=rec.path, + source=MemorySource(rec.source), + start_line=rec.start_line, + end_line=rec.end_line, + text=rec.text, + hash=rec.hash, + embedding=rec.embedding, + ) + for rec in records + ] + + # ------------------------------------------------------------------ + # Search + # ------------------------------------------------------------------ + + async def vector_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform cosine-similarity vector search over in-memory embeddings.""" + if not self.vector_enabled or not query: + return [] + + query_embedding = await self.get_embedding(query) + if not query_embedding: + return [] + + source_values = {s.value for s in sources} if sources else None + results = [] + for rec in self._chunks.values(): + if source_values and rec.source not in source_values: + continue + if not rec.embedding: + continue + + similarity = cosine_similarity(query_embedding, rec.embedding) + results.append( + MemorySearchResult( + path=rec.path, + start_line=rec.start_line, + end_line=rec.end_line, + score=similarity, + snippet=rec.text, + source=MemorySource(rec.source), + raw_metric=1.0 - similarity, # distance equivalent + ), + ) + + results.sort(key=lambda r: r.score, reverse=True) + return results[:limit] + + async def keyword_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform keyword/full-text search via Python substring matching.""" + if not self.fts_enabled or not query: + return [] + + words = query.split() + if not words: + return [] + + query_lower = query.lower() + words_lower = [w.lower() for w in words] + n_words = len(words) + + source_values = {s.value for s in sources} if sources else None + results = [] + for rec in self._chunks.values(): + if source_values and rec.source not in source_values: + continue + + text_lower = rec.text.lower() + match_count = sum(1 for w in words_lower if w in text_lower) + if match_count == 0: + continue + + base_score = match_count / n_words + # Bonus for full phrase match (multi-word queries only) + phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 + score = min(1.0, base_score + phrase_bonus) + + results.append( + MemorySearchResult( + path=rec.path, + start_line=rec.start_line, + end_line=rec.end_line, + score=score, + snippet=rec.text, + source=MemorySource(rec.source), + ), + ) + + results.sort(key=lambda r: r.score, reverse=True) + return results[:limit] + + async def hybrid_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + vector_weight: float = 0.7, + candidate_multiplier: float = 3.0, + ) -> list[MemorySearchResult]: + """Perform hybrid search combining vector and keyword search. + + Args: + query: Search query text + limit: Maximum number of results + sources: Optional list of sources to filter + vector_weight: Weight for vector search results (0.0-1.0). + Keyword weight = 1.0 - vector_weight. + candidate_multiplier: Multiplier for candidate pool size. + + Returns: + List of search results sorted by combined relevance score + """ + assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" + + candidates = min(200, max(1, int(limit * candidate_multiplier))) + text_weight = 1.0 - vector_weight + + if self.vector_enabled and self.fts_enabled: + keyword_results = await self.keyword_search(query, candidates, sources) + vector_results = await self.vector_search(query, candidates, sources) + + logger.info("\n=== Vector Search Results ===") + for i, r in enumerate(vector_results[:10], 1): + snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet + logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") + + logger.info("\n=== Keyword Search Results ===") + for i, r in enumerate(keyword_results[:10], 1): + snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet + logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") + + if not keyword_results: + return vector_results[:limit] + elif not vector_results: + return keyword_results[:limit] + else: + merged = self._merge_hybrid_results( + vector=vector_results, + keyword=keyword_results, + vector_weight=vector_weight, + text_weight=text_weight, + ) + + logger.info("\n=== Merged Hybrid Results ===") + for i, r in enumerate(merged[:10], 1): + snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet + logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") + + return merged[:limit] + elif self.vector_enabled: + return await self.vector_search(query, limit, sources) + elif self.fts_enabled: + return await self.keyword_search(query, limit, sources) + else: + return [] + + @staticmethod + def _merge_hybrid_results( + vector: list[MemorySearchResult], + keyword: list[MemorySearchResult], + vector_weight: float, + text_weight: float, + ) -> list[MemorySearchResult]: + """Merge vector and keyword search results with weighted scoring.""" + merged: dict[str, MemorySearchResult] = {} + + for result in vector: + result.score = result.score * vector_weight + merged[result.merge_key] = result + + for result in keyword: + key = result.merge_key + if key in merged: + merged[key].score += result.score * text_weight + else: + result.score = result.score * text_weight + merged[key] = result + + results = list(merged.values()) + results.sort(key=lambda r: r.score, reverse=True) + return results + + async def clear_all(self) -> None: + """Clear all indexed data from memory and disk.""" + self._chunks.clear() + self._files.clear() + await self._save_chunks() + await self._save_metadata() + logger.info(f"Cleared all data from LocalMemoryStore '{self.store_name}'") diff --git a/reme/core/memory_store/sqlite_memory_store.py b/reme/core/memory_store/sqlite_memory_store.py index e5aee948..b3ee67b9 100644 --- a/reme/core/memory_store/sqlite_memory_store.py +++ b/reme/core/memory_store/sqlite_memory_store.py @@ -958,3 +958,4 @@ class SqliteMemoryStore(BaseMemoryStore): if self.conn: self.conn.close() self.conn = None + await super().close() diff --git a/tests/test_embedding_cache.py b/tests/test_embedding_cache.py index 94b6673e..5cc27e53 100644 --- a/tests/test_embedding_cache.py +++ b/tests/test_embedding_cache.py @@ -16,6 +16,8 @@ Usage: # pylint: disable=C0413 import asyncio +import shutil +import tempfile from typing import List from reme.core.utils import load_env @@ -42,50 +44,56 @@ async def test_cache_basic_functionality(): print("Test 1: Basic Cache Functionality") print(f"{'='*60}") - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=100, - max_retries=2, - raise_exception=True, - ) + temp_dir = tempfile.mkdtemp() + try: + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=100, + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, + ) - test_text = "Hello, this is a test sentence for embedding cache." + test_text = "Hello, this is a test sentence for embedding cache." - print(f"Input text: {test_text}") + print(f"Input text: {test_text}") + print(f"Cache directory: {temp_dir}") - # First call - should be a cache miss - print("\n1️⃣ First embedding call (cold cache):") - embedding1 = await model.get_embedding(test_text) - stats1 = model.get_cache_stats() + # First call - should be a cache miss + print("\n1️⃣ First embedding call (cold cache):") + embedding1 = await model.get_embedding(test_text) + stats1 = model.get_cache_stats() - print(f" Embedding dimension: {len(embedding1)}") - print(f" Cache size: {stats1['cache_size']}") - print(f" Cache hits: {stats1['cache_hits']}") - print(f" Cache misses: {stats1['cache_misses']}") - print(f" Hit rate: {stats1['hit_rate']:.2%}") + print(f" Embedding dimension: {len(embedding1)}") + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache hits: {stats1['cache_hits']}") + print(f" Cache misses: {stats1['cache_misses']}") + print(f" Hit rate: {stats1['hit_rate']:.2%}") - assert len(embedding1) == 1024, "Embedding dimension mismatch" - assert stats1["cache_misses"] == 1, "Should have 1 cache miss" - assert stats1["cache_hits"] == 0, "Should have 0 cache hits" - assert stats1["cache_size"] == 1, "Cache should have 1 entry" + assert len(embedding1) == 1024, "Embedding dimension mismatch" + assert stats1["cache_misses"] == 1, "Should have 1 cache miss" + assert stats1["cache_hits"] == 0, "Should have 0 cache hits" + assert stats1["cache_size"] == 1, "Cache should have 1 entry" - # Second call with same text - should be a cache hit - print("\n2️⃣ Second embedding call (same text):") - embedding2 = await model.get_embedding(test_text) - stats2 = model.get_cache_stats() + # Second call with same text - should be a cache hit + print("\n2️⃣ Second embedding call (same text):") + embedding2 = await model.get_embedding(test_text) + stats2 = model.get_cache_stats() - print(f" Cache hits: {stats2['cache_hits']}") - print(f" Cache misses: {stats2['cache_misses']}") - print(f" Hit rate: {stats2['hit_rate']:.2%}") + print(f" Cache hits: {stats2['cache_hits']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Hit rate: {stats2['hit_rate']:.2%}") - assert embedding1 == embedding2, "Cached embedding should be identical" - assert stats2["cache_hits"] == 1, "Should have 1 cache hit" - assert stats2["cache_misses"] == 1, "Should still have 1 cache miss" - assert stats2["hit_rate"] == 0.5, "Hit rate should be 50%" + assert embedding1 == embedding2, "Cached embedding should be identical" + assert stats2["cache_hits"] == 1, "Should have 1 cache hit" + assert stats2["cache_misses"] == 1, "Should still have 1 cache miss" + assert stats2["hit_rate"] == 0.5, "Hit rate should be 50%" - await model.close() - print("\n✓ PASSED: Basic cache functionality works correctly") + await model.close() + print("\n✓ PASSED: Basic cache functionality works correctly") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def test_batch_cache_efficiency(): @@ -94,54 +102,59 @@ async def test_batch_cache_efficiency(): print("Test 2: Batch Cache Efficiency") print(f"{'='*60}") - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=1000, - max_retries=2, - raise_exception=True, - ) + temp_dir = tempfile.mkdtemp() + try: + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=1000, + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, + ) - texts = get_test_texts() + texts = get_test_texts() - # Create a list with duplicates - texts_with_duplicates = texts + texts[:3] # 5 unique + 3 duplicates = 8 total + # Create a list with duplicates + texts_with_duplicates = texts + texts[:3] # 5 unique + 3 duplicates = 8 total - print(f"Processing {len(texts_with_duplicates)} texts (5 unique + 3 duplicates)") + print(f"Processing {len(texts_with_duplicates)} texts (5 unique + 3 duplicates)") - # First batch - print("\n1️⃣ First batch (cold cache):") - embeddings1 = await model.get_embeddings(texts) - stats1 = model.get_cache_stats() + # First batch + print("\n1️⃣ First batch (cold cache):") + embeddings1 = await model.get_embeddings(texts) + stats1 = model.get_cache_stats() - print(f" Embeddings generated: {len(embeddings1)}") - print(f" Cache size: {stats1['cache_size']}") - print(f" Cache misses: {stats1['cache_misses']}") - print(f" Cache hits: {stats1['cache_hits']}") + print(f" Embeddings generated: {len(embeddings1)}") + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache misses: {stats1['cache_misses']}") + print(f" Cache hits: {stats1['cache_hits']}") - assert len(embeddings1) == len(texts), "Embeddings count mismatch" - assert stats1["cache_size"] == len(texts), f"Cache should have {len(texts)} entries" - assert stats1["cache_misses"] == len(texts), "All should be cache misses" + assert len(embeddings1) == len(texts), "Embeddings count mismatch" + assert stats1["cache_size"] == len(texts), f"Cache should have {len(texts)} entries" + assert stats1["cache_misses"] == len(texts), "All should be cache misses" - # Second batch with duplicates - print("\n2️⃣ Second batch (with duplicates):") - embeddings2 = await model.get_embeddings(texts_with_duplicates) - stats2 = model.get_cache_stats() + # Second batch with duplicates + print("\n2️⃣ Second batch (with duplicates):") + embeddings2 = await model.get_embeddings(texts_with_duplicates) + stats2 = model.get_cache_stats() - print(f" Embeddings generated: {len(embeddings2)}") - print(f" Cache hits: {stats2['cache_hits']}") - print(f" Cache misses: {stats2['cache_misses']}") - print(f" Hit rate: {stats2['hit_rate']:.2%}") + print(f" Embeddings generated: {len(embeddings2)}") + print(f" Cache hits: {stats2['cache_hits']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Hit rate: {stats2['hit_rate']:.2%}") - assert len(embeddings2) == len(texts_with_duplicates), "Embeddings count mismatch" - assert stats2["cache_hits"] >= 3, "Should have at least 3 cache hits from duplicates" + assert len(embeddings2) == len(texts_with_duplicates), "Embeddings count mismatch" + assert stats2["cache_hits"] >= 3, "Should have at least 3 cache hits from duplicates" - # Verify embeddings are identical for duplicated texts - for i in range(3): - assert embeddings2[i] == embeddings2[len(texts) + i], f"Duplicate {i} should have identical embedding" + # Verify embeddings are identical for duplicated texts + for i in range(3): + assert embeddings2[i] == embeddings2[len(texts) + i], f"Duplicate {i} should have identical embedding" - await model.close() - print("\n✓ PASSED: Batch cache efficiently handles duplicates") + await model.close() + print("\n✓ PASSED: Batch cache efficiently handles duplicates") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def test_cache_lru_eviction(): @@ -150,61 +163,66 @@ async def test_cache_lru_eviction(): print("Test 3: LRU Cache Eviction") print(f"{'='*60}") - # Create model with small cache size - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=3, # Small cache for testing eviction - max_retries=2, - raise_exception=True, - ) - - texts = get_test_texts()[:5] # Use 5 texts, cache size is 3 - - print(f"Cache size limit: {model.max_cache_size}") - print(f"Number of unique texts: {len(texts)}") - - # Fill cache beyond capacity - print("\n1️⃣ Filling cache with 5 texts (capacity = 3):") - for i, text in enumerate(texts): - await model.get_embedding(text) - stats = model.get_cache_stats() - print( - f" After text {i+1}: cache_size={stats['cache_size']}, " - f"hits={stats['cache_hits']}, misses={stats['cache_misses']}", + temp_dir = tempfile.mkdtemp() + try: + # Create model with small cache size + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=3, # Small cache for testing eviction + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, ) - final_stats = model.get_cache_stats() - assert final_stats["cache_size"] <= 3, "Cache size should not exceed max_cache_size" - assert final_stats["cache_misses"] == 5, "Should have 5 cache misses for 5 unique texts" + texts = get_test_texts()[:5] # Use 5 texts, cache size is 3 - # Access the most recent entries - should be cache hits - print("\n2️⃣ Accessing recent entries (should be cached):") - recent_texts = texts[-3:] # Last 3 texts should still be in cache + print(f"Cache size limit: {model.max_cache_size}") + print(f"Number of unique texts: {len(texts)}") - for i, text in enumerate(recent_texts): - await model.get_embedding(text) - stats = model.get_cache_stats() - print(f" Text {len(texts) - 3 + i + 1}: hits={stats['cache_hits']}") + # Fill cache beyond capacity + print("\n1️⃣ Filling cache with 5 texts (capacity = 3):") + for i, text in enumerate(texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print( + f" After text {i+1}: cache_size={stats['cache_size']}, " + f"hits={stats['cache_hits']}, misses={stats['cache_misses']}", + ) - final_stats = model.get_cache_stats() - assert final_stats["cache_hits"] == 3, "Should have 3 cache hits for recent entries" + final_stats = model.get_cache_stats() + assert final_stats["cache_size"] <= 3, "Cache size should not exceed max_cache_size" + assert final_stats["cache_misses"] == 5, "Should have 5 cache misses for 5 unique texts" - # Access oldest entries - should be cache misses (evicted) - print("\n3️⃣ Accessing oldest entries (should be evicted):") - old_texts = texts[:2] # First 2 texts should have been evicted + # Access the most recent entries - should be cache hits + print("\n2️⃣ Accessing recent entries (should be cached):") + recent_texts = texts[-3:] # Last 3 texts should still be in cache - before_misses = final_stats["cache_misses"] - for i, text in enumerate(old_texts): - await model.get_embedding(text) - stats = model.get_cache_stats() - print(f" Text {i + 1}: misses={stats['cache_misses']}") + for i, text in enumerate(recent_texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print(f" Text {len(texts) - 3 + i + 1}: hits={stats['cache_hits']}") - final_stats = model.get_cache_stats() - assert final_stats["cache_misses"] == before_misses + 2, "Should have 2 more cache misses for evicted entries" + final_stats = model.get_cache_stats() + assert final_stats["cache_hits"] == 3, "Should have 3 cache hits for recent entries" - await model.close() - print("\n✓ PASSED: LRU eviction works correctly") + # Access oldest entries - should be cache misses (evicted) + print("\n3️⃣ Accessing oldest entries (should be evicted):") + old_texts = texts[:2] # First 2 texts should have been evicted + + before_misses = final_stats["cache_misses"] + for i, text in enumerate(old_texts): + await model.get_embedding(text) + stats = model.get_cache_stats() + print(f" Text {i + 1}: misses={stats['cache_misses']}") + + final_stats = model.get_cache_stats() + assert final_stats["cache_misses"] == before_misses + 2, "Should have 2 more cache misses for evicted entries" + + await model.close() + print("\n✓ PASSED: LRU eviction works correctly") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def test_cache_stats_and_clear(): @@ -213,49 +231,54 @@ async def test_cache_stats_and_clear(): print("Test 4: Cache Statistics and Clearing") print(f"{'='*60}") - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=100, - max_retries=2, - raise_exception=True, - ) + temp_dir = tempfile.mkdtemp() + try: + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=100, + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, + ) - texts = get_test_texts() + texts = get_test_texts() - # Generate some cache activity - print("\n1️⃣ Generating cache activity:") - await model.get_embeddings(texts) - await model.get_embeddings(texts[:3]) # Repeat first 3 + # Generate some cache activity + print("\n1️⃣ Generating cache activity:") + await model.get_embeddings(texts) + await model.get_embeddings(texts[:3]) # Repeat first 3 - stats = model.get_cache_stats() - print(f" Cache size: {stats['cache_size']}") - print(f" Max cache size: {stats['max_cache_size']}") - print(f" Cache hits: {stats['cache_hits']}") - print(f" Cache misses: {stats['cache_misses']}") - print(f" Hit rate: {stats['hit_rate']:.2%}") + stats = model.get_cache_stats() + print(f" Cache size: {stats['cache_size']}") + print(f" Max cache size: {stats['max_cache_size']}") + print(f" Cache hits: {stats['cache_hits']}") + print(f" Cache misses: {stats['cache_misses']}") + print(f" Hit rate: {stats['hit_rate']:.2%}") - assert stats["cache_size"] > 0, "Cache should not be empty" - assert stats["cache_hits"] >= 3, "Should have at least 3 cache hits" - assert "hit_rate" in stats, "Stats should include hit_rate" + assert stats["cache_size"] > 0, "Cache should not be empty" + assert stats["cache_hits"] >= 3, "Should have at least 3 cache hits" + assert "hit_rate" in stats, "Stats should include hit_rate" - # Clear cache - print("\n2️⃣ Clearing cache:") - model.clear_cache() - stats_after_clear = model.get_cache_stats() + # Clear cache + print("\n2️⃣ Clearing cache:") + model.clear_cache() + stats_after_clear = model.get_cache_stats() - print(f" Cache size after clear: {stats_after_clear['cache_size']}") - print(f" Hits after clear: {stats_after_clear['cache_hits']}") - print(f" Misses after clear: {stats_after_clear['cache_misses']}") - print(f" Hit rate after clear: {stats_after_clear['hit_rate']:.2%}") + print(f" Cache size after clear: {stats_after_clear['cache_size']}") + print(f" Hits after clear: {stats_after_clear['cache_hits']}") + print(f" Misses after clear: {stats_after_clear['cache_misses']}") + print(f" Hit rate after clear: {stats_after_clear['hit_rate']:.2%}") - assert stats_after_clear["cache_size"] == 0, "Cache should be empty after clear" - assert stats_after_clear["cache_hits"] == 0, "Hits should be reset" - assert stats_after_clear["cache_misses"] == 0, "Misses should be reset" - assert stats_after_clear["hit_rate"] == 0.0, "Hit rate should be 0" + assert stats_after_clear["cache_size"] == 0, "Cache should be empty after clear" + assert stats_after_clear["cache_hits"] == 0, "Hits should be reset" + assert stats_after_clear["cache_misses"] == 0, "Misses should be reset" + assert stats_after_clear["hit_rate"] == 0.0, "Hit rate should be 0" - await model.close() - print("\n✓ PASSED: Cache statistics and clearing work correctly") + await model.close() + print("\n✓ PASSED: Cache statistics and clearing work correctly") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def test_cache_disabled(): @@ -264,40 +287,45 @@ async def test_cache_disabled(): print("Test 5: Cache Disabled") print(f"{'='*60}") - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=0, # Disable cache - max_retries=2, - raise_exception=True, - ) + temp_dir = tempfile.mkdtemp() + try: + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=0, # Disable cache + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, + ) - test_text = "Test text with cache disabled" + test_text = "Test text with cache disabled" - print(f"Cache size limit: {model.max_cache_size} (disabled)") - print(f"Input text: {test_text}") + print(f"Cache size limit: {model.max_cache_size} (disabled)") + print(f"Input text: {test_text}") - # Call twice with same text - print("\n1️⃣ First call:") - embedding1 = await model.get_embedding(test_text) - stats1 = model.get_cache_stats() - print(f" Cache size: {stats1['cache_size']}") - print(f" Cache misses: {stats1['cache_misses']}") + # Call twice with same text + print("\n1️⃣ First call:") + embedding1 = await model.get_embedding(test_text) + stats1 = model.get_cache_stats() + print(f" Cache size: {stats1['cache_size']}") + print(f" Cache misses: {stats1['cache_misses']}") - print("\n2️⃣ Second call (same text):") - embedding2 = await model.get_embedding(test_text) - stats2 = model.get_cache_stats() - print(f" Cache size: {stats2['cache_size']}") - print(f" Cache misses: {stats2['cache_misses']}") - print(f" Cache hits: {stats2['cache_hits']}") + print("\n2️⃣ Second call (same text):") + embedding2 = await model.get_embedding(test_text) + stats2 = model.get_cache_stats() + print(f" Cache size: {stats2['cache_size']}") + print(f" Cache misses: {stats2['cache_misses']}") + print(f" Cache hits: {stats2['cache_hits']}") - assert stats2["cache_size"] == 0, "Cache should remain empty when disabled" - assert stats2["cache_misses"] == 2, "Both calls should be cache misses" - assert stats2["cache_hits"] == 0, "Should have no cache hits when disabled" - assert embedding1 == embedding2, "Embeddings should still be consistent" + assert stats2["cache_size"] == 0, "Cache should remain empty when disabled" + assert stats2["cache_misses"] == 2, "Both calls should be cache misses" + assert stats2["cache_hits"] == 0, "Should have no cache hits when disabled" + assert embedding1 == embedding2, "Embeddings should still be consistent" - await model.close() - print("\n✓ PASSED: Cache correctly disabled when max_cache_size=0") + await model.close() + print("\n✓ PASSED: Cache correctly disabled when max_cache_size=0") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def test_cache_performance_demo(): @@ -306,58 +334,63 @@ async def test_cache_performance_demo(): print("Test 6: Cache Performance Demo") print(f"{'='*60}") - model = OpenAIEmbeddingModel( - model_name="text-embedding-v4", - dimensions=1024, - max_cache_size=1000, - max_retries=2, - raise_exception=True, - ) + temp_dir = tempfile.mkdtemp() + try: + model = OpenAIEmbeddingModel( + model_name="text-embedding-v4", + dimensions=1024, + max_cache_size=1000, + max_retries=2, + raise_exception=True, + cache_dir=temp_dir, + ) - texts = get_test_texts() + texts = get_test_texts() - # Create a realistic workload with many repeated queries - workload = texts * 3 # 15 queries total, 5 unique + # Create a realistic workload with many repeated queries + workload = texts * 3 # 15 queries total, 5 unique - print(f"\nProcessing {len(workload)} queries ({len(texts)} unique texts)") - print("This simulates a realistic scenario with repeated queries\n") + print(f"\nProcessing {len(workload)} queries ({len(texts)} unique texts)") + print("This simulates a realistic scenario with repeated queries\n") - # Process all queries - for i, text in enumerate(workload, 1): - await model.get_embedding(text) - if i % 5 == 0: # Report every 5 queries - stats = model.get_cache_stats() - print( - f"After {i:2d} queries: hits={stats['cache_hits']:2d}, " - f"misses={stats['cache_misses']:2d}, " - f"hit_rate={stats['hit_rate']:5.1%}", - ) + # Process all queries + for i, text in enumerate(workload, 1): + await model.get_embedding(text) + if i % 5 == 0: # Report every 5 queries + stats = model.get_cache_stats() + print( + f"After {i:2d} queries: hits={stats['cache_hits']:2d}, " + f"misses={stats['cache_misses']:2d}, " + f"hit_rate={stats['hit_rate']:5.1%}", + ) - final_stats = model.get_cache_stats() - total_requests = final_stats["cache_hits"] + final_stats["cache_misses"] + final_stats = model.get_cache_stats() + total_requests = final_stats["cache_hits"] + final_stats["cache_misses"] - print(f"\n{'─'*60}") - print("📊 Final Statistics:") - print(f"{'─'*60}") - print(f" Total queries: {total_requests}") - print(f" Unique texts: {len(texts)}") - print(f" Cache hits: {final_stats['cache_hits']}") - print(f" Cache misses: {final_stats['cache_misses']}") - print(f" Hit rate: {final_stats['hit_rate']:.1%}") - print(f" Cache size: {final_stats['cache_size']}/{final_stats['max_cache_size']}") - print(f"{'─'*60}") - print( - f"💰 API calls saved: {final_stats['cache_hits']} out of {total_requests} " - f"({final_stats['cache_hits']/total_requests*100:.1f}%)", - ) - print(f"{'─'*60}") + print(f"\n{'─'*60}") + print("📊 Final Statistics:") + print(f"{'─'*60}") + print(f" Total queries: {total_requests}") + print(f" Unique texts: {len(texts)}") + print(f" Cache hits: {final_stats['cache_hits']}") + print(f" Cache misses: {final_stats['cache_misses']}") + print(f" Hit rate: {final_stats['hit_rate']:.1%}") + print(f" Cache size: {final_stats['cache_size']}/{final_stats['max_cache_size']}") + print(f"{'─'*60}") + print( + f"💰 API calls saved: {final_stats['cache_hits']} out of {total_requests} " + f"({final_stats['cache_hits']/total_requests*100:.1f}%)", + ) + print(f"{'─'*60}") - assert final_stats["cache_hits"] == 10, "Should have 10 cache hits (2 repeats × 5 texts)" - assert final_stats["cache_misses"] == 5, "Should have 5 cache misses (5 unique texts)" - assert final_stats["hit_rate"] > 0.6, "Hit rate should be > 60%" + assert final_stats["cache_hits"] == 10, "Should have 10 cache hits (2 repeats × 5 texts)" + assert final_stats["cache_misses"] == 5, "Should have 5 cache misses (5 unique texts)" + assert final_stats["hit_rate"] > 0.6, "Hit rate should be > 60%" - await model.close() - print("\n✓ PASSED: Cache provides significant performance improvement") + await model.close() + print("\n✓ PASSED: Cache provides significant performance improvement") + finally: + shutil.rmtree(temp_dir, ignore_errors=True) async def main(): diff --git a/tests/test_memory_store.py b/tests/test_memory_store.py index a5937b10..9a7e961f 100644 --- a/tests/test_memory_store.py +++ b/tests/test_memory_store.py @@ -1,12 +1,14 @@ # pylint: disable=too-many-lines """Unified test suite for memory store implementations. -This module provides comprehensive test coverage for SqliteMemoryStore, ChromaMemoryStore -and future memory store implementations. Tests can be run for specific stores or all implementations. +This module provides comprehensive test coverage for SqliteMemoryStore, ChromaMemoryStore, +LocalMemoryStore and future memory store implementations. Tests can be run for specific stores +or all implementations. Usage: python test_memory_store.py --sqlite # Test SqliteMemoryStore only python test_memory_store.py --chroma # Test ChromaMemoryStore only + python test_memory_store.py --local # Test LocalMemoryStore only python test_memory_store.py --all # Test all memory stores """ @@ -15,6 +17,7 @@ import asyncio import hashlib import shutil import time +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import List @@ -24,6 +27,7 @@ from reme.core.embedding import OpenAIEmbeddingModel from reme.core.enumeration.memory_source import MemorySource from reme.core.memory_store.base_memory_store import BaseMemoryStore from reme.core.memory_store.chroma_memory_store import ChromaMemoryStore +from reme.core.memory_store.local_memory_store import LocalMemoryStore from reme.core.memory_store.sqlite_memory_store import SqliteMemoryStore from reme.core.schema.file_metadata import FileMetadata from reme.core.schema.memory_chunk import MemoryChunk @@ -50,6 +54,10 @@ class TestConfig: CHROMA_DB_PATH = "./test_memory_store_chroma" CHROMA_FTS_ENABLED = True + # LocalMemoryStore settings + LOCAL_DB_PATH = "./test_memory_store_local" + LOCAL_FTS_ENABLED = True + # Embedding model settings EMBEDDING_MODEL_NAME = "text-embedding-v4" EMBEDDING_DIMENSIONS = 64 @@ -190,6 +198,8 @@ def get_store_type(store: BaseMemoryStore) -> str: return "sqlite" elif isinstance(store, ChromaMemoryStore): return "chroma" + elif isinstance(store, LocalMemoryStore): + return "local" else: raise ValueError(f"Unknown memory store type: {type(store)}") @@ -211,6 +221,8 @@ def create_memory_store(store_type: str) -> BaseMemoryStore: dimensions=config.EMBEDDING_DIMENSIONS, ) + thread_pool = ThreadPoolExecutor() + if store_type == "sqlite": return SqliteMemoryStore( store_name=config.NAME, @@ -218,6 +230,7 @@ def create_memory_store(store_type: str) -> BaseMemoryStore: embedding_model=embedding_model, vec_ext_path=config.SQLITE_VEC_EXT_PATH, fts_enabled=config.SQLITE_FTS_ENABLED, + thread_pool=thread_pool, ) elif store_type == "chroma": return ChromaMemoryStore( @@ -225,6 +238,15 @@ def create_memory_store(store_type: str) -> BaseMemoryStore: db_path=config.CHROMA_DB_PATH, embedding_model=embedding_model, fts_enabled=config.CHROMA_FTS_ENABLED, + thread_pool=thread_pool, + ) + elif store_type == "local": + return LocalMemoryStore( + store_name=config.NAME, + db_path=config.LOCAL_DB_PATH, + embedding_model=embedding_model, + fts_enabled=config.LOCAL_FTS_ENABLED, + thread_pool=thread_pool, ) else: raise ValueError(f"Unknown store type: {store_type}") @@ -260,6 +282,14 @@ async def test_start_store(store: BaseMemoryStore, _store_name: str): assert store.chunks_collection is not None, "ChromaDB collection should exist" logger.info(f"✓ ChromaDB collection created: {store.collection_name}") + # Verify LocalMemoryStore initialized (access internals for test assertions) + if isinstance(store, LocalMemoryStore): + # pylint: disable=protected-access + assert store._started, "LocalMemoryStore should be marked as started" + assert isinstance(store._chunks, dict), "Chunks index should be a dict" + assert isinstance(store._files, dict), "Files index should be a dict" + logger.info(f"✓ LocalMemoryStore ready (chunks file: {store._chunks_file})") + async def test_upsert_file(store: BaseMemoryStore, _store_name: str) -> tuple[FileMetadata, List[MemoryChunk]]: """Test file and chunks insertion.""" @@ -976,6 +1006,19 @@ async def cleanup_store(store: BaseMemoryStore, store_type: str): metadata_file.unlink() logger.info(f"✓ Cleaned up metadata file: {metadata_file}") + # Clean up LocalMemoryStore JSON persistence files + if store_type == "local": + config = TestConfig() + db_dir = Path(config.LOCAL_DB_PATH) + if db_dir.exists(): + shutil.rmtree(db_dir) + logger.info(f"✓ Cleaned up directory: {db_dir}") + for suffix in ("_chunks.json", "_file_metadata.json"): + json_file = db_dir.parent / f"{config.NAME}{suffix}" + if json_file.exists(): + json_file.unlink() + logger.info(f"✓ Cleaned up file: {json_file}") + logger.info("✓ Cleanup completed") except Exception as e: logger.error(f"Cleanup error: {e}") @@ -993,6 +1036,7 @@ async def main(): Examples: python test_memory_store.py --sqlite # Test SqliteMemoryStore only python test_memory_store.py --chroma # Test ChromaMemoryStore only + python test_memory_store.py --local # Test LocalMemoryStore only python test_memory_store.py --all # Test all memory stores """, ) @@ -1006,6 +1050,11 @@ Examples: action="store_true", help="Test ChromaMemoryStore", ) + parser.add_argument( + "--local", + action="store_true", + help="Test LocalMemoryStore", + ) parser.add_argument( "--all", action="store_true", @@ -1021,6 +1070,7 @@ Examples: stores_to_test = [ ("sqlite", "SqliteMemoryStore"), ("chroma", "ChromaMemoryStore"), + ("local", "LocalMemoryStore"), ] else: # Build list based on individual flags @@ -1028,15 +1078,18 @@ Examples: stores_to_test.append(("sqlite", "SqliteMemoryStore")) if args.chroma: stores_to_test.append(("chroma", "ChromaMemoryStore")) + if args.local: + stores_to_test.append(("local", "LocalMemoryStore")) if not stores_to_test: # Default to all memory stores if no argument provided stores_to_test = [ ("sqlite", "SqliteMemoryStore"), ("chroma", "ChromaMemoryStore"), + ("local", "LocalMemoryStore"), ] print("No memory store specified, defaulting to test all memory stores") - print("Use --sqlite or --chroma to test specific ones\n") + print("Use --sqlite, --chroma, or --local to test specific ones\n") # Run tests for each memory store for store_type, store_name in stores_to_test: From 48b1136707ba7fa9e88a6e0664367f165d754e83 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 20 Feb 2026 00:29:38 +0800 Subject: [PATCH 2/3] refactor(memory_store): remove async executor for file operations and add metadata caching --- reme/core/memory_store/chroma_memory_store.py | 95 ++++++++----------- reme/core/memory_store/local_memory_store.py | 33 +------ tests/test_memory_store.py | 2 +- 3 files changed, 43 insertions(+), 87 deletions(-) diff --git a/reme/core/memory_store/chroma_memory_store.py b/reme/core/memory_store/chroma_memory_store.py index abb81a42..d5ea1737 100644 --- a/reme/core/memory_store/chroma_memory_store.py +++ b/reme/core/memory_store/chroma_memory_store.py @@ -48,6 +48,7 @@ class ChromaMemoryStore(BaseMemoryStore): self.chunks_collection: "chromadb.Collection | None" = None # Initialize metadata file path (db_path and store_name are set by base class) self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json" + self._metadata_cache: dict[str, dict[str, FileMetadata]] = {} @property def collection_name(self) -> str: @@ -64,10 +65,7 @@ class ChromaMemoryStore(BaseMemoryStore): return {} try: - data = await self._run_sync_in_executor( - self._metadata_file.read_text, - encoding="utf-8", - ) + data = self._metadata_file.read_text(encoding="utf-8") metadata_dict = json.loads(data) # Convert dict to FileMetadata objects @@ -104,11 +102,7 @@ class ChromaMemoryStore(BaseMemoryStore): } data = json.dumps(metadata_dict, indent=2, ensure_ascii=False) - await self._run_sync_in_executor( - self._metadata_file.write_text, - data, - encoding="utf-8", - ) + self._metadata_file.write_text(data, encoding="utf-8") logger.debug(f"Saved file metadata to {self._metadata_file}") except Exception as e: logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") @@ -121,8 +115,7 @@ class ChromaMemoryStore(BaseMemoryStore): self.db_path.mkdir(parents=True, exist_ok=True) # Initialize persistent ChromaDB client - self.client = await self._run_sync_in_executor( - chromadb.PersistentClient, + self.client = chromadb.PersistentClient( path=str(self.db_path), settings=Settings( anonymized_telemetry=False, @@ -132,12 +125,14 @@ class ChromaMemoryStore(BaseMemoryStore): # Get or create the chunks collection # ChromaDB uses cosine distance by default for similarity - self.chunks_collection = await self._run_sync_in_executor( - self.client.get_or_create_collection, + self.chunks_collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) + # Load metadata into cache + self._metadata_cache = await self._load_metadata() + logger.info(f"ChromaDB initialized with collection: {self.collection_name}") logger.info(f"File metadata will be persisted to: {self._metadata_file}") @@ -181,70 +176,59 @@ class ChromaMemoryStore(BaseMemoryStore): ) # Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function) - await self._run_sync_in_executor( - self.chunks_collection.upsert, + self.chunks_collection.upsert( ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas, ) - # Store file metadata to disk - metadata = await self._load_metadata() - if source.value not in metadata: - metadata[source.value] = {} - metadata[source.value][file_meta.path] = FileMetadata( + # Update file metadata in cache + if source.value not in self._metadata_cache: + self._metadata_cache[source.value] = {} + self._metadata_cache[source.value][file_meta.path] = FileMetadata( hash=file_meta.hash, mtime_ms=file_meta.mtime_ms, size=file_meta.size, path=file_meta.path, chunk_count=len(chunks), ) - await self._save_metadata(metadata) async def delete_file(self, path: str, source: MemorySource) -> None: """Delete file and all its chunks.""" # Query for all chunks with this path and source - results = await self._run_sync_in_executor( - self.chunks_collection.get, + results = self.chunks_collection.get( where={"$and": [{"path": path}, {"source": source.value}]}, include=[], ) if results["ids"]: - await self._run_sync_in_executor( - self.chunks_collection.delete, + self.chunks_collection.delete( ids=results["ids"], ) - # Remove from file metadata - metadata = await self._load_metadata() - if source.value in metadata: - metadata[source.value].pop(path, None) - await self._save_metadata(metadata) + # Remove from file metadata cache + if source.value in self._metadata_cache: + self._metadata_cache[source.value].pop(path, None) async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: """Delete specific chunks for a file.""" if not chunk_ids: return - await self._run_sync_in_executor( - self.chunks_collection.delete, + self.chunks_collection.delete( ids=chunk_ids, ) - # Update chunk count in file metadata - metadata = await self._load_metadata() - for source_meta in metadata.values(): + # Update chunk count in file metadata cache + for source_meta in self._metadata_cache.values(): if path in source_meta: # Recalculate chunk count - results = await self._run_sync_in_executor( - self.chunks_collection.get, + results = self.chunks_collection.get( where={"path": path}, include=[], ) source_meta[path].chunk_count = len(results["ids"]) - await self._save_metadata(metadata) break async def upsert_chunks( @@ -282,8 +266,7 @@ class ChromaMemoryStore(BaseMemoryStore): ) # Always pass embeddings to prevent default embedding function - await self._run_sync_in_executor( - self.chunks_collection.upsert, + self.chunks_collection.upsert( ids=ids, documents=documents, embeddings=embeddings, @@ -292,10 +275,9 @@ class ChromaMemoryStore(BaseMemoryStore): async def list_files(self, source: MemorySource) -> list[str]: """List all indexed files for a source.""" - metadata = await self._load_metadata() - if source.value not in metadata: + if source.value not in self._metadata_cache: return [] - return list(metadata[source.value].keys()) + return list(self._metadata_cache[source.value].keys()) async def get_file_metadata( self, @@ -303,10 +285,9 @@ class ChromaMemoryStore(BaseMemoryStore): source: MemorySource, ) -> FileMetadata | None: """Get file metadata with chunk count.""" - metadata = await self._load_metadata() - if source.value not in metadata: + if source.value not in self._metadata_cache: return None - return metadata[source.value].get(path) + return self._metadata_cache[source.value].get(path) async def get_file_chunks( self, @@ -314,8 +295,7 @@ class ChromaMemoryStore(BaseMemoryStore): source: MemorySource, ) -> list[MemoryChunk]: """Get all chunks for a file.""" - results = await self._run_sync_in_executor( - self.chunks_collection.get, + results = self.chunks_collection.get( where={"$and": [{"path": path}, {"source": source.value}]}, include=["documents", "embeddings", "metadatas"], ) @@ -364,8 +344,7 @@ class ChromaMemoryStore(BaseMemoryStore): where_filter = {"source": {"$in": [s.value for s in sources]}} # Perform vector search - results = await self._run_sync_in_executor( - self.chunks_collection.query, + results = self.chunks_collection.query( query_embeddings=[query_embedding], n_results=limit, where=where_filter, @@ -445,8 +424,7 @@ class ChromaMemoryStore(BaseMemoryStore): where_document = {"$or": [{"$contains": w} for w in word_variants_list]} # Get all matching documents - results = await self._run_sync_in_executor( - self.chunks_collection.get, + results = self.chunks_collection.get( where=where_filter, where_document=where_document, include=["documents", "metadatas"], @@ -589,23 +567,26 @@ class ChromaMemoryStore(BaseMemoryStore): async def clear_all(self) -> None: """Clear all indexed data.""" # Delete and recreate the collection - await self._run_sync_in_executor( - self.client.delete_collection, + self.client.delete_collection( name=self.collection_name, ) - self.chunks_collection = await self._run_sync_in_executor( - self.client.get_or_create_collection, + self.chunks_collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) - # Clear file metadata on disk + # Clear file metadata cache and disk + self._metadata_cache = {} await self._save_metadata({}) logger.info(f"Cleared all data from ChromaDB collection: {self.collection_name}") async def close(self) -> None: """Close ChromaDB client and release resources.""" + # Persist metadata cache to disk before closing + if self._metadata_cache: + await self._save_metadata(self._metadata_cache) + # ChromaDB PersistentClient handles persistence automatically self.client = None self.chunks_collection = None diff --git a/reme/core/memory_store/local_memory_store.py b/reme/core/memory_store/local_memory_store.py index 12863f16..b7ecd51b 100644 --- a/reme/core/memory_store/local_memory_store.py +++ b/reme/core/memory_store/local_memory_store.py @@ -63,10 +63,7 @@ class LocalMemoryStore(BaseMemoryStore): if not self._chunks_file.exists(): return try: - data = await self._run_sync_in_executor( - self._chunks_file.read_text, - encoding="utf-8", - ) + data = self._chunks_file.read_text(encoding="utf-8") self._chunks = {} for line in data.strip().split("\n"): if not line: @@ -96,11 +93,7 @@ class LocalMemoryStore(BaseMemoryStore): } lines.append(json.dumps(chunk_dict, ensure_ascii=False)) data = "\n".join(lines) - await self._run_sync_in_executor( - self._chunks_file.write_text, - data, - encoding="utf-8", - ) + self._chunks_file.write_text(data, encoding="utf-8") logger.debug(f"Saved {len(self._chunks)} chunks to {self._chunks_file}") except Exception as e: logger.error(f"Failed to save chunks to {self._chunks_file}: {e}") @@ -110,10 +103,7 @@ class LocalMemoryStore(BaseMemoryStore): if not self._metadata_file.exists(): return try: - data = await self._run_sync_in_executor( - self._metadata_file.read_text, - encoding="utf-8", - ) + data = self._metadata_file.read_text(encoding="utf-8") raw: dict = json.loads(data) self._files = { source: {path: FileMetadata(**meta) for path, meta in files.items()} for source, files in raw.items() @@ -138,11 +128,7 @@ class LocalMemoryStore(BaseMemoryStore): for path, meta in files.items() } data = json.dumps(raw, indent=2, ensure_ascii=False) - await self._run_sync_in_executor( - self._metadata_file.write_text, - data, - encoding="utf-8", - ) + self._metadata_file.write_text(data, encoding="utf-8") logger.debug(f"Saved file metadata to {self._metadata_file}") except Exception as e: logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") @@ -216,9 +202,6 @@ class LocalMemoryStore(BaseMemoryStore): chunk_count=len(chunks), ) - await self._save_chunks() - await self._save_metadata() - async def delete_file(self, path: str, source: MemorySource) -> None: """Delete file and all its chunks.""" to_delete = [cid for cid, rec in self._chunks.items() if rec.path == path and rec.source == source.value] @@ -228,9 +211,6 @@ class LocalMemoryStore(BaseMemoryStore): if source.value in self._files: self._files[source.value].pop(path, None) - await self._save_chunks() - await self._save_metadata() - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: """Delete specific chunks for a file.""" if not chunk_ids: @@ -244,9 +224,6 @@ class LocalMemoryStore(BaseMemoryStore): if path in source_meta: source_meta[path].chunk_count = sum(1 for rec in self._chunks.values() if rec.path == path) - await self._save_chunks() - await self._save_metadata() - async def upsert_chunks( self, chunks: list[MemoryChunk], @@ -272,8 +249,6 @@ class LocalMemoryStore(BaseMemoryStore): updated_at=now, ) - await self._save_chunks() - # ------------------------------------------------------------------ # Read operations # ------------------------------------------------------------------ diff --git a/tests/test_memory_store.py b/tests/test_memory_store.py index 9a7e961f..b8d6e0d3 100644 --- a/tests/test_memory_store.py +++ b/tests/test_memory_store.py @@ -1013,7 +1013,7 @@ async def cleanup_store(store: BaseMemoryStore, store_type: str): if db_dir.exists(): shutil.rmtree(db_dir) logger.info(f"✓ Cleaned up directory: {db_dir}") - for suffix in ("_chunks.json", "_file_metadata.json"): + for suffix in ("_chunks.jsonl", "_file_metadata.json"): json_file = db_dir.parent / f"{config.NAME}{suffix}" if json_file.exists(): json_file.unlink() From 793b8f7d595014ef719b3e25a689b1d458b76fba Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 20 Feb 2026 00:34:39 +0800 Subject: [PATCH 3/3] config: switch memory store backend from local to chroma --- reme/config/cli.yaml | 4 ++-- reme/config/fs.yaml | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/reme/config/cli.yaml b/reme/config/cli.yaml index 2f259869..b71988c9 100644 --- a/reme/config/cli.yaml +++ b/reme/config/cli.yaml @@ -22,8 +22,8 @@ embedding_models: memory_stores: default: -# backend: chroma - backend: local + backend: chroma +# backend: local db_name: reme.db store_name: reme embedding_model: default diff --git a/reme/config/fs.yaml b/reme/config/fs.yaml index 2d92c7bc..d9da8f1f 100644 --- a/reme/config/fs.yaml +++ b/reme/config/fs.yaml @@ -18,8 +18,8 @@ embedding_models: memory_stores: default: # backend: sqlite -# backend: chroma - backend: local + backend: chroma +# backend: local db_name: reme.db store_name: reme embedding_model: default