From ced4af681fefa1a88110bb79458d5451bf1655af Mon Sep 17 00:00:00 2001 From: Bartok Date: Tue, 23 Jun 2026 07:22:35 -0400 Subject: [PATCH] fix(claude-plugin): rerank auto-injected memory context by default (#5690) --- integrations/mem0-plugin/scripts/_search.py | 19 ++++++ .../mem0-plugin/scripts/file_context.py | 3 +- .../mem0-plugin/scripts/on_bash_output.sh | 7 ++- .../mem0-plugin/scripts/on_user_prompt.sh | 7 ++- integrations/mem0-plugin/tests/test_search.py | 61 +++++++++++++++++++ 5 files changed, 90 insertions(+), 7 deletions(-) diff --git a/integrations/mem0-plugin/scripts/_search.py b/integrations/mem0-plugin/scripts/_search.py index 8f32862f2..5a19cfd41 100644 --- a/integrations/mem0-plugin/scripts/_search.py +++ b/integrations/mem0-plugin/scripts/_search.py @@ -7,12 +7,31 @@ All pre-fetch hooks use this instead of duplicating urllib boilerplate. from __future__ import annotations import json +import os import urllib.request SEARCH_URL = "https://api.mem0.ai/v3/memories/search/" SEARCH_TIMEOUT = 5 +def should_rerank() -> bool: + """Whether auto-injection searches should request Platform reranking. + + The REST search endpoint does not rerank when ``rerank`` is omitted, so + auto-injected context is ordered by raw vector similarity and the single + most relevant memory can fall outside the injected top_k window. We default + reranking ON for the hook-driven injection path (the extra ~150-200ms is + well within the hook's curl budget) and let users opt out via MEM0_RERANK. + + MEM0_RERANK is read case-insensitively; ``0``, ``false``, ``no``, and + ``off`` disable reranking. Anything else (including unset) enables it. + """ + raw = os.environ.get("MEM0_RERANK") + if raw is None: + return True + return raw.strip().lower() not in ("0", "false", "no", "off", "") + + def _do_search(api_key: str, payload: dict) -> list[dict]: body = json.dumps(payload).encode() req = urllib.request.Request( diff --git a/integrations/mem0-plugin/scripts/file_context.py b/integrations/mem0-plugin/scripts/file_context.py index 8256ae6fa..a1ebf2735 100644 --- a/integrations/mem0-plugin/scripts/file_context.py +++ b/integrations/mem0-plugin/scripts/file_context.py @@ -23,7 +23,7 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from _formatting import TYPE_ICONS, format_age from _identity import resolve_api_key, resolve_user_id from _project import resolve_project_id -from _search import search_memories +from _search import search_memories, should_rerank FILE_READ_GATE_MIN_BYTES = 1500 MAX_RESULTS = 5 @@ -93,6 +93,7 @@ def search_file_context( api_key, user_id, project_id, query, top_k=MAX_RESULTS, threshold=0.3, global_search=global_search, + rerank=should_rerank(), ) results = results[:MAX_RESULTS] diff --git a/integrations/mem0-plugin/scripts/on_bash_output.sh b/integrations/mem0-plugin/scripts/on_bash_output.sh index 698729f73..dc8621fc9 100755 --- a/integrations/mem0-plugin/scripts/on_bash_output.sh +++ b/integrations/mem0-plugin/scripts/on_bash_output.sh @@ -77,15 +77,16 @@ RESULTS=$(PYTHONPATH="$SCRIPT_DIR" MEM0_SEARCH_QUERY="$ERROR_QUERY" MEM0_SEARCH_ python3 -c " import os, sys sys.path.insert(0, os.environ.get('PYTHONPATH', '.')) -from _search import search_memories, format_results_for_context +from _search import search_memories, format_results_for_context, should_rerank api_key = os.environ.get('MEM0_API_KEY', '') user_id = os.environ.get('MEM0_SEARCH_USER', 'default') project_id = os.environ.get('MEM0_PROJECT_ID', 'unknown') query = os.environ.get('MEM0_SEARCH_QUERY', '') +rerank = should_rerank() -r1 = search_memories(api_key, user_id, project_id, query, metadata_type='anti_pattern', top_k=3) -r2 = search_memories(api_key, user_id, project_id, query, metadata_type='bug_fix', top_k=3) +r1 = search_memories(api_key, user_id, project_id, query, metadata_type='anti_pattern', top_k=3, rerank=rerank) +r2 = search_memories(api_key, user_id, project_id, query, metadata_type='bug_fix', top_k=3, rerank=rerank) seen = set() combined = [] diff --git a/integrations/mem0-plugin/scripts/on_user_prompt.sh b/integrations/mem0-plugin/scripts/on_user_prompt.sh index a65af96c4..3f48fc01f 100755 --- a/integrations/mem0-plugin/scripts/on_user_prompt.sh +++ b/integrations/mem0-plugin/scripts/on_user_prompt.sh @@ -120,14 +120,15 @@ if [ -n "$HAS_RESUME" ]; then RESUME_RESULTS=$(PYTHONPATH="$SCRIPT_DIR" MEM0_SEARCH_USER="$USER_ID" python3 -c " import os, sys sys.path.insert(0, os.environ.get('PYTHONPATH', '.')) -from _search import search_memories, format_results_for_context +from _search import search_memories, format_results_for_context, should_rerank api_key = os.environ.get('MEM0_API_KEY', '') user_id = os.environ.get('MEM0_SEARCH_USER', 'default') project_id = os.environ.get('MEM0_PROJECT_ID', 'unknown') +rerank = should_rerank() -state = search_memories(api_key, user_id, project_id, 'session state current task', metadata_type='session_state', top_k=3) -decisions = search_memories(api_key, user_id, project_id, 'recent decisions and learnings', metadata_type='decision', top_k=3) +state = search_memories(api_key, user_id, project_id, 'session state current task', metadata_type='session_state', top_k=3, rerank=rerank) +decisions = search_memories(api_key, user_id, project_id, 'recent decisions and learnings', metadata_type='decision', top_k=3, rerank=rerank) all_r = state + decisions seen = set() diff --git a/integrations/mem0-plugin/tests/test_search.py b/integrations/mem0-plugin/tests/test_search.py index fd9f91de7..ccf1437d3 100644 --- a/integrations/mem0-plugin/tests/test_search.py +++ b/integrations/mem0-plugin/tests/test_search.py @@ -101,6 +101,67 @@ def test_search_memories_no_api_key_returns_empty(): assert results == [] +def test_search_memories_omits_rerank_by_default(): + """Regression for #5684: rerank must not be sent unless requested.""" + from _search import search_memories + + captured_body = {} + + def mock_urlopen(req, timeout=None): + captured_body.update(json.loads(req.data.decode())) + resp = MagicMock() + resp.read.return_value = json.dumps({"results": []}).encode() + resp.__enter__ = lambda s: s + resp.__exit__ = MagicMock(return_value=False) + return resp + + with patch("urllib.request.urlopen", side_effect=mock_urlopen): + search_memories("key", "user", "proj", "query") + + assert "rerank" not in captured_body + + +def test_search_memories_forwards_rerank_true(): + """Regression for #5684: rerank=True must reach the request body so the + REST endpoint actually reranks (it does not rerank when omitted).""" + from _search import search_memories + + captured_body = {} + + def mock_urlopen(req, timeout=None): + captured_body.update(json.loads(req.data.decode())) + resp = MagicMock() + resp.read.return_value = json.dumps({"results": []}).encode() + resp.__enter__ = lambda s: s + resp.__exit__ = MagicMock(return_value=False) + return resp + + with patch("urllib.request.urlopen", side_effect=mock_urlopen): + search_memories("key", "user", "proj", "query", rerank=True) + + assert captured_body.get("rerank") is True + + +def test_should_rerank_defaults_true(monkeypatch): + """Regression for #5684: auto-injection reranks by default.""" + from _search import should_rerank + + monkeypatch.delenv("MEM0_RERANK", raising=False) + assert should_rerank() is True + + +def test_should_rerank_opt_out_values(monkeypatch): + from _search import should_rerank + + for falsey in ("0", "false", "False", "NO", "off", ""): + monkeypatch.setenv("MEM0_RERANK", falsey) + assert should_rerank() is False, falsey + + for truthy in ("1", "true", "yes", "on"): + monkeypatch.setenv("MEM0_RERANK", truthy) + assert should_rerank() is True, truthy + + def test_format_results_for_context(): from _search import format_results_for_context