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
This commit is contained in:
jinliyl 2026-03-10 18:19:00 +08:00 • committed by GitHub
parent 083ed6a137
commit 8b1698451e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 192 additions and 62 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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: |
上述消息是要整合到现有摘要中的新对话消息,这些消息在<previous-summary>标签中提供。
以上消息是需要整合到现有摘要中的新对话内容,现有摘要位于<previous-summary>标签中。
update_user_message_suffix_zh: |
用新信息更新现有的结构化摘要。规则: