97 lines
3.5 KiB
Python
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"
|