mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
feat(store): add ChromaDB memory store implementation with hybrid search
This commit is contained in:
parent
4b414b18de
commit
52d21392f2
13 changed files with 906 additions and 207 deletions
|
|
@ -20,7 +20,7 @@ __all__ = [
|
|||
"ReMeFs",
|
||||
]
|
||||
|
||||
__version__ = "0.3.0.0a6"
|
||||
__version__ = "0.3.0.0a7"
|
||||
|
||||
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
611
reme/core/memory_store/chroma_memory_store.py
Normal file
611
reme/core/memory_store/chroma_memory_store.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue