feat(store): add ChromaDB memory store implementation with hybrid search

This commit is contained in:
jinli.yl 2026-02-12 20:34:50 +08:00
parent 4b414b18de
commit 52d21392f2
13 changed files with 906 additions and 207 deletions

View file

@ -20,7 +20,7 @@ __all__ = [
"ReMeFs",
]
__version__ = "0.3.0.0a6"
__version__ = "0.3.0.0a7"
"""

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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