fix(claude-plugin): rerank auto-injected memory context by default (#5690)

This commit is contained in:
Bartok
2026-06-23 07:22:35 -04:00
committed by GitHub
parent 565db27121
commit ced4af681f
5 changed files with 90 additions and 7 deletions
@@ -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(
@@ -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]
@@ -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 = []
@@ -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()
@@ -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