mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 14:12:58 +00:00
fix(agents): cite the pages the graph tool read
The graph tool records every page read_entity_pages returns in retrieved_docs, but agents only ever collected internal_search's. A turn answered from graph pages therefore emitted no sources: on the multi-hop e2e run, a question answered from two read_entity_pages calls came back with none. _search_tool_docs gathers both tools' documents the way the executor caches them, and both the classic/agentic collector and the research agent's per-step citations use it. Checked against the real executor and a real graph: a graph-only turn now cites the pages it read.
This commit is contained in:
1 parent
877609dfb6
commit
b32d27c912
4 files changed
+59
-17
No files matched your search
+21
-7
@@ -955,16 +955,30 @@ class BaseAgent(ABC):
|
||||
)
|
||||
self.retrieved_docs = scrubbed
|
||||
|
||||
def _collect_internal_sources(self) -> None:
|
||||
"""Merge the cached InternalSearchTool's docs into ``retrieved_docs``,
|
||||
deduped, preserving any pre-fetched docs so a mixed-exposure agent cites
|
||||
both pre-fetched and tool-retrieved sources (not just the tool's)."""
|
||||
def _search_tool_docs(self) -> List[Dict]:
|
||||
"""Documents this run's search tools read: internal search and the graph tool.
|
||||
|
||||
Both record what they surface in ``retrieved_docs``; a page read from
|
||||
the graph carries the answer as much as a search hit does, so both are
|
||||
cited. Tools are looked up the way the executor caches them.
|
||||
"""
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
from docsgpt.agents.tools.internal_search import INTERNAL_TOOL_ID
|
||||
|
||||
executor = getattr(self, "tool_executor", None)
|
||||
loaded = getattr(executor, "_loaded_tools", None) or {}
|
||||
tool = loaded.get(f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}")
|
||||
if not (tool and getattr(tool, "retrieved_docs", None)):
|
||||
docs: List[Dict] = []
|
||||
for name, tool_id in (("internal_search", INTERNAL_TOOL_ID), ("graph_search", GRAPH_TOOL_ID)):
|
||||
tool = loaded.get(f"{name}:{tool_id}:{self.user or ''}")
|
||||
docs.extend(getattr(tool, "retrieved_docs", None) or [])
|
||||
return docs
|
||||
|
||||
def _collect_internal_sources(self) -> None:
|
||||
"""Merge the search tools' docs into ``retrieved_docs``, deduped,
|
||||
preserving any pre-fetched docs so a mixed-exposure agent cites both
|
||||
pre-fetched and tool-retrieved sources (not just the tools')."""
|
||||
tool_docs = self._search_tool_docs()
|
||||
if not tool_docs:
|
||||
return
|
||||
|
||||
def _key(d):
|
||||
@@ -974,7 +988,7 @@ class BaseAgent(ABC):
|
||||
|
||||
merged = list(self.retrieved_docs or [])
|
||||
seen = {_key(d) for d in merged}
|
||||
for doc in tool.retrieved_docs:
|
||||
for doc in tool_docs:
|
||||
k = _key(doc)
|
||||
if k not in seen:
|
||||
seen.add(k)
|
||||
|
||||
@@ -7,10 +7,7 @@ from typing import Dict, Generator, List, Optional
|
||||
from docsgpt.agents.base import BaseAgent
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
from docsgpt.agents.tools.graph_search import add_graph_search_tool
|
||||
from docsgpt.agents.tools.internal_search import (
|
||||
INTERNAL_TOOL_ID,
|
||||
add_internal_search_tool,
|
||||
)
|
||||
from docsgpt.agents.tools.internal_search import add_internal_search_tool
|
||||
from docsgpt.agents.tools.wiki import add_wiki_tool
|
||||
from docsgpt.agents.tools.think import THINK_TOOL_ENTRY, THINK_TOOL_ID
|
||||
from docsgpt.logging import LogContext
|
||||
@@ -622,12 +619,9 @@ class ResearchAgent(BaseAgent):
|
||||
return messages, search_returned_empty
|
||||
|
||||
def _collect_step_sources(self):
|
||||
"""Collect sources from InternalSearchTool and register with CitationManager."""
|
||||
cache_key = f"internal_search:{INTERNAL_TOOL_ID}:{self.user or ''}"
|
||||
tool = self.tool_executor._loaded_tools.get(cache_key)
|
||||
if tool and hasattr(tool, "retrieved_docs"):
|
||||
for doc in tool.retrieved_docs:
|
||||
self.citations.add(doc)
|
||||
"""Register the search tools' docs (internal search and graph pages) with CitationManager."""
|
||||
for doc in self._search_tool_docs():
|
||||
self.citations.add(doc)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Phase 3: Synthesis
|
||||
|
||||
@@ -313,6 +313,26 @@ class TestClassicAgentSearchExposure:
|
||||
"Tool Doc",
|
||||
]
|
||||
|
||||
def test_collect_internal_sources_includes_graph_pages(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
# Pages the graph tool read carry the answer as much as search hits do,
|
||||
# so they are cited the same way.
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
|
||||
retriever_config = {"source": {"active_docs": ["b"]}}
|
||||
agent = ClassicAgent(retriever_config=retriever_config, **agent_base_params)
|
||||
search = Mock()
|
||||
search.retrieved_docs = [{"text": "Found", "title": "Search Doc", "source": "b"}]
|
||||
graph = Mock()
|
||||
graph.retrieved_docs = [{"text": "Quill is a store.", "title": "quill.md", "source": "b"}]
|
||||
user = agent.user or ""
|
||||
agent.tool_executor._loaded_tools[f"internal_search:{INTERNAL_TOOL_ID}:{user}"] = search
|
||||
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{user}"] = graph
|
||||
|
||||
agent._collect_internal_sources()
|
||||
assert [d["title"] for d in agent.retrieved_docs] == ["Search Doc", "quill.md"]
|
||||
|
||||
def test_collect_internal_sources_dedupes(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
|
||||
@@ -779,6 +779,20 @@ class TestCollectStepSources:
|
||||
|
||||
assert len(agent.citations.citations) == 2
|
||||
|
||||
def test_collects_pages_the_graph_tool_read(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
from docsgpt.agents.tools.graph_search import GRAPH_TOOL_ID
|
||||
|
||||
agent = ResearchAgent(**agent_base_params)
|
||||
graph = Mock()
|
||||
graph.retrieved_docs = [{"source": "s3", "title": "quill.md", "text": "Quill"}]
|
||||
agent.tool_executor._loaded_tools[f"graph_search:{GRAPH_TOOL_ID}:{agent.user or ''}"] = graph
|
||||
|
||||
agent._collect_step_sources()
|
||||
|
||||
assert len(agent.citations.citations) == 1
|
||||
|
||||
def test_no_tool_no_error(
|
||||
self, agent_base_params, mock_llm_creator, mock_llm_handler_creator
|
||||
):
|
||||
|
||||
Reference in new issue
Block a user