mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-19 11:06:18 +00:00
* 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>
579 lines
24 KiB
Python
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"
|