mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 10:13:06 +00:00
CI has no pgvector, so every live graph test skips there and the queries this branch added ran in no CI job at all. These pin what holds without a database: the four new read queries bind every value (the entity name comes from an LLM tool call) and map their rows; empty input runs no query; a failed query returns nothing and releases its connection. The graph tool's plumbing, the sources_have_graph gate and the hybrid path's vector ranking get the same. Also drops a redundant chained comparison flagged by code scanning.
174 lines
6.1 KiB
Python
174 lines
6.1 KiB
Python
"""The graph retriever's default path, end to end through ``_graph_docs_for_source``.
|
|
|
|
The shipped defaults — seed from entities, walk the passages, blend with vector
|
|
search — are the configuration that measured best, so they are what most graph
|
|
sources run. This drives that whole path with a store that returns real values,
|
|
and checks each per-source option actually switches its stage off.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from docsgpt.retriever.graph_rag import GraphRAGRetriever
|
|
from docsgpt.storage.db.source_config import RetrievalConfig
|
|
|
|
TEXTS = {
|
|
"c-alder": "Alder streams audit events to Quill.",
|
|
"c-quill": "Quill is compacted every six hours.",
|
|
}
|
|
VECTOR_ONLY = "A passage only plain vector search found."
|
|
|
|
|
|
class _Store:
|
|
"""A two-entity chain: the question matches Alder, the answer is on Quill."""
|
|
|
|
def __init__(self):
|
|
self.calls: list[str] = []
|
|
|
|
def search_nodes_by_embedding(self, source_id, query_embedding, k=10):
|
|
return [{"id": "alder", "name": "Alder", "distance": 0.1}]
|
|
|
|
def get_subgraph(self, source_id, node_ids, hops=1):
|
|
return {
|
|
"nodes": [{"id": "alder", "doc_freq": 1}, {"id": "quill", "doc_freq": 1}],
|
|
"edges": [{"src_node_id": "alder", "dst_node_id": "quill", "weight": 1.0}],
|
|
}
|
|
|
|
def get_chunk_ids_for_nodes(self, source_id, node_ids):
|
|
return {"alder": ["c-alder"], "quill": ["c-quill"]}
|
|
|
|
def chunk_similarities(self, source_id, chunk_ids, query_embedding):
|
|
self.calls.append("chunk_similarities")
|
|
return {"c-alder": 0.9, "c-quill": 0.2}
|
|
|
|
def get_chunk_texts(self, source_id, chunk_ids):
|
|
return {
|
|
c: {"text": TEXTS[c], "metadata": {"title": c}}
|
|
for c in chunk_ids
|
|
if c in TEXTS
|
|
}
|
|
|
|
|
|
def _retriever(per_source=None):
|
|
"""A retriever without its constructor (which builds a ClassicRAG)."""
|
|
retriever = object.__new__(GraphRAGRetriever)
|
|
retriever.chunks = 3
|
|
retriever.base_chunks = None
|
|
retriever.doc_token_limit = 50000
|
|
retriever.vectorstores = ["src"]
|
|
retriever.per_source_retrieval = per_source or {}
|
|
retriever.vector_calls = 0
|
|
|
|
def _vector_ranking(source_id, query_embedding):
|
|
retriever.vector_calls += 1
|
|
return [(VECTOR_ONLY, {"title": "vector"})]
|
|
|
|
retriever._vector_ranking = _vector_ranking
|
|
return retriever
|
|
|
|
|
|
def _texts(docs):
|
|
return [doc["text"] for doc in docs]
|
|
|
|
|
|
class TestDefaultPath:
|
|
def test_walks_passages_and_blends_in_vector_hits(self):
|
|
store = _Store()
|
|
retriever = _retriever()
|
|
|
|
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
|
|
|
# The answer sits one edge away from the seed: the walk reached it.
|
|
assert TEXTS["c-quill"] in _texts(docs)
|
|
# A hit only vector search found is blended in, not lost.
|
|
assert VECTOR_ONLY in _texts(docs)
|
|
assert store.calls == ["chunk_similarities"]
|
|
assert retriever.vector_calls == 1
|
|
|
|
|
|
class TestPerSourceOptions:
|
|
def test_passage_walk_can_be_switched_off(self):
|
|
store = _Store()
|
|
retriever = _retriever(
|
|
{"src": RetrievalConfig(chunks=3, graph={"passage_nodes": False})}
|
|
)
|
|
|
|
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
|
|
|
assert "chunk_similarities" not in store.calls
|
|
assert TEXTS["c-quill"] in _texts(docs)
|
|
|
|
def test_vector_blending_can_be_switched_off(self):
|
|
store = _Store()
|
|
retriever = _retriever(
|
|
{"src": RetrievalConfig(chunks=3, graph={"blend_vector": False})}
|
|
)
|
|
|
|
docs = retriever._graph_docs_for_source(store, "src", [0.1, 0.2])
|
|
|
|
assert retriever.vector_calls == 0
|
|
assert VECTOR_ONLY not in _texts(docs)
|
|
|
|
|
|
class TestVectorRanking:
|
|
"""The vector half of the blend, keyed on chunk text since hits carry no id."""
|
|
|
|
class _VectorStore:
|
|
def __init__(self, hits=None, error=None):
|
|
self.hits = hits or []
|
|
self.error = error
|
|
self.searched = None
|
|
self.closed = False
|
|
|
|
def search(self, question, k, query_vector=None):
|
|
self.searched = (question, k, query_vector)
|
|
if self.error:
|
|
raise self.error
|
|
return self.hits
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
@staticmethod
|
|
def _real_retriever(monkeypatch, store):
|
|
from types import SimpleNamespace
|
|
|
|
retriever = object.__new__(GraphRAGRetriever)
|
|
retriever.chunks = 3
|
|
retriever._classic = SimpleNamespace(_get_rephrased_question=lambda: "where does Alder stream?")
|
|
monkeypatch.setattr(
|
|
"docsgpt.vectorstore.vector_creator.VectorCreator.create_vectorstore",
|
|
lambda *args, **kwargs: store,
|
|
)
|
|
return retriever
|
|
|
|
def test_object_and_dict_hits_become_text_and_metadata(self, monkeypatch):
|
|
from types import SimpleNamespace
|
|
|
|
store = self._VectorStore(
|
|
hits=[
|
|
SimpleNamespace(page_content="Alder streams to Quill.", metadata={"title": "alder.md"}),
|
|
{"text": "Quill is compacted every six hours.", "metadata": {"title": "quill.md"}},
|
|
{"page_content": "A passage without metadata."},
|
|
{"metadata": {"title": "no text"}},
|
|
]
|
|
)
|
|
retriever = self._real_retriever(monkeypatch, store)
|
|
|
|
ranked = retriever._vector_ranking("src", [0.1, 0.2])
|
|
|
|
assert ranked == [
|
|
("Alder streams to Quill.", {"title": "alder.md"}),
|
|
("Quill is compacted every six hours.", {"title": "quill.md"}),
|
|
("A passage without metadata.", {}),
|
|
]
|
|
# The rephrased question and the shared query vector, with room to fuse.
|
|
assert store.searched == ("where does Alder stream?", 20, [0.1, 0.2])
|
|
assert store.closed
|
|
|
|
def test_a_failed_search_ranks_nothing_and_still_closes_the_store(self, monkeypatch):
|
|
store = self._VectorStore(error=RuntimeError("pgvector down"))
|
|
retriever = self._real_retriever(monkeypatch, store)
|
|
|
|
assert retriever._vector_ranking("src", [0.1]) == []
|
|
assert store.closed
|