fix: minor graph rag improvements

This commit is contained in:
Alex committed 2026-06-23 20:10:41 +01:00
1 parent ae3e79bf8f
commit d7bbfcfe17
24 files changed
+1048 -277

No files matched your search

+4
View File
@@ -1153,10 +1153,14 @@ class EnableSourceGraphRAG(Resource):
return make_response(jsonify({"success": False}), 400)
try:
from application.worker import (
_reset_graph_for_source,
_source_updated_at,
graph_extraction_key,
)
# Drop any prior graph so each enable rebuilds from scratch rather
# than no-opping against an already-``done`` checkpoint.
_reset_graph_for_source(resolved_source_id)
task = extract_graph.delay(
resolved_source_id,
owner,
+24 -1
View File
@@ -137,8 +137,31 @@ def convert_source_to_wiki(self, source_id, user, idempotency_key=None):
return resp
def _emit_graph_poison_event(task_name, bound):
"""Publish a terminal ``graph.extract.failed`` when the poison-guard trips.
The guard returns before the worker runs, so the worker's own failed event
never fires — without this the build UI spins forever.
"""
user = bound.get("user")
source_id = bound.get("source_id")
if not user or not source_id:
return
from application.events.publisher import publish_user_event
publish_user_event(
user,
"graph.extract.failed",
{
"source_id": str(source_id),
"error": "Graph extraction stopped after repeated failures.",
},
scope={"kind": "source", "id": str(source_id)},
)
@celery.task(**DURABLE_TASK)
@with_idempotency(task_name="extract_graph")
@with_idempotency(task_name="extract_graph", on_poison=_emit_graph_poison_event)
def extract_graph(self, source_id, user, idempotency_key=None):
from application.worker import extract_graph_worker
+2 -2
View File
@@ -1,4 +1,4 @@
"""GraphRAG feature package (flag-gated, pgvector-only per D29)."""
"""GraphRAG feature package (flag-gated, pgvector-only)."""
from __future__ import annotations
@@ -6,5 +6,5 @@ from application.core.settings import settings
def graphrag_available() -> bool:
"""Return True when GraphRAG is enabled and the store is pgvector (D29)."""
"""Return True when GraphRAG is enabled and the store is pgvector."""
return settings.GRAPHRAG_ENABLED and settings.VECTOR_STORE == "pgvector"
+98 -75
View File
@@ -1,4 +1,4 @@
"""Ingest-time GraphRAG extraction pipeline (D28; pgvector-only per D29).
"""Ingest-time GraphRAG extraction pipeline (pgvector-only).
Turns a source's chunks into the per-source knowledge graph held by
``GraphStore``: each chunk is sent through a schema-constrained LLM extraction
@@ -6,7 +6,7 @@ Turns a source's chunks into the per-source knowledge graph held by
are added between resolved endpoints, and chunk links are recorded so retrieval
can join ``graph_node_chunks`` back to the retrievable chunk ids.
Cost controls (D28): gleanings off (exactly one ``.gen()`` per chunk), a hard
Cost controls: gleanings off (exactly one ``.gen()`` per chunk), a hard
chunk cap, a resumable ``graph_ingest_progress`` checkpoint (an idempotent retry
never re-bills), and concat-merge of entity descriptions (no LLM summary pass).
@@ -21,7 +21,7 @@ from __future__ import annotations
import json
import logging
import re
from typing import Any, Dict, List, Optional
from typing import Any, Callable, Dict, List, Optional
from application.core.settings import settings
from application.llm.llm_creator import LLMCreator
@@ -151,6 +151,7 @@ def extract_graph_for_source(
*,
config: SourceConfig,
request_id: Optional[str] = None,
progress_cb: Optional[Callable[[Dict[str, int]], None]] = None,
) -> Dict[str, int]:
"""Build the per-source graph from its chunks via per-chunk LLM extraction.
@@ -160,6 +161,9 @@ def extract_graph_for_source(
reported under ``skipped_over_cap``. A malformed response or an LLM error on
a single chunk marks it ``failed`` and continues — the pipeline never crashes.
Each chunk is written in a single transaction with one batched embedding
call (entity + relationship-endpoint names together).
Args:
source_id: The source whose graph is being built.
user: Owner id for token-usage attribution (``None`` skips attribution).
@@ -167,6 +171,8 @@ def extract_graph_for_source(
retrievable id (``doc_id``/``chunk_id``/``id``) and text.
config: The source's parsed ``SourceConfig`` (graph knobs).
request_id: Originating request id stamped on the extraction LLM.
progress_cb: Optional callback invoked after each processed chunk with
``{current, total, nodes, edges}`` for progress reporting.
Returns:
A summary ``{nodes, edges, chunks_processed, skipped_over_cap,
@@ -199,70 +205,47 @@ def extract_graph_for_source(
edges = 0
chunks_processed = 0
failed_chunks = 0
total = len(to_process)
def _report():
if progress_cb is None:
return
try:
progress_cb(
{
"current": chunks_processed + failed_chunks,
"total": total,
"nodes": nodes,
"edges": edges,
}
)
except Exception as exc:
logger.debug("graph progress callback failed: %s", exc)
for chunk, chunk_id in to_process:
text = _chunk_text(chunk)
if not text:
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
_report()
continue
extracted = _extract_chunk(llm, text)
if extracted is None:
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
continue
try:
entities = [
e for e in extracted["entities"]
if isinstance(e, dict) and str(e.get("name", "")).strip()
]
names = [str(e["name"]).strip() for e in entities]
name_embeddings = (
embedding.embed_documents(names) if names else []
entities = _build_entities(extracted["entities"])
relationships = _build_relationships(extracted["relationships"])
name_embeddings = _embed_names(embedding, entities, relationships)
chunk_nodes, chunk_edges = store.apply_chunk(
source_id, chunk_id, entities, relationships, name_embeddings
)
node_ids: Dict[str, str] = {}
for entity, name_embedding in zip(entities, name_embeddings):
name = str(entity["name"]).strip()
normalized_name = name.lower()
node_id = store.upsert_node(
source_id=source_id,
name=name,
normalized_name=normalized_name,
type=str(entity.get("type") or "") or None,
description=str(entity.get("description") or "") or None,
name_embedding=name_embedding,
)
node_ids[normalized_name] = node_id
nodes += 1
for node_id in node_ids.values():
store.link_node_chunk(source_id, node_id, chunk_id)
for rel in extracted["relationships"]:
if not isinstance(rel, dict):
continue
src_id = _resolve_endpoint(
store, source_id, rel.get("source"), node_ids, embedding
)
dst_id = _resolve_endpoint(
store, source_id, rel.get("target"), node_ids, embedding
)
if src_id is None or dst_id is None:
continue
store.add_edge(
source_id=source_id,
src_node_id=src_id,
dst_node_id=dst_id,
type=str(rel.get("type") or "") or None,
description=str(rel.get("description") or "") or None,
weight=_coerce_weight(rel.get("weight", 1.0)),
source_chunk_ids=[chunk_id],
)
edges += 1
nodes += chunk_nodes
edges += chunk_edges
store.mark_chunk(source_id, chunk_id, "done")
chunks_processed += 1
except Exception as exc:
@@ -273,6 +256,7 @@ def extract_graph_for_source(
)
store.mark_chunk(source_id, chunk_id, "failed")
failed_chunks += 1
_report()
try:
store.set_node_degrees(source_id)
@@ -288,28 +272,67 @@ def extract_graph_for_source(
}
def _resolve_endpoint(
store,
source_id: str,
name: Any,
node_ids: Dict[str, str],
def _build_entities(raw_entities: Any) -> List[Dict[str, Any]]:
"""Normalize the LLM's entity dicts (drop nameless ones)."""
entities = []
for e in raw_entities:
if not isinstance(e, dict):
continue
name = str(e.get("name", "")).strip()
if not name:
continue
entities.append(
{
"name": name,
"normalized_name": name.lower(),
"type": str(e.get("type") or "") or None,
"description": str(e.get("description") or "") or None,
}
)
return entities
def _build_relationships(raw_relationships: Any) -> List[Dict[str, Any]]:
"""Normalize the LLM's relationship dicts (endpoints kept as raw names)."""
relationships = []
for rel in raw_relationships:
if not isinstance(rel, dict):
continue
relationships.append(
{
"source": rel.get("source"),
"target": rel.get("target"),
"type": str(rel.get("type") or "") or None,
"description": str(rel.get("description") or "") or None,
"weight": _coerce_weight(rel.get("weight", 1.0)),
}
)
return relationships
def _embed_names(
embedding,
) -> Optional[str]:
"""Resolve a relationship endpoint to a node id, upserting if unseen this chunk."""
if name is None:
return None
clean = str(name).strip()
if not clean:
return None
normalized_name = clean.lower()
if normalized_name in node_ids:
return node_ids[normalized_name]
name_embedding = embedding.embed_documents([clean])[0]
node_id = store.upsert_node(
source_id=source_id,
name=clean,
normalized_name=normalized_name,
name_embedding=name_embedding,
)
node_ids[normalized_name] = node_id
return node_id
entities: List[Dict[str, Any]],
relationships: List[Dict[str, Any]],
) -> Dict[str, List[float]]:
"""Embed every distinct name in a chunk (entities + endpoints) in one call.
Returns a ``normalized_name -> embedding`` map. One batched ``embed_documents``
per chunk instead of a call per relationship endpoint.
"""
name_by_norm: Dict[str, str] = {}
for entity in entities:
name_by_norm.setdefault(entity["normalized_name"], entity["name"])
for rel in relationships:
for endpoint in (rel.get("source"), rel.get("target")):
if endpoint is None:
continue
clean = str(endpoint).strip()
if clean:
name_by_norm.setdefault(clean.lower(), clean)
if not name_by_norm:
return {}
norms = list(name_by_norm.keys())
vectors = embedding.embed_documents([name_by_norm[n] for n in norms])
return {norm: vector for norm, vector in zip(norms, vectors)}
+209 -68
View File
@@ -1,6 +1,6 @@
"""Per-source knowledge-graph store co-located with the pgvector ``documents`` table.
GraphRAG is pgvector-only (D29): the graph tables live in the same DB as the
GraphRAG is pgvector-only: the graph tables live in the same DB as the
pgvector store and are created on-demand (``CREATE TABLE IF NOT EXISTS`` +
``CREATE EXTENSION IF NOT EXISTS vector``), mirroring
``PGVectorStore._ensure_table_exists`` rather than going through app-DB Alembic.
@@ -214,6 +214,63 @@ class GraphStore:
finally:
cursor.close()
def _upsert_node(
self,
cursor,
source_id: str,
name: str,
normalized_name: str,
type: Optional[str] = None,
description: Optional[str] = None,
name_embedding: Optional[List[float]] = None,
) -> str:
"""Upsert a node on an open cursor (no commit). Returns the node id.
On conflict the description is concatenated (de-duped), ``doc_freq`` is
incremented, the type is refreshed if previously empty, and the
embedding is refreshed when provided.
"""
node_id = str(uuid.uuid4())
cursor.execute(
"""
INSERT INTO graph_nodes
(id, source_id, name, normalized_name, type, description,
doc_freq, name_embedding)
VALUES (%s, %s, %s, %s, %s, %s, 1, %s)
ON CONFLICT (source_id, normalized_name) DO UPDATE SET
description = CASE
WHEN EXCLUDED.description IS NULL
OR EXCLUDED.description = '' THEN graph_nodes.description
WHEN graph_nodes.description IS NULL
OR graph_nodes.description = '' THEN EXCLUDED.description
WHEN position(EXCLUDED.description IN graph_nodes.description) > 0
THEN graph_nodes.description
ELSE graph_nodes.description || ' ' || EXCLUDED.description
END,
type = CASE
WHEN graph_nodes.type IS NULL
OR graph_nodes.type = '' THEN EXCLUDED.type
ELSE graph_nodes.type
END,
name = COALESCE(graph_nodes.name, EXCLUDED.name),
doc_freq = graph_nodes.doc_freq + 1,
name_embedding = COALESCE(
EXCLUDED.name_embedding, graph_nodes.name_embedding
)
RETURNING id;
""",
(
node_id,
source_id,
name,
normalized_name,
type,
description,
name_embedding,
),
)
return str(cursor.fetchone()[0])
def upsert_node(
self,
source_id: str,
@@ -229,51 +286,15 @@ class GraphStore:
incremented, the type is refreshed if previously empty, and the
embedding is refreshed when provided. Returns the node id either way.
"""
node_id = str(uuid.uuid4())
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO graph_nodes
(id, source_id, name, normalized_name, type, description,
doc_freq, name_embedding)
VALUES (%s, %s, %s, %s, %s, %s, 1, %s)
ON CONFLICT (source_id, normalized_name) DO UPDATE SET
description = CASE
WHEN EXCLUDED.description IS NULL
OR EXCLUDED.description = '' THEN graph_nodes.description
WHEN graph_nodes.description IS NULL
OR graph_nodes.description = '' THEN EXCLUDED.description
WHEN position(EXCLUDED.description IN graph_nodes.description) > 0
THEN graph_nodes.description
ELSE graph_nodes.description || ' ' || EXCLUDED.description
END,
type = CASE
WHEN graph_nodes.type IS NULL
OR graph_nodes.type = '' THEN EXCLUDED.type
ELSE graph_nodes.type
END,
name = COALESCE(graph_nodes.name, EXCLUDED.name),
doc_freq = graph_nodes.doc_freq + 1,
name_embedding = COALESCE(
EXCLUDED.name_embedding, graph_nodes.name_embedding
)
RETURNING id;
""",
(
node_id,
source_id,
name,
normalized_name,
type,
description,
name_embedding,
),
returned_id = self._upsert_node(
cursor, source_id, name, normalized_name, type, description,
name_embedding,
)
returned_id = cursor.fetchone()[0]
conn.commit()
return str(returned_id)
return returned_id
except Exception as e:
conn.rollback()
logging.error(f"Error upserting node: {e}")
@@ -281,6 +302,43 @@ class GraphStore:
finally:
cursor.close()
def _add_edge(
self,
cursor,
source_id: str,
src_node_id: str,
dst_node_id: str,
type: Optional[str] = None,
description: Optional[str] = None,
weight: float = 1.0,
source_chunk_ids: Optional[List[str]] = None,
) -> str:
"""Insert an edge on an open cursor (no commit, no degree bump).
Callers that batch many edges run ``set_node_degrees`` once afterwards
instead of bumping degree per edge.
"""
edge_id = str(uuid.uuid4())
cursor.execute(
"""
INSERT INTO graph_edges
(id, source_id, src_node_id, dst_node_id, type, description,
weight, source_chunk_ids)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s);
""",
(
edge_id,
source_id,
src_node_id,
dst_node_id,
type,
description,
weight,
Jsonb(source_chunk_ids or []),
),
)
return edge_id
def add_edge(
self,
source_id: str,
@@ -292,27 +350,12 @@ class GraphStore:
source_chunk_ids: Optional[List[str]] = None,
) -> str:
"""Insert an edge and bump the degree of both endpoints. Returns its id."""
edge_id = str(uuid.uuid4())
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO graph_edges
(id, source_id, src_node_id, dst_node_id, type, description,
weight, source_chunk_ids)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s);
""",
(
edge_id,
source_id,
src_node_id,
dst_node_id,
type,
description,
weight,
Jsonb(source_chunk_ids or []),
),
edge_id = self._add_edge(
cursor, source_id, src_node_id, dst_node_id, type, description,
weight, source_chunk_ids,
)
cursor.execute(
"UPDATE graph_nodes SET degree = degree + 1 "
@@ -328,18 +371,22 @@ class GraphStore:
finally:
cursor.close()
def _link_node_chunk(self, cursor, source_id: str, node_id: str, chunk_id: str):
"""Link a node to a chunk on an open cursor (no commit)."""
cursor.execute(
"""
INSERT INTO graph_node_chunks (source_id, node_id, chunk_id)
VALUES (%s, %s, %s)
ON CONFLICT (source_id, node_id, chunk_id) DO NOTHING;
""",
(source_id, node_id, str(chunk_id)),
)
def link_node_chunk(self, source_id: str, node_id: str, chunk_id: str):
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO graph_node_chunks (source_id, node_id, chunk_id)
VALUES (%s, %s, %s)
ON CONFLICT (source_id, node_id, chunk_id) DO NOTHING;
""",
(source_id, node_id, str(chunk_id)),
)
self._link_node_chunk(cursor, source_id, node_id, chunk_id)
conn.commit()
except Exception as e:
conn.rollback()
@@ -348,6 +395,100 @@ class GraphStore:
finally:
cursor.close()
def apply_chunk(
self,
source_id: str,
chunk_id: str,
entities: List[Dict[str, Any]],
relationships: List[Dict[str, Any]],
name_embeddings: Dict[str, List[float]],
) -> tuple[int, int]:
"""Write one chunk's extracted entities and relationships in one transaction.
``entities`` are ``{name, normalized_name, type, description}`` dicts;
each is upserted and linked to ``chunk_id``. ``relationships`` are
``{source, target, type, description, weight}`` dicts keyed by entity
name; an endpoint not among the chunk's entities is upserted edge-only
(not linked to the chunk), mirroring the per-call path.
``name_embeddings`` maps ``normalized_name`` to its embedding. Degrees
are not bumped here — the caller runs ``set_node_degrees`` once at the
end. Returns ``(nodes_upserted, edges_added)``.
"""
conn = self._get_connection()
cursor = conn.cursor()
node_ids: Dict[str, str] = {}
edges_added = 0
try:
for entity in entities:
normalized_name = entity["normalized_name"]
node_id = self._upsert_node(
cursor,
source_id,
entity["name"],
normalized_name,
entity.get("type"),
entity.get("description"),
name_embeddings.get(normalized_name),
)
node_ids[normalized_name] = node_id
self._link_node_chunk(cursor, source_id, node_id, chunk_id)
for rel in relationships:
src_id = self._resolve_endpoint(
cursor, source_id, rel.get("source"), node_ids, name_embeddings
)
dst_id = self._resolve_endpoint(
cursor, source_id, rel.get("target"), node_ids, name_embeddings
)
if src_id is None or dst_id is None:
continue
self._add_edge(
cursor,
source_id,
src_id,
dst_id,
type=rel.get("type"),
description=rel.get("description"),
weight=float(rel.get("weight") or 1.0),
source_chunk_ids=[chunk_id],
)
edges_added += 1
conn.commit()
return len(entities), edges_added
except Exception:
conn.rollback()
raise
finally:
cursor.close()
def _resolve_endpoint(
self,
cursor,
source_id: str,
name: Any,
node_ids: Dict[str, str],
name_embeddings: Dict[str, List[float]],
) -> Optional[str]:
"""Resolve a relationship endpoint to a node id, upserting if unseen this chunk."""
if name is None:
return None
clean = str(name).strip()
if not clean:
return None
normalized_name = clean.lower()
if normalized_name in node_ids:
return node_ids[normalized_name]
node_id = self._upsert_node(
cursor,
source_id,
clean,
normalized_name,
name_embedding=name_embeddings.get(normalized_name),
)
node_ids[normalized_name] = node_id
return node_id
def get_node_by_normalized(
self, source_id: str, normalized_name: str
) -> Optional[Dict[str, Any]]:
@@ -382,7 +523,7 @@ class GraphStore:
conn.rollback()
def count_nodes(self, source_id: str) -> int:
"""Number of nodes for a source. Zero drives the ClassicRAG fallback (G5)."""
"""Number of nodes for a source. Zero drives the ClassicRAG fallback."""
conn = self._get_connection()
cursor = conn.cursor()
try:
+5 -28
View File
@@ -1,9 +1,9 @@
import logging
import os
from application.core.settings import settings
from application.llm.llm_creator import LLMCreator
from application.retriever.base import BaseRetriever
from application.retriever.labels import labels_from_metadata
from application.utils import num_tokens_from_string
from application.vectorstore.vector_creator import VectorCreator
@@ -223,38 +223,15 @@ class ClassicRAG(BaseRetriever):
page_content = doc.get("text", doc.get("page_content", ""))
metadata = doc.get("metadata", {})
title = metadata.get(
"title", metadata.get("post_title", page_content)
labels = labels_from_metadata(
metadata, page_content, vectorstore_id
)
if not isinstance(title, str):
title = str(title)
title = title.split("/")[-1]
filename = (
metadata.get("filename")
or metadata.get("file_name")
or metadata.get("source")
)
if isinstance(filename, str):
filename = os.path.basename(filename) or filename
else:
filename = title
if not filename:
filename = title
source_path = metadata.get("source") or vectorstore_id
doc_text_with_header = f"{filename}\n{page_content}"
doc_text_with_header = f"{labels['filename']}\n{page_content}"
doc_tokens = num_tokens_from_string(doc_text_with_header)
if cumulative_tokens + doc_tokens < token_budget:
all_docs.append(
{
"title": title,
"text": page_content,
"source": source_path,
"filename": filename,
}
)
all_docs.append({"text": page_content, **labels})
cumulative_tokens += doc_tokens
if cumulative_tokens >= token_budget:
+17 -35
View File
@@ -1,20 +1,19 @@
"""GraphRAG local retriever — Personalized PageRank over a per-source graph (D27/D31).
"""GraphRAG local retriever — Personalized PageRank over a per-source graph.
Local mode only (v1): rephrased query -> entity-name NN seeds -> bounded 1-2-hop
fetch -> networkx Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked
by landed PPR mass -> shared token budget. No LLM call at query time beyond the
(optional, reused) rephrase.
Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx
Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR
mass -> shared token budget. No LLM call at query time beyond the (optional,
reused) rephrase.
Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the
``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the
token-budget loop, and the per-source fallback when a source has no graph (D31).
token-budget loop, and the per-source fallback when a source has no graph.
"""
from __future__ import annotations
import logging
import math
import os
from typing import Any, Dict, List
import networkx as nx
@@ -24,6 +23,7 @@ from application.graphrag import graphrag_available
from application.graphrag.store import GraphStore
from application.retriever.base import BaseRetriever
from application.retriever.classic_rag import ClassicRAG
from application.retriever.labels import labels_from_metadata
from application.utils import num_tokens_from_string
from application.vectorstore.base import EmbeddingsSingleton
@@ -36,31 +36,6 @@ def _idf(doc_freq: Any) -> float:
return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0)
def _labels_from_metadata(
metadata: Dict[str, Any], text: str, source_id: str
) -> Dict[str, str]:
"""Derive ``title``/``source``/``filename`` from chunk metadata as ClassicRAG does."""
metadata = metadata or {}
title = metadata.get("title", metadata.get("post_title", text))
if not isinstance(title, str):
title = str(title)
title = title.split("/")[-1]
filename = (
metadata.get("filename")
or metadata.get("file_name")
or metadata.get("source")
)
if isinstance(filename, str):
filename = os.path.basename(filename) or filename
else:
filename = title
if not filename:
filename = title
source_path = metadata.get("source") or source_id
return {"title": title, "source": source_path, "filename": filename}
class GraphRAGRetriever(BaseRetriever):
"""Per-source PPR retriever; falls back to ClassicRAG when a source has no graph."""
@@ -165,7 +140,14 @@ class GraphRAGRetriever(BaseRetriever):
return []
seed_ids = [row["id"] for row in seed_rows]
seeds = {row["id"]: 1.0 - float(row.get("distance") or 0.0) for row in seed_rows}
# Clamp to >= 0: cosine distance can exceed 1 (negative similarity) for
# some embedding backends, and networkx pagerank produces garbage on
# negative personalization (and ZeroDivisionError when the weights sum
# to ~0). All-zero collapses to uniform PPR via the None guard below.
seeds = {
row["id"]: max(0.0, 1.0 - float(row.get("distance") or 0.0))
for row in seed_rows
}
subgraph = store.get_subgraph(source_id, seed_ids, hops=SUBGRAPH_HOPS)
node_scores = self._ppr_scores(subgraph, seeds)
@@ -185,7 +167,7 @@ class GraphRAGRetriever(BaseRetriever):
text = chunk.get("text") if chunk else None
if not text:
continue
labels = _labels_from_metadata(chunk.get("metadata"), text, source_id)
labels = labels_from_metadata(chunk.get("metadata"), text, source_id)
doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}")
if cumulative_tokens + doc_tokens >= token_budget:
break
@@ -194,7 +176,7 @@ class GraphRAGRetriever(BaseRetriever):
return docs
def _classic_for_source(self, source_id) -> List[Dict[str, Any]]:
"""Reuse the composed ClassicRAG to retrieve one source's chunks (D31)."""
"""Reuse the composed ClassicRAG to retrieve one source's chunks."""
original = self._classic.vectorstores
original_overrides = self._classic.per_source_retrieval
try:
+38
View File
@@ -0,0 +1,38 @@
"""Shared chunk-label derivation for retrievers."""
from __future__ import annotations
import os
from typing import Any, Dict
def labels_from_metadata(
metadata: Dict[str, Any], text: str, fallback_source: str
) -> Dict[str, str]:
"""Derive ``title``/``source``/``filename`` from a chunk's metadata.
Falls back to the chunk text for the title and to ``fallback_source`` (the
vectorstore/source id) when metadata carries no source. Used by both
ClassicRAG and GraphRAG so citation labels stay identical across retrievers.
"""
metadata = metadata or {}
title = metadata.get("title", metadata.get("post_title", text))
if not isinstance(title, str):
title = str(title)
title = title.split("/")[-1]
filename = (
metadata.get("filename")
or metadata.get("file_name")
or metadata.get("source")
)
if isinstance(filename, str):
filename = os.path.basename(filename) or filename
else:
filename = title
if not filename:
filename = title
source = metadata.get("source") or fallback_source
return {"title": title, "source": source, "filename": filename}
+2 -2
View File
@@ -50,7 +50,7 @@ class PreScreenConfig(BaseModel):
class GraphConfig(BaseModel):
"""Ingest-time GraphRAG extraction knobs (D28; pgvector-only per D29).
"""Ingest-time GraphRAG extraction knobs (pgvector-only).
``extraction_model`` None reuses the instance default model
(``LLM_PROVIDER``/``LLM_NAME``); ``max_chunks`` None falls back to the
@@ -162,7 +162,7 @@ class SourceConfig(BaseModel):
return new_config
def graph_enabled(self) -> dict:
"""Return a config dict flipped to GraphRAG mode (D28).
"""Return a config dict flipped to GraphRAG mode.
Sets ``kind="graphrag"`` so ingest paths run graph extraction and
``retrieval.retriever="graphrag"`` so the Dispatcher routes queries to
+93 -9
View File
@@ -78,13 +78,41 @@ def _source_updated_at(source) -> str:
return str(stamp) if stamp is not None else ""
def _reset_graph_for_source(source_id) -> None:
"""Drop a source's existing graph so a re-enable/re-ingest rebuilds from scratch.
Clears nodes, edges, node→chunk links and the ingest checkpoint. Run at the
enqueue site (not inside the worker) so a broker redelivery of an interrupted
build still resumes from its checkpoint rather than restarting from zero.
"""
from application.graphrag.store import GraphStore
GraphStore().delete_by_source(str(source_id))
def _publish_graph_event(user, source_id, event_type, payload) -> None:
"""Publish a graph-extraction SSE event, scoped to the source. Never raises."""
if not user:
return
try:
publish_user_event(
user,
event_type,
payload,
scope={"kind": "source", "id": str(source_id)},
)
except Exception as e:
logging.debug(f"Failed to publish graph event {event_type}: {e}")
def _maybe_enqueue_graph_extraction(cfg, source_id, user):
"""Enqueue graph extraction after embed for a graphrag source (D28).
"""Reset and re-enqueue graph extraction after embed for a graphrag source.
The graph lights up asynchronously once chunks are embedded, so ClassicRAG
works immediately. A no-op for non-graphrag sources or when GraphRAG is
unavailable. The enqueue is isolated so a broker hiccup can never fail an
otherwise-successful ingest.
unavailable. The prior graph is cleared first so a re-ingest rebuilds rather
than accumulating stale nodes. The work is isolated so a broker hiccup can
never fail an otherwise-successful ingest.
"""
if cfg.kind != "graphrag":
return
@@ -99,6 +127,7 @@ def _maybe_enqueue_graph_extraction(cfg, source_id, user):
with db_readonly() as conn:
source = SourcesRepository(conn).get_any(source_id, user)
_reset_graph_for_source(source_id)
key = graph_extraction_key(source_id, _source_updated_at(source))
extract_graph.delay(source_id, user, idempotency_key=key)
except Exception as e:
@@ -2344,13 +2373,16 @@ def convert_source_to_wiki_worker(self, source_id, user):
def extract_graph_worker(self, source_id, user):
"""Build a graphrag source's knowledge graph from its embedded chunks (D28).
"""Build a graphrag source's knowledge graph from its embedded chunks.
Loads the source, fetches its chunks from the vector store, and runs the
per-chunk LLM extraction pipeline. The chunks carry ``doc_id`` + ``text``,
matching the retrievable ids the extraction pipeline links against. No-ops
cleanly when GraphRAG is unavailable or the source has no chunks yet.
Streams ``graph.extract.progress`` SSE events as chunks are processed and a
terminal ``graph.extract.completed``/``graph.extract.failed`` on exit.
Args:
self: Celery task instance.
source_id: Source whose graph is being built.
@@ -2381,10 +2413,62 @@ def extract_graph_worker(self, source_id, user):
)
chunks = store.get_chunks() or []
return extract_graph_for_source(
source_id,
total = len(chunks)
# Throttle: at most ~20 progress events regardless of chunk count.
step = max(1, total // 20)
def _progress(info):
current = int(info.get("current", 0))
if current and current % step != 0 and current != info.get("total"):
return
_publish_graph_event(
user,
source_id,
"graph.extract.progress",
{
"source_id": source_id,
"current": current,
"total": int(info.get("total", total)),
"nodes": int(info.get("nodes", 0)),
"edges": int(info.get("edges", 0)),
},
)
_publish_graph_event(
user,
chunks,
config=cfg,
request_id=getattr(self.request, "id", None),
source_id,
"graph.extract.progress",
{
"source_id": source_id,
"current": 0,
"total": total,
"nodes": 0,
"edges": 0,
},
)
try:
summary = extract_graph_for_source(
source_id,
user,
chunks,
config=cfg,
request_id=getattr(self.request, "id", None),
progress_cb=_progress,
)
except Exception as e:
_publish_graph_event(
user,
source_id,
"graph.extract.failed",
{"source_id": source_id, "error": str(e)[:1024]},
)
raise
_publish_graph_event(
user,
source_id,
"graph.extract.completed",
{"source_id": source_id, **summary},
)
return summary
+4
View File
@@ -40,6 +40,10 @@ const KNOWN_TYPES: ReadonlySet<string> = new Set([
// TeamNotificationToast via selectRecentEvents.
'team.member_added',
'resource.shared',
// GraphRAG extraction progress (worker.py); consumed by graphBuildSlice.
'graph.extract.progress',
'graph.extract.completed',
'graph.extract.failed',
]);
/**
+2
View File
@@ -152,6 +152,7 @@
"graphrag": {
"badge": "GraphRAG",
"building": "Building graph…",
"buildingPct": "Building graph… {{pct}}%",
"enable": {
"action": "Enable GraphRAG",
"title": "Enable GraphRAG",
@@ -162,6 +163,7 @@
"estimate": "~1 LLM call per chunk · ~{{lo}}–{{hi}} total tokens (about half input, half output); output varies with how entity-dense the document is.",
"confirm": "Enable GraphRAG",
"inProgress": "Building graph…",
"inProgressPct": "Building graph… {{pct}}%",
"summaryNodes_one": "{{count}} entity",
"summaryNodes_other": "{{count}} entities",
"summaryEdges_one": "{{count}} relationship",
+94 -25
View File
@@ -1,6 +1,6 @@
import { useEffect, useRef, useState } from 'react';
import { useTranslation } from 'react-i18next';
import { useSelector } from 'react-redux';
import { useDispatch, useSelector } from 'react-redux';
import userService from '../api/services/userService';
import Spinner from '../components/Spinner';
@@ -8,7 +8,9 @@ import { Button } from '../components/ui/button';
import { Modal } from '../components/ui/modal';
import { ActiveState, Doc } from '../models/misc';
import { selectToken } from '../preferences/preferenceSlice';
import type { AppDispatch, RootState } from '../store';
import { clearGraphBuild, selectGraphBuilds } from './graphBuildSlice';
import {
GraphRAGSummary,
estimateGraphTokens,
@@ -17,6 +19,19 @@ import {
} from './graphragEnableUtils';
const POLL_INTERVAL_MS = 2000;
// Fallback poll backstop. SSE (graphBuildSlice) is the primary driver; the poll
// only covers SSE-off deployments. Bounded so a dead/stuck backend can't spin
// forever (consecutive transient errors), with a hard attempt cap on top.
const MAX_POLL_ATTEMPTS = 150;
const MAX_CONSECUTIVE_ERRORS = 5;
const ZERO_SUMMARY: GraphRAGSummary = {
nodes: 0,
edges: 0,
chunksProcessed: 0,
skippedOverCap: 0,
failedChunks: 0,
};
type Phase = 'confirm' | 'building' | 'summary' | 'error';
@@ -35,6 +50,12 @@ export default function EnableGraphRAGModal({
}: EnableGraphRAGModalProps) {
const { t } = useTranslation();
const token = useSelector(selectToken);
const dispatch = useDispatch<AppDispatch>();
const sourceId = document?.id;
const build = useSelector((state: RootState) =>
sourceId ? selectGraphBuilds(state)[sourceId] : undefined,
);
const estimate = estimateGraphTokens(
Number(document?.tokens) || 0,
@@ -46,6 +67,8 @@ export default function EnableGraphRAGModal({
const [summary, setSummary] = useState<GraphRAGSummary | null>(null);
const pollTimer = useRef<number | null>(null);
const isActiveRef = useRef(true);
// Guards against the SSE path and the poll fallback both resolving.
const resolvedRef = useRef(false);
const clearPoll = () => {
if (pollTimer.current !== null) {
@@ -62,6 +85,7 @@ export default function EnableGraphRAGModal({
useEffect(() => {
if (modalState === 'ACTIVE') {
isActiveRef.current = true;
resolvedRef.current = false;
setPhase('confirm');
setError(null);
setSummary(null);
@@ -74,56 +98,84 @@ export default function EnableGraphRAGModal({
setModalState('INACTIVE');
};
const succeed = (result: GraphRAGSummary) => {
if (resolvedRef.current) return;
resolvedRef.current = true;
clearPoll();
setSummary(result);
setPhase('summary');
};
const fail = (message?: string) => {
if (resolvedRef.current) return;
resolvedRef.current = true;
clearPoll();
setError(message ?? t('settings.sources.graphrag.enable.errors.generic'));
setPhase('error');
};
const poll = (taskId: string) => {
// Primary resolution path: the build's terminal SSE event lands in the slice.
useEffect(() => {
if (phase !== 'building' || !build) return;
if (build.status === 'completed') succeed(build.summary ?? ZERO_SUMMARY);
else if (build.status === 'failed') fail(build.error);
}, [build, phase]);
// Fallback poll: bounded so it can't loop forever on a dead/stuck backend.
const poll = (taskId: string, attempt: number, consecutiveErrors: number) => {
if (attempt > MAX_POLL_ATTEMPTS) return; // SSE may still resolve
pollTaskOnce(userService, taskId, token)
.then((result) => {
if (!isActiveRef.current) return;
if (result.status === 'pending') {
if (!isActiveRef.current || resolvedRef.current) return;
if (result.status === 'done') {
succeed(result.summary);
return;
}
if (result.status === 'failed') {
fail(result.message);
return;
}
if (result.status === 'error') {
if (consecutiveErrors + 1 >= MAX_CONSECUTIVE_ERRORS) {
fail();
return;
}
pollTimer.current = window.setTimeout(
() => poll(taskId),
() => poll(taskId, attempt + 1, consecutiveErrors + 1),
POLL_INTERVAL_MS,
);
return;
}
if (result.status === 'done') {
setSummary(result.summary);
setPhase('summary');
onEnabled();
return;
}
fail(result.message);
// pending — reset the error streak and keep polling.
pollTimer.current = window.setTimeout(
() => poll(taskId, attempt + 1, 0),
POLL_INTERVAL_MS,
);
})
.catch(() => {
if (!isActiveRef.current) return;
if (!isActiveRef.current || resolvedRef.current) return;
fail();
});
};
const handleEnable = async () => {
if (!document?.id) return;
if (!sourceId) return;
resolvedRef.current = false;
// Reset any prior build record so this build's progress events flow.
dispatch(clearGraphBuild(sourceId));
setPhase('building');
setError(null);
const start = await startGraphRAG(userService, document.id, token);
setSummary(null);
const start = await startGraphRAG(userService, sourceId, token);
if (!isActiveRef.current) return;
if (start.status === 'enabled') {
setSummary({
nodes: 0,
edges: 0,
chunksProcessed: 0,
skippedOverCap: 0,
failedChunks: 0,
});
setPhase('summary');
onEnabled();
succeed(ZERO_SUMMARY);
return;
}
if (start.status === 'task') {
onEnabled();
poll(start.taskId);
poll(start.taskId, 0, 0);
return;
}
if (start.status === 'forbidden') {
@@ -139,6 +191,11 @@ export default function EnableGraphRAGModal({
fail(start.message);
};
const progressPct =
build && build.status === 'building' && build.total > 0
? Math.min(100, Math.round((build.current / build.total) * 100))
: null;
return (
<Modal
open={modalState === 'ACTIVE'}
@@ -198,8 +255,20 @@ export default function EnableGraphRAGModal({
<div className="flex flex-col items-center gap-3 py-6">
<Spinner size="medium" />
<p className="text-muted-foreground text-sm">
{t('settings.sources.graphrag.enable.inProgress')}
{progressPct !== null
? t('settings.sources.graphrag.enable.inProgressPct', {
pct: progressPct,
})
: t('settings.sources.graphrag.enable.inProgress')}
</p>
{progressPct !== null && (
<div className="bg-muted h-1.5 w-48 overflow-hidden rounded-full">
<div
className="bg-foreground/70 h-full rounded-full transition-all"
style={{ width: `${progressPct}%` }}
/>
</div>
)}
</div>
)}
+52 -27
View File
@@ -57,6 +57,7 @@ import WikiViewer from '../components/WikiViewer';
import GraphView from '../components/GraphView';
import ConvertToWikiModal from './ConvertToWikiModal';
import EnableGraphRAGModal from './EnableGraphRAGModal';
import { clearGraphBuild, selectGraphBuilds } from './graphBuildSlice';
import SourceConfigModal from './SourceConfigModal';
type SourceMenuOption = {
@@ -133,9 +134,9 @@ export default function Sources({
const [graphRAGAvailable, setGraphRAGAvailable] = useState<boolean>(false);
const [hybridAvailable, setHybridAvailable] = useState<boolean>(false);
const [availableModels, setAvailableModels] = useState<Model[]>([]);
const [buildingGraphIds, setBuildingGraphIds] = useState<Set<string>>(
() => new Set(),
);
// Graph-build progress is SSE-driven (graphBuildSlice), so the "building"
// badge survives closing the modal and reflects the real backend state.
const graphBuilds = useSelector(selectGraphBuilds);
const [syncMenuState, setSyncMenuState] = useState<{
isOpen: boolean;
docId: string | null;
@@ -481,6 +482,19 @@ export default function Sources({
refreshDocs(undefined, 1, rowsPerPage);
}, [debouncedSearchTerm]);
// When a graph build reaches a terminal state via SSE, refresh the list so
// the badge reflects the final state, then drop the entry. The modal (a
// child) captures its summary in its own effect first, so clearing here
// doesn't race its summary view.
useEffect(() => {
const terminal = Object.entries(graphBuilds).filter(
([, b]) => b.status === 'completed' || b.status === 'failed',
);
if (terminal.length === 0) return;
terminal.forEach(([sourceId]) => dispatch(clearGraphBuild(sourceId)));
refreshDocs(undefined, currentPage, rowsPerPage);
}, [graphBuilds, dispatch, refreshDocs, currentPage, rowsPerPage]);
useEffect(() => {
let cancelled = false;
userService
@@ -772,18 +786,39 @@ export default function Sources({
{t('settings.sources.ingestProcessing')}
</span>
)}
{document.config?.kind === 'graphrag' && (
<span className="bg-muted-foreground/10 text-muted-foreground flex items-center gap-1 rounded-full px-2 py-0.5 text-xs leading-[16px] font-medium">
<Network
size={11}
strokeWidth={2}
aria-hidden="true"
/>
{document.id && buildingGraphIds.has(document.id)
? t('settings.sources.graphrag.building')
: t('settings.sources.graphrag.badge')}
</span>
)}
{document.config?.kind === 'graphrag' &&
(() => {
const build = document.id
? graphBuilds[document.id]
: undefined;
const isBuilding = build?.status === 'building';
const pct =
isBuilding && build.total > 0
? Math.min(
100,
Math.round(
(build.current / build.total) * 100,
),
)
: null;
return (
<span className="bg-muted-foreground/10 text-muted-foreground flex items-center gap-1 rounded-full px-2 py-0.5 text-xs leading-[16px] font-medium">
<Network
size={11}
strokeWidth={2}
aria-hidden="true"
/>
{isBuilding
? pct !== null
? t(
'settings.sources.graphrag.buildingPct',
{ pct },
)
: t('settings.sources.graphrag.building')
: t('settings.sources.graphrag.badge')}
</span>
);
})()}
<div className="flex items-center gap-2">
<img
src={CalendarIcon}
@@ -906,23 +941,13 @@ export default function Sources({
setModalState={(state) => {
setGraphRAGModalState(state);
if (state === 'INACTIVE') {
const closedId = documentToGraphRAG?.id;
if (closedId) {
setBuildingGraphIds((prev) => {
const next = new Set(prev);
next.delete(closedId);
return next;
});
}
setDocumentToGraphRAG(null);
}
}}
document={documentToGraphRAG}
onEnabled={() => {
const enabledId = documentToGraphRAG?.id;
if (enabledId) {
setBuildingGraphIds((prev) => new Set(prev).add(enabledId));
}
// The "building" badge is now driven by SSE progress events; just
// refresh so the source flips to graphrag kind in the list.
refreshDocs(undefined, currentPage, rowsPerPage);
}}
/>
@@ -0,0 +1,104 @@
import { describe, expect, it } from 'vitest';
import { sseEventReceived } from '../notifications/notificationsSlice';
import reducer, { clearGraphBuild } from './graphBuildSlice';
const ev = (type: string, id: string, payload: Record<string, unknown>) =>
sseEventReceived({
type,
ts: '',
user_id: 'u',
topic: 't',
scope: { kind: 'source', id },
payload,
} as never);
describe('graphBuildSlice', () => {
it('ignores non-graph events', () => {
const state = reducer(undefined, ev('source.ingest.progress', 's1', {}));
expect(state.builds).toEqual({});
});
it('tracks building progress keyed by source id', () => {
const state = reducer(
undefined,
ev('graph.extract.progress', 's1', {
current: 3,
total: 10,
nodes: 5,
edges: 2,
}),
);
expect(state.builds.s1).toEqual({
status: 'building',
current: 3,
total: 10,
nodes: 5,
edges: 2,
});
});
it('records a completed summary on the terminal event', () => {
const state = reducer(
undefined,
ev('graph.extract.completed', 's1', {
nodes: 12,
edges: 7,
chunks_processed: 30,
skipped_over_cap: 2,
failed_chunks: 1,
}),
);
expect(state.builds.s1.status).toBe('completed');
expect(state.builds.s1.summary).toEqual({
nodes: 12,
edges: 7,
chunksProcessed: 30,
skippedOverCap: 2,
failedChunks: 1,
});
});
it('records a failure with its error', () => {
const state = reducer(
undefined,
ev('graph.extract.failed', 's1', { error: 'boom' }),
);
expect(state.builds.s1).toMatchObject({ status: 'failed', error: 'boom' });
});
it('does not let a late progress event resurrect a terminal state', () => {
let state = reducer(
undefined,
ev('graph.extract.completed', 's1', { nodes: 1 }),
);
state = reducer(
state,
ev('graph.extract.progress', 's1', { current: 1, total: 2 }),
);
expect(state.builds.s1.status).toBe('completed');
});
it('clears an entry on acknowledgement', () => {
let state = reducer(
undefined,
ev('graph.extract.completed', 's1', { nodes: 1 }),
);
state = reducer(state, clearGraphBuild('s1'));
expect(state.builds.s1).toBeUndefined();
});
it('drops events without a scope id', () => {
const state = reducer(
undefined,
sseEventReceived({
type: 'graph.extract.progress',
ts: '',
user_id: 'u',
topic: 't',
payload: { current: 1, total: 2 },
} as never),
);
expect(state.builds).toEqual({});
});
});
+109
View File
@@ -0,0 +1,109 @@
import { createSlice, PayloadAction } from '@reduxjs/toolkit';
import { sseEventReceived } from '../notifications/notificationsSlice';
import { RootState } from '../store';
import { GraphRAGSummary } from './graphragEnableUtils';
export type GraphBuildStatus = 'building' | 'completed' | 'failed';
export interface GraphBuild {
status: GraphBuildStatus;
current: number;
total: number;
nodes: number;
edges: number;
summary?: GraphRAGSummary;
error?: string;
}
interface GraphBuildState {
/** Per-source graph-extraction state, driven entirely by SSE events. */
builds: Record<string, GraphBuild>;
}
const initialState: GraphBuildState = { builds: {} };
function toNumber(value: unknown): number {
const n = Number(value);
return Number.isFinite(n) ? n : 0;
}
export const graphBuildSlice = createSlice({
name: 'graphBuild',
initialState,
reducers: {
// Drop a source's build entry once the UI has acknowledged a terminal
// state (e.g. after refreshing the source list), so a stale completed/
// failed record doesn't linger.
clearGraphBuild: (state, action: PayloadAction<string>) => {
delete state.builds[action.payload];
},
},
extraReducers: (builder) => {
builder.addCase(sseEventReceived, (state, action) => {
const e = action.payload;
if (!e.type.startsWith('graph.extract.')) return;
const scopeId =
typeof e.scope?.id === 'string' && e.scope.id.length > 0
? e.scope.id
: undefined;
if (!scopeId) return;
const payload = (e.payload || {}) as Record<string, unknown>;
switch (e.type) {
case 'graph.extract.progress': {
const prev = state.builds[scopeId];
// A progress event must never resurrect a terminal state (a delayed
// or replayed frame arriving after completed/failed).
if (prev && prev.status !== 'building') break;
state.builds[scopeId] = {
status: 'building',
current: toNumber(payload.current),
total: toNumber(payload.total),
nodes: toNumber(payload.nodes),
edges: toNumber(payload.edges),
};
break;
}
case 'graph.extract.completed': {
const summary: GraphRAGSummary = {
nodes: toNumber(payload.nodes),
edges: toNumber(payload.edges),
chunksProcessed: toNumber(payload.chunks_processed),
skippedOverCap: toNumber(payload.skipped_over_cap),
failedChunks: toNumber(payload.failed_chunks),
};
state.builds[scopeId] = {
status: 'completed',
current: toNumber(payload.chunks_processed),
total: toNumber(payload.chunks_processed),
nodes: summary.nodes,
edges: summary.edges,
summary,
};
break;
}
case 'graph.extract.failed': {
state.builds[scopeId] = {
status: 'failed',
current: 0,
total: 0,
nodes: 0,
edges: 0,
error:
typeof payload.error === 'string' ? payload.error : undefined,
};
break;
}
default:
break;
}
});
},
});
export const { clearGraphBuild } = graphBuildSlice.actions;
export const selectGraphBuilds = (state: RootState) => state.graphBuild.builds;
export default graphBuildSlice.reducer;
@@ -165,9 +165,15 @@ describe('pollTaskOnce', () => {
});
});
it('returns pending on a non-ok response', async () => {
it('returns error (transient) on a non-ok response', async () => {
const getTaskStatus = vi.fn().mockResolvedValue(jsonResponse(503, {}));
const result = await pollTaskOnce({ getTaskStatus }, 't-1', 'tok');
expect(result).toEqual({ status: 'pending' });
expect(result).toEqual({ status: 'error' });
});
it('returns error (transient) on a network throw', async () => {
const getTaskStatus = vi.fn().mockRejectedValue(new Error('offline'));
const result = await pollTaskOnce({ getTaskStatus }, 't-1', 'tok');
expect(result).toEqual({ status: 'error' });
});
});
+5 -2
View File
@@ -47,6 +47,9 @@ export type GraphRAGStart =
export type GraphRAGPoll =
| { status: 'pending' }
// Transient (network throw / non-2xx). Distinct from 'pending' so a caller
// can bound consecutive errors instead of polling a dead backend forever.
| { status: 'error' }
| { status: 'done'; summary: GraphRAGSummary }
| { status: 'failed'; message?: string };
@@ -113,9 +116,9 @@ export async function pollTaskOnce(
try {
response = await service.getTaskStatus(taskId, token);
} catch {
return { status: 'pending' };
return { status: 'error' };
}
if (!response.ok) return { status: 'pending' };
if (!response.ok) return { status: 'error' };
const data = await response.json().catch(() => ({}));
return interpretTaskStatus(data?.status, data?.result);
}
+2
View File
@@ -15,6 +15,7 @@ import {
prefListenerMiddleware,
prefSlice,
} from './preferences/preferenceSlice';
import graphBuildReducer from './settings/graphBuildSlice';
import teamsReducer from './teams/teamsSlice';
import uploadReducer from './upload/uploadSlice';
@@ -78,6 +79,7 @@ const store = configureStore({
notifications: notificationsReducer,
schedules: schedulesReducer,
teams: teamsReducer,
graphBuild: graphBuildReducer,
},
middleware: (getDefaultMiddleware) =>
getDefaultMiddleware().concat(
@@ -111,6 +111,8 @@ class TestEnableSourceGraphRAG:
"application.api.user.sources.routes.graphrag_available",
return_value=True,
), patch(
"application.worker._reset_graph_for_source",
) as mock_reset, patch(
"application.api.user.sources.routes.extract_graph.delay",
return_value=fake_task,
) as mock_extract, app.test_request_context(
@@ -129,6 +131,8 @@ class TestEnableSourceGraphRAG:
assert cfg.kind == "graphrag"
assert cfg.retrieval.retriever == "graphrag"
# Each enable wipes any prior graph so it rebuilds from scratch.
mock_reset.assert_called_once_with(sid)
mock_extract.assert_called_once()
assert mock_extract.call_args.args[0] == sid
assert mock_extract.call_args.args[1] == user
+45
View File
@@ -161,6 +161,51 @@ class TestGraphStoreLive:
finally:
store.delete_by_source(source_id)
def test_apply_chunk_writes_nodes_links_and_edges(self, store, source_id):
"""One transactional write: entities linked to the chunk, edges added,
and a bare relationship endpoint upserted but not chunk-linked."""
try:
entities = [
{"name": "Ada", "normalized_name": "ada", "type": "person",
"description": "mathematician"},
{"name": "Engine", "normalized_name": "engine", "type": "machine",
"description": None},
]
relationships = [
{"source": "Ada", "target": "Engine", "type": "designed",
"description": "Ada designed the Engine", "weight": 2.0},
# 'Babbage' is only an endpoint — upserted edge-only.
{"source": "Babbage", "target": "Engine", "type": "built",
"description": None, "weight": 1.0},
]
name_embeddings = {
"ada": [0.1] * store._embedding_dim(),
"engine": [0.2] * store._embedding_dim(),
"babbage": [0.3] * store._embedding_dim(),
}
nodes, edges = store.apply_chunk(
source_id, "c1", entities, relationships, name_embeddings
)
assert nodes == 2 # only entities are counted
assert edges == 2
ada = store.get_node_by_normalized(source_id, "ada")
engine = store.get_node_by_normalized(source_id, "engine")
babbage = store.get_node_by_normalized(source_id, "babbage")
assert ada is not None and engine is not None
assert babbage is not None # endpoint upserted
mapping = store.get_chunk_ids_for_nodes(
source_id, [ada["id"], engine["id"], babbage["id"]]
)
assert mapping[ada["id"]] == ["c1"]
assert mapping[engine["id"]] == ["c1"]
# Bare endpoint is not linked to the chunk.
assert babbage["id"] not in mapping
finally:
store.delete_by_source(source_id)
def test_self_loop_degree_agrees_across_paths(self, store, source_id):
"""``add_edge``'s incremental +1 and ``set_node_degrees`` recompute must
agree on a self-loop (count it once)."""
+28
View File
@@ -142,6 +142,34 @@ class TestGraphRAGHappyPath:
assert texts.index("near") < texts.index("far")
assert docs[0].keys() == {"title", "text", "source", "filename"}
@patch("application.retriever.graph_rag.num_tokens_from_string", return_value=10)
@patch("application.retriever.graph_rag.GraphStore")
@patch("application.retriever.graph_rag.graphrag_available", return_value=True)
def test_seed_distance_over_one_is_clamped(
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed
):
# One seed at cosine distance > 1 (negative similarity) => raw weight
# 1 - 1.5 < 0. Paired with a positive seed the personalization sums to
# ~0, which makes networkx pagerank raise ZeroDivisionError. Clamping
# each weight to >= 0 keeps the personalization a valid distribution.
nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}]
edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}]
node_chunks = {"n1": ["c1"], "n2": ["c2"]}
chunk_texts = {"c1": "a", "c2": "b"}
seed_rows = [
{"id": "n1", "distance": 0.5},
{"id": "n2", "distance": 1.5},
]
store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows)
mock_store_cls.return_value = store
rag = _make_retriever(chunks=2)
# Call the PPR path directly: _get_data would swallow a raise and fall
# back to ClassicRAG, hiding the regression.
docs = rag._graph_docs_for_source(store, "src1")
assert len(docs) >= 1
@patch("application.retriever.graph_rag.num_tokens_from_string", return_value=10)
@patch("application.retriever.graph_rag.GraphStore")
@patch("application.retriever.graph_rag.graphrag_available", return_value=True)
+55
View File
@@ -115,3 +115,58 @@ class TestExtractGraphWorker:
extract.assert_called_once()
assert extract.call_args.args[2] == []
assert result["chunks_processed"] == 0
def test_publishes_completed_event(
self, pg_conn, patch_worker_db, task_self, monkeypatch
):
from application import worker
source_id = _seed_source(pg_conn)
_patch_store(monkeypatch, [{"doc_id": "c1", "text": "alpha"}])
monkeypatch.setattr(
"application.graphrag.graphrag_available", lambda: True
)
monkeypatch.setattr(
"application.graphrag.extraction.extract_graph_for_source",
MagicMock(return_value={"nodes": 1, "edges": 0, "chunks_processed": 1}),
)
events = []
monkeypatch.setattr(
worker, "publish_user_event",
lambda user, etype, payload, **kw: events.append((etype, payload)),
)
worker.extract_graph_worker(task_self, source_id, "alice")
types = [e[0] for e in events]
assert "graph.extract.progress" in types
assert types[-1] == "graph.extract.completed"
assert events[-1][1]["nodes"] == 1
def test_publishes_failed_event_on_error(
self, pg_conn, patch_worker_db, task_self, monkeypatch
):
from application import worker
source_id = _seed_source(pg_conn)
_patch_store(monkeypatch, [{"doc_id": "c1", "text": "alpha"}])
monkeypatch.setattr(
"application.graphrag.graphrag_available", lambda: True
)
def _boom(*a, **kw):
raise RuntimeError("extraction blew up")
monkeypatch.setattr(
"application.graphrag.extraction.extract_graph_for_source", _boom
)
events = []
monkeypatch.setattr(
worker, "publish_user_event",
lambda user, etype, payload, **kw: events.append((etype, payload)),
)
with pytest.raises(RuntimeError):
worker.extract_graph_worker(task_self, source_id, "alice")
assert "graph.extract.failed" in [e[0] for e in events]
+44 -1
View File
@@ -1,4 +1,4 @@
"""Ingest paths enqueue graph extraction after embed for graphrag sources (D28).
"""Ingest paths enqueue graph extraction after embed for graphrag sources.
Exercises ``remote_worker`` as the representative ingest path: the remote
loader, embedding pipeline, and ``upload_index`` are mocked, so the only thing
@@ -44,6 +44,10 @@ def _mock_remote_pipeline(monkeypatch):
monkeypatch.setattr(
worker, "upload_index", lambda full_path, file_data: None
)
# Reset constructs a real GraphStore (pgvector); not under test here.
monkeypatch.setattr(
worker, "_reset_graph_for_source", lambda *a, **kw: None
)
def _patch_delay(monkeypatch):
@@ -160,6 +164,39 @@ class TestRemoteWorkerEnqueuesGraphExtraction:
assert _run() == _run()
def test_resets_graph_before_enqueue(
self, task_self, pg_conn, patch_worker_db, monkeypatch,
_mock_remote_pipeline,
):
"""A re-ingest clears the prior graph before re-enqueuing extraction."""
from application import worker
monkeypatch.setattr(
"application.graphrag.graphrag_available", lambda: True
)
reset = MagicMock(name="reset_graph")
monkeypatch.setattr(worker, "_reset_graph_for_source", reset)
delay = _patch_delay(monkeypatch)
config = {"kind": "graphrag", "retrieval": {"retriever": "graphrag"}}
sid = _seed_source(pg_conn, "bob", config)
worker.remote_worker(
task_self,
{"urls": ["http://example.com"]},
"graph-remote",
"bob",
"crawler",
directory="temp",
retriever="classic",
operation_mode="upload",
config=config,
source_id=sid,
)
reset.assert_called_once_with(sid)
delay.assert_called_once()
def test_classic_source_does_not_enqueue(
self, task_self, pg_conn, patch_worker_db, monkeypatch,
_mock_remote_pipeline,
@@ -230,6 +267,9 @@ class TestEnqueueIsolatesBrokerFailures:
monkeypatch.setattr(
"application.graphrag.graphrag_available", lambda: True
)
monkeypatch.setattr(
worker, "_reset_graph_for_source", lambda *a, **kw: None
)
def _boom(*a, **kw):
raise RuntimeError("broker down")
@@ -255,6 +295,9 @@ class TestEnqueueIsolatesBrokerFailures:
monkeypatch.setattr(
"application.graphrag.graphrag_available", lambda: True
)
monkeypatch.setattr(
worker, "_reset_graph_for_source", lambda *a, **kw: None
)
delay = _patch_delay(monkeypatch)
def _boom(self, source_id, user_id):