"""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 '' 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"