From 8b1698451e96b99c71ffe157f02e1f2e64061d16 Mon Sep 17 00:00:00 2001 From: jinliyl <6469360+jinliyl@users.noreply.github.com> Date: Tue, 10 Mar 2026 18:19:00 +0800 Subject: [PATCH] fin(emb_dim) (#150) * fix(logger): handle file logging configuration errors gracefully * fix(embedding): handle embedding dimension mismatches and improve caching * fix(file-watcher): clear file store on changes to prevent stale data * refactor(logger): update logger implementation and fix message translation * refactor(logger): update logger configuration and add documentation --- reme/__init__.py | 2 +- reme/core/embedding/base_embedding_model.py | 92 +++++++++++++++---- reme/core/file_store/chroma_file_store.py | 44 +++++++-- reme/core/file_store/local_file_store.py | 24 ++++- reme/core/file_watcher/delta_file_watcher.py | 1 + reme/core/file_watcher/full_file_watcher.py | 2 + reme/core/utils/logger_utils.py | 44 +++++---- reme/core/utils/std_logger.py | 43 +++++---- .../file_based/components/compactor.yaml | 2 +- 9 files changed, 192 insertions(+), 62 deletions(-) diff --git a/reme/__init__.py b/reme/__init__.py index c75dfbcc..41537822 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -6,7 +6,7 @@ from . import extension from . import memory from .reme import ReMe -__version__ = "0.3.0.6b2" +__version__ = "0.3.0.6b3" __all__ = [ "config", diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py index 78a91b8b..2b60e8e6 100644 --- a/reme/core/embedding/base_embedding_model.py +++ b/reme/core/embedding/base_embedding_model.py @@ -29,7 +29,7 @@ class BaseEmbeddingModel(ABC): api_key: str | None = None, base_url: str | None = None, model_name: str = "", - dimensions: int | None = 1024, + dimensions: int = 1024, use_dimensions: bool = False, max_batch_size: int = 10, max_retries: int = 3, @@ -99,7 +99,34 @@ class BaseEmbeddingModel(ABC): """Truncate a list of texts to max_input_length.""" return [self._truncate_text(text) for text in texts] - def _get_cache_key(self, text: str) -> str: + def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]: + """Validate and adjust embedding dimensions to match expected dimensions. + + Args: + embedding: The embedding vector to validate + + Returns: + Embedding vector adjusted to match self.dimensions + """ + actual_len = len(embedding) + if actual_len == self.dimensions: + return embedding + + elif actual_len < self.dimensions: + logger.warning( + f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, " + f"padding with zeros", + ) + return embedding + [0.0] * (self.dimensions - actual_len) + + else: + logger.warning( + f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, " + f"truncating to {self.dimensions}", + ) + return embedding[: self.dimensions] + + def _get_cache_key(self, text: str, dimensions: int) -> str: """Generate a cache key by hashing text + model_name + dimensions. This ensures that the same text produces different cache keys when @@ -107,12 +134,13 @@ class BaseEmbeddingModel(ABC): Args: text: Input text to hash + dimensions: Vector dimensions of the embeddings Returns: SHA256 hash combining text, model name, and dimensions """ # Combine text, model_name, and dimensions to create unique cache key - cache_string = f"{text}|{self.model_name}|{self.dimensions}" + cache_string = f"{text}|{self.model_name}|{dimensions}" return hashlib.sha256(cache_string.encode("utf-8")).hexdigest() def _get_cache_file_path(self) -> Path: @@ -164,6 +192,13 @@ class BaseEmbeddingModel(ABC): if cache_key in self._embedding_cache: continue + if len(embedding) != self.dimensions: + logger.warning( + f"Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got {len(embedding)}", + ) + continue + # Respect max_cache_size during loading if len(self._embedding_cache) >= self.max_cache_size: logger.info( @@ -204,6 +239,12 @@ class BaseEmbeddingModel(ABC): try: with open(cache_file, "w", encoding="utf-8") as f: for cache_key, embedding in self._embedding_cache.items(): + if len(embedding) != self.dimensions: + logger.warning( + f"Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got {len(embedding)}", + ) + continue cache_entry = {cache_key: embedding} f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n") @@ -223,16 +264,27 @@ class BaseEmbeddingModel(ABC): if not self.enable_cache: return None - cache_key = self._get_cache_key(text) + cache_key = self._get_cache_key(text, self.dimensions) if cache_key in self._embedding_cache: + embeddings: list[float] = self._embedding_cache[cache_key] + + # Validate embedding dimensions match expected dimensions + if len(embeddings) != self.dimensions: + logger.warning( + f"Cached embedding dimensions mismatch: expected {self.dimensions}, " + f"got {len(embeddings)}. Removing invalid cache entry.", + ) + del self._embedding_cache[cache_key] + self._cache_misses += 1 + return None + # Move to end (most recently used) self._embedding_cache.move_to_end(cache_key) self._cache_hits += 1 text_preview = text[:50] + "..." if len(text) > 50 else text - logger.info( - f"Cache hit for text: '{text_preview}' (hits: {self._cache_hits}, misses: {self._cache_misses})", - ) - return self._embedding_cache[cache_key] + logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})") + return embeddings + self._cache_misses += 1 return None @@ -249,9 +301,15 @@ class BaseEmbeddingModel(ABC): if self.max_cache_size <= 0: return - cache_key = self._get_cache_key(text) + cache_key = self._get_cache_key(text, self.dimensions) + if len(embedding) != self.dimensions: + logger.warning( + f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, " + f"expected {self.dimensions}, got real length {len(embedding)}", + ) + return - # Remove oldest entry if cache is full + # Remove the oldest entry if cache is full if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache: self._embedding_cache.popitem(last=False) @@ -299,7 +357,7 @@ class BaseEmbeddingModel(ABC): for i in range(self.max_retries): try: result = await self._get_embeddings([truncated_text], **kwargs) - embedding = result[0] + embedding = self._validate_and_adjust_embedding(result[0]) # Store in cache self._put_to_cache(truncated_text, embedding) return embedding @@ -345,8 +403,9 @@ class BaseEmbeddingModel(ABC): if batch_embeddings: # Store results and cache them for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): - results[orig_idx] = embedding - self._put_to_cache(text, embedding) + adjusted_embedding = self._validate_and_adjust_embedding(embedding) + results[orig_idx] = adjusted_embedding + self._put_to_cache(text, adjusted_embedding) break except Exception as e: logger.error(f"Model {self.model_name} batch failed: {e}") @@ -371,7 +430,7 @@ class BaseEmbeddingModel(ABC): for i in range(self.max_retries): try: result = self._get_embeddings_sync([truncated_text], **kwargs) - embedding = result[0] + embedding = self._validate_and_adjust_embedding(result[0]) # Store in cache self._put_to_cache(truncated_text, embedding) return embedding @@ -417,8 +476,9 @@ class BaseEmbeddingModel(ABC): if batch_embeddings: # Store results and cache them for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): - results[orig_idx] = embedding - self._put_to_cache(text, embedding) + adjusted_embedding = self._validate_and_adjust_embedding(embedding) + results[orig_idx] = adjusted_embedding + self._put_to_cache(text, adjusted_embedding) break except Exception as exc: logger.error(f"Model {self.model_name} batch failed: {exc}") diff --git a/reme/core/file_store/chroma_file_store.py b/reme/core/file_store/chroma_file_store.py index 390a2bc4..d87f5fef 100644 --- a/reme/core/file_store/chroma_file_store.py +++ b/reme/core/file_store/chroma_file_store.py @@ -1,6 +1,7 @@ """ChromaDB storage backend for file store.""" import json +import random import time from pathlib import Path @@ -355,12 +356,41 @@ class ChromaFileStore(BaseFileStore): where_filter = {"source": {"$in": [s.value for s in sources]}} # Perform vector search - results = self.chunks_collection.query( - query_embeddings=[query_embedding], - n_results=limit, - where=where_filter, - include=["documents", "metadatas", "distances"], - ) + try: + results = self.chunks_collection.query( + query_embeddings=[query_embedding], + n_results=limit, + where=where_filter, + include=["documents", "metadatas", "distances"], + ) + except Exception as e: + logger.error(f"Vector search failed: {e}, falling back to random results") + # Fallback: get some documents without vector search and assign random scores + try: + fallback_results = self.chunks_collection.get( + where=where_filter, + limit=limit, + include=["documents", "metadatas"], + ) + search_results = [] + if fallback_results["ids"]: + for i, _ in enumerate(fallback_results["ids"]): + metadata = fallback_results["metadatas"][i] + search_results.append( + MemorySearchResult( + path=metadata["path"], + start_line=metadata["start_line"], + end_line=metadata["end_line"], + score=random.uniform(0.3, 0.7), # Random score in middle range + snippet=fallback_results["documents"][i], + source=MemorySource(metadata["source"]), + raw_metric=None, + ), + ) + return search_results + except Exception as fallback_e: + logger.error(f"Fallback search also failed: {fallback_e}") + return [] search_results = [] if results["ids"] and results["ids"][0]: @@ -430,7 +460,7 @@ class ChromaFileStore(BaseFileStore): # 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]} + where_document: dict = {"$contains": word_variants_list[0]} else: where_document = {"$or": [{"$contains": w} for w in word_variants_list]} diff --git a/reme/core/file_store/local_file_store.py b/reme/core/file_store/local_file_store.py index 37df17f4..38757d7a 100644 --- a/reme/core/file_store/local_file_store.py +++ b/reme/core/file_store/local_file_store.py @@ -259,6 +259,8 @@ class LocalFileStore(BaseFileStore): if not query_embedding: return [] + expected_dim = self.embedding_dim + # Collect candidate chunks with embeddings candidates = [ chunk for chunk in self._chunks.values() if (not sources or chunk.source in sources) and chunk.embedding @@ -267,9 +269,29 @@ class LocalFileStore(BaseFileStore): if not candidates: return [] + # Validate and fix chunk embedding dimensions + valid_embeddings = [] + for chunk in candidates: + emb = chunk.embedding + emb_len = len(emb) + if emb_len != expected_dim: + if emb_len < expected_dim: + emb = emb + [0.0] * (expected_dim - emb_len) + logger.warning( + f"Chunk embedding dimension {emb_len} < expected {expected_dim}, " + f"padded with zeros (chunk_id={chunk.id})", + ) + else: + emb = emb[:expected_dim] + logger.warning( + f"Chunk embedding dimension {emb_len} > expected {expected_dim}, " + f"truncated to {expected_dim} (chunk_id={chunk.id})", + ) + valid_embeddings.append(emb) + # Build embedding matrix and compute similarities in batch query_array = np.array([query_embedding]) # Shape: (1, emb_size) - chunk_embeddings = np.array([chunk.embedding for chunk in candidates]) # Shape: (n, emb_size) + chunk_embeddings = np.array(valid_embeddings) # Shape: (n, emb_size) similarities = batch_cosine_similarity(query_array, chunk_embeddings)[0] # Shape: (n,) # Build results diff --git a/reme/core/file_watcher/delta_file_watcher.py b/reme/core/file_watcher/delta_file_watcher.py index 6148bd07..f35c9b2f 100644 --- a/reme/core/file_watcher/delta_file_watcher.py +++ b/reme/core/file_watcher/delta_file_watcher.py @@ -141,6 +141,7 @@ class DeltaFileWatcher(BaseFileWatcher): async def _on_changes(self, changes: set[tuple[Change, str]]): """Handle file changes with incremental synchronization.""" self.dirty = True + await self.file_store.clear_all() for change_type, path in changes: if change_type == Change.added: diff --git a/reme/core/file_watcher/full_file_watcher.py b/reme/core/file_watcher/full_file_watcher.py index 2d365f23..c49a94fa 100644 --- a/reme/core/file_watcher/full_file_watcher.py +++ b/reme/core/file_watcher/full_file_watcher.py @@ -44,6 +44,8 @@ class FullFileWatcher(BaseFileWatcher): async def _on_changes(self, changes: set[tuple[Change, str]]): """Handle file changes with full synchronization""" self.dirty = True + await self.file_store.clear_all() + for change_type, path in changes: if change_type in [Change.added, Change.modified]: file_meta = await self._build_file_metadata(path) diff --git a/reme/core/utils/logger_utils.py b/reme/core/utils/logger_utils.py index 22819db8..512c0a42 100644 --- a/reme/core/utils/logger_utils.py +++ b/reme/core/utils/logger_utils.py @@ -18,26 +18,6 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool # Remove default handler to avoid duplicate logs logger.remove() - # Ensure the logging directory exists - os.makedirs(log_dir, exist_ok=True) - - # Generate filename based on the current timestamp - # Use dashes instead of colons for Windows compatibility - current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - log_filename = f"{current_ts}.log" - log_filepath = os.path.join(log_dir, log_filename) - - # Configure file-based logging with rotation and compression - logger.add( - log_filepath, - level=level, - rotation="00:00", - retention="7 days", - compression="zip", - encoding="utf-8", - format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}", - ) - # Configure colorized standard output logging if enabled if log_to_console: logger.add( @@ -46,3 +26,27 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}", colorize=True, ) + + # Try to configure file-based logging (skip if permission denied) + try: + # Ensure the logging directory exists + os.makedirs(log_dir, exist_ok=True) + + # Generate filename based on the current timestamp + # Use dashes instead of colons for Windows compatibility + current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + log_filename = f"{current_ts}.log" + log_filepath = os.path.join(log_dir, log_filename) + + # Configure file-based logging with rotation and compression + logger.add( + log_filepath, + level=level, + rotation="00:00", + retention="7 days", + compression="zip", + encoding="utf-8", + format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}", + ) + except Exception as e: + logger.error(f"Error configuring file logging: {e}") diff --git a/reme/core/utils/std_logger.py b/reme/core/utils/std_logger.py index e2de0e9c..dbdf1908 100644 --- a/reme/core/utils/std_logger.py +++ b/reme/core/utils/std_logger.py @@ -38,7 +38,7 @@ class CustomFormatter(logging.Formatter): return super().format(record) -def get_logger( +def get_loggerv2( name: str = "reme", log_dir: str = "logs", level: str = "INFO", @@ -80,22 +80,26 @@ def get_logger( # Configure file logging if log_to_file: - os.makedirs(log_dir, exist_ok=True) - current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - log_filename = f"{log_file_prefix}_{current_ts}.log" - log_filepath = os.path.join(log_dir, log_filename) + try: + os.makedirs(log_dir, exist_ok=True) + current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + log_filename = f"{log_file_prefix}_{current_ts}.log" + log_filepath = os.path.join(log_dir, log_filename) - file_handler = TimedRotatingFileHandler( - log_filepath, - when=rotation, - interval=1, - backupCount=retention_days, - encoding="utf-8", - ) - file_handler.setLevel(getattr(logging, level.upper(), logging.INFO)) - file_handler.setFormatter(CustomFormatter(log_format, colorize=False)) - file_handler.suffix = "%Y-%m-%d" - logger.addHandler(file_handler) + file_handler = TimedRotatingFileHandler( + log_filepath, + when=rotation, + interval=1, + backupCount=retention_days, + encoding="utf-8", + ) + file_handler.setLevel(getattr(logging, level.upper(), logging.INFO)) + file_handler.setFormatter(CustomFormatter(log_format, colorize=False)) + file_handler.suffix = "%Y-%m-%d" + logger.addHandler(file_handler) + + except Exception as e: + logger.error(f"Error configuring file logging: {e}") # Configure console logging if log_to_console: @@ -107,3 +111,10 @@ def get_logger( # Cache logger _loggers[name] = logger return logger + + +def get_logger(): + """Get a configured logger instance using loguru.""" + from loguru import logger + + return logger diff --git a/reme/memory/file_based/components/compactor.yaml b/reme/memory/file_based/components/compactor.yaml index 2e3b15a8..83c4f952 100644 --- a/reme/memory/file_based/components/compactor.yaml +++ b/reme/memory/file_based/components/compactor.yaml @@ -119,7 +119,7 @@ update_user_message_suffix: | Keep each section concise. Preserve exact file paths, function names, and error messages. update_user_message_prefix_zh: | - 上述消息是要整合到现有摘要中的新对话消息,这些消息在标签中提供。 + 以上消息是需要整合到现有摘要中的新对话内容,现有摘要位于标签中。 update_user_message_suffix_zh: | 用新信息更新现有的结构化摘要。规则: