deer-flow/backend/tests/test_memory_relevance_retrieval.py
wd_pan cc27730348
feat(memory): add opt-in relevance-aware retrieval ranking (#5251)
* feat(memory): add opt-in relevance-aware retrieval ranking

Add a deterministic, network-free lexical relevance strategy for DeerMem
(issue #4495): memory_search ranks every fact in scope by idf-weighted
token overlap combined with confidence, with optional greedy-MMR diversity
against near-duplicate facts; prompt injection ranks facts against the
current-turn query threaded from DynamicContextMiddleware through the new
optional `query` keyword on MemoryManager.get_context/aget_context.

Defaults preserve the legacy confidence-only behavior exactly; no prompt,
storage-format, or vector/embedding-dependency changes.

Refs #4495
Signed-off-by: pwd11 <fvdsrc@163.com>

* fix(memory): bound relevance retrieval and apply review feedback

Bound tokenization and index shared stems, preserve mixed CJK tokens, warm jieba, and align missing confidence with legacy injection. Cache MMR token sets and stop selection at result or injection budgets. Document retrieval-adapter precedence and add regression coverage. Refs #4495.

Signed-off-by: pwd11 <fvdsrc@163.com>

* fix(memory): preserve backend compatibility and normalize relevance

Signed-off-by: pwd11 <fvdsrc@163.com>

* fix(memory): omit absent query hints and share injection IDF

Signed-off-by: pwd11 <fvdsrc@163.com>

* test(memory): retain timeout mock until injection worker exits

Signed-off-by: pwd11 <fvdsrc@163.com>

* docs(agents): drop root guidance compaction

Signed-off-by: pwd11 <fvdsrc@163.com>

* fix(memory): validate token prefixes and preserve upload queries

---------

Signed-off-by: pwd11 <fvdsrc@163.com>
Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
2026-09-18 11:32:04 +08:00

579 lines
24 KiB
Python

"""Tests for the optional relevance-aware retrieval strategy (issue #4495).
The strategy is opt-in via DeerMem-private config
(``retrieval_relevance_enabled``) and must never change the default
confidence-based behavior. Coverage:
- deterministic lexical relevance + confidence scoring;
- greedy MMR diversity selection;
- ``DeerMem.search`` relevance mode (including related facts without a
literal substring match);
- prompt-injection fact ordering under a query;
- the DynamicContextMiddleware -> ``_get_memory_context`` query wiring.
"""
from types import SimpleNamespace
import pytest
from langchain_core.messages import HumanMessage
from deerflow.agents.memory.backends.deermem.deer_mem import DeerMem
from deerflow.agents.memory.backends.deermem.deermem.config import DeerMemConfig
from deerflow.agents.memory.backends.deermem.deermem.core.relevance import (
build_idf,
diversify,
lexical_relevance,
rank_facts,
tokenize,
)
from deerflow.agents.middlewares.dynamic_context_middleware import DynamicContextMiddleware
def _make_fact(content: str, category: str = "context", confidence: float = 0.7) -> dict:
return {
"id": f"fact_test_{hash(content) & 0xFFFFFFFF:08x}",
"content": content,
"category": category,
"confidence": confidence,
"createdAt": "2026-07-09T00:00:00Z",
"source": "test",
}
def _deer_mem_with_facts(facts: list[dict], backend_config: dict | None = None) -> DeerMem:
"""Build a DeerMem whose updater returns the given facts (no disk I/O)."""
mgr = DeerMem(backend_config=backend_config)
mgr._updater = SimpleNamespace(get_memory_data=lambda agent_name=None, *, user_id=None: {"facts": facts})
return mgr
# ---------------------------------------------------------------------------
# Lexical relevance scoring
# ---------------------------------------------------------------------------
class TestLexicalRelevance:
def test_missing_confidence_defaults_to_zero(self):
missing = {"content": "unrelated first"}
low = _make_fact("unrelated second", confidence=0.1)
assert rank_facts([missing, low], "python")[0] is low
def test_optional_segmenter_receives_bounded_input(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import relevance
seen = []
def cut(text):
seen.append(len(text))
yield from ("token" for _ in range(10000))
monkeypatch.setattr(relevance, "_jieba_available", True)
monkeypatch.setattr(relevance, "jieba", SimpleNamespace(cut=cut), raising=False)
assert len(tokenize("word" * 10000)) == 128
assert seen == [4096]
def test_mixed_cjk_without_jieba(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import relevance
monkeypatch.setattr(relevance, "_jieba_available", False)
assert {"python", "我喜", "喜欢", "编程"} <= set(tokenize("我喜欢Python编程"))
assert {"你好", "世界"} <= set(tokenize("你好 世界"))
assert lexical_relevance("数据库升级", "Python数据库迁移") > 0
@pytest.mark.parametrize("confidence", [None, "invalid", float("nan"), float("inf")])
def test_invalid_confidence_does_not_outrank_low_confidence(self, confidence):
invalid = _make_fact("unrelated first", confidence=confidence)
low = _make_fact("unrelated second", confidence=0.1)
assert rank_facts([invalid, low], "python")[0] is low
def test_bounded_tokens(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import relevance
monkeypatch.setattr(relevance, "_jieba_available", False)
assert len(tokenize("word " * 10000)) <= 128
assert len(tokenize("数据库迁移" * 10000)) <= 128
def test_query_tokenized_once_per_ranking(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import relevance
original = relevance.tokenize
queries = []
def counted(text):
if text == "database migration":
queries.append(text)
return original(text)
monkeypatch.setattr(relevance, "tokenize", counted)
rank_facts([_make_fact(f"python fact {i}") for i in range(100)], "database migration")
assert len(queries) == 1
def test_overlapping_content_scores_higher_than_unrelated(self):
query = "database migration"
related = lexical_relevance(query, "Migrations are managed with alembic and a PostgreSQL database")
unrelated = lexical_relevance(query, "User prefers cooking Italian food on weekends")
assert related > unrelated
def test_zero_for_no_overlap(self):
assert lexical_relevance("python", "User lives in Beijing") == 0.0
def test_case_insensitive(self):
assert lexical_relevance("PYTHON", "User prefers Python") > 0.0
def test_substring_signal_without_word_boundaries(self):
"""CJK / unsegmented content: containment still contributes relevance."""
assert lexical_relevance("Python", "我喜欢Python编程") > 0.0
def test_empty_query_scores_zero(self):
assert lexical_relevance("", "anything") == 0.0
assert lexical_relevance(" ", "anything") == 0.0
class TestIdf:
def test_common_tokens_are_downweighted(self):
corpus = [
tokenize("database migration conventions"),
tokenize("database backup schedule"),
tokenize("database replica lag"),
tokenize("the database is used everywhere"),
]
idf = build_idf(corpus)
assert idf["migration"] > idf["database"]
class TestRankFacts:
def test_combines_relevance_and_confidence(self):
facts = [
_make_fact("User prefers concise answers", confidence=0.95),
_make_fact("Migrations are managed with alembic", confidence=0.5),
]
ranked = rank_facts(facts, "database migration", relevance_weight=0.7)
assert ranked[0]["content"] == "Migrations are managed with alembic"
def test_pure_confidence_when_relevance_weight_is_zero(self):
facts = [
_make_fact("Low", confidence=0.2),
_make_fact("High", confidence=0.9),
]
ranked = rank_facts(facts, "high", relevance_weight=0.0)
assert [f["content"] for f in ranked] == ["High", "Low"]
def test_does_not_mutate_input(self):
facts = [
_make_fact("Migrations are managed with alembic", confidence=0.5),
_make_fact("User prefers concise answers", confidence=0.95),
]
snapshot = [dict(f) for f in facts]
rank_facts(facts, "database migration", relevance_weight=0.7)
assert facts == snapshot
# ---------------------------------------------------------------------------
# Diversity
# ---------------------------------------------------------------------------
class TestDiversify:
def test_incremental_penalties_match_reference_mmr(self):
scored = [(0.9 - (i % 4) * 0.1, _make_fact(f"database {i % 3} fact {i % 5}")) for i in range(20)]
remaining = list(scored)
expected = []
def penalty(fact):
left = set(tokenize(fact["content"]))
return max((len(left & set(tokenize(picked["content"]))) / len(left | set(tokenize(picked["content"]))) for picked in expected), default=0.0)
while remaining:
index = max(range(len(remaining)), key=lambda i: remaining[i][0] - 0.5 * penalty(remaining[i][1]))
expected.append(remaining.pop(index)[1])
for limit in (0, 1, 5, len(scored), len(scored) + 1):
assert diversify(scored, similarity_weight=0.5, limit=limit) == expected[:limit]
def test_limit_preserves_full_prefix(self):
from deerflow.agents.memory.backends.deermem.deermem.core.relevance import order_facts_for_query
facts = [_make_fact(text) for text in ["database migrations", "database migration", "python testing", "Italian cooking"]]
full = order_facts_for_query(facts, "database", diversity_weight=0.5)
assert order_facts_for_query(facts, "database", diversity_weight=0.5, limit=2) == full[:2]
assert order_facts_for_query(facts, "database", diversity_weight=0.5, limit=0) == []
def test_tokenization_is_linear(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import relevance
calls = []
original = relevance.tokenize
def counted(text):
calls.append(text)
return original(text)
monkeypatch.setattr(relevance, "tokenize", counted)
scored = [(0.7, _make_fact(f"database fact {i}")) for i in range(30)]
diversify(scored, similarity_weight=0.5, limit=5)
assert len(calls) <= len(scored)
def test_promotes_distinct_fact_over_near_duplicate(self):
facts = [
_make_fact("Use ruff for linting"),
_make_fact("Use ruff for linting"),
_make_fact("Deploys go through GitHub Actions"),
]
ranked = rank_facts(facts, "linting", relevance_weight=0.7)
scored = [(1.0 - index * 0.1, fact) for index, fact in enumerate(ranked)]
picked = diversify(scored, similarity_weight=0.5, limit=2)
contents = [fact["content"] for fact in picked]
assert contents[0] == "Use ruff for linting"
assert "Deploys go through GitHub Actions" in contents
assert len(contents) == 2
def test_identity_when_similarity_weight_is_zero(self):
facts = [
_make_fact("Use ruff for linting"),
_make_fact("Deploys go through GitHub Actions"),
]
ranked = rank_facts(facts, "linting", relevance_weight=0.7)
scored = [(1.0 - index * 0.1, fact) for index, fact in enumerate(ranked)]
picked = diversify(scored, similarity_weight=0.0)
assert [f["content"] for f in picked] == [f["content"] for f in ranked]
# ---------------------------------------------------------------------------
# Config defaults
# ---------------------------------------------------------------------------
class TestRelevanceConfig:
def test_defaults_keep_legacy_behavior(self):
config = DeerMemConfig()
assert config.retrieval_relevance_enabled is False
assert config.retrieval_relevance_weight == 0.5
assert config.retrieval_diversity_weight == 0.0
def test_backend_config_accepts_new_knobs(self):
config = DeerMemConfig.from_backend_config(
{
"retrieval_relevance_enabled": True,
"retrieval_relevance_weight": 0.8,
"retrieval_diversity_weight": 0.4,
}
)
assert config.retrieval_relevance_enabled is True
assert config.retrieval_relevance_weight == 0.8
assert config.retrieval_diversity_weight == 0.4
# ---------------------------------------------------------------------------
# DeerMem.search with relevance mode
# ---------------------------------------------------------------------------
class TestRelevanceSearch:
def test_search_passes_top_k_to_mmr(self, monkeypatch):
from deerflow.agents.memory.backends.deermem import deer_mem
facts = [_make_fact(f"database fact {i}") for i in range(30)]
original = deer_mem.order_facts_for_query
limits = []
def ranked(*args, **kwargs):
limits.append(kwargs.get("limit"))
return original(*args, **kwargs)
monkeypatch.setattr(deer_mem, "order_facts_for_query", ranked)
mgr = _deer_mem_with_facts(facts, {"retrieval_relevance_enabled": True, "retrieval_diversity_weight": 0.5, "retrieval_adapter": ""})
assert len(mgr.search("database", top_k=3)) == 3
assert limits == [3]
@pytest.mark.parametrize("enabled", [False, True])
@pytest.mark.parametrize("counting", ["char", "tiktoken"])
def test_warms_segmenter_only_when_enabled(self, monkeypatch, enabled, counting):
from deerflow.agents.memory.backends.deermem import deer_mem
calls = []
monkeypatch.setattr(deer_mem, "warm_tokenizer", lambda: calls.append("jieba"))
monkeypatch.setattr(deer_mem, "warm_tiktoken_cache", lambda: calls.append("tiktoken") or True)
mgr = _deer_mem_with_facts([], {"retrieval_relevance_enabled": enabled, "token_counting": counting, "retrieval_adapter": ""})
assert mgr.warm() is True
assert calls == (["jieba"] if enabled else []) + (["tiktoken"] if counting == "tiktoken" else [])
def test_returns_related_fact_without_literal_substring(self):
facts = [
_make_fact("Database migrations are handled with alembic", "project", 0.4),
_make_fact("User prefers concise answers", "preference", 0.9),
]
mgr = _deer_mem_with_facts(
facts,
backend_config={
"retrieval_relevance_enabled": True,
"retrieval_adapter": "",
"retrieval_relevance_weight": 0.7,
},
)
results = mgr.search("how do I add a database migration", top_k=5)
assert results[0]["content"] == "Database migrations are handled with alembic"
assert len(results) == 2 # every fact in scope competes, not only substring matches
def test_relevance_outweighs_confidence(self):
facts = [
_make_fact("User prefers concise answers", "preference", 0.9),
_make_fact("Migrations are managed with alembic", "project", 0.4),
]
mgr = _deer_mem_with_facts(
facts,
backend_config={
"retrieval_relevance_enabled": True,
"retrieval_adapter": "",
"retrieval_relevance_weight": 0.7,
},
)
results = mgr.search("database migration", top_k=5)
assert results[0]["content"] == "Migrations are managed with alembic"
def test_respects_category_filter_and_top_k(self):
facts = [_make_fact(f"Database fact {index}", "project", 0.5) for index in range(6)] + [_make_fact("Unrelated preference", "preference", 0.9)]
mgr = _deer_mem_with_facts(
facts,
backend_config={"retrieval_relevance_enabled": True, "retrieval_adapter": ""},
)
results = mgr.search("database", top_k=3, category="project")
assert len(results) == 3
assert all(fact["category"] == "project" for fact in results)
def test_diversity_dedups_near_duplicates(self):
facts = [
_make_fact("Use ruff for linting", confidence=0.9),
_make_fact("Use ruff for linting", confidence=0.8),
_make_fact("CI lints on every pull request", confidence=0.7),
]
mgr = _deer_mem_with_facts(
facts,
backend_config={
"retrieval_relevance_enabled": True,
"retrieval_adapter": "",
"retrieval_diversity_weight": 0.5,
},
)
# Both "lints" and "linting" extend this complete query token, so
# the test isolates diversity rather than arbitrary shared stems.
results = mgr.search("lint", top_k=2)
assert len(results) == 2
assert "CI lints on every pull request" in [fact["content"] for fact in results]
def test_legacy_behavior_unchanged_when_disabled(self):
facts = [
_make_fact("Fact A", confidence=0.3),
_make_fact("Fact B", confidence=0.9),
]
mgr = _deer_mem_with_facts(facts) # default config
results = mgr.search("Fact", top_k=5)
assert [fact["confidence"] for fact in results] == [0.9, 0.3]
def test_legacy_empty_result_without_substring_match_when_disabled(self):
facts = [_make_fact("The project uses PostgreSQL for persistence")]
mgr = _deer_mem_with_facts(facts)
assert mgr.search("database migration", top_k=5) == []
# ---------------------------------------------------------------------------
# Prompt injection with query-aware ranking
# ---------------------------------------------------------------------------
class TestInjectionRelevance:
def test_diversification_stops_at_budget_and_preserves_guaranteed_pool(self, monkeypatch):
from deerflow.agents.memory.backends.deermem.deermem.core import prompt
original = prompt.iter_diversify
picked = []
def counted(*args, **kwargs):
for fact in original(*args, **kwargs):
picked.append(fact)
yield fact
monkeypatch.setattr(prompt, "iter_diversify", counted)
facts = [_make_fact(f"database fact {i}", confidence=0.9) for i in range(100)]
facts.append(_make_fact("Always ask before deleting files", category="correction", confidence=0.1))
result = prompt.format_memory_for_injection(
{"facts": facts},
query="database",
relevance_weight=0.7,
diversity_weight=0.5,
**self._injection_args(max_tokens=40, guaranteed_categories=["correction"], guaranteed_token_budget=20),
)
assert "Always ask before deleting files" in result
assert "database fact" in result
assert len(picked) < 10
def _injection_args(self, **overrides):
args = {
"max_tokens": 300,
"use_tiktoken": False,
"guaranteed_categories": None,
"guaranteed_token_budget": 500,
}
args.update(overrides)
return args
def test_relevance_reranks_facts_under_token_budget(self):
from deerflow.agents.memory.backends.deermem.deermem.core.prompt import (
format_memory_for_injection,
)
facts = [
_make_fact("User prefers concise answers", "preference", 0.95),
_make_fact("Migrations are managed with alembic", "project", 0.4),
]
memory_data = {"facts": facts}
legacy = format_memory_for_injection(
memory_data,
**self._injection_args(max_tokens=20),
)
relevance = format_memory_for_injection(
memory_data,
query="how do I add a database migration",
relevance_weight=0.7,
**self._injection_args(max_tokens=20),
)
assert "concise answers" in legacy
assert "alembic" in relevance
assert "alembic" not in legacy
def test_query_none_preserves_legacy_order(self):
from deerflow.agents.memory.backends.deermem.deermem.core.prompt import (
format_memory_for_injection,
)
facts = [
_make_fact("User prefers concise answers", "preference", 0.95),
_make_fact("Migrations are managed with alembic", "project", 0.4),
_make_fact("User lives in Beijing", "personal", 0.8),
]
memory_data = {"facts": facts}
legacy = format_memory_for_injection(memory_data, **self._injection_args())
with_query_none = format_memory_for_injection(memory_data, query=None, relevance_weight=0.7, **self._injection_args())
assert legacy == with_query_none
class TestGetContextQuery:
def test_get_context_uses_query_when_enabled(self):
facts = [
_make_fact("User prefers concise answers", "preference", 0.95),
_make_fact("Migrations are managed with alembic", "project", 0.4),
]
mgr = _deer_mem_with_facts(
facts,
backend_config={"retrieval_relevance_enabled": True, "retrieval_relevance_weight": 0.7},
)
body = mgr.get_context("user-1", agent_name="assistant", query="how do I add a database migration")
assert "alembic" in body
def test_get_context_without_query_keeps_confidence_order(self):
facts = [
_make_fact("User prefers concise answers", "preference", 0.95),
_make_fact("Migrations are managed with alembic", "project", 0.4),
]
enabled = _deer_mem_with_facts(
facts,
backend_config={"retrieval_relevance_enabled": True},
)
disabled = _deer_mem_with_facts(facts)
assert enabled.get_context("user-1", agent_name="assistant") == disabled.get_context("user-1", agent_name="assistant")
# ---------------------------------------------------------------------------
# Middleware wiring
# ---------------------------------------------------------------------------
class TestMiddlewareQueryWiring:
@pytest.mark.parametrize("multimodal", [False, True])
@pytest.mark.parametrize("user_text", ["Use my PostgreSQL preferences to analyze these reports.", "", "PostgreSQL " * 200], ids=["request", "attachment_only", "bounded_request"])
def test_upload_context_does_not_replace_original_query(self, monkeypatch, tmp_path, multimodal, user_text):
from unittest import mock
from deerflow.agents.middlewares.uploads_middleware import UploadsMiddleware
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
uploads = UploadsMiddleware(base_dir=str(tmp_path))
files = [{"filename": f"report-{i}.csv", "size": 1024, "path": f"/mnt/user-data/uploads/report-{i}.csv", "extension": ".csv"} for i in range(5)]
monkeypatch.setattr(uploads, "_files_from_kwargs", lambda *_: files)
content = [{"type": "text", "text": user_text}] if multimodal else user_text
runtime = SimpleNamespace(context={})
update = uploads.before_agent({"messages": [HumanMessage(content=content, id="msg-1")]}, runtime)
uploaded_message = update["messages"][0]
assert uploaded_message.additional_kwargs[ORIGINAL_USER_CONTENT_KEY] == user_text
with mock.patch("deerflow.agents.lead_agent.prompt._get_memory_context", return_value="") as get_context:
DynamicContextMiddleware().before_agent({"messages": [uploaded_message]}, runtime)
get_context.assert_called_once()
assert get_context.call_args.kwargs["query"] == (user_text.strip()[:1000] or None)
def test_invalid_original_content_metadata_uses_message_text(self):
from deerflow.agents.middlewares.dynamic_context_middleware import _derive_injection_query
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
message = HumanMessage(content="database migration", additional_kwargs={ORIGINAL_USER_CONTENT_KEY: ["not a string"]})
assert _derive_injection_query(message) == "database migration"
def test_first_turn_passes_current_query_to_memory_context(self):
from unittest import mock
mw = DynamicContextMiddleware()
state = {
"messages": [
HumanMessage(content="how do I add a database migration", id="msg-1"),
]
}
with (
mock.patch(
"deerflow.agents.lead_agent.prompt._get_memory_context",
return_value="",
) as get_context,
mock.patch("deerflow.agents.middlewares.dynamic_context_middleware.datetime") as mock_dt,
):
mock_dt.now.return_value.strftime.return_value = "2026-05-08, Friday"
mw.before_agent(state, SimpleNamespace(context={}))
get_context.assert_called_once()
assert get_context.call_args.kwargs.get("query") == "how do I add a database migration"
def test_multimodal_content_yields_text_query(self):
from unittest import mock
mw = DynamicContextMiddleware()
state = {
"messages": [
HumanMessage(
content=[
{"type": "text", "text": "how do I "},
{"type": "text", "text": "add a database migration"},
],
id="msg-1",
),
]
}
with (
mock.patch(
"deerflow.agents.lead_agent.prompt._get_memory_context",
return_value="",
) as get_context,
mock.patch("deerflow.agents.middlewares.dynamic_context_middleware.datetime") as mock_dt,
):
mock_dt.now.return_value.strftime.return_value = "2026-05-08, Friday"
mw.before_agent(state, SimpleNamespace(context={}))
assert get_context.call_args.kwargs.get("query") == "how do I add a database migration"