deerflow-code/offline-backend-20260512/backend/packages/harness/tests/test_rag_references.py
2026-09-07 18:24:55 +08:00

97 lines
3.5 KiB
Python

"""Tests for RAG citation batch construction."""
from __future__ import annotations
import json
from deerflow.runtime.references import (
build_latest_reference_batch,
build_latest_reference_prompt,
build_reference_batches,
)
def _tool_message(call_id: str, query: str, start: int) -> dict:
return {
"type": "tool",
"name": "web_search",
"tool_call_id": call_id,
"content": json.dumps(
{
"query": query,
"results": [
{
"title": f"{query} result {start}",
"content": f"{query} content {start}",
"url": f"https://example.test/{start}",
},
{
"title": f"{query} result {start + 1}",
"content": f"{query} content {start + 1}",
"url": f"https://example.test/{start + 1}",
},
],
},
ensure_ascii=False,
),
}
def test_latest_reference_prompt_uses_one_combined_batch_for_multi_search_turn() -> None:
messages = [
{"type": "human", "id": "human-1", "content": "search two terms"},
{
"type": "ai",
"id": "ai-search-1",
"content": "",
"tool_calls": [{"id": "call-1", "name": "web_search", "args": {"query": "Trump"}}],
},
_tool_message("call-1", "Trump", 1),
{
"type": "ai",
"id": "ai-search-2",
"content": "",
"tool_calls": [{"id": "call-2", "name": "web_search", "args": {"query": "White House"}}],
},
_tool_message("call-2", "White House", 3),
]
batch = build_latest_reference_batch(messages)
assert batch is not None
assert batch["source_count"] == 4
assert [source["index"] for source in batch["sources"]] == [1, 2, 3, 4]
assert [source["title"] for source in batch["sources"]] == [
"Trump result 1",
"Trump result 2",
"White House result 3",
"White House result 4",
]
prompt, prompt_batch = build_latest_reference_prompt(messages)
assert prompt_batch == batch
assert '<reference_batch source_count="4">' in prompt
assert "[1] Trump result 1" in prompt
assert "[4] White House result 4" in prompt
assert "重新分批" in prompt
def test_reference_batches_keep_separate_user_turns_but_merge_searches_inside_turn() -> None:
messages = [
{"type": "human", "id": "human-1", "content": "first"},
{"type": "ai", "id": "ai-1", "content": "", "tool_calls": [{"id": "c1", "name": "web_search"}]},
_tool_message("c1", "first", 1),
{"type": "ai", "id": "answer-1", "content": "answer [1]"},
{"type": "human", "id": "human-2", "content": "second"},
{"type": "ai", "id": "ai-2", "content": "", "tool_calls": [{"id": "c2", "name": "web_search"}]},
_tool_message("c2", "second-a", 3),
{"type": "ai", "id": "ai-3", "content": "", "tool_calls": [{"id": "c3", "name": "web_search"}]},
_tool_message("c3", "second-b", 5),
{"type": "ai", "id": "answer-2", "content": "answer [1][4]"},
]
batches = build_reference_batches(messages)
assert [len(batch["sources"]) for batch in batches] == [2, 4]
assert [source["index"] for source in batches[1]["sources"]] == [1, 2, 3, 4]
assert batches[0]["assistant_message_id"] == "answer-1"
assert batches[1]["assistant_message_id"] == "answer-2"