From 0a0ad89b221fe850765304f03923f1a4fd20c144 Mon Sep 17 00:00:00 2001 From: Wenqiang Wei <46308778+endxxxx@users.noreply.github.com> Date: Wed, 15 Jul 2026 19:29:31 +0800 Subject: [PATCH 1/5] refactor: summarize search and embedding logs --- src/memos/api/handlers/search_handler.py | 10 +- src/memos/embedders/ark.py | 3 +- src/memos/embedders/base.py | 55 +++++++++++ src/memos/embedders/ollama.py | 3 +- src/memos/embedders/sentence_transformer.py | 3 +- src/memos/embedders/universal_api.py | 33 +------ src/memos/log.py | 72 ++++++++++++++ .../tree_text_memory/retrieve/searcher.py | 10 +- src/memos/multi_mem_cube/single_cube.py | 7 +- tests/embedders/test_base.py | 78 ++++++++++++++- tests/test_log.py | 95 +++++++++++++++++++ 11 files changed, 321 insertions(+), 48 deletions(-) diff --git a/src/memos/api/handlers/search_handler.py b/src/memos/api/handlers/search_handler.py index 03e6977ad..2f170a8a0 100644 --- a/src/memos/api/handlers/search_handler.py +++ b/src/memos/api/handlers/search_handler.py @@ -16,7 +16,7 @@ from memos.api.handlers.formatters_handler import rerank_knowledge_mem from memos.api.product_models import APISearchRequest, SearchResponse from memos.dream.contextualization import CONTEXT_MEMORY_TYPE -from memos.log import get_logger +from memos.log import get_logger, summarize_search_request, summarize_search_results from memos.memories.textual.tree_text_memory.retrieve.retrieve_utils import ( cosine_similarity_matrix, ) @@ -77,7 +77,10 @@ def handle_search_memories(self, search_req: APISearchRequest) -> SearchResponse Returns: SearchResponse with formatted results """ - self.logger.info(f"[SearchHandler] Search Req is: {search_req}") + self.logger.info( + "[SearchHandler] Search request summary: %s", + summarize_search_request(search_req), + ) # Use deepcopy to avoid modifying the original request object search_req_local = copy.deepcopy(search_req) @@ -137,7 +140,8 @@ def handle_search_memories(self, search_req: APISearchRequest) -> SearchResponse results = hooked_results self.logger.info( - f"[SearchHandler] Final search results: count={len(results)} results={results}" + "[SearchHandler] Final search result summary: %s", + summarize_search_results(results), ) return SearchResponse( diff --git a/src/memos/embedders/ark.py b/src/memos/embedders/ark.py index a8b47e200..aa183fcdc 100644 --- a/src/memos/embedders/ark.py +++ b/src/memos/embedders/ark.py @@ -1,6 +1,6 @@ from memos.configs.embedder import ArkEmbedderConfig from memos.dependency import require_python_package -from memos.embedders.base import BaseEmbedder +from memos.embedders.base import BaseEmbedder, log_embedding_call from memos.log import get_logger @@ -35,6 +35,7 @@ def __init__(self, config: ArkEmbedderConfig): # Initialize ark client self.client = Ark(api_key=self.config.api_key, base_url=self.config.api_base) + @log_embedding_call def embed(self, texts: list[str]) -> list[list[float]]: """ Generate embeddings for the given texts. diff --git a/src/memos/embedders/base.py b/src/memos/embedders/base.py index e46611d1a..4cc8e6c85 100644 --- a/src/memos/embedders/base.py +++ b/src/memos/embedders/base.py @@ -1,8 +1,63 @@ +import functools import re +import time from abc import ABC, abstractmethod +from collections.abc import Callable +from typing import Any, TypeVar, cast from memos.configs.embedder import BaseEmbedderConfig +from memos.log import get_logger, text_hash + + +logger = get_logger(__name__) +EmbeddingCallable = TypeVar("EmbeddingCallable", bound=Callable[..., Any]) + + +def log_embedding_call(func: EmbeddingCallable) -> EmbeddingCallable: + """Log embedding request dimensions and timing without text or vectors.""" + + @functools.wraps(func) + def wrapper(self, texts, *args, **kwargs): + normalized_texts = [texts] if isinstance(texts, str) else list(texts or []) + text_lengths = [len(str(text or "")) for text in normalized_texts] + config = getattr(self, "config", None) + model = getattr(config, "model_name_or_path", None) or "unknown" + backup_model = getattr(config, "backup_model_name_or_path", None) or "none" + backup_enabled = bool(getattr(self, "use_backup_client", False)) + started_at = time.perf_counter() + status = "success" + error_type = None + try: + return func(self, texts, *args, **kwargs) + except Exception as exc: + status = "failed" + error_type = type(exc).__name__ + raise + finally: + elapsed_ms = (time.perf_counter() - started_at) * 1000 + log_message = ( + "Embedding request model=%s backup_model=%s backup_enabled=%s " + "batch_size=%d total_chars=%d max_chars=%d text_hash=%s " + "elapsed_ms=%.2f status=%s" + ) + log_values = ( + model, + backup_model, + backup_enabled, + len(normalized_texts), + sum(text_lengths), + max(text_lengths, default=0), + text_hash(normalized_texts), + elapsed_ms, + status, + ) + if error_type is None: + logger.info(log_message, *log_values) + else: + logger.info(log_message + " error_type=%s", *log_values, error_type) + + return cast("EmbeddingCallable", wrapper) def _count_tokens_for_embedding(text: str) -> int: diff --git a/src/memos/embedders/ollama.py b/src/memos/embedders/ollama.py index dfd8e230d..789dfb9c0 100644 --- a/src/memos/embedders/ollama.py +++ b/src/memos/embedders/ollama.py @@ -1,7 +1,7 @@ from ollama import Client from memos.configs.embedder import OllamaEmbedderConfig -from memos.embedders.base import BaseEmbedder +from memos.embedders.base import BaseEmbedder, log_embedding_call from memos.log import get_logger @@ -57,6 +57,7 @@ def _ensure_model_exists(self): except Exception as e: logger.warning(f"Could not verify model existence: {e}") + @log_embedding_call def embed(self, texts: list[str]) -> list[list[float]]: """ Generate embeddings for the given texts. diff --git a/src/memos/embedders/sentence_transformer.py b/src/memos/embedders/sentence_transformer.py index de086cb49..7b1cad146 100644 --- a/src/memos/embedders/sentence_transformer.py +++ b/src/memos/embedders/sentence_transformer.py @@ -1,6 +1,6 @@ from memos.configs.embedder import SenTranEmbedderConfig from memos.dependency import require_python_package -from memos.embedders.base import BaseEmbedder +from memos.embedders.base import BaseEmbedder, log_embedding_call from memos.log import get_logger @@ -32,6 +32,7 @@ def __init__(self, config: SenTranEmbedderConfig): # Get embedding dimensions from the model self.config.embedding_dims = self.model.get_sentence_embedding_dimension() + @log_embedding_call def embed(self, texts: list[str]) -> list[list[float]]: """ Generate embeddings for the given texts. diff --git a/src/memos/embedders/universal_api.py b/src/memos/embedders/universal_api.py index 66ed4e62c..24a022eae 100644 --- a/src/memos/embedders/universal_api.py +++ b/src/memos/embedders/universal_api.py @@ -1,32 +1,17 @@ import asyncio import os -import time from openai import AzureOpenAI as AzureClient from openai import OpenAI as OpenAIClient from memos.configs.embedder import UniversalAPIEmbedderConfig -from memos.embedders.base import BaseEmbedder +from memos.embedders.base import BaseEmbedder, log_embedding_call from memos.log import get_logger -from memos.utils import timed_with_status logger = get_logger(__name__) -def _embedding_log_extra_args(embedder: "UniversalAPIEmbedder", texts: list[str] | str) -> dict: - text_items = [texts] if isinstance(texts, str) else texts - return { - "model_name_or_path": getattr( - embedder.config, "model_name_or_path", "text-embedding-3-large" - ), - "backup_model_name_or_path": getattr(embedder.config, "backup_model_name_or_path", None), - "use_backup_client": getattr(embedder, "use_backup_client", False), - "text_len": len(text_items), - "text_content": text_items, - } - - def _sanitize_unicode(text: str) -> str: """ Remove Unicode surrogates and other problematic characters. @@ -71,10 +56,7 @@ def __init__(self, config: UniversalAPIEmbedderConfig): else None, ) - @timed_with_status( - log_prefix="model_timed_embedding", - log_extra_args=_embedding_log_extra_args, - ) + @log_embedding_call def embed(self, texts: list[str]) -> list[list[float]]: if isinstance(texts, str): texts = [texts] @@ -82,7 +64,6 @@ def embed(self, texts: list[str]) -> list[list[float]]: texts = [_sanitize_unicode(t) for t in texts] # Truncate texts if max_tokens is configured texts = self._truncate_texts(texts) - logger.info(f"Embeddings request with input: {texts}") if self.provider == "openai" or self.provider == "azure": try: @@ -92,18 +73,17 @@ async def _create_embeddings(): input=texts, ) - init_time = time.time() response = asyncio.run( asyncio.wait_for( _create_embeddings(), timeout=int(os.getenv("MOS_EMBEDDER_TIMEOUT", 5)) ) ) - logger.info(f"Embeddings request succeeded with {time.time() - init_time} seconds") return [r.embedding for r in response.data] except Exception as e: if self.use_backup_client: logger.warning( - f"Embeddings request ended with {type(e).__name__} error: {e}, try backup client" + "Embedding request failed error_type=%s; trying backup client", + type(e).__name__, ) try: @@ -117,17 +97,12 @@ async def _create_embeddings_backup(): input=texts, ) - init_time = time.time() response = asyncio.run( asyncio.wait_for( _create_embeddings_backup(), timeout=int(os.getenv("MOS_EMBEDDER_TIMEOUT", 5)), ) ) - logger.info( - f"Backup embeddings request succeeded with {time.time() - init_time} seconds" - ) - logger.info(f"Backup embeddings request response: {response}") return [r.embedding for r in response.data] except Exception as e: raise ValueError(f"Backup embeddings request ended with error: {e}") from e diff --git a/src/memos/log.py b/src/memos/log.py index c18bd2118..0e521df76 100644 --- a/src/memos/log.py +++ b/src/memos/log.py @@ -1,13 +1,17 @@ import atexit +import hashlib import logging import os import threading import time +from collections import Counter +from collections.abc import Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from logging.config import dictConfig from pathlib import Path from sys import stdout +from typing import Any import requests @@ -31,6 +35,74 @@ _LOGGING_CONFIGURED_PID: int | None = None +def text_hash(texts: str | Sequence[str]) -> str: + """Return a stable short hash without retaining the source text.""" + normalized_texts = [texts] if isinstance(texts, str) else texts + digest = hashlib.sha256() + for text in normalized_texts: + encoded = str(text or "").encode("utf-8", errors="replace") + digest.update(len(encoded).to_bytes(8, byteorder="big")) + digest.update(encoded) + return digest.hexdigest()[:16] + + +def summarize_search_request(request: Any) -> dict[str, Any]: + """Summarize search controls without logging query or user content.""" + query = str(getattr(request, "query", "") or "") + mode = getattr(request, "mode", None) + readable_cube_ids = getattr(request, "readable_cube_ids", None) or [] + views = getattr(request, "include_memory_view", None) or [] + return { + "query_chars": len(query), + "query_hash": text_hash(query), + "mode": getattr(mode, "value", mode), + "readable_cube_count": len(readable_cube_ids), + "views": list(views), + "top_k": getattr(request, "memory_limit_number", None), + "dedup": getattr(request, "dedup", None), + "rerank": getattr(request, "rerank", None), + } + + +def summarize_search_results(results: Mapping[str, Any]) -> dict[str, Any]: + """Count result buckets and items without serializing result payloads.""" + bucket_counts: dict[str, int] = {} + item_counts: dict[str, int] = {} + for key, value in results.items(): + if not isinstance(value, list): + continue + bucket_counts[key] = len(value) + item_counts[key] = sum( + len(bucket.get("memories") or []) + if isinstance(bucket, Mapping) and isinstance(bucket.get("memories"), list) + else 1 + for bucket in value + ) + return { + "bucket_counts": bucket_counts, + "item_counts": item_counts, + "total_items": sum(item_counts.values()), + } + + +def summarize_textual_memories(memories: Sequence[Any]) -> dict[str, Any]: + """Count textual memories by type without reading their content or metadata values.""" + type_counts: Counter[str] = Counter() + for item in memories: + metadata = ( + item.get("metadata") if isinstance(item, Mapping) else getattr(item, "metadata", None) + ) + if isinstance(metadata, Mapping): + memory_type = metadata.get("memory_type") + else: + memory_type = getattr(metadata, "memory_type", None) + type_counts[str(memory_type or "unknown")] += 1 + return { + "total_items": len(memories), + "type_counts": dict(type_counts), + } + + def _setup_logfile() -> Path: """ensure the logger filepath is in place diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py index cd27d92a1..5c9cce78e 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py @@ -9,7 +9,7 @@ from memos.embedders.factory import OllamaEmbedder from memos.graph_dbs.factory import Neo4jGraphDB from memos.llms.factory import AzureLLM, OllamaLLM, OpenAILLM -from memos.log import get_logger +from memos.log import get_logger, summarize_textual_memories from memos.memories.textual.item import SearchedTreeNodeTextualMemoryMetadata, TextualMemoryItem from memos.memories.textual.tree_text_memory.retrieve.bm25_util import EnhancedBM25 from memos.memories.textual.tree_text_memory.retrieve.retrieve_utils import ( @@ -277,13 +277,7 @@ def search( dedup=dedup, ) - logger.info(f"[SEARCH] Done. Total {len(final_results)} results.") - res_results = "" - for _num_i, result in enumerate(final_results): - res_results += "\n" + ( - result.id + "|" + result.metadata.memory_type + "|" + result.memory - ) - logger.info(f"[SEARCH] Results. {res_results}") + logger.info("[SEARCH] Result summary: %s", summarize_textual_memories(final_results)) return final_results @timed diff --git a/src/memos/multi_mem_cube/single_cube.py b/src/memos/multi_mem_cube/single_cube.py index f84fc60e1..27399259a 100644 --- a/src/memos/multi_mem_cube/single_cube.py +++ b/src/memos/multi_mem_cube/single_cube.py @@ -12,7 +12,7 @@ format_memory_item, post_process_textual_mem, ) -from memos.log import get_logger +from memos.log import get_logger, summarize_search_request, summarize_search_results from memos.mem_reader.utils import parse_keep_filter_response from memos.mem_scheduler.schemas.message_schemas import ScheduleMessageItem from memos.mem_scheduler.schemas.task_schemas import ( @@ -104,7 +104,7 @@ def search_memories(self, search_req: APISearchRequest) -> dict[str, Any]: mem_cube_id=self.cube_id, session_id=search_req.session_id or "default_session", ) - self.logger.info(f"Search Req is: {search_req}") + self.logger.info("Search request summary: %s", summarize_search_request(search_req)) memories_result: MOSSearchResult = { "text_mem": [], @@ -129,8 +129,7 @@ def search_memories(self, search_req: APISearchRequest) -> dict[str, Any]: self.cube_id, ) - self.logger.info(f"Search memories result: {memories_result}") - self.logger.info(f"Search {len(memories_result)} memories.") + self.logger.info("Search result summary: %s", summarize_search_results(memories_result)) return memories_result @timed diff --git a/tests/embedders/test_base.py b/tests/embedders/test_base.py index 029d3a4c2..47a05864d 100644 --- a/tests/embedders/test_base.py +++ b/tests/embedders/test_base.py @@ -1,6 +1,82 @@ -from memos.embedders.base import BaseEmbedder +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from memos.embedders.base import BaseEmbedder, log_embedding_call from tests.utils import check_module_base_class def test_base_embedder_class(): check_module_base_class(BaseEmbedder) + + +def test_log_embedding_call_records_safe_structured_summary(): + class StubEmbedder: + config = SimpleNamespace(model_name_or_path="embedding-model") + + @log_embedding_call + def embed(self, texts: list[str]) -> list[list[float]]: + return [[0.1, 0.2] for _ in texts] + + private_texts = ["private first text", "private second text"] + with patch("memos.embedders.base.logger") as mock_logger: + result = StubEmbedder().embed(private_texts) + + assert result == [[0.1, 0.2], [0.1, 0.2]] + log_args = mock_logger.info.call_args.args + rendered = log_args[0] % log_args[1:] + assert "model=embedding-model" in rendered + assert "batch_size=2" in rendered + assert "total_chars=37" in rendered + assert "max_chars=19" in rendered + assert "text_hash=" in rendered + assert "elapsed_ms=" in rendered + assert "status=success" in rendered + assert "private first text" not in rendered + assert "private second text" not in rendered + assert "0.1" not in rendered + + +def test_log_embedding_call_records_error_type_without_exception_content(): + class FailingEmbedder: + config = SimpleNamespace(model_name_or_path="embedding-model") + + @log_embedding_call + def embed(self, texts: list[str]) -> list[list[float]]: + raise ValueError(f"failed to embed {texts}") + + with ( + patch("memos.embedders.base.logger") as mock_logger, + pytest.raises(ValueError, match="private failing text"), + ): + FailingEmbedder().embed(["private failing text"]) + + log_args = mock_logger.info.call_args.args + rendered = log_args[0] % log_args[1:] + assert "status=failed" in rendered + assert "error_type=ValueError" in rendered + assert "private failing text" not in rendered + + +def test_log_embedding_call_records_backup_model_without_text_content(): + class StubEmbedder: + config = SimpleNamespace( + model_name_or_path="primary-model", + backup_model_name_or_path="backup-model", + ) + use_backup_client = True + + @log_embedding_call + def embed(self, texts: list[str]) -> list[list[float]]: + return [[0.1] for _ in texts] + + with patch("memos.embedders.base.logger") as mock_logger: + StubEmbedder().embed(["private input"]) + + log_args = mock_logger.info.call_args.args + rendered = log_args[0] % log_args[1:] + assert "model=primary-model" in rendered + assert "backup_model=backup-model" in rendered + assert "backup_enabled=True" in rendered + assert "private input" not in rendered diff --git a/tests/test_log.py b/tests/test_log.py index 5387adc34..5b72fb9c6 100644 --- a/tests/test_log.py +++ b/tests/test_log.py @@ -1,6 +1,8 @@ import logging import os +from types import SimpleNamespace + from dotenv import load_dotenv from memos import log @@ -59,3 +61,96 @@ def getpid(): log.get_logger("child") assert len(calls) == 2 + + +def test_summarize_search_request_uses_metadata_and_query_hash_only(): + request = SimpleNamespace( + query="private query content", + user_id="private-user", + mode="fast", + readable_cube_ids=["cube-a", "cube-b"], + include_memory_view=["detail_factual", "preference"], + memory_limit_number=6, + dedup="mmr", + rerank=True, + ) + + summary = log.summarize_search_request(request) + + assert summary["query_chars"] == len(request.query) + assert len(summary["query_hash"]) == 16 + assert summary["mode"] == "fast" + assert summary["readable_cube_count"] == 2 + assert summary["views"] == ["detail_factual", "preference"] + assert summary["top_k"] == 6 + assert summary["dedup"] == "mmr" + assert summary["rerank"] is True + assert request.query not in str(summary) + assert request.user_id not in str(summary) + + +def test_summarize_search_results_only_counts_buckets_and_items(): + results = { + "text_mem": [ + { + "cube_id": "cube-a", + "memories": [ + { + "memory": "private memory value", + "metadata": { + "embedding": [0.1] * 256, + "properties": {"secret": 1}, + }, + }, + {"memory": "another private value"}, + ], + } + ], + "pref_mem": [{"cube_id": "cube-a", "memories": [{"memory": "private preference"}]}], + "skill_mem": [], + "pref_note": "private preference note", + } + + summary = log.summarize_search_results(results) + + assert summary == { + "bucket_counts": {"text_mem": 1, "pref_mem": 1, "skill_mem": 0}, + "item_counts": {"text_mem": 2, "pref_mem": 1, "skill_mem": 0}, + "total_items": 3, + } + rendered = str(summary) + assert "private memory value" not in rendered + assert "0.1" not in rendered + assert "secret" not in rendered + assert "private preference note" not in rendered + assert len(rendered) <= len(str(results)) * 0.2 + + +def test_summarize_textual_memories_only_counts_types(): + memories = [ + SimpleNamespace( + memory="private long-term memory", + metadata=SimpleNamespace( + memory_type="LongTermMemory", + embedding=[0.1, 0.2], + ), + ), + SimpleNamespace( + memory="private user memory", + metadata=SimpleNamespace( + memory_type="UserMemory", + properties={"secret": True}, + ), + ), + ] + + summary = log.summarize_textual_memories(memories) + + assert summary == { + "total_items": 2, + "type_counts": {"LongTermMemory": 1, "UserMemory": 1}, + } + rendered = str(summary) + assert "private long-term memory" not in rendered + assert "0.1" not in rendered + assert "secret" not in rendered From 54ad5e4dbadbf356176a0194cfb69366bf3fd422 Mon Sep 17 00:00:00 2001 From: Wenqiang Wei <46308778+endxxxx@users.noreply.github.com> Date: Wed, 15 Jul 2026 20:22:03 +0800 Subject: [PATCH 2/5] perf: cache and coalesce embedding requests --- src/memos/embedders/cache.py | 229 +++++++++++++++++++++++++++++++++ src/memos/embedders/factory.py | 7 +- tests/embedders/test_cache.py | 185 ++++++++++++++++++++++++++ 3 files changed, 420 insertions(+), 1 deletion(-) create mode 100644 src/memos/embedders/cache.py create mode 100644 tests/embedders/test_cache.py diff --git a/src/memos/embedders/cache.py b/src/memos/embedders/cache.py new file mode 100644 index 000000000..963d2d725 --- /dev/null +++ b/src/memos/embedders/cache.py @@ -0,0 +1,229 @@ +import os +import threading + +from collections import Counter +from concurrent.futures import Future +from typing import Any + +from cachetools import LRUCache, TTLCache + +from memos.context.context import get_current_trace_id +from memos.embedders.base import BaseEmbedder +from memos.exceptions import EmbedderError +from memos.log import get_logger + + +logger = get_logger(__name__) + + +_OPTIMIZATION_ENABLED_ENV = "MEMOS_EMBEDDING_OPTIMIZATION_ENABLED" +_CACHE_TTL_ENV = "MEMOS_EMBEDDING_CACHE_TTL_SECONDS" +_CACHE_MAX_SIZE_ENV = "MEMOS_EMBEDDING_CACHE_MAX_SIZE" +_REQUEST_CACHE_TTL_ENV = "MEMOS_EMBEDDING_REQUEST_CACHE_TTL_SECONDS" +_REQUEST_CACHE_MAX_REQUESTS_ENV = "MEMOS_EMBEDDING_REQUEST_CACHE_MAX_REQUESTS" + +_DEFAULT_CACHE_TTL_SECONDS = 30.0 +_DEFAULT_CACHE_MAX_SIZE = 4096 +_DEFAULT_REQUEST_CACHE_TTL_SECONDS = 60.0 +_DEFAULT_REQUEST_CACHE_MAX_REQUESTS = 1024 + +_INVALID_REQUEST_IDS = {None, "", "trace-id"} + +CachedVector = tuple[float, ...] + + +def _env_enabled(name: str, default: bool = False) -> bool: + raw = os.getenv(name) + if raw is None: + return default + return raw.strip().lower() in {"1", "true", "yes", "y", "on"} + + +def embedding_optimization_enabled() -> bool: + return _env_enabled(_OPTIMIZATION_ENABLED_ENV) + + +def _env_float(name: str, default: float, minimum: float = 0.0) -> float: + raw = os.getenv(name) + if raw is None: + return default + try: + return max(minimum, float(raw)) + except ValueError: + logger.warning("Invalid %s=%r; using default %s", name, raw, default) + return default + + +def _env_int(name: str, default: int, minimum: int = 1) -> int: + raw = os.getenv(name) + if raw is None: + return default + try: + return max(minimum, int(raw)) + except ValueError: + logger.warning("Invalid %s=%r; using default %s", name, raw, default) + return default + + +class CachingEmbedder(BaseEmbedder): + """Add exact caches and singleflight coordination to an embedder.""" + + def __init__(self, backend: BaseEmbedder): + self._backend = backend + self.config = backend.config + self._lock = threading.RLock() + self._inflight: dict[str, Future[CachedVector]] = {} + self._stats: Counter[str] = Counter() + + cache_ttl = _env_float(_CACHE_TTL_ENV, _DEFAULT_CACHE_TTL_SECONDS) + cache_max_size = _env_int(_CACHE_MAX_SIZE_ENV, _DEFAULT_CACHE_MAX_SIZE) + self._cache: TTLCache[str, CachedVector] | None = ( + TTLCache(maxsize=cache_max_size, ttl=cache_ttl) if cache_ttl > 0 else None + ) + + request_cache_ttl = _env_float(_REQUEST_CACHE_TTL_ENV, _DEFAULT_REQUEST_CACHE_TTL_SECONDS) + request_cache_max_requests = _env_int( + _REQUEST_CACHE_MAX_REQUESTS_ENV, _DEFAULT_REQUEST_CACHE_MAX_REQUESTS + ) + self._request_cache_text_limit = cache_max_size + self._request_caches: TTLCache[str, LRUCache[str, CachedVector]] | None = ( + TTLCache(maxsize=request_cache_max_requests, ttl=request_cache_ttl) + if request_cache_ttl > 0 + else None + ) + + def __getattr__(self, name: str) -> Any: + return getattr(self._backend, name) + + def embed(self, texts: list[str]) -> list[list[float]]: + if not embedding_optimization_enabled(): + return self._backend.embed(texts) + if isinstance(texts, str): + texts = [texts] + if not texts: + return [] + + unique_texts = list(dict.fromkeys(texts)) + request_id = get_current_trace_id() + resolved: dict[str, CachedVector] = {} + owned: dict[str, Future[CachedVector]] = {} + waiting: dict[str, Future[CachedVector]] = {} + local_stats = Counter( + { + "batch_dedup_hits": len(texts) - len(unique_texts), + "request_hits": 0, + "ttl_hits": 0, + "singleflight_joins": 0, + "misses": 0, + } + ) + + with self._lock: + request_cache = self._get_request_cache(request_id) + for text in unique_texts: + if request_cache is not None and text in request_cache: + resolved[text] = request_cache[text] + local_stats["request_hits"] += 1 + elif self._cache is not None and text in self._cache: + resolved[text] = self._cache[text] + if request_cache is not None: + request_cache[text] = resolved[text] + local_stats["ttl_hits"] += 1 + elif text in self._inflight: + waiting[text] = self._inflight[text] + local_stats["singleflight_joins"] += 1 + else: + future: Future[CachedVector] = Future() + self._inflight[text] = future + owned[text] = future + local_stats["misses"] += 1 + + self._stats.update(local_stats) + + if owned: + owned_texts = list(owned) + with self._lock: + self._stats["backend_calls"] += 1 + self._stats["backend_texts"] += len(owned_texts) + try: + computed = self._backend.embed(owned_texts) + if len(computed) != len(owned_texts): + raise EmbedderError( + "Embedding backend returned a different number of vectors than texts" + ) + computed_vectors = [tuple(vector) for vector in computed] + except Exception as exc: + with self._lock: + self._stats["backend_errors"] += 1 + for text, future in owned.items(): + self._inflight.pop(text, None) + future.set_exception(exc) + raise + + with self._lock: + for text, vector in zip(owned_texts, computed_vectors, strict=True): + resolved[text] = vector + if self._cache is not None: + self._cache[text] = vector + if request_cache is not None: + request_cache[text] = vector + self._inflight.pop(text, None) + owned[text].set_result(vector) + + for text, future in waiting.items(): + vector = future.result() + resolved[text] = vector + if request_cache is not None: + with self._lock: + request_cache[text] = vector + + logger.info( + "embedding cache summary batch_size=%d unique_texts=%d " + "request_hits=%d ttl_hits=%d batch_dedup_hits=%d " + "singleflight_joins=%d misses=%d", + len(texts), + len(unique_texts), + local_stats["request_hits"], + local_stats["ttl_hits"], + local_stats["batch_dedup_hits"], + local_stats["singleflight_joins"], + local_stats["misses"], + ) + return [list(resolved[text]) for text in texts] + + def _get_request_cache(self, request_id: str | None) -> LRUCache[str, CachedVector] | None: + if request_id in _INVALID_REQUEST_IDS or self._request_caches is None: + return None + request_cache = self._request_caches.get(request_id) + if request_cache is None: + request_cache = LRUCache(maxsize=self._request_cache_text_limit) + self._request_caches[request_id] = request_cache + return request_cache + + def cache_info(self) -> dict[str, int]: + with self._lock: + info = dict(self._stats) + for key in ( + "batch_dedup_hits", + "request_hits", + "ttl_hits", + "singleflight_joins", + "misses", + "backend_calls", + "backend_texts", + "backend_errors", + ): + info.setdefault(key, 0) + info["ttl_cache_size"] = len(self._cache) if self._cache is not None else 0 + info["request_cache_count"] = ( + len(self._request_caches) if self._request_caches is not None else 0 + ) + info["inflight"] = len(self._inflight) + return info + + def clear_cache(self) -> None: + with self._lock: + if self._cache is not None: + self._cache.clear() + if self._request_caches is not None: + self._request_caches.clear() diff --git a/src/memos/embedders/factory.py b/src/memos/embedders/factory.py index be14db9e2..4f9d1bda8 100644 --- a/src/memos/embedders/factory.py +++ b/src/memos/embedders/factory.py @@ -3,6 +3,7 @@ from memos.configs.embedder import EmbedderConfigFactory from memos.embedders.ark import ArkEmbedder from memos.embedders.base import BaseEmbedder +from memos.embedders.cache import CachingEmbedder, embedding_optimization_enabled from memos.embedders.ollama import OllamaEmbedder from memos.embedders.sentence_transformer import SenTranEmbedder from memos.embedders.universal_api import UniversalAPIEmbedder @@ -18,6 +19,7 @@ class EmbedderFactory(BaseEmbedder): "ark": ArkEmbedder, "universal_api": UniversalAPIEmbedder, } + cacheable_backends: ClassVar[set[str]] = {"ollama", "ark", "universal_api"} @classmethod @singleton_factory() @@ -26,4 +28,7 @@ def from_config(cls, config_factory: EmbedderConfigFactory) -> BaseEmbedder: if backend not in cls.backend_to_class: raise ValueError(f"Invalid backend: {backend}") embedder_class = cls.backend_to_class[backend] - return embedder_class(config_factory.config) + embedder = embedder_class(config_factory.config) + if backend in cls.cacheable_backends and embedding_optimization_enabled(): + return CachingEmbedder(embedder) + return embedder diff --git a/tests/embedders/test_cache.py b/tests/embedders/test_cache.py new file mode 100644 index 000000000..2baff6780 --- /dev/null +++ b/tests/embedders/test_cache.py @@ -0,0 +1,185 @@ +import threading +import time + +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from memos.configs.embedder import EmbedderConfigFactory +from memos.context.context import RequestContext, set_request_context +from memos.embedders.cache import CachingEmbedder +from memos.embedders.factory import EmbedderFactory + + +@pytest.fixture(autouse=True) +def clear_request_context(): + set_request_context(None) + yield + set_request_context(None) + + +def _backend(side_effect=None): + backend = MagicMock() + backend.config = SimpleNamespace(model_name_or_path="embedding-model") + backend.embed.side_effect = side_effect or ( + lambda texts: [[float(len(text)), float(index)] for index, text in enumerate(texts)] + ) + return backend + + +def _enable_optimization(monkeypatch, ttl_seconds="0"): + monkeypatch.setenv("MEMOS_EMBEDDING_OPTIMIZATION_ENABLED", "true") + monkeypatch.setenv("MEMOS_EMBEDDING_CACHE_TTL_SECONDS", ttl_seconds) + monkeypatch.setenv("MEMOS_EMBEDDING_CACHE_MAX_SIZE", "32") + monkeypatch.setenv("MEMOS_EMBEDDING_REQUEST_CACHE_TTL_SECONDS", "60") + monkeypatch.setenv("MEMOS_EMBEDDING_REQUEST_CACHE_MAX_REQUESTS", "8") + + +def test_disabled_cache_preserves_backend_batch(monkeypatch): + monkeypatch.setenv("MEMOS_EMBEDDING_OPTIMIZATION_ENABLED", "false") + backend = _backend() + embedder = CachingEmbedder(backend) + + first = embedder.embed(["same", "same"]) + second = embedder.embed(["same", "same"]) + + assert first == second + assert backend.embed.call_count == 2 + backend.embed.assert_called_with(["same", "same"]) + + +def test_enabled_cache_treats_string_as_one_text(monkeypatch): + _enable_optimization(monkeypatch) + backend = _backend() + embedder = CachingEmbedder(backend) + + result = embedder.embed("whole query") + + assert result == [[11.0, 0.0]] + backend.embed.assert_called_once_with(["whole query"]) + + +def test_reuses_duplicate_texts_within_request(monkeypatch): + _enable_optimization(monkeypatch) + backend = _backend() + embedder = CachingEmbedder(backend) + set_request_context(RequestContext(trace_id="request-a")) + + first = embedder.embed(["same", "same", "other"]) + first[0][0] = 999.0 + second = embedder.embed(["same"]) + + backend.embed.assert_called_once_with(["same", "other"]) + assert first[1] == [4.0, 0.0] + assert second == [[4.0, 0.0]] + assert embedder.cache_info()["request_hits"] == 1 + + +def test_short_ttl_cache_reuses_text_across_requests(monkeypatch): + _enable_optimization(monkeypatch, ttl_seconds="60") + backend = _backend() + embedder = CachingEmbedder(backend) + + set_request_context(RequestContext(trace_id="request-a")) + first = embedder.embed(["same"]) + set_request_context(RequestContext(trace_id="request-b")) + second = embedder.embed(["same"]) + + assert first == second + backend.embed.assert_called_once_with(["same"]) + assert embedder.cache_info()["ttl_hits"] == 1 + + +def test_partial_cache_hit_preserves_input_order(monkeypatch): + _enable_optimization(monkeypatch, ttl_seconds="60") + backend = _backend() + embedder = CachingEmbedder(backend) + + assert embedder.embed(["cached"]) == [[6.0, 0.0]] + result = embedder.embed(["new-a", "cached", "new-b", "new-a"]) + + assert result == [ + [5.0, 0.0], + [6.0, 0.0], + [5.0, 1.0], + [5.0, 0.0], + ] + assert backend.embed.call_args_list[-1].args[0] == ["new-a", "new-b"] + + +def test_request_cache_does_not_leak_when_ttl_disabled(monkeypatch): + _enable_optimization(monkeypatch) + backend = _backend() + embedder = CachingEmbedder(backend) + + set_request_context(RequestContext(trace_id="request-a")) + embedder.embed(["same"]) + set_request_context(RequestContext(trace_id="request-b")) + embedder.embed(["same"]) + + assert backend.embed.call_count == 2 + + +def test_concurrent_same_text_uses_single_backend_call(monkeypatch): + _enable_optimization(monkeypatch, ttl_seconds="60") + backend_started = threading.Event() + release_backend = threading.Event() + + def blocking_embed(texts): + backend_started.set() + assert release_backend.wait(timeout=2) + return [[1.0, 2.0] for _ in texts] + + backend = _backend(blocking_embed) + embedder = CachingEmbedder(backend) + + with ThreadPoolExecutor(max_workers=2) as executor: + first = executor.submit(embedder.embed, ["same"]) + assert backend_started.wait(timeout=2) + second = executor.submit(embedder.embed, ["same"]) + + deadline = time.monotonic() + 2 + while embedder.cache_info()["singleflight_joins"] != 1: + assert time.monotonic() < deadline + time.sleep(0.01) + release_backend.set() + + assert first.result(timeout=2) == [[1.0, 2.0]] + assert second.result(timeout=2) == [[1.0, 2.0]] + + backend.embed.assert_called_once_with(["same"]) + + +def test_backend_failure_is_not_cached(monkeypatch): + _enable_optimization(monkeypatch, ttl_seconds="60") + backend = _backend() + backend.embed.side_effect = [ValueError("temporary"), [[1.0, 2.0]]] + embedder = CachingEmbedder(backend) + + with pytest.raises(ValueError, match="temporary"): + embedder.embed(["same"]) + + assert embedder.embed(["same"]) == [[1.0, 2.0]] + assert backend.embed.call_count == 2 + + +def test_factory_wraps_remote_embedder_when_enabled(monkeypatch): + _enable_optimization(monkeypatch) + backend = _backend() + monkeypatch.setitem(EmbedderFactory.backend_to_class, "ollama", lambda _config: backend) + config = EmbedderConfigFactory.model_validate( + { + "backend": "ollama", + "config": { + "model_name_or_path": "cache-test-model", + "api_base": "http://cache-test.invalid", + }, + } + ) + + embedder = EmbedderFactory.from_config(config) + + assert isinstance(embedder, CachingEmbedder) + assert embedder.config is backend.config From 372861b7cb4af2c518e5087e746dcbeff222b67d Mon Sep 17 00:00:00 2001 From: Wenqiang Wei <46308778+endxxxx@users.noreply.github.com> Date: Thu, 16 Jul 2026 19:04:18 +0800 Subject: [PATCH 3/5] fix: summarize cosine reranker logs --- src/memos/reranker/cosine_local.py | 10 +++++++- tests/reranker/test_local_rerankers.py | 33 ++++++++++++++++++++++++++ 2 files changed, 42 insertions(+), 1 deletion(-) diff --git a/src/memos/reranker/cosine_local.py b/src/memos/reranker/cosine_local.py index 318cd744a..140074b1d 100644 --- a/src/memos/reranker/cosine_local.py +++ b/src/memos/reranker/cosine_local.py @@ -98,5 +98,13 @@ def get_weight(it: TextualMemoryItem) -> float: chosen = {it.id for it, _ in top_items} remain = [(it, -1.0) for it in graph_results if it.id not in chosen] top_items.extend(remain[: top_k - len(top_items)]) - logger.info(f"CosineLocalReranker rerank result: {top_items[:1]}") + top_score = round(top_items[0][1], 6) if top_items else None + logger.info( + "CosineLocalReranker rerank result: input_count=%s embedded_count=%s " + "output_count=%s top_score=%s", + len(graph_results), + len(items_with_emb), + len(top_items), + top_score, + ) return top_items diff --git a/tests/reranker/test_local_rerankers.py b/tests/reranker/test_local_rerankers.py index cf6a1fb04..0bb157db0 100644 --- a/tests/reranker/test_local_rerankers.py +++ b/tests/reranker/test_local_rerankers.py @@ -1,3 +1,5 @@ +import logging + from memos.memories.textual.item import TextualMemoryItem, TreeNodeTextualMemoryMetadata from memos.reranker.cosine_local import CosineLocalReranker from memos.reranker.noop import NoopReranker @@ -77,3 +79,34 @@ def test_cosine_local_reranker_fills_missing_embeddings_with_negative_score(): assert ranked[0][0] == embedded assert ranked[1] == (missing, -1.0) + + +def test_cosine_local_reranker_logs_summary_without_candidate_payload(caplog): + private_memory = "private candidate payload" + private_embedding_value = 0.123456789 + item = _memory_item( + "00000000-0000-0000-0000-000000000001", + private_memory, + embedding=[private_embedding_value] * 128, + ) + + with caplog.at_level(logging.INFO): + CosineLocalReranker().rerank( + "private query", + [item], + top_k=1, + query_embedding=[private_embedding_value] * 128, + ) + + message = next( + record.getMessage() + for record in caplog.records + if "CosineLocalReranker rerank result" in record.getMessage() + ) + assert "input_count=1" in message + assert "embedded_count=1" in message + assert "output_count=1" in message + assert "top_score=" in message + assert private_memory not in message + assert str(private_embedding_value) not in message + assert "embedding" not in message From 1d956685c4fa697f5b563e1e9d89caf559de8fd7 Mon Sep 17 00:00:00 2001 From: Wenqiang Wei <46308778+endxxxx@users.noreply.github.com> Date: Wed, 22 Jul 2026 16:17:42 +0800 Subject: [PATCH 4/5] perf: optimize PolarDB retrieval round trips --- src/memos/graph_dbs/polardb.py | 71 +++++++++++-- .../tree_text_memory/retrieve/recall.py | 100 +++++++++++++++++- tests/graph_dbs/test_search_return_fields.py | 37 +++++++ tests/memories/textual/test_tree_retriever.py | 56 ++++++++++ 4 files changed, 252 insertions(+), 12 deletions(-) diff --git a/src/memos/graph_dbs/polardb.py b/src/memos/graph_dbs/polardb.py index 33a79aa75..bf74fbb8b 100644 --- a/src/memos/graph_dbs/polardb.py +++ b/src/memos/graph_dbs/polardb.py @@ -1,4 +1,5 @@ import json +import os import random import textwrap import threading @@ -20,6 +21,29 @@ logger = get_logger(__name__) +def _build_lightweight_return_columns(return_fields: list[str]) -> str: + columns = [] + for field in return_fields: + expression = ( + "embedding" + if field == "embedding" + else f"ag_catalog.agtype_access_operator(properties, '\"{field}\"'::agtype)" + ) + columns.append(f"{expression} AS return_{field}") + return "".join(f",\n {column}" for column in columns) + + +def _decode_lightweight_return_value(value: Any) -> Any: + if hasattr(value, "value"): + value = value.value + if isinstance(value, str): + try: + return json.loads(value) + except (json.JSONDecodeError, TypeError): + return value + return value + + def _compose_node(item: dict[str, Any]) -> tuple[str, str, dict[str, Any]]: node_id = item["id"] memory = item["memory"] @@ -101,6 +125,8 @@ def escape_sql_string(value: str) -> str: class PolarDBGraphDB(BaseGraphDB): """PolarDB-based implementation using Apache AGE graph database extension.""" + supports_lightweight_vector_search = True + @require_python_package( import_name="psycopg2", install_command="pip install psycopg2-binary", @@ -1828,6 +1854,7 @@ def search_by_embedding( filter: dict | None = None, knowledgebase_ids: list[str] | None = None, return_fields: list[str] | None = None, + light_weight_mode: bool = False, **kwargs, ) -> list[dict]: logger.info( @@ -1841,6 +1868,14 @@ def search_by_embedding( knowledgebase_ids, return_fields, ) + validated_return_fields = [ + field for field in self._validate_return_fields(return_fields) if field != "id" + ] + properties_projection = "NULL::agtype AS properties" if light_weight_mode else "properties" + lightweight_return_columns = ( + _build_lightweight_return_columns(validated_return_fields) if light_weight_mode else "" + ) + start_time = time.perf_counter() where_clauses = [] if scope: @@ -1888,10 +1923,10 @@ def search_by_embedding( set hnsw.ef_search = 100;set hnsw.iterative_scan = relaxed_order; WITH t AS ( SELECT id, - properties, + {properties_projection}, timeline, ag_catalog.agtype_access_operator(properties, '"id"'::agtype) AS old_id, - (embedding <=> %s::vector(1024)) AS scope_distance + (embedding <=> %s::vector(1024)) AS scope_distance{lightweight_return_columns} FROM "{self.db_name}_graph"."Memory" {where_clause} ORDER BY scope_distance ASC @@ -1916,7 +1951,19 @@ def search_by_embedding( else: pass - logger.info(" search_by_embedding query: %s", query) + if os.getenv("POLARDB_LOG_EMBEDDING_QUERY", "false").lower() == "true": + logger.info(" search_by_embedding query: %s", query) + else: + logger.info( + "search_by_embedding query omitted query_len=%d vector_dim=%d " + "user_name=%s top_k=%s scope=%s status=%s", + len(query), + len(vector) if vector else 0, + user_name, + top_k, + scope, + status, + ) with self._get_connection() as conn, conn.cursor() as cursor: if params: @@ -1926,8 +1973,11 @@ def search_by_embedding( results = cursor.fetchall() output = [] for row in results: - if len(row) < 5: - logger.warning(f"Row has {len(row)} columns, expected 5. Row: {row}") + expected_columns = 6 + len(validated_return_fields) if light_weight_mode else 5 + if len(row) < expected_columns: + logger.warning( + "Row has %d columns, expected at least %d", len(row), expected_columns + ) continue oldid = row[3] # old_id score = row[4] # scope @@ -1938,9 +1988,16 @@ def search_by_embedding( score_val = (score_val + 1) / 2 # align to neo4j, Normalized Cosine Score if threshold is None or score_val >= threshold: item = {"id": id_val, "score": score_val} - if return_fields: + if light_weight_mode: + for offset, field in enumerate(validated_return_fields, start=5): + item[field] = _decode_lightweight_return_value(row[offset]) + elif validated_return_fields: properties = row[1] # properties column - item.update(self._extract_fields_from_properties(properties, return_fields)) + item.update( + self._extract_fields_from_properties( + properties, validated_return_fields + ) + ) output.append(item) elapsed_time = (time.perf_counter() - start_time) * 1000.0 logger.info( diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/recall.py b/src/memos/memories/textual/tree_text_memory/retrieve/recall.py index dd90b8932..416fb7014 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/recall.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/recall.py @@ -1,4 +1,5 @@ import concurrent.futures +import os from memos.context.context import ContextThreadPoolExecutor from memos.embedders.factory import OllamaEmbedder @@ -11,6 +12,49 @@ logger = get_logger(__name__) +_POLARDB_LIGHTWEIGHT_SEARCH_ENV = "MEMOS_POLARDB_LIGHTWEIGHT_SEARCH_ENABLED" +_LIGHTWEIGHT_VECTOR_RETURN_FIELDS = ( + "memory", + "memory_type", + "user_name", + "user_id", + "session_id", + "status", + "is_fast", + "evolve_to", + "version", + "history", + "working_binding", + "type", + "key", + "confidence", + "source", + "tags", + "visibility", + "updated_at", + "info", + "internal_info", + "message_ids", + "covered_history", + "sources", + "created_at", + "usage", + "background", + "file_ids", + "event_time", + "event_location", + "event_roles", + "preference_type", + "dialog_id", + "original_text", + "preference", + "mem_cube_id", +) + + +def _env_flag_enabled(name: str) -> bool: + return os.getenv(name, "false").strip().lower() in {"1", "true", "yes", "on"} + class GraphMemoryRetriever: """ @@ -340,7 +384,21 @@ def _vector_recall( if not query_embedding: return [] + use_lightweight_search = ( + _env_flag_enabled(_POLARDB_LIGHTWEIGHT_SEARCH_ENV) + and getattr(self.graph_store, "supports_lightweight_vector_search", False) is True + ) + lightweight_return_fields = list(_LIGHTWEIGHT_VECTOR_RETURN_FIELDS) + if self.include_embedding: + lightweight_return_fields.append("embedding") + def search_single(vec, search_priority=None, search_filter=None): + search_kwargs = {} + if use_lightweight_search: + search_kwargs = { + "light_weight_mode": True, + "return_fields": lightweight_return_fields, + } return ( self.graph_store.search_by_embedding( vector=vec, @@ -351,6 +409,7 @@ def search_single(vec, search_priority=None, search_filter=None): search_filter=search_priority, filter=search_filter, user_name=user_name, + **search_kwargs, ) or [] ) @@ -394,17 +453,32 @@ def search_path_b(): return [] # merge and deduplicate, keeping highest score per ID - id_to_score = {} + id_to_hit = {} for r in all_hits: rid = r.get("id") if rid: rid = str(rid).strip("\"'") score = r.get("score", 0.0) - if rid not in id_to_score or score > id_to_score[rid]: - id_to_score[rid] = score + if rid not in id_to_hit or score > id_to_hit[rid].get("score", 0.0): + id_to_hit[rid] = {**r, "id": rid, "score": score} # Sort IDs by score (descending) to preserve ranking - sorted_ids = sorted(id_to_score.keys(), key=lambda x: id_to_score[x], reverse=True) + sorted_ids = sorted( + id_to_hit.keys(), + key=lambda result_id: id_to_hit[result_id].get("score", 0.0), + reverse=True, + ) + + if use_lightweight_search: + lightweight_hits = [id_to_hit[result_id] for result_id in sorted_ids] + if all( + hit.get("memory") is not None and hit.get("memory_type") is not None + for hit in lightweight_hits + ): + return [self._memory_item_from_vector_hit(hit) for hit in lightweight_hits] + logger.warning( + "Lightweight vector search returned incomplete fields; falling back to get_nodes" + ) node_dicts = ( self.graph_store.get_nodes( @@ -434,11 +508,27 @@ def search_path_b(): # Inject similarity score as relativity if "metadata" not in node: node["metadata"] = {} - node["metadata"]["relativity"] = id_to_score.get(rid, 0.0) + node["metadata"]["relativity"] = id_to_hit.get(rid, {}).get("score", 0.0) ordered_nodes.append(node) return [TextualMemoryItem.from_dict(n) for n in ordered_nodes] + @staticmethod + def _memory_item_from_vector_hit(hit: dict) -> TextualMemoryItem: + metadata = { + key: value + for key, value in hit.items() + if key not in {"id", "memory", "score"} and value is not None + } + metadata["relativity"] = hit.get("score", 0.0) + return TextualMemoryItem.from_dict( + { + "id": hit["id"], + "memory": hit["memory"], + "metadata": metadata, + } + ) + def _bm25_recall( self, query: str, diff --git a/tests/graph_dbs/test_search_return_fields.py b/tests/graph_dbs/test_search_return_fields.py index fc95d5a81..1ed0f6173 100644 --- a/tests/graph_dbs/test_search_return_fields.py +++ b/tests/graph_dbs/test_search_return_fields.py @@ -227,6 +227,43 @@ def fake_connection(): assert "n.memory_type = 'DreamDiary'" in query assert ids == ["node-1"] + def test_lightweight_embedding_search_projects_only_requested_fields(self, polardb_instance): + polardb_instance.db_name = "test_db" + polardb_instance.config = {"user_name": "default_user"} + polardb_instance._build_user_name_and_kb_ids_conditions_sql = MagicMock(return_value=[]) + polardb_instance._build_filter_conditions_sql = MagicMock(return_value=[]) + + cursor = MagicMock() + cursor.fetchall.return_value = [ + (1, None, None, '"node-1"', 0.2, '"hello"', '["tag-a"]', 0.8) + ] + conn = MagicMock() + conn.cursor.return_value.__enter__.return_value = cursor + + @contextmanager + def fake_connection(): + yield conn + + polardb_instance._get_connection = fake_connection + + results = polardb_instance.search_by_embedding( + vector=[0.1] * 1024, + user_name="cube-a", + return_fields=["memory", "tags"], + light_weight_mode=True, + ) + + query = cursor.execute.call_args[0][0] + assert "NULL::agtype AS properties" in query + assert "AS return_memory" in query + assert "AS return_tags" in query + assert results == [ + {"id": "node-1", "score": pytest.approx(0.6), "memory": "hello", "tags": ["tag-a"]} + ] + + def test_polardb_declares_lightweight_vector_search_support(self, polardb_instance): + assert polardb_instance.supports_lightweight_vector_search is True + class TestFieldNameValidation: """Tests for _validate_return_fields injection prevention.""" diff --git a/tests/memories/textual/test_tree_retriever.py b/tests/memories/textual/test_tree_retriever.py index 5f97c911a..bc4adee64 100644 --- a/tests/memories/textual/test_tree_retriever.py +++ b/tests/memories/textual/test_tree_retriever.py @@ -80,6 +80,62 @@ def test_vector_recall_combines_and_dedups(retriever, mock_graph_store): assert all(isinstance(r, TextualMemoryItem) for r in results) +def test_vector_recall_uses_lightweight_hits_without_get_nodes( + retriever, mock_graph_store, monkeypatch +): + monkeypatch.setenv("MEMOS_POLARDB_LIGHTWEIGHT_SEARCH_ENABLED", "true") + mock_graph_store.supports_lightweight_vector_search = True + node_id = str(uuid.uuid4()) + mock_graph_store.search_by_embedding.return_value = [ + { + "id": node_id, + "score": 0.91, + "memory": "remembered content", + "memory_type": "LongTermMemory", + "user_name": "cube-a", + "key": "topic", + "tags": ["tag-a"], + } + ] + + results = retriever._vector_recall([[0.1] * 5], "LongTermMemory", top_k=5, user_name="cube-a") + + assert len(results) == 1 + assert results[0].id == node_id + assert results[0].memory == "remembered content" + assert results[0].metadata.memory_type == "LongTermMemory" + assert results[0].metadata.user_name == "cube-a" + assert results[0].metadata.relativity == pytest.approx(0.91) + mock_graph_store.get_nodes.assert_not_called() + search_kwargs = mock_graph_store.search_by_embedding.call_args.kwargs + assert search_kwargs["light_weight_mode"] is True + assert "memory" in search_kwargs["return_fields"] + assert "memory_type" in search_kwargs["return_fields"] + assert "user_name" in search_kwargs["return_fields"] + assert "embedding" not in search_kwargs["return_fields"] + + +def test_vector_recall_falls_back_when_lightweight_hit_lacks_required_fields( + retriever, mock_graph_store, monkeypatch +): + monkeypatch.setenv("MEMOS_POLARDB_LIGHTWEIGHT_SEARCH_ENABLED", "true") + mock_graph_store.supports_lightweight_vector_search = True + node_id = str(uuid.uuid4()) + mock_graph_store.search_by_embedding.return_value = [{"id": node_id, "score": 0.8}] + mock_graph_store.get_nodes.return_value = [ + { + "id": node_id, + "memory": "fallback content", + "metadata": {"memory_type": "LongTermMemory"}, + } + ] + + results = retriever._vector_recall([[0.1] * 5], "LongTermMemory", top_k=5) + + assert results[0].memory == "fallback content" + mock_graph_store.get_nodes.assert_called_once() + + def test_retrieve_merges_graph_and_vector(retriever, mock_graph_store): parsed_goal = ParsedTaskGoal(keys=["k"], tags=["t"]) From 9841b87287d58fe67facf861572f6bdf2b6649d1 Mon Sep 17 00:00:00 2001 From: Wenqiang Wei <46308778+endxxxx@users.noreply.github.com> Date: Wed, 22 Jul 2026 16:19:43 +0800 Subject: [PATCH 5/5] perf: prune MMR candidates by memory bucket --- src/memos/api/handlers/search_handler.py | 89 ++++++++++++++++++++ tests/api/test_mmr_candidate_pruning.py | 100 +++++++++++++++++++++++ 2 files changed, 189 insertions(+) create mode 100644 tests/api/test_mmr_candidate_pruning.py diff --git a/src/memos/api/handlers/search_handler.py b/src/memos/api/handlers/search_handler.py index 2f170a8a0..363e0b253 100644 --- a/src/memos/api/handlers/search_handler.py +++ b/src/memos/api/handlers/search_handler.py @@ -32,6 +32,8 @@ _ENV_CONTEXT_RECALL = "MEMOS_DREAM_CONTEXT_RECALL" _ENV_CONTEXT_RECALL_TOP_K = "MEMOS_DREAM_CONTEXT_RECALL_TOP_K" _DEFAULT_CONTEXT_RECALL_TOP_K = 2 +_ENV_MMR_CANDIDATE_PRUNING = "MEMOS_MMR_CANDIDATE_PRUNING_ENABLED" +_MMR_CANDIDATE_MULTIPLIER = 2 def _env_enabled(name: str, default: str = "off") -> bool: @@ -352,6 +354,12 @@ def _mmr_dedup_text_memories( if not text_buckets and not pref_buckets: return results + self._prune_mmr_candidates_by_bucket( + results, + text_top_k=text_top_k, + pref_top_k=pref_top_k, + ) + # Flatten all memories with their type and scores # flat structure: (memory_type, bucket_idx, mem, score) flat: list[tuple[str, int, dict[str, Any], float]] = [] @@ -556,6 +564,87 @@ def _mmr_dedup_text_memories( return results + def _prune_mmr_candidates_by_bucket( + self, + results: dict[str, Any], + *, + text_top_k: int, + pref_top_k: int, + ) -> dict[str, Any]: + """Keep at most twice the final quota for each memory type in each bucket.""" + if not _env_enabled(_ENV_MMR_CANDIDATE_PRUNING, "off"): + return results + + total_before = 0 + total_after = 0 + bucket_counts: list[tuple[int, int]] = [] + + for result_key, target_top_k in ( + ("text_mem", text_top_k), + ("pref_mem", pref_top_k), + ): + buckets = results.get(result_key) + if not isinstance(buckets, list): + continue + + candidate_limit = max(0, int(target_top_k)) * _MMR_CANDIDATE_MULTIPLIER + for bucket in buckets: + memories = bucket.get("memories") if isinstance(bucket, dict) else None + if not isinstance(memories, list): + continue + + before_count = len(memories) + total_before += before_count + memories_by_type: dict[str, list[dict[str, Any]]] = {} + for memory in memories: + if not isinstance(memory, dict): + continue + metadata = memory.get("metadata") + memory_type = ( + metadata.get("memory_type") if isinstance(metadata, dict) else None + ) + memories_by_type.setdefault(str(memory_type or result_key), []).append(memory) + + selected: list[dict[str, Any]] = [] + if candidate_limit > 0: + for typed_memories in memories_by_type.values(): + selected.extend( + sorted( + typed_memories, + key=self._mmr_candidate_score, + reverse=True, + )[:candidate_limit] + ) + selected.sort(key=self._mmr_candidate_score, reverse=True) + + bucket["memories"] = selected + if "total_nodes" in bucket: + bucket["total_nodes"] = len(selected) + total_after += len(selected) + bucket_counts.append((before_count, len(selected))) + + self.logger.info( + "[SearchHandler] MMR candidate pruning: multiplier=%s before=%s after=%s " + "dropped=%s bucket_counts=%s", + _MMR_CANDIDATE_MULTIPLIER, + total_before, + total_after, + total_before - total_after, + bucket_counts, + ) + return results + + @staticmethod + def _mmr_candidate_score(memory: dict[str, Any]) -> float: + metadata = memory.get("metadata") + if not isinstance(metadata, dict): + return 0.0 + score = metadata.get("score", metadata.get("relativity", 0.0)) + try: + return float(score or 0.0) + except (TypeError, ValueError): + return 0.0 + @staticmethod def _is_unrelated( index: int, diff --git a/tests/api/test_mmr_candidate_pruning.py b/tests/api/test_mmr_candidate_pruning.py new file mode 100644 index 000000000..f25734e08 --- /dev/null +++ b/tests/api/test_mmr_candidate_pruning.py @@ -0,0 +1,100 @@ +from unittest.mock import Mock + +from memos.api.handlers.search_handler import SearchHandler + + +def _memory(memory_id: str, memory_type: str, score: float) -> dict: + return { + "id": memory_id, + "memory": memory_id, + "metadata": {"memory_type": memory_type, "relativity": score}, + } + + +def _bucket(*memories: dict) -> dict: + return {"cube_id": "cube-a", "memories": list(memories), "total_nodes": len(memories)} + + +def test_mmr_candidate_pruning_is_disabled_by_default(monkeypatch): + monkeypatch.delenv("MEMOS_MMR_CANDIDATE_PRUNING_ENABLED", raising=False) + handler = SearchHandler.__new__(SearchHandler) + handler.logger = Mock() + results = { + "text_mem": [ + _bucket( + _memory("text-1", "LongTermMemory", 0.9), + _memory("text-2", "LongTermMemory", 0.8), + _memory("text-3", "LongTermMemory", 0.7), + ) + ] + } + + handler._prune_mmr_candidates_by_bucket(results, text_top_k=1, pref_top_k=1) + + assert [item["id"] for item in results["text_mem"][0]["memories"]] == [ + "text-1", + "text-2", + "text-3", + ] + + +def test_mmr_candidate_pruning_limits_each_memory_type_and_bucket(monkeypatch): + monkeypatch.setenv("MEMOS_MMR_CANDIDATE_PRUNING_ENABLED", "true") + handler = SearchHandler.__new__(SearchHandler) + handler.logger = Mock() + results = { + "text_mem": [ + _bucket( + _memory("long-1", "LongTermMemory", 0.90), + _memory("long-2", "LongTermMemory", 0.70), + _memory("long-3", "LongTermMemory", 0.50), + _memory("user-1", "UserMemory", 0.85), + _memory("user-2", "UserMemory", 0.65), + _memory("user-3", "UserMemory", 0.45), + ) + ], + "pref_mem": [ + _bucket( + _memory("pref-1", "PreferenceMemory", 0.88), + _memory("pref-2", "PreferenceMemory", 0.68), + _memory("pref-3", "PreferenceMemory", 0.48), + ) + ], + } + + handler._prune_mmr_candidates_by_bucket(results, text_top_k=1, pref_top_k=1) + + assert {item["id"] for item in results["text_mem"][0]["memories"]} == { + "long-1", + "long-2", + "user-1", + "user-2", + } + assert [item["id"] for item in results["pref_mem"][0]["memories"]] == [ + "pref-1", + "pref-2", + ] + assert results["text_mem"][0]["total_nodes"] == 4 + assert results["pref_mem"][0]["total_nodes"] == 2 + + +def test_mmr_dedup_prunes_before_embedding_extraction(monkeypatch): + monkeypatch.setenv("MEMOS_MMR_CANDIDATE_PRUNING_ENABLED", "true") + handler = SearchHandler.__new__(SearchHandler) + handler.logger = Mock() + handler._extract_embeddings = Mock(return_value=[[1.0, 0.0], [0.0, 1.0]]) + results = { + "text_mem": [ + _bucket( + _memory("text-1", "LongTermMemory", 0.9), + _memory("text-2", "LongTermMemory", 0.8), + _memory("text-3", "LongTermMemory", 0.7), + ) + ], + "pref_mem": [], + } + + handler._mmr_dedup_text_memories(results, text_top_k=1, pref_top_k=1) + + embedded_memories = handler._extract_embeddings.call_args.args[0] + assert [item["id"] for item in embedded_memories] == ["text-1", "text-2"]