Compare commits

Author SHA1 Message Date
Dark-Alex-17 b91f738209 docs: updated the configuratino examples for graph-based RAG
CI / All (ubuntu-latest) (push) Failing after 24s
CI / All (macos-latest) (push) Has been cancelled
CI / All (windows-latest) (push) Has been cancelled
2026-07-13 16:55:18 -06:00
Dark-Alex-17 4f0dae9b49 feat: fully functional graph-based RAG
CI / All (ubuntu-latest) (push) Failing after 26s
CI / All (macos-latest) (push) Has been cancelled
CI / All (windows-latest) (push) Has been cancelled
2026-07-13 16:50:07 -06:00
Dark-Alex-17 deb673ebc9 fmt: applied some formatting changes 2026-07-13 16:07:19 -06:00
4 changed files with 556 additions and 64 deletions
+1 -1
View File
@@ -199,7 +199,7 @@ rag_chunk_size: null # Defines the size of chunks for document proce
rag_chunk_overlap: null # Defines the overlap between chunks rag_chunk_overlap: null # Defines the overlap between chunks
rag_extractor_model: null # LLM model for graph-based entity/relationship extraction; when set, enables a graph RAG signal alongside vector and BM25 rag_extractor_model: null # LLM model for graph-based entity/relationship extraction; when set, enables a graph RAG signal alongside vector and BM25
rag_extractor_prompt: null # Custom extraction prompt template; must contain __CHUNK__ placeholder; defaults to built-in prompt when null rag_extractor_prompt: null # Custom extraction prompt template; must contain __CHUNK__ placeholder; defaults to built-in prompt when null
rag_graph_hops: 1 # Number of hops to expand from matched entities at query time (1 = direct neighbors; increase for denser graphs) rag_graph_hops: 1 # Number of hops to expand from matched entities at query time (0 = seed nodes only; 1 = direct neighbors; increase for denser graphs)
# Defines the query structure using variables like __CONTEXT__, __SOURCES__, and __INPUT__ to tailor searches to specific needs # Defines the query structure using variables like __CONTEXT__, __SOURCES__, and __INPUT__ to tailor searches to specific needs
rag_template: | rag_template: |
Answer the query based on the context while respecting the rules. (user query, some textual context and rules, all inside xml tags) Answer the query based on the context while respecting the rules. (user query, some textual context and rules, all inside xml tags)
+1 -1
View File
@@ -227,7 +227,7 @@ nodes:
reranker_model: null # Optional reranker for hybrid-search results reranker_model: null # Optional reranker for hybrid-search results
extractor_model: null # Optional chat model for graph-based entity/relationship extraction; enables graph RAG signal when set extractor_model: null # Optional chat model for graph-based entity/relationship extraction; enables graph RAG signal when set
extractor_prompt: null # Optional custom extraction prompt; must contain __CHUNK__ placeholder; uses built-in prompt when null extractor_prompt: null # Optional custom extraction prompt; must contain __CHUNK__ placeholder; uses built-in prompt when null
graph_hops: 1 # Graph expansion depth at query time (1 = direct neighbors; increase for denser knowledge graphs) graph_hops: 1 # Graph expansion depth at query time (0 = seed nodes only; 1 = direct neighbors; increase for denser knowledge graphs)
batch_size: 100 # Optional embedding-request batch size batch_size: 100 # Optional embedding-request batch size
state_updates: # {{output}} = { context: <str>, sources: [<path>, ...] } state_updates: # {{output}} = { context: <str>, sources: [<path>, ...] }
context: "{{output.context}}" # writes `context` -> `reducers.context = concat` context: "{{output.context}}" # writes `context` -> `reducers.context = concat`
+503 -21
View File
@@ -2,12 +2,21 @@ use super::DocumentId;
use crate::client::*; use crate::client::*;
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use indexmap::IndexMap; use indexmap::{IndexMap, IndexSet};
use petgraph::Direction; use petgraph::Direction;
use petgraph::graph::NodeIndex; use petgraph::graph::NodeIndex;
use petgraph::stable_graph::StableGraph; use petgraph::stable_graph::StableGraph;
use petgraph::visit::EdgeRef;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::HashSet; use std::collections::{HashMap, HashSet};
/// Heuristic upper bound on chunk size before warning the user that the
/// extraction LLM call may be truncated. Not a hard limit.
const MAX_CHUNK_CHARS: usize = 24_000;
/// Maximum number of nodes the BFS may visit during a single graph_search.
/// Keeps the synchronous traversal bounded on dense graphs.
pub const MAX_GRAPH_NODES: usize = 500;
const EXTRACTION_PROMPT: &str = r#"Extract entities and relationships from the following text chunk. const EXTRACTION_PROMPT: &str = r#"Extract entities and relationships from the following text chunk.
@@ -89,16 +98,27 @@ impl Default for KnowledgeGraph {
impl KnowledgeGraph { impl KnowledgeGraph {
pub fn merge(&mut self, doc_id: DocumentId, result: ExtractionResult) { pub fn merge(&mut self, doc_id: DocumentId, result: ExtractionResult) {
let mut chunk_nodes: Vec<u32> = vec![]; let mut chunk_nodes: IndexSet<u32> = IndexSet::new();
for extracted in &result.entities { for extracted in &result.entities {
let key = extracted.name.to_lowercase(); let key = extracted.name.to_lowercase();
let normalized_type = extracted.entity_type.to_uppercase();
let node_raw = if let Some(&existing) = self.entity_index.get(&key) { let node_raw = if let Some(&existing) = self.entity_index.get(&key) {
let idx = NodeIndex::new(existing as usize);
if self.graph.contains_node(idx) {
let node = &mut self.graph[idx];
if node.entity_type == "OTHER" && normalized_type != "OTHER" {
node.entity_type = normalized_type;
}
if node.description.is_none() {
node.description = extracted.description.clone();
}
}
existing existing
} else { } else {
let entity = Entity { let entity = Entity {
name: extracted.name.clone(), name: extracted.name.clone(),
entity_type: extracted.entity_type.clone(), entity_type: normalized_type,
description: extracted.description.clone(), description: extracted.description.clone(),
}; };
let idx = self.graph.add_node(entity); let idx = self.graph.add_node(entity);
@@ -106,7 +126,7 @@ impl KnowledgeGraph {
self.entity_index.insert(key, raw); self.entity_index.insert(key, raw);
raw raw
}; };
chunk_nodes.push(node_raw); chunk_nodes.insert(node_raw);
} }
for extracted in &result.relationships { for extracted in &result.relationships {
@@ -118,11 +138,14 @@ impl KnowledgeGraph {
) { ) {
let from_idx = NodeIndex::new(from_raw as usize); let from_idx = NodeIndex::new(from_raw as usize);
let to_idx = NodeIndex::new(to_raw as usize); let to_idx = NodeIndex::new(to_raw as usize);
// Avoid duplicate edges let already_exists = self
if !self.graph.contains_edge(from_idx, to_idx) { .graph
.edges_connecting(from_idx, to_idx)
.any(|e| e.weight().relation_type == extracted.relation_type);
if !already_exists {
let rel = Relationship { let rel = Relationship {
relation_type: extracted.relation_type.clone(), relation_type: extracted.relation_type.clone(),
weight: extracted.weight.unwrap_or(1.0), weight: extracted.weight.unwrap_or(1.0).clamp(0.0, 1.0),
}; };
self.graph.add_edge(from_idx, to_idx, rel); self.graph.add_edge(from_idx, to_idx, rel);
} }
@@ -158,6 +181,10 @@ impl KnowledgeGraph {
.filter(|raw| !still_used.contains(raw)) .filter(|raw| !still_used.contains(raw))
.collect(); .collect();
if to_remove.is_empty() {
return;
}
for raw in to_remove { for raw in to_remove {
let idx = NodeIndex::new(raw as usize); let idx = NodeIndex::new(raw as usize);
if self.graph.contains_node(idx) { if self.graph.contains_node(idx) {
@@ -166,6 +193,57 @@ impl KnowledgeGraph {
self.entity_index.swap_remove(&name); self.entity_index.swap_remove(&name);
} }
} }
self.compact();
}
/// Rebuild the internal graph with consecutive node indices. Eliminates
/// the null tombstone slots that petgraph's StableGraph accumulates after
/// repeated `remove_node` calls, keeping serialized YAML size in check.
fn compact(&mut self) {
let mut new_graph: StableGraph<Entity, Relationship> = StableGraph::new();
let mut old_to_new: HashMap<u32, u32> = HashMap::new();
for &old_raw in self.entity_index.values() {
let old_idx = NodeIndex::new(old_raw as usize);
if self.graph.contains_node(old_idx) {
let entity = self.graph[old_idx].clone();
let new_idx = new_graph.add_node(entity);
old_to_new.insert(old_raw, new_idx.index() as u32);
}
}
for edge_idx in self.graph.edge_indices() {
if let Some((from, to)) = self.graph.edge_endpoints(edge_idx) {
let from_raw = from.index() as u32;
let to_raw = to.index() as u32;
if let (Some(&new_from), Some(&new_to)) =
(old_to_new.get(&from_raw), old_to_new.get(&to_raw))
{
let rel = self.graph[edge_idx].clone();
new_graph.add_edge(
NodeIndex::new(new_from as usize),
NodeIndex::new(new_to as usize),
rel,
);
}
}
}
for raw in self.entity_index.values_mut() {
if let Some(&new_raw) = old_to_new.get(raw) {
*raw = new_raw;
}
}
for node_raws in self.document_entities.values_mut() {
*node_raws = node_raws
.iter()
.filter_map(|raw| old_to_new.get(raw).copied())
.collect();
}
self.graph = new_graph;
} }
pub fn build_node_to_docs(&self) -> IndexMap<u32, Vec<DocumentId>> { pub fn build_node_to_docs(&self) -> IndexMap<u32, Vec<DocumentId>> {
@@ -179,30 +257,79 @@ impl KnowledgeGraph {
map map
} }
pub fn expand_neighbors(&self, seed_nodes: &[u32], hops: usize) -> Vec<u32> { /// BFS from seed nodes with weight-decayed scoring.
let mut expanded: indexmap::IndexSet<u32> = seed_nodes.iter().copied().collect(); ///
let mut frontier: Vec<u32> = seed_nodes.to_vec(); /// Seed node scores are provided by the caller (typically token-overlap
/// ratios). Each neighbor's score is `edge_weight * parent_score`, so
/// strongly-connected neighbors rank higher and weakly-connected ones
/// naturally contribute less. Traversal is capped at `MAX_GRAPH_NODES`
/// total nodes; the highest-scored frontier nodes are expanded first so
/// the budget is spent on the most relevant entities.
///
/// Returns a map of raw node index → score (includes seed nodes).
pub fn expand_neighbors_scored(
&self,
seed_scores: &[(u32, f32)],
hops: usize,
) -> IndexMap<u32, f32> {
let mut node_scores: IndexMap<u32, f32> = IndexMap::new();
for &(raw, score) in seed_scores {
node_scores.insert(raw, score);
}
let mut frontier: Vec<(u32, f32)> = seed_scores.to_vec();
for _ in 0..hops { for _ in 0..hops {
let mut next_frontier: Vec<u32> = vec![]; if node_scores.len() >= MAX_GRAPH_NODES {
for &raw in &frontier { break;
let idx = NodeIndex::new(raw as usize); }
if self.graph.contains_node(idx) {
for dir in [Direction::Outgoing, Direction::Incoming] { frontier.sort_unstable_by(|a, b| {
for neighbor in self.graph.neighbors_directed(idx, dir) { b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)
let n = neighbor.index() as u32; });
if expanded.insert(n) {
next_frontier.push(n); let mut next_frontier: Vec<(u32, f32)> = vec![];
'nodes: for (raw, parent_score) in &frontier {
let idx = NodeIndex::new(*raw as usize);
if !self.graph.contains_node(idx) {
continue;
}
for dir in [Direction::Outgoing, Direction::Incoming] {
for edge_ref in self.graph.edges_directed(idx, dir) {
let neighbor_idx = match dir {
Direction::Outgoing => edge_ref.target(),
Direction::Incoming => edge_ref.source(),
};
let neighbor_raw = neighbor_idx.index() as u32;
let candidate = edge_ref.weight().weight * parent_score;
match node_scores.entry(neighbor_raw) {
indexmap::map::Entry::Vacant(e) => {
e.insert(candidate);
next_frontier.push((neighbor_raw, candidate));
} }
indexmap::map::Entry::Occupied(mut e) => {
if candidate > *e.get() {
*e.get_mut() = candidate;
}
}
}
if node_scores.len() >= MAX_GRAPH_NODES {
break 'nodes;
} }
} }
} }
} }
frontier = next_frontier; frontier = next_frontier;
if frontier.is_empty() { if frontier.is_empty() {
break; break;
} }
} }
expanded.into_iter().collect()
node_scores
} }
} }
@@ -213,6 +340,14 @@ pub async fn extract_entities(
chunk: &str, chunk: &str,
prompt_template: Option<&str>, prompt_template: Option<&str>,
) -> Result<ExtractionResult> { ) -> Result<ExtractionResult> {
if chunk.len() > MAX_CHUNK_CHARS {
warn!(
"Entity extraction chunk is {} chars (heuristic limit: {}); \
the LLM response may be truncated",
chunk.len(),
MAX_CHUNK_CHARS
);
}
let template = prompt_template.unwrap_or(EXTRACTION_PROMPT); let template = prompt_template.unwrap_or(EXTRACTION_PROMPT);
let prompt = template.replace("__CHUNK__", chunk); let prompt = template.replace("__CHUNK__", chunk);
let mut messages = vec![Message::new( let mut messages = vec![Message::new(
@@ -250,3 +385,350 @@ pub async fn extract_entities(
serde_json::from_str::<ExtractionResult>(&json) serde_json::from_str::<ExtractionResult>(&json)
.context("Failed to parse entity extraction JSON") .context("Failed to parse entity extraction JSON")
} }
#[cfg(test)]
mod tests {
use super::*;
fn entity(name: &str, entity_type: &str) -> ExtractedEntity {
ExtractedEntity {
name: name.to_string(),
entity_type: entity_type.to_string(),
description: None,
}
}
fn rel(from: &str, to: &str, rel_type: &str, weight: f32) -> ExtractedRelationship {
ExtractedRelationship {
from: from.to_string(),
to: to.to_string(),
relation_type: rel_type.to_string(),
weight: Some(weight),
}
}
fn doc(id: usize) -> DocumentId {
DocumentId(id)
}
fn extraction(
entities: Vec<ExtractedEntity>,
rels: Vec<ExtractedRelationship>,
) -> ExtractionResult {
ExtractionResult {
entities,
relationships: rels,
}
}
#[test]
fn merge_deduplicates_by_lowercase_name() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![
entity("Python", "TECHNOLOGY"),
entity("python", "TECHNOLOGY"),
],
vec![],
),
);
assert_eq!(kg.entity_index.len(), 1);
assert_eq!(kg.graph.node_count(), 1);
}
#[test]
fn merge_chunk_nodes_no_duplicate_doc_entries() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(1),
extraction(
vec![
entity("Python", "TECHNOLOGY"),
entity("python", "TECHNOLOGY"),
],
vec![],
),
);
let count = kg.document_entities.get(&1).map(|v| v.len()).unwrap_or(0);
assert_eq!(
count, 1,
"duplicate entity in one chunk should produce one doc_entity entry"
);
}
#[test]
fn merge_normalizes_entity_type_to_uppercase() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(vec![entity("Django", "technology")], vec![]),
);
let raw = kg.entity_index["django"];
assert_eq!(
kg.graph[NodeIndex::new(raw as usize)].entity_type,
"TECHNOLOGY"
);
}
#[test]
fn merge_promotes_type_from_other_to_specific() {
let mut kg = KnowledgeGraph::default();
kg.merge(doc(0), extraction(vec![entity("Python", "OTHER")], vec![]));
kg.merge(
doc(1),
extraction(vec![entity("Python", "TECHNOLOGY")], vec![]),
);
let raw = kg.entity_index["python"];
assert_eq!(
kg.graph[NodeIndex::new(raw as usize)].entity_type,
"TECHNOLOGY"
);
}
#[test]
fn merge_does_not_demote_specific_type_to_other() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(vec![entity("Python", "TECHNOLOGY")], vec![]),
);
kg.merge(doc(1), extraction(vec![entity("Python", "OTHER")], vec![]));
let raw = kg.entity_index["python"];
assert_eq!(
kg.graph[NodeIndex::new(raw as usize)].entity_type,
"TECHNOLOGY"
);
}
#[test]
fn merge_allows_multiple_relation_types_between_same_pair() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![
entity("Python", "TECHNOLOGY"),
entity("Django", "TECHNOLOGY"),
],
vec![rel("Python", "Django", "implements", 0.9)],
),
);
kg.merge(
doc(1),
extraction(
vec![
entity("Python", "TECHNOLOGY"),
entity("Django", "TECHNOLOGY"),
],
vec![rel("Python", "Django", "uses", 0.8)],
),
);
let from_idx = NodeIndex::new(kg.entity_index["python"] as usize);
let to_idx = NodeIndex::new(kg.entity_index["django"] as usize);
let count = kg.graph.edges_connecting(from_idx, to_idx).count();
assert_eq!(
count, 2,
"two different relation types should produce two edges"
);
}
#[test]
fn merge_deduplicates_same_relation_type() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("A", "B", "uses", 1.0)],
),
);
kg.merge(
doc(1),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("A", "B", "uses", 0.5)],
),
);
let from_idx = NodeIndex::new(kg.entity_index["a"] as usize);
let to_idx = NodeIndex::new(kg.entity_index["b"] as usize);
let count = kg.graph.edges_connecting(from_idx, to_idx).count();
assert_eq!(
count, 1,
"same relation type should not create a duplicate edge"
);
}
#[test]
fn remove_documents_preserves_entity_shared_across_docs() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![entity("Python", "TECHNOLOGY"), entity("A", "CONCEPT")],
vec![],
),
);
kg.merge(
doc(1),
extraction(
vec![entity("Python", "TECHNOLOGY"), entity("B", "CONCEPT")],
vec![],
),
);
kg.remove_documents(&[doc(0)]);
assert!(
kg.entity_index.contains_key("python"),
"shared entity should survive"
);
assert!(
!kg.entity_index.contains_key("a"),
"exclusive entity should be removed"
);
assert!(
kg.entity_index.contains_key("b"),
"other doc's entity should survive"
);
}
#[test]
fn remove_documents_noop_on_empty_slice() {
let mut kg = KnowledgeGraph::default();
kg.merge(doc(0), extraction(vec![entity("X", "CONCEPT")], vec![]));
kg.remove_documents(&[]);
assert_eq!(kg.entity_index.len(), 1);
}
#[test]
fn remove_documents_compacts_graph() {
let mut kg = KnowledgeGraph::default();
// doc 0: A, B with an edge
kg.merge(
doc(0),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("A", "B", "uses", 1.0)],
),
);
// doc 1: C only
kg.merge(doc(1), extraction(vec![entity("C", "CONCEPT")], vec![]));
kg.remove_documents(&[doc(0)]);
assert_eq!(kg.graph.node_count(), 1);
let c_raw = kg.entity_index["c"];
assert_eq!(
c_raw, 0,
"compacted graph should give surviving node index 0"
);
let refs = kg.document_entities.get(&1).cloned().unwrap_or_default();
assert_eq!(refs, vec![0u32]);
}
#[test]
fn expand_zero_hops_returns_seeds_only() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("A", "B", "uses", 0.9)],
),
);
let a_raw = kg.entity_index["a"];
let result = kg.expand_neighbors_scored(&[(a_raw, 1.0)], 0);
assert_eq!(result.len(), 1);
assert_eq!(result[&a_raw], 1.0);
}
#[test]
fn expand_one_hop_decays_score_by_edge_weight() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("A", "B", "uses", 0.8)],
),
);
let a_raw = kg.entity_index["a"];
let b_raw = kg.entity_index["b"];
let result = kg.expand_neighbors_scored(&[(a_raw, 1.0)], 1);
assert_eq!(result.len(), 2);
assert_eq!(result[&a_raw], 1.0);
let b_score = result[&b_raw];
assert!(
(b_score - 0.8).abs() < 1e-6,
"neighbor score should be edge_weight * parent_score = 0.8, got {b_score}"
);
}
#[test]
fn expand_incoming_edges_also_traversed() {
let mut kg = KnowledgeGraph::default();
// Edge goes B → A; seeding A should still discover B via incoming edge
kg.merge(
doc(0),
extraction(
vec![entity("A", "CONCEPT"), entity("B", "CONCEPT")],
vec![rel("B", "A", "uses", 0.7)],
),
);
let a_raw = kg.entity_index["a"];
let b_raw = kg.entity_index["b"];
let result = kg.expand_neighbors_scored(&[(a_raw, 1.0)], 1);
assert!(
result.contains_key(&b_raw),
"B should be reachable via incoming edge from A"
);
let b_score = result[&b_raw];
assert!((b_score - 0.7).abs() < 1e-6);
}
#[test]
fn expand_picks_best_path_score() {
let mut kg = KnowledgeGraph::default();
// A(0.5) → C(0.9): score 0.45; B(1.0) → C(0.4): score 0.40 — A→C path wins.
kg.merge(
doc(0),
extraction(
vec![
entity("A", "CONCEPT"),
entity("B", "CONCEPT"),
entity("C", "CONCEPT"),
],
vec![rel("A", "C", "uses", 0.9), rel("B", "C", "uses", 0.4)],
),
);
let a_raw = kg.entity_index["a"];
let b_raw = kg.entity_index["b"];
let c_raw = kg.entity_index["c"];
let seeds = vec![(a_raw, 0.5f32), (b_raw, 1.0f32)];
let result = kg.expand_neighbors_scored(&seeds, 1);
let c_score = result[&c_raw];
// Best path: B(1.0) * 0.4 = 0.4, A(0.5) * 0.9 = 0.45 → should be 0.45
assert!(
(c_score - 0.45).abs() < 1e-6,
"C score should reflect best path (0.45), got {c_score}"
);
}
#[test]
fn build_node_to_docs_maps_shared_entity_to_multiple_docs() {
let mut kg = KnowledgeGraph::default();
kg.merge(
doc(0),
extraction(vec![entity("Python", "TECHNOLOGY")], vec![]),
);
kg.merge(
doc(1),
extraction(vec![entity("Python", "TECHNOLOGY")], vec![]),
);
let n2d = kg.build_node_to_docs();
let raw = kg.entity_index["python"];
let docs = &n2d[&raw];
assert!(docs.contains(&DocumentId(0)));
assert!(docs.contains(&DocumentId(1)));
}
}
+51 -41
View File
@@ -25,6 +25,8 @@ use std::{
}; };
use tokio::time::sleep; use tokio::time::sleep;
const BM25_SEED_SCORE: f32 = 0.5;
const RAG_TEMPLATE: &str = r#"Answer the query based on the context while respecting the rules. (user query, some textual context and rules, all inside xml tags) const RAG_TEMPLATE: &str = r#"Answer the query based on the context while respecting the rules. (user query, some textual context and rules, all inside xml tags)
<context> <context>
@@ -752,14 +754,14 @@ impl Rag {
bail!("No RAG files"); bail!("No RAG files");
} }
if self.data.extractor_model.is_some() if !new_doc_contents.is_empty()
&& !new_doc_contents.is_empty()
&& let Some(extractor_model_id) = self.data.extractor_model.clone() && let Some(extractor_model_id) = self.data.extractor_model.clone()
{ {
match Model::retrieve_model(&self.app_config, &extractor_model_id, ModelType::Chat) { match Model::retrieve_model(&self.app_config, &extractor_model_id, ModelType::Chat) {
Ok(model) => match self.create_embeddings_client(model) { Ok(model) => match self.create_embeddings_client(model) {
Ok(client) => { Ok(client) => {
let total = new_doc_contents.len(); let total = new_doc_contents.len();
let mut failures = 0usize;
for (i, (doc_id, content)) in new_doc_contents.into_iter().enumerate() { for (i, (doc_id, content)) in new_doc_contents.into_iter().enumerate() {
progress( progress(
&spinner, &spinner,
@@ -774,14 +776,21 @@ impl Rag {
{ {
Ok(result) => self.data.knowledge_graph.merge(doc_id, result), Ok(result) => self.data.knowledge_graph.merge(doc_id, result),
Err(e) => { Err(e) => {
debug!("Entity extraction failed for doc {doc_id:?}: {e}") warn!("Entity extraction failed for doc {doc_id:?}: {e}");
failures += 1;
} }
} }
} }
if failures > 0 {
progress(
&spinner,
format!("Entity extraction: {failures}/{total} chunks failed"),
);
}
} }
Err(e) => debug!("Failed to create extractor client: {e}"), Err(e) => warn!("Failed to create extractor client: {e}"),
}, },
Err(e) => debug!("Extractor model not found: {e}"), Err(e) => warn!("Extractor model not found: {e}"),
} }
} }
@@ -930,9 +939,31 @@ impl Rag {
if kg.entity_index.is_empty() { if kg.entity_index.is_empty() {
return vec![]; return vec![];
} }
let query_lower = query.to_lowercase();
let mut seed_nodes: Vec<u32> = kg let query_lower = query.to_lowercase();
let query_tokens: Vec<&str> = query_lower.split_whitespace().collect();
let token_count = query_tokens.len().max(1);
let score_node = |raw: u32| -> f32 {
let idx = NodeIndex::new(raw as usize);
if !kg.graph.contains_node(idx) {
return 0.0;
}
let entity = &kg.graph[idx];
let combined = format!(
"{} {}",
entity.name,
entity.description.as_deref().unwrap_or("")
)
.to_lowercase();
query_tokens
.iter()
.filter(|t| combined.contains(*t))
.count() as f32
/ token_count as f32
};
let mut seed_scores: Vec<(u32, f32)> = kg
.entity_index .entity_index
.iter() .iter()
.filter(|(name, _)| { .filter(|(name, _)| {
@@ -946,52 +977,31 @@ impl Rag {
.any(|token| token.trim_matches(|c: char| !c.is_alphanumeric()) == name_str) .any(|token| token.trim_matches(|c: char| !c.is_alphanumeric()) == name_str)
} }
}) })
.map(|(_, &raw)| raw) .map(|(_, &raw)| (raw, score_node(raw).max(BM25_SEED_SCORE)))
.collect(); .collect();
if seed_nodes.is_empty() { if seed_scores.is_empty() {
let bm25_results = self.bm25.search(query, top_k * 2); let bm25_results = self.bm25.search(query, top_k * 2);
'outer: for result in bm25_results { 'outer: for result in bm25_results {
if let Some(node_raws) = kg.document_entities.get(&result.document.id.0) { if let Some(node_raws) = kg.document_entities.get(&result.document.id.0) {
seed_nodes.extend(node_raws.iter().copied()); for &raw in node_raws {
if seed_nodes.len() >= top_k { seed_scores.push((raw, BM25_SEED_SCORE));
break 'outer; if seed_scores.len() >= top_k {
break 'outer;
}
} }
} }
} }
} }
if seed_nodes.is_empty() { if seed_scores.is_empty() {
return vec![]; return vec![];
} }
let hops = self.data.graph_hops.unwrap_or(1); let hops = self.data.graph_hops.unwrap_or(1);
let expanded = kg.expand_neighbors(&seed_nodes, hops); let mut scored: Vec<(u32, f32)> = kg
.expand_neighbors_scored(&seed_scores, hops)
let query_tokens: Vec<&str> = query_lower.split_whitespace().collect();
let token_count = query_tokens.len().max(1);
let mut scored: Vec<(u32, f32)> = expanded
.into_iter() .into_iter()
.map(|raw| {
let idx = NodeIndex::new(raw as usize);
let score = if kg.graph.contains_node(idx) {
let entity = &kg.graph[idx];
let combined = format!(
"{} {}",
entity.name,
entity.description.as_deref().unwrap_or("")
)
.to_lowercase();
query_tokens
.iter()
.filter(|t| combined.contains(*t))
.count() as f32
/ token_count as f32
} else {
0.0
};
(raw, score)
})
.collect(); .collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal)); scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal));
@@ -1349,11 +1359,11 @@ fn set_chunk_size(model: &Model) -> Result<usize> {
fn set_graph_hops(default_value: usize) -> Result<usize> { fn set_graph_hops(default_value: usize) -> Result<usize> {
let value = Text::new("Set graph expansion hops:") let value = Text::new("Set graph expansion hops:")
.with_default(&default_value.to_string()) .with_default(&default_value.to_string())
.with_help_message("Number of hops to expand from matched entities (1 = direct neighbors, 2 = neighbors of neighbors)") .with_help_message("Number of hops to expand from matched entities (0 = seed nodes only, 1 = direct neighbors, 2 = neighbors of neighbors)")
.with_validator(move |text: &str| { .with_validator(move |text: &str| {
let out = match text.parse::<usize>() { let out = match text.parse::<usize>() {
Ok(v) if v >= 1 => Validation::Valid, Ok(_) => Validation::Valid,
_ => Validation::Invalid("Must be an integer >= 1".into()), _ => Validation::Invalid("Must be a non-negative integer".into()),
}; };
Ok(out) Ok(out)
}) })