From 52d21392f22983ba291224c23ef68fd1b796df41 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 12 Feb 2026 20:34:50 +0800 Subject: [PATCH] feat(store): add ChromaDB memory store implementation with hybrid search --- reme/__init__.py | 2 +- reme/config/fs.yaml | 4 +- reme/core/context/service_context.py | 19 +- reme/core/memory_store/__init__.py | 5 +- reme/core/memory_store/base_memory_store.py | 54 ++ reme/core/memory_store/chroma_memory_store.py | 611 ++++++++++++++++++ reme/core/memory_store/sqlite_memory_store.py | 103 ++- reme/core/schema/service_config.py | 1 + reme/reme_cli.py | 4 +- reme/reme_fs.py | 75 +-- reme/tool/fs/fs_memory_search.py | 112 +--- tests/test_fs_memory_search.py | 4 +- tests/test_memory_store.py | 119 ++-- 13 files changed, 906 insertions(+), 207 deletions(-) create mode 100644 reme/core/memory_store/chroma_memory_store.py diff --git a/reme/__init__.py b/reme/__init__.py index 31a75142..9f738f72 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -20,7 +20,7 @@ __all__ = [ "ReMeFs", ] -__version__ = "0.3.0.0a6" +__version__ = "0.3.0.0a7" """ diff --git a/reme/config/fs.yaml b/reme/config/fs.yaml index daa59c6a..dc07e41d 100644 --- a/reme/config/fs.yaml +++ b/reme/config/fs.yaml @@ -17,7 +17,9 @@ embedding_models: memory_stores: default: - backend: sqlite +# backend: sqlite + backend: chroma + db_name: reme.db store_name: reme embedding_model: default fts_enabled: true diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py index 295f16fc..ced8fa5f 100644 --- a/reme/core/context/service_context.py +++ b/reme/core/context/service_context.py @@ -89,7 +89,8 @@ class ServiceContext(BaseContext): logger.info(f"ReMe Config: {service_config.model_dump_json()}") if self.service_config.working_dir: - Path(self.service_config.working_dir).mkdir(parents=True, exist_ok=True) + self.working_path = Path(self.service_config.working_dir) + self.working_path.mkdir(parents=True, exist_ok=True) if self.service_config.enable_logo: print_logo(service_config=self.service_config) @@ -198,8 +199,12 @@ class ServiceContext(BaseContext): else: # Extract config dict and replace special fields with actual instances config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict["embedding_model"] = self.embedding_models[config.embedding_model] - config_dict["thread_pool"] = self.thread_pool + config_dict.update( + { + "embedding_model": self.embedding_models[config.embedding_model], + "thread_pool": self.thread_pool, + }, + ) self.vector_stores[name] = R.vector_stores[config.backend](**config_dict) await self.vector_stores[name].create_collection(config.collection_name) @@ -209,7 +214,13 @@ class ServiceContext(BaseContext): else: # Extract config dict and replace embedding_model string with actual instance config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict["embedding_model"] = self.embedding_models[config.embedding_model] + config_dict.update( + { + "embedding_model": self.embedding_models[config.embedding_model], + "thread_pool": self.thread_pool, + "db_path": self.working_path / config.db_name, + }, + ) self.memory_stores[name] = R.memory_stores[config.backend](**config_dict) await self.memory_stores[name].start() diff --git a/reme/core/memory_store/__init__.py b/reme/core/memory_store/__init__.py index 2226e1b0..ab7ba78c 100644 --- a/reme/core/memory_store/__init__.py +++ b/reme/core/memory_store/__init__.py @@ -1,16 +1,19 @@ """Memory store module for persistent memory management. This module provides storage backends for memory chunks and file metadata, -including SQLite-based implementations with vector and full-text search. +including SQLite-based and ChromaDB-based implementations with vector and full-text search. """ from .base_memory_store import BaseMemoryStore +from .chroma_memory_store import ChromaMemoryStore from .sqlite_memory_store import SqliteMemoryStore from ..context import R __all__ = [ "BaseMemoryStore", + "ChromaMemoryStore", "SqliteMemoryStore", ] R.memory_stores.register("sqlite")(SqliteMemoryStore) +R.memory_stores.register("chroma")(ChromaMemoryStore) diff --git a/reme/core/memory_store/base_memory_store.py b/reme/core/memory_store/base_memory_store.py index 069fe820..01121114 100644 --- a/reme/core/memory_store/base_memory_store.py +++ b/reme/core/memory_store/base_memory_store.py @@ -1,7 +1,12 @@ """Base storage interface for memory manager.""" +import asyncio import re from abc import ABC, abstractmethod +from concurrent.futures import ThreadPoolExecutor +from functools import partial +from pathlib import Path +from typing import Callable from ..embedding import BaseEmbeddingModel from ..enumeration import MemorySource @@ -14,6 +19,8 @@ class BaseMemoryStore(ABC): def __init__( self, store_name: str, + db_path: str | Path, + thread_pool: ThreadPoolExecutor, embedding_model: BaseEmbeddingModel, vector_enabled: bool = False, fts_enabled: bool = True, @@ -30,6 +37,8 @@ class BaseMemoryStore(ABC): raise ValueError("At least one of vector_enabled or fts_enabled must be True.") self.store_name: str = store_name + self.db_path: Path = Path(db_path) + self.thread_pool: ThreadPoolExecutor = thread_pool self.embedding_model: BaseEmbeddingModel = embedding_model self.vector_enabled: bool = vector_enabled self.fts_enabled: bool = fts_enabled @@ -40,22 +49,43 @@ class BaseMemoryStore(ABC): """Get the embedding model's dimensionality.""" return self.embedding_model.dimensions + def _get_mock_embedding(self) -> list[float]: + """Generate a zero vector based on embedding model dimensions.""" + return [0.0] * self.embedding_dim + async def get_embedding(self, query: str, **kwargs) -> list[float]: """Get embedding for a single query string.""" + if not self.vector_enabled: + return self._get_mock_embedding() return await self.embedding_model.get_embedding(query, **kwargs) async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]: """Get embeddings for a batch of query strings.""" + if not self.vector_enabled: + return [self._get_mock_embedding() for _ in queries] return await self.embedding_model.get_embeddings(queries, **kwargs) async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: """Generate and populate embedding field for a single MemoryChunk object.""" + if not self.vector_enabled: + chunk.embedding = self._get_mock_embedding() + return chunk return await self.embedding_model.get_chunk_embedding(chunk, **kwargs) async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: """Generate and populate embedding fields for a batch of MemoryChunk objects.""" + if not self.vector_enabled: + mock_embedding = self._get_mock_embedding() + for chunk in chunks: + chunk.embedding = mock_embedding.copy() + return chunks return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs) + async def _run_sync_in_executor(self, sync_func: Callable, *args, **kwargs): + """Run a synchronous function in the context-defined thread pool executor.""" + loop = asyncio.get_running_loop() + return await loop.run_in_executor(self.thread_pool, partial(sync_func, *args, **kwargs)) # noqa + @abstractmethod async def start(self): """Initialize the storage backend.""" @@ -124,6 +154,30 @@ class BaseMemoryStore(ABC): List of search results sorted by relevance """ + @abstractmethod + 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. + candidates = limit * candidate_multiplier + + Returns: + List of search results sorted by combined relevance score + """ + @abstractmethod async def clear_all(self): """Clear all indexed data.""" diff --git a/reme/core/memory_store/chroma_memory_store.py b/reme/core/memory_store/chroma_memory_store.py new file mode 100644 index 00000000..186cdfdf --- /dev/null +++ b/reme/core/memory_store/chroma_memory_store.py @@ -0,0 +1,611 @@ +"""ChromaDB storage backend for memory index.""" + +import json +import time +from pathlib import Path + +from loguru import logger + +from .base_memory_store import BaseMemoryStore +from ..enumeration import MemorySource +from ..schema import FileMetadata, MemoryChunk, MemorySearchResult + +try: + import chromadb + from chromadb.config import Settings + + CHROMADB_AVAILABLE = True +except ImportError: + CHROMADB_AVAILABLE = False + chromadb = None + Settings = None + + +class ChromaMemoryStore(BaseMemoryStore): + """ChromaDB memory storage with vector and full-text search. + + Inherits embedding methods from BaseMemoryStore: + - get_chunk_embedding / get_chunk_embeddings (async) + - get_embedding / get_embeddings (async) + + Provides ChromaDB-backed persistent storage with: + - Vector similarity search (native ChromaDB) + - Full-text search (via ChromaDB where_document filter) + - Efficient chunk and file metadata management + """ + + def __init__( + self, + **kwargs, + ): + if not CHROMADB_AVAILABLE: + raise ImportError( + "chromadb package is required for ChromaMemoryStore. Install it with: pip install chromadb", + ) + + super().__init__(**kwargs) + self.client: "chromadb.ClientAPI | None" = None + 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" + + @property + def collection_name(self) -> str: + """Get the name of the ChromaDB collection for this store.""" + return f"chunks_{self.store_name}" + + async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]: + """Load file metadata from disk. + + Returns: + Dictionary mapping source -> path -> FileMetadata + """ + if not self._metadata_file.exists(): + return {} + + try: + data = await self._run_sync_in_executor( + self._metadata_file.read_text, + encoding="utf-8", + ) + metadata_dict = json.loads(data) + + # Convert dict to FileMetadata objects + result = {} + for source, files in metadata_dict.items(): + result[source] = {} + for path, meta in files.items(): + result[source][path] = FileMetadata(**meta) + + logger.debug(f"Loaded file metadata from {self._metadata_file}") + return result + except Exception as e: + logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") + return {} + + async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None: + """Save file metadata to disk. + + Args: + metadata: Dictionary mapping source -> path -> FileMetadata + """ + try: + # Convert FileMetadata objects to dict for JSON serialization + metadata_dict = {} + for source, files in metadata.items(): + metadata_dict[source] = {} + for path, meta in files.items(): + metadata_dict[source][path] = { + "path": meta.path, + "hash": meta.hash, + "mtime_ms": meta.mtime_ms, + "size": meta.size, + "chunk_count": meta.chunk_count, + } + + 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", + ) + 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}") + + async def start(self) -> None: + """Initialize ChromaDB client and collection.""" + if self.client is not None: + return + + self.db_path.mkdir(parents=True, exist_ok=True) + + # Initialize persistent ChromaDB client + self.client = await self._run_sync_in_executor( + chromadb.PersistentClient, + path=str(self.db_path), + settings=Settings( + anonymized_telemetry=False, + allow_reset=True, + ), + ) + + # 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, + name=self.collection_name, + metadata={"hnsw:space": "cosine"}, + ) + + logger.info(f"ChromaDB initialized with collection: {self.collection_name}") + logger.info(f"File metadata will be persisted to: {self._metadata_file}") + + 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 + + # Delete existing chunks for this file first + await self.delete_file(file_meta.path, source) + + # Batch generate embeddings for all chunks + # (base class returns mock embeddings when vector_enabled=False) + chunks = await self.get_chunk_embeddings(chunks) + + # Prepare data for ChromaDB batch upsert + ids = [] + documents = [] + embeddings = [] + metadatas = [] + + now = int(time.time() * 1000) + for chunk in chunks: + ids.append(chunk.id) + documents.append(chunk.text) + embeddings.append(chunk.embedding) + metadatas.append( + { + "path": file_meta.path, + "source": source.value, + "start_line": chunk.start_line, + "end_line": chunk.end_line, + "hash": chunk.hash, + "updated_at": now, + }, + ) + + # Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function) + await self._run_sync_in_executor( + 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( + 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, + where={"$and": [{"path": path}, {"source": source.value}]}, + include=[], + ) + + if results["ids"]: + await self._run_sync_in_executor( + 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) + + 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, + ids=chunk_ids, + ) + + # Update chunk count in file metadata + metadata = await self._load_metadata() + for source_meta in metadata.values(): + if path in source_meta: + # Recalculate chunk count + results = await self._run_sync_in_executor( + 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( + self, + chunks: list[MemoryChunk], + source: MemorySource, + ) -> None: + """Insert or update specific chunks without affecting other chunks.""" + if not chunks: + return + + # Batch generate embeddings for all chunks + # (base class returns mock embeddings when vector_enabled=False) + chunks = await self.get_chunk_embeddings(chunks) + + ids = [] + documents = [] + embeddings = [] + metadatas = [] + + now = int(time.time() * 1000) + for chunk in chunks: + ids.append(chunk.id) + documents.append(chunk.text) + embeddings.append(chunk.embedding) + metadatas.append( + { + "path": chunk.path, + "source": source.value, + "start_line": chunk.start_line, + "end_line": chunk.end_line, + "hash": chunk.hash, + "updated_at": now, + }, + ) + + # Always pass embeddings to prevent default embedding function + await self._run_sync_in_executor( + self.chunks_collection.upsert, + ids=ids, + documents=documents, + embeddings=embeddings, + metadatas=metadatas, + ) + + 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: + return [] + return list(metadata[source.value].keys()) + + async def get_file_metadata( + self, + path: str, + source: MemorySource, + ) -> FileMetadata | None: + """Get file metadata with chunk count.""" + metadata = await self._load_metadata() + if source.value not in metadata: + return None + return metadata[source.value].get(path) + + async def get_file_chunks( + self, + path: str, + source: MemorySource, + ) -> list[MemoryChunk]: + """Get all chunks for a file.""" + results = await self._run_sync_in_executor( + self.chunks_collection.get, + where={"$and": [{"path": path}, {"source": source.value}]}, + include=["documents", "embeddings", "metadatas"], + ) + + chunks = [] + for i, chunk_id in enumerate(results["ids"]): + metadata = results["metadatas"][i] + chunks.append( + MemoryChunk( + id=chunk_id, + path=metadata["path"], + source=MemorySource(metadata["source"]), + start_line=metadata["start_line"], + end_line=metadata["end_line"], + text=results["documents"][i], + hash=metadata["hash"], + embedding=results["embeddings"][i] if results["embeddings"] is not None else None, + ), + ) + + # Sort by start_line + chunks.sort(key=lambda c: c.start_line) + return chunks + + async def vector_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform vector similarity search.""" + if not self.vector_enabled or not query: + return [] + + # Get query embedding + query_embedding = await self.get_embedding(query) + if not query_embedding: + return [] + + # Build where filter for sources + where_filter = None + if sources: + if len(sources) == 1: + where_filter = {"source": sources[0].value} + else: + where_filter = {"source": {"$in": [s.value for s in sources]}} + + # Perform vector search + results = await self._run_sync_in_executor( + self.chunks_collection.query, + query_embeddings=[query_embedding], + n_results=limit, + where=where_filter, + include=["documents", "metadatas", "distances"], + ) + + search_results = [] + if results["ids"] and results["ids"][0]: + for i, _ in enumerate(results["ids"][0]): + metadata = results["metadatas"][0][i] + distance = results["distances"][0][i] + + # Convert cosine distance to similarity score + # Cosine distance range is [0, 2], convert to [1, 0] score + score = max(0.0, 1.0 - distance / 2.0) + + search_results.append( + MemorySearchResult( + path=metadata["path"], + start_line=metadata["start_line"], + end_line=metadata["end_line"], + score=score, + snippet=results["documents"][0][i], + source=MemorySource(metadata["source"]), + raw_metric=distance, + ), + ) + + # Sort by score descending + search_results.sort(key=lambda r: r.score, reverse=True) + return search_results + + async def keyword_search( + self, + query: str, + limit: int, + sources: list[MemorySource] | None = None, + ) -> list[MemorySearchResult]: + """Perform keyword/full-text search. + + ChromaDB supports where_document filter for text matching. + Note: ChromaDB's $contains is case-sensitive, so we generate multiple + case variants (original, lowercase, capitalized) for each word to + improve recall while maintaining case-insensitive scoring. + """ + if not self.fts_enabled or not query: + return [] + + # Normalize whitespace and split into words + words = query.split() + if not words: + return [] + + # Generate case variants for each word to handle case-sensitive $contains + # Include: original, lowercase, and capitalized forms + word_variants = set() + for word in words: + word_variants.add(word) # original + word_variants.add(word.lower()) # lowercase + word_variants.add(word.capitalize()) # Capitalized + word_variants.add(word.upper()) # UPPERCASE + word_variants_list = list(word_variants) + + # Build where filter for sources + where_filter = None + if sources: + if len(sources) == 1: + where_filter = {"source": sources[0].value} + else: + where_filter = {"source": {"$in": [s.value for s in sources]}} + + # ChromaDB where_document uses $contains for substring matching (case-sensitive) + # Use multiple case variants to improve recall + if len(word_variants_list) == 1: + where_document = {"$contains": word_variants_list[0]} + else: + 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, + where=where_filter, + where_document=where_document, + include=["documents", "metadatas"], + ) + + search_results = [] + query_lower = query.lower() + words_lower = [w.lower() for w in words] # lowercase words for scoring + n_words = len(words) + + for i, _ in enumerate(results["ids"]): + metadata = results["metadatas"][i] + text = results["documents"][i] + text_lower = text.lower() + + # Calculate relevance score based on word matches + match_count = sum(1 for w in words_lower if w in text_lower) + base_score = match_count / n_words + + # Bonus for full phrase match (only applies to multi-word queries) + phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 + # Scale base_score and add phrase bonus, max score is 1.0 + score = min(1.0, base_score + phrase_bonus) + + search_results.append( + MemorySearchResult( + path=metadata["path"], + start_line=metadata["start_line"], + end_line=metadata["end_line"], + score=score, + snippet=text, + source=MemorySource(metadata["source"]), + ), + ) + + # Sort by score descending and limit results + search_results.sort(key=lambda r: r.score, reverse=True) + return search_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 + + # Perform search based on enabled backends + 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) + + # Log original vector results + 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}") + + # Log original keyword results + 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, + ) + + # Log merged results + 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: + vector_results = await self.vector_search(query, limit, sources) + return vector_results + elif self.fts_enabled: + keyword_results = await self.keyword_search(query, limit, sources) + return keyword_results + 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] = {} + + # Process vector results + for result in vector: + result.score = result.score * vector_weight + merged[result.merge_key] = result + + # Process keyword results + 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 + + # Sort by score and return + 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.""" + # Delete and recreate the collection + await self._run_sync_in_executor( + self.client.delete_collection, + name=self.collection_name, + ) + self.chunks_collection = await self._run_sync_in_executor( + self.client.get_or_create_collection, + name=self.collection_name, + metadata={"hnsw:space": "cosine"}, + ) + + # Clear file metadata on disk + 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.""" + # ChromaDB PersistentClient handles persistence automatically + self.client = None + self.chunks_collection = None diff --git a/reme/core/memory_store/sqlite_memory_store.py b/reme/core/memory_store/sqlite_memory_store.py index 1c6a2287..e5aee948 100644 --- a/reme/core/memory_store/sqlite_memory_store.py +++ b/reme/core/memory_store/sqlite_memory_store.py @@ -27,9 +27,8 @@ class SqliteMemoryStore(BaseMemoryStore): - Efficient chunk and file metadata management """ - def __init__(self, db_path: str = ".reme/memory.db", vec_ext_path: str = "", **kwargs): + def __init__(self, vec_ext_path: str = "", **kwargs): super().__init__(**kwargs) - self.db_path = db_path self.vec_ext_path = vec_ext_path self.conn: sqlite3.Connection | None = None @@ -831,6 +830,106 @@ class SqliteMemoryStore(BaseMemoryStore): finally: cursor.close() + 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 + + # Perform search based on enabled backends + 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) + + # Log original vector results + 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}") + + # Log original keyword results + 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, + ) + + # Log merged results + 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: + vector_results = await self.vector_search(query, limit, sources) + return vector_results + elif self.fts_enabled: + keyword_results = await self.keyword_search(query, limit, sources) + return keyword_results + 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] = {} + + # Process vector results + for result in vector: + result.score = result.score * vector_weight + merged[result.merge_key] = result + + # Process keyword results + 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 + + # Sort by score and return + results = list(merged.values()) + results.sort(key=lambda r: r.score, reverse=True) + return results + async def clear_all(self): """Clear all indexed data.""" cursor = self.conn.cursor() diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py index 766d1fd3..0893197b 100644 --- a/reme/core/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -83,6 +83,7 @@ class MemoryStoreConfig(BaseModel): model_config = ConfigDict(extra="allow") backend: str = Field(default="sqlite") + db_name: str = Field(default="reme.db") store_name: str = Field(default="reme") embedding_model: str = Field(default="default") diff --git a/reme/reme_cli.py b/reme/reme_cli.py index 476cdf2e..6a429774 100644 --- a/reme/reme_cli.py +++ b/reme/reme_cli.py @@ -42,8 +42,8 @@ class ReMeCli(ReMeFs): fs_cli = FsCli( tools=[ FsMemorySearch( - hybrid_vector_weight=self.hybrid_vector_weight, - hybrid_candidate_multiplier=self.hybrid_candidate_multiplier, + vector_weight=self.vector_weight, + candidate_multiplier=self.candidate_multiplier, ), BashTool(cwd=self.working_dir), LsTool(cwd=self.working_dir), diff --git a/reme/reme_fs.py b/reme/reme_fs.py index 33c335d6..ae597964 100644 --- a/reme/reme_fs.py +++ b/reme/reme_fs.py @@ -29,31 +29,18 @@ class ReMeFs(Application): log_to_console: bool = True, llm_api_key: str | None = None, llm_base_url: str | None = None, - default_llm_name: str | None = None, - default_llm_config: dict | None = None, embedding_api_key: str | None = None, embedding_base_url: str | None = None, - default_embedding_model_name: str | None = None, + default_llm_config: dict | None = None, default_embedding_model_config: dict | None = None, - default_store_name: str = "reme", - vector_enabled: bool = False, - fts_enabled: bool = True, default_memory_store_config: dict | None = None, - token_counter_backend: str = "base", default_token_counter_config: dict | None = None, - watch_paths: list[str] | None = None, - suffix_filters: list[str] | None = None, - recursive: bool = False, - debounce: int = 500, - chunk_tokens: int = 1000, - chunk_overlap: int = 100, - scan_on_start: bool = True, default_file_watcher_config: dict | None = None, context_window_tokens: int = 128000, reserve_tokens: int = 36000, keep_recent_tokens: int = 20000, - hybrid_vector_weight: float = 0.7, - hybrid_candidate_multiplier: float = 3.0, + vector_weight: float = 0.7, + candidate_multiplier: float = 3.0, **kwargs, ): """Initialize ReMe with config.""" @@ -63,54 +50,20 @@ class ReMeFs(Application): memory_path.mkdir(parents=True, exist_ok=True) self.working_dir: str = str(working_path.absolute()) - default_llm_config = default_llm_config or {} - if default_llm_name: - default_llm_config["model_name"] = default_llm_name - - default_embedding_model_config = default_embedding_model_config or {} - if default_embedding_model_name: - default_embedding_model_config["model_name"] = default_embedding_model_name - - default_memory_store_config = default_memory_store_config or {} - default_memory_store_config.update( - { - "store_name": default_store_name, - "vector_enabled": vector_enabled, - "fts_enabled": fts_enabled, - }, - ) - - default_token_counter_config = default_token_counter_config or {} - default_token_counter_config.update( - { - "backend": token_counter_backend, - }, - ) - default_file_watcher_config = default_file_watcher_config or {} - default_file_watcher_config.update( - { - "watch_paths": watch_paths - or [ - str(working_path / "MEMORY.md"), - str(working_path / "memory.md"), - str(memory_path), - ], - "suffix_filters": suffix_filters or [".md"], - "recursive": recursive, - "debounce": debounce, - "chunk_tokens": chunk_tokens, - "chunk_overlap": chunk_overlap, - "scan_on_start": scan_on_start, - }, - ) - + if not default_file_watcher_config.get("watch_paths", None): + default_file_watcher_config["watch_paths"] = [ + str(working_path / "MEMORY.md"), + str(working_path / "memory.md"), + str(memory_path), + ] super().__init__( *args, llm_api_key=llm_api_key, llm_base_url=llm_base_url, embedding_api_key=embedding_api_key, embedding_base_url=embedding_base_url, + working_dir=working_dir, config_path=config_path, enable_logo=enable_logo, log_to_console=log_to_console, @@ -126,8 +79,8 @@ class ReMeFs(Application): self.context_window_tokens: int = context_window_tokens self.reserve_tokens: int = reserve_tokens self.keep_recent_tokens: int = keep_recent_tokens - self.hybrid_vector_weight: float = hybrid_vector_weight - self.hybrid_candidate_multiplier: float = hybrid_candidate_multiplier + self.vector_weight: float = vector_weight + self.candidate_multiplier: float = candidate_multiplier async def context_check(self, messages: list[Message | dict]) -> dict: """Check if messages exceed context limits.""" @@ -194,8 +147,8 @@ class ReMeFs(Application): Search results as formatted string """ search_tool = FsMemorySearch( - hybrid_vector_weight=self.hybrid_vector_weight, - hybrid_candidate_multiplier=self.hybrid_candidate_multiplier, + vector_weight=self.vector_weight, + candidate_multiplier=self.candidate_multiplier, ) return await search_tool.call( query=query, diff --git a/reme/tool/fs/fs_memory_search.py b/reme/tool/fs/fs_memory_search.py index 012dc0d8..3edd344b 100644 --- a/reme/tool/fs/fs_memory_search.py +++ b/reme/tool/fs/fs_memory_search.py @@ -2,10 +2,8 @@ import json -from loguru import logger - from reme.core.enumeration import MemorySource -from reme.core.schema import MemorySearchResult, ToolCall +from reme.core.schema import ToolCall from .base_fs_tool import BaseFsTool @@ -17,21 +15,19 @@ class FsMemorySearch(BaseFsTool): sources: list[MemorySource] | None = None, min_score: float = 0.1, max_results: int = 5, - hybrid_vector_weight: float = 0.7, - hybrid_candidate_multiplier: float = 3.0, + vector_weight: float = 0.7, + candidate_multiplier: float = 3.0, **kwargs, ): """Initialize memory search tool.""" - assert ( - 0.0 <= hybrid_vector_weight <= 1.0 - ), f"hybrid_vector_weight must be between 0 and 1, got {hybrid_vector_weight}" + assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" kwargs.setdefault("name", "memory_search") super().__init__(**kwargs) self.sources = sources or [MemorySource.MEMORY] self.min_score = min_score self.max_results = max_results - self.hybrid_vector_weight = hybrid_vector_weight - self.hybrid_candidate_multiplier = hybrid_candidate_multiplier + self.vector_weight = vector_weight + self.candidate_multiplier = candidate_multiplier def _build_tool_call(self) -> ToolCall: return ToolCall( @@ -68,93 +64,17 @@ class FsMemorySearch(BaseFsTool): query: str = self.context.query.strip() min_score = self.context.get("min_score", self.min_score) max_results = self.context.get("max_results", self.max_results) - candidates = min(200, max(1, int(max_results * self.hybrid_candidate_multiplier))) - vector_enabled = self.memory_store.vector_enabled - fts_enabled = self.memory_store.fts_enabled + # Use hybrid_search from memory_store + results = await self.memory_store.hybrid_search( + query=query, + limit=max_results, + sources=self.sources, + vector_weight=self.vector_weight, + candidate_multiplier=self.candidate_multiplier, + ) - # Perform search based on enabled backends - if vector_enabled and fts_enabled: - keyword_results = await self._search_keyword(query, candidates) - vector_results = await self._search_vector(query, candidates) - - # Log original vector results - 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}") - - # Log original keyword results - 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: - results = [r for r in vector_results if r.score >= min_score][:max_results] - elif not vector_results: - results = [r for r in keyword_results if r.score >= min_score][:max_results] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=self.hybrid_vector_weight, - text_weight=1.0 - self.hybrid_vector_weight, - ) - - # Log merged results - 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}") - - results = [r for r in merged if r.score >= min_score][:max_results] - elif vector_enabled: - vector_results = await self._search_vector(query, candidates) - results = [r for r in vector_results if r.score >= min_score][:max_results] - elif fts_enabled: - keyword_results = await self._search_keyword(query, candidates) - results = [r for r in keyword_results if r.score >= min_score][:max_results] - else: - results = [] + # Filter by min_score + results = [r for r in results if r.score >= min_score] return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False) - - async def _search_vector(self, query: str, limit: int) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - return await self.memory_store.vector_search(query, limit, sources=self.sources) - - async def _search_keyword(self, query: str, limit: int) -> list[MemorySearchResult]: - """Perform keyword/FTS search.""" - if not self.memory_store.fts_enabled: - return [] - return await self.memory_store.keyword_search(query, limit, sources=self.sources) - - @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] = {} - - # Process vector results - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - # Process keyword results - 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 - - # Sort by score and return - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results diff --git a/tests/test_fs_memory_search.py b/tests/test_fs_memory_search.py index d2a6256f..c6de7e44 100644 --- a/tests/test_fs_memory_search.py +++ b/tests/test_fs_memory_search.py @@ -611,7 +611,7 @@ async def test_memory_search_hybrid_mode(): "fts_enabled": True, }, search_params={ - "hybrid_vector_weight": 0.7, + "vector_weight": 0.7, }, ) await reme_fs_hybrid.start() @@ -672,7 +672,7 @@ async def test_memory_search_hybrid_mode(): "fts_enabled": True, }, search_params={ - "hybrid_vector_weight": vec_weight, + "vector_weight": vec_weight, }, ) await reme_fs_weights.start() diff --git a/tests/test_memory_store.py b/tests/test_memory_store.py index df45d77d..a5937b10 100644 --- a/tests/test_memory_store.py +++ b/tests/test_memory_store.py @@ -1,11 +1,12 @@ # pylint: disable=too-many-lines """Unified test suite for memory store implementations. -This module provides comprehensive test coverage for SqliteMemoryStore and future -memory store implementations. Tests can be run for specific stores or all 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. Usage: python test_memory_store.py --sqlite # Test SqliteMemoryStore only + python test_memory_store.py --chroma # Test ChromaMemoryStore only python test_memory_store.py --all # Test all memory stores """ @@ -22,6 +23,7 @@ from loguru import logger 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.sqlite_memory_store import SqliteMemoryStore from reme.core.schema.file_metadata import FileMetadata from reme.core.schema.memory_chunk import MemoryChunk @@ -44,6 +46,10 @@ class TestConfig: SQLITE_VEC_EXT_PATH = "" # Empty string to use default vec0/sqlite_vec/vector0 SQLITE_FTS_ENABLED = True + # ChromaMemoryStore settings + CHROMA_DB_PATH = "./test_memory_store_chroma" + CHROMA_FTS_ENABLED = True + # Embedding model settings EMBEDDING_MODEL_NAME = "text-embedding-v4" EMBEDDING_DIMENSIONS = 64 @@ -178,10 +184,12 @@ def get_store_type(store: BaseMemoryStore) -> str: store: Memory store instance Returns: - str: Type identifier ("sqlite", etc.) + str: Type identifier ("sqlite", "chroma", etc.) """ if isinstance(store, SqliteMemoryStore): return "sqlite" + elif isinstance(store, ChromaMemoryStore): + return "chroma" else: raise ValueError(f"Unknown memory store type: {type(store)}") @@ -190,7 +198,7 @@ def create_memory_store(store_type: str) -> BaseMemoryStore: """Create a memory store instance based on type. Args: - store_type: Type of memory store ("sqlite", etc.) + store_type: Type of memory store ("sqlite", "chroma", etc.) Returns: BaseMemoryStore: Initialized memory store instance @@ -211,6 +219,13 @@ def create_memory_store(store_type: str) -> BaseMemoryStore: vec_ext_path=config.SQLITE_VEC_EXT_PATH, fts_enabled=config.SQLITE_FTS_ENABLED, ) + elif store_type == "chroma": + return ChromaMemoryStore( + store_name=config.NAME, + db_path=config.CHROMA_DB_PATH, + embedding_model=embedding_model, + fts_enabled=config.CHROMA_FTS_ENABLED, + ) else: raise ValueError(f"Unknown store type: {store_type}") @@ -239,6 +254,12 @@ async def test_start_store(store: BaseMemoryStore, _store_name: str): assert store.chunks_table_name in tables, f"{store.chunks_table_name} table should exist" logger.info("✓ Required tables created") + # Verify ChromaDB collection created + if isinstance(store, ChromaMemoryStore): + assert store.client is not None, "ChromaDB client should be initialized" + assert store.chunks_collection is not None, "ChromaDB collection should exist" + logger.info(f"✓ ChromaDB collection created: {store.collection_name}") + async def test_upsert_file(store: BaseMemoryStore, _store_name: str) -> tuple[FileMetadata, List[MemoryChunk]]: """Test file and chunks insertion.""" @@ -414,11 +435,9 @@ async def test_vector_search(store: BaseMemoryStore, _store_name: str): """Test vector similarity search.""" logger.info("=" * 20 + " VECTOR SEARCH TEST " + "=" * 20) - # Check if vector search is available (SQLite-specific) - if isinstance(store, SqliteMemoryStore) and not store.vector_available: - logger.warning("⚠ Vector extension not available, skipping vector search tests") - logger.info(" Install sqlite-vec extension to enable vector search") - logger.info(" See: https://github.com/asg017/sqlite-vec") + # Check if vector search is enabled + if not store.vector_enabled: + logger.warning("⚠ Vector search not enabled, skipping vector search tests") return # Search for AI-related content @@ -447,9 +466,9 @@ async def test_vector_search_with_source_filter(store: BaseMemoryStore, _store_n """Test vector search with source filtering.""" logger.info("=" * 20 + " VECTOR SEARCH WITH SOURCE FILTER TEST " + "=" * 20) - # Check if vector search is available (SQLite-specific) - if isinstance(store, SqliteMemoryStore) and not store.vector_available: - logger.warning("⚠ Vector extension not available, skipping vector search with source filter test") + # Check if vector search is enabled + if not store.vector_enabled: + logger.warning("⚠ Vector search not enabled, skipping vector search with source filter test") return query = "sales data analysis" @@ -503,9 +522,9 @@ async def test_keyword_search(store: BaseMemoryStore, _store_name: str): """Test full-text keyword search.""" logger.info("=" * 20 + " KEYWORD SEARCH TEST " + "=" * 20) - # Check if FTS is available - if isinstance(store, SqliteMemoryStore) and not store.fts_available: - logger.info("⊘ Skipped: FTS not available") + # Check if FTS is enabled + if not store.fts_enabled: + logger.info("⊘ Skipped: FTS not enabled") return query = "neural networks" @@ -535,9 +554,9 @@ async def test_keyword_search_with_source_filter(store: BaseMemoryStore, _store_ """Test keyword search with source filtering.""" logger.info("=" * 20 + " KEYWORD SEARCH WITH SOURCE FILTER TEST " + "=" * 20) - # Check if FTS is available - if isinstance(store, SqliteMemoryStore) and not store.fts_available: - logger.info("⊘ Skipped: FTS not available") + # Check if FTS is enabled + if not store.fts_enabled: + logger.info("⊘ Skipped: FTS not enabled") return query = "data" @@ -581,9 +600,9 @@ async def test_keyword_search_special_chars(store: BaseMemoryStore, _store_name: """Test keyword search with special characters like ?, *, etc.""" logger.info("=" * 20 + " KEYWORD SEARCH SPECIAL CHARS TEST " + "=" * 20) - # Check if FTS is available - if isinstance(store, SqliteMemoryStore) and not store.fts_available: - logger.info("⊘ Skipped: FTS not available") + # Check if FTS is enabled + if not store.fts_enabled: + logger.info("⊘ Skipped: FTS not enabled") return # Test various queries with special characters @@ -681,16 +700,17 @@ async def test_concurrent_searches(store: BaseMemoryStore, _store_name: str): "data analysis techniques", ] - # Concurrent vector searches - search_tasks = [store.vector_search(q, limit=3) for q in queries] - results = await asyncio.gather(*search_tasks) + # Concurrent vector searches (if available) + if store.vector_enabled: + search_tasks = [store.vector_search(q, limit=3) for q in queries] + results = await asyncio.gather(*search_tasks) - logger.info(f"✓ Completed {len(results)} concurrent vector searches") - for i, (query, result) in enumerate(zip(queries, results), 1): - logger.info(f" Query {i}: '{query}' -> {len(result)} results") + logger.info(f"✓ Completed {len(results)} concurrent vector searches") + for i, (query, result) in enumerate(zip(queries, results), 1): + logger.info(f" Query {i}: '{query}' -> {len(result)} results") # Concurrent keyword searches (if available) - if isinstance(store, SqliteMemoryStore) and store.fts_available: + if store.fts_enabled: keyword_tasks = [store.keyword_search(q, limit=3) for q in queries] keyword_results = await asyncio.gather(*keyword_tasks) logger.info(f"✓ Completed {len(keyword_results)} concurrent keyword searches") @@ -796,15 +816,17 @@ async def test_edge_cases(store: BaseMemoryStore, _store_name: str): logger.info("✓ Handled unicode and emoji") # Test 5: Search with empty query - try: - results = await store.vector_search("", limit=5) - logger.info(f"✓ Empty query returned {len(results)} results") - except Exception as e: - logger.info(f"⊘ Empty query not supported: {e}") + if store.vector_enabled: + try: + results = await store.vector_search("", limit=5) + logger.info(f"✓ Empty query returned {len(results)} results") + except Exception as e: + logger.info(f"⊘ Empty query not supported: {e}") # Test 6: Very high limit - results = await store.vector_search("test", limit=1000) - logger.info(f"✓ High limit search returned {len(results)} results") + if store.vector_enabled: + results = await store.vector_search("test", limit=1000) + logger.info(f"✓ High limit search returned {len(results)} results") # Test 7: Non-existent file non_existent_meta = await store.get_file_metadata("non_existent_file.txt", MemorySource.MEMORY) @@ -924,7 +946,7 @@ async def cleanup_store(store: BaseMemoryStore, store_type: str): Args: store: Memory store instance - store_type: Type of memory store ("sqlite", etc.) + store_type: Type of memory store ("sqlite", "chroma", etc.) """ logger.info("=" * 20 + " CLEANUP " + "=" * 20) @@ -941,6 +963,19 @@ async def cleanup_store(store: BaseMemoryStore, store_type: str): shutil.rmtree(db_dir) logger.info(f"✓ Cleaned up directory: {db_dir}") + # Clean up local directory if ChromaMemoryStore + if store_type == "chroma": + config = TestConfig() + db_dir = Path(config.CHROMA_DB_PATH) + if db_dir.exists(): + shutil.rmtree(db_dir) + logger.info(f"✓ Cleaned up directory: {db_dir}") + # Also clean up the metadata file + metadata_file = db_dir.parent / f"{config.NAME}_file_metadata.json" + if metadata_file.exists(): + metadata_file.unlink() + logger.info(f"✓ Cleaned up metadata file: {metadata_file}") + logger.info("✓ Cleanup completed") except Exception as e: logger.error(f"Cleanup error: {e}") @@ -957,6 +992,7 @@ async def main(): epilog=""" Examples: python test_memory_store.py --sqlite # Test SqliteMemoryStore only + python test_memory_store.py --chroma # Test ChromaMemoryStore only python test_memory_store.py --all # Test all memory stores """, ) @@ -965,6 +1001,11 @@ Examples: action="store_true", help="Test SqliteMemoryStore", ) + parser.add_argument( + "--chroma", + action="store_true", + help="Test ChromaMemoryStore", + ) parser.add_argument( "--all", action="store_true", @@ -979,19 +1020,23 @@ Examples: if args.all: stores_to_test = [ ("sqlite", "SqliteMemoryStore"), + ("chroma", "ChromaMemoryStore"), ] else: # Build list based on individual flags if args.sqlite: stores_to_test.append(("sqlite", "SqliteMemoryStore")) + if args.chroma: + stores_to_test.append(("chroma", "ChromaMemoryStore")) if not stores_to_test: # Default to all memory stores if no argument provided stores_to_test = [ ("sqlite", "SqliteMemoryStore"), + ("chroma", "ChromaMemoryStore"), ] print("No memory store specified, defaulting to test all memory stores") - print("Use --sqlite to test specific ones\n") + print("Use --sqlite or --chroma to test specific ones\n") # Run tests for each memory store for store_type, store_name in stores_to_test: