#![cfg(feature = "embeddings")]
use crate::db::models::CodeElement;
use crate::embeddings::{build_blob, Embedder};
use crate::graph::traversal::is_indexer_noise;
use crate::graph::GraphEngine;
use std::collections::{HashMap, HashSet, VecDeque};
pub const FUNCTION_TARGET_TYPES: &[&str] = &["function", "method", "constructor"];
pub const UPPER_TYPES: &[&str] = &[
"class",
"struct",
"interface",
"trait",
"module",
"file",
"document",
"doc_section",
"workflow",
"workflow_step",
"decision_point",
"failure_mode",
"domain_entity",
"service",
"api_endpoint",
"data_store",
"known_issue",
"playbook",
"playbook_step",
"team_knowledge",
];
pub fn is_function_target(element_type: &str) -> bool {
FUNCTION_TARGET_TYPES.contains(&element_type)
}
pub fn is_upper_type(element_type: &str) -> bool {
UPPER_TYPES.contains(&element_type)
}
pub struct DownwardRule {
pub hops: u32,
pub edge_types: &'static [&'static str],
pub fanout_cap: usize,
}
const WORKFLOW_EDGES: &[&str] = &[
"has_step",
"next_step",
"branches_to",
"implemented_by",
"entry_point_of",
"step_in_process",
"has_failure_mode",
];
const STEP_EDGES: &[&str] = &[
"next_step",
"branches_to",
"implemented_by",
"handled_by_playbook",
"has_failure_mode",
"resolved_by_playbook",
];
const CONCEPT_EDGES: &[&str] = &[
"owns_concept",
"implements_concept",
"exposes_endpoint",
"reads_from",
"writes_to",
"documents_concept",
"has_known_issue",
];
const ISSUE_EDGES: &[&str] = &[
"has_known_issue",
"resolved_by_playbook",
"documents_concept",
];
const CLASS_DOWN_EDGES: &[&str] = &["contains", "defines", "has_method", "has_property"];
const FILE_DOWN_EDGES: &[&str] = &[
"contains",
"defines",
"imports",
"references",
"tested_by",
"documented_by",
];
const DOC_DOWN_EDGES: &[&str] = &["references", "documented_by"];
const DOC_FALLBACK_EDGES: &[&str] = &["documented_by", "documents_concept"];
pub fn downward_rule_for(element_type: &str) -> DownwardRule {
match element_type {
"class" | "struct" | "interface" | "trait" | "module" => DownwardRule {
hops: 1,
edge_types: CLASS_DOWN_EDGES,
fanout_cap: 12,
},
"file" => DownwardRule {
hops: 1,
edge_types: FILE_DOWN_EDGES,
fanout_cap: 12,
},
"document" | "doc_section" => DownwardRule {
hops: 1,
edge_types: DOC_DOWN_EDGES,
fanout_cap: 10,
},
"workflow" => DownwardRule {
hops: 2,
edge_types: WORKFLOW_EDGES,
fanout_cap: 15,
},
"workflow_step" | "decision_point" | "failure_mode" => DownwardRule {
hops: 1,
edge_types: STEP_EDGES,
fanout_cap: 12,
},
"domain_entity" | "service" | "api_endpoint" | "data_store" => DownwardRule {
hops: 2,
edge_types: CONCEPT_EDGES,
fanout_cap: 12,
},
"known_issue" | "playbook" | "playbook_step" | "team_knowledge" => DownwardRule {
hops: 1,
edge_types: ISSUE_EDGES,
fanout_cap: 8,
},
_ => DownwardRule {
hops: 1,
edge_types: DOC_FALLBACK_EDGES,
fanout_cap: 5,
},
}
}
pub const GLOBAL_FUNCTION_CAP: usize = 80;
#[derive(Debug, Clone)]
pub struct UpperSeed {
pub qualified_name: String,
pub element_type: String,
pub name: String,
}
impl UpperSeed {
pub fn new(qualified_name: impl Into<String>, element_type: impl Into<String>) -> Self {
Self::with_name(qualified_name, element_type, "")
}
pub fn with_name(
qualified_name: impl Into<String>,
element_type: impl Into<String>,
name: impl Into<String>,
) -> Self {
let qualified_name = qualified_name.into();
let name = name.into();
let display_name = if name.is_empty() {
derive_display_name(&qualified_name)
} else {
name
};
Self {
qualified_name,
element_type: element_type.into(),
name: display_name,
}
}
}
fn derive_display_name(qualified_name: &str) -> String {
if let Some(gid) = qualified_name.strip_prefix("ontology://") {
let parts: Vec<&str> = gid.split(':').collect();
if parts.len() >= 5 {
let id = parts[parts.len() - 2];
if !id.is_empty() {
return id.to_string();
}
}
}
qualified_name
.rsplit(['/', ':'])
.next()
.filter(|s| !s.is_empty())
.unwrap_or(qualified_name)
.to_string()
}
#[derive(Debug, Clone)]
pub struct DiscoveredFunction {
pub qualified_name: String,
pub element_type: String,
pub file_path: String,
pub env: String,
pub via_upper: String,
pub via_upper_type: String,
pub via_upper_name: String,
pub via_edge: String,
pub hop: u32,
pub composite_score: Option<f32>,
}
pub fn traverse_to_functions(
graph: &GraphEngine,
upper_seeds: &[UpperSeed],
env: Option<&str>,
) -> Result<Vec<DiscoveredFunction>, Box<dyn std::error::Error>> {
let mut discovered: Vec<DiscoveredFunction> = Vec::new();
let mut seen: HashMap<String, usize> = HashMap::new();
let mut total = 0usize;
for upper in upper_seeds {
if total >= GLOBAL_FUNCTION_CAP {
break;
}
let rule = downward_rule_for(&upper.element_type);
let seed_fanout_cap = rule.fanout_cap.min(GLOBAL_FUNCTION_CAP.saturating_sub(total));
let mut found_for_this_seed = 0usize;
let mut visited: HashSet<String> = HashSet::new();
visited.insert(upper.qualified_name.clone());
let mut frontier: VecDeque<(String, u32, String)> = VecDeque::new();
frontier.push_back((upper.qualified_name.clone(), 0, "seed".to_string()));
while let Some((current, hop, _via_edge_into_current)) = frontier.pop_front() {
if found_for_this_seed >= seed_fanout_cap || total >= GLOBAL_FUNCTION_CAP {
break;
}
if hop >= rule.hops {
continue;
}
let outgoing = graph.get_relationships(¤t).unwrap_or_default();
let incoming = graph.get_relationships_for_target(¤t).unwrap_or_default();
for rel in outgoing.iter().chain(incoming.iter()) {
if found_for_this_seed >= seed_fanout_cap || total >= GLOBAL_FUNCTION_CAP {
break;
}
if !rule.edge_types.contains(&rel.rel_type.as_str()) {
continue;
}
if let Some(wanted) = env {
if rel.env != wanted {
continue;
}
}
let neighbor = if rel.source_qualified == current {
rel.target_qualified.clone()
} else {
rel.source_qualified.clone()
};
if !visited.insert(neighbor.clone()) {
continue;
}
let Some(element) = graph.find_element(&neighbor).ok().flatten() else {
continue;
};
if is_indexer_noise(&element.element_type) {
continue;
}
let next_hop = hop + 1;
let via_edge = rel.rel_type.clone();
if is_function_target(&element.element_type) {
record_function(
&mut discovered,
&mut seen,
&mut total,
&element,
&upper.qualified_name,
&upper.element_type,
&upper.name,
&via_edge,
next_hop,
);
found_for_this_seed += 1;
continue;
}
if next_hop < rule.hops {
frontier.push_back((neighbor, next_hop, via_edge));
}
}
}
if found_for_this_seed == 0 && needs_code_refs_fallback(&upper.element_type) {
resolve_code_refs_fallback(
graph,
upper,
env,
seed_fanout_cap,
&mut discovered,
&mut seen,
&mut total,
)?;
}
}
Ok(discovered)
}
fn needs_code_refs_fallback(element_type: &str) -> bool {
matches!(
element_type,
"domain_entity" | "service" | "api_endpoint" | "data_store"
| "workflow" | "workflow_step" | "known_issue" | "playbook"
| "team_knowledge"
)
}
fn resolve_code_refs_fallback(
graph: &GraphEngine,
upper: &UpperSeed,
env: Option<&str>,
fanout_cap: usize,
discovered: &mut Vec<DiscoveredFunction>,
seen: &mut HashMap<String, usize>,
total: &mut usize,
) -> Result<(), Box<dyn std::error::Error>> {
let Some(element) = graph.find_element(&upper.qualified_name)? else {
return Ok(());
};
let code_refs = element
.metadata
.get("code_refs")
.and_then(|v| v.as_array())
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(String::from))
.collect::<Vec<_>>()
});
let Some(code_refs) = code_refs else {
return Ok(());
};
if code_refs.is_empty() {
return Ok(());
}
let mut matched = 0usize;
for raw_ref in &code_refs {
if matched >= fanout_cap || *total >= GLOBAL_FUNCTION_CAP {
break;
}
let candidates: Vec<CodeElement> = if let Ok(Some(el)) = graph.find_element(raw_ref.trim()) {
vec![el]
} else if let Some((file_part, sym_part)) = raw_ref.trim().split_once("::") {
let per_ref_cap = (fanout_cap - matched).min(80);
graph
.find_elements_by_file_path_prefix(file_part, per_ref_cap)?
.into_iter()
.filter(|e| {
let sym_lower = sym_part.to_lowercase();
e.name.to_lowercase() == sym_lower
|| e
.qualified_name
.to_lowercase()
.ends_with(&format!("::{}", sym_lower))
})
.collect()
} else {
let per_ref_cap = (fanout_cap - matched).min(80);
graph.find_elements_by_file_path_prefix(raw_ref.trim(), per_ref_cap)?
};
for el in candidates {
if matched >= fanout_cap || *total >= GLOBAL_FUNCTION_CAP {
break;
}
if let Some(wanted) = env {
if el.env != wanted {
continue;
}
}
if !is_function_target(&el.element_type) {
continue;
}
record_function(
discovered,
seen,
total,
&el,
&upper.qualified_name,
&upper.element_type,
&upper.name,
"code_ref",
1,
);
matched += 1;
}
}
Ok(())
}
fn record_function(
discovered: &mut Vec<DiscoveredFunction>,
seen: &mut HashMap<String, usize>,
total: &mut usize,
element: &CodeElement,
via_upper: &str,
via_upper_type: &str,
via_upper_name: &str,
via_edge: &str,
hop: u32,
) {
if let Some(&idx) = seen.get(&element.qualified_name) {
if hop < discovered[idx].hop {
discovered[idx] = DiscoveredFunction {
qualified_name: element.qualified_name.clone(),
element_type: element.element_type.clone(),
file_path: element.file_path.clone(),
env: element.env.clone(),
via_upper: via_upper.to_string(),
via_upper_type: via_upper_type.to_string(),
via_upper_name: via_upper_name.to_string(),
via_edge: via_edge.to_string(),
hop,
composite_score: None,
};
}
return;
}
seen.insert(element.qualified_name.clone(), discovered.len());
discovered.push(DiscoveredFunction {
qualified_name: element.qualified_name.clone(),
element_type: element.element_type.clone(),
file_path: element.file_path.clone(),
env: element.env.clone(),
via_upper: via_upper.to_string(),
via_upper_type: via_upper_type.to_string(),
via_upper_name: via_upper_name.to_string(),
via_edge: via_edge.to_string(),
hop,
composite_score: None,
});
*total += 1;
}
pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b.iter())
.map(|(x, y)| x * y)
.sum()
}
pub fn composite_text(upper_name: &str, func_blob: &str) -> String {
if func_blob.is_empty() {
upper_name.to_string()
} else {
format!("{upper_name}\n{func_blob}")
}
}
pub fn score_functions(
graph: &GraphEngine,
query_vec: &[f32],
functions: Vec<DiscoveredFunction>,
embedder: &Embedder,
) -> Result<Vec<DiscoveredFunction>, Box<dyn std::error::Error>> {
if functions.is_empty() {
return Ok(functions);
}
let mut blob_cache: HashMap<String, String> = HashMap::new();
let qns: Vec<String> = functions
.iter()
.map(|f| f.qualified_name.clone())
.collect();
for qn in &qns {
if blob_cache.contains_key(qn) {
continue;
}
let blob = graph
.find_element(qn)
.ok()
.flatten()
.and_then(|el| build_blob(&el))
.unwrap_or_default();
blob_cache.insert(qn.clone(), blob);
}
let composite_texts: Vec<String> = functions
.iter()
.map(|f| composite_text(&f.via_upper_name, blob_cache.get(&f.qualified_name).map(|s| s.as_str()).unwrap_or("")))
.collect();
let borrowed: Vec<String> = composite_texts; let vectors = embedder.embed(&borrowed)?;
let mut scored = functions;
for (func, vec) in scored.iter_mut().zip(vectors.iter()) {
let raw = cosine(query_vec, vec);
let clamped = raw.clamp(-1.0, 1.0);
func.composite_score = Some(clamped);
}
scored.sort_by(|a, b| {
b.composite_score
.unwrap_or(-1.0)
.partial_cmp(&a.composite_score.unwrap_or(-1.0))
.unwrap_or(std::cmp::Ordering::Equal)
});
Ok(scored)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn downward_rule_for_class_is_one_hop_via_containment() {
for t in ["class", "struct", "interface", "trait", "module"] {
let r = downward_rule_for(t);
assert_eq!(r.hops, 1, "{t} should be 1 hop");
assert!(r.edge_types.contains(&"contains"), "{t} should allow contains");
assert!(r.edge_types.contains(&"defines"), "{t} should allow defines");
assert!(r.edge_types.contains(&"has_method"), "{t} should allow has_method");
}
}
#[test]
fn downward_rule_for_workflow_is_two_hops() {
let r = downward_rule_for("workflow");
assert_eq!(r.hops, 2);
assert!(r.edge_types.contains(&"has_step"));
assert!(r.edge_types.contains(&"implemented_by"));
}
#[test]
fn downward_rule_for_concept_two_hops_with_code_refs_fallback() {
for t in ["domain_entity", "service", "api_endpoint", "data_store"] {
let r = downward_rule_for(t);
assert_eq!(r.hops, 2, "{t} should be 2 hops");
assert!(needs_code_refs_fallback(t), "{t} should need code_refs fallback");
}
}
#[test]
fn downward_rule_for_step_one_hop_via_implemented_by() {
for t in ["workflow_step", "decision_point", "failure_mode"] {
let r = downward_rule_for(t);
assert_eq!(r.hops, 1);
assert!(r.edge_types.contains(&"implemented_by"), "{t} should allow implemented_by");
}
}
#[test]
fn downward_rule_default_uses_doc_edges() {
let r = downward_rule_for("some_unknown_type");
assert_eq!(r.hops, 1);
assert_eq!(r.fanout_cap, 5);
assert!(r.edge_types.contains(&"documented_by"));
}
#[test]
fn is_function_target_matches_function_types() {
assert!(is_function_target("function"));
assert!(is_function_target("method"));
assert!(is_function_target("constructor"));
assert!(!is_function_target("class"));
assert!(!is_function_target("struct"));
assert!(!is_function_target("file"));
}
#[test]
fn is_upper_type_matches_upper_types() {
assert!(is_upper_type("class"));
assert!(is_upper_type("workflow"));
assert!(is_upper_type("domain_entity"));
assert!(is_upper_type("document"));
assert!(!is_upper_type("function"));
assert!(!is_upper_type("method"));
}
#[test]
fn indexer_noise_types_are_filtered_by_is_indexer_noise() {
assert!(is_indexer_noise("unknown"));
assert!(is_indexer_noise("environment"));
assert!(!is_indexer_noise("function"));
}
#[test]
fn cosine_identical_vectors_is_one() {
let a = vec![0.6, 0.8, 0.0];
let c = cosine(&a, &a);
assert!((c - 1.0).abs() < 1e-5, "identical unit vectors: cosine ~ 1, got {c}");
}
#[test]
fn cosine_orthogonal_vectors_is_zero() {
let a = vec![1.0, 0.0];
let b = vec![0.0, 1.0];
let c = cosine(&a, &b);
assert!(c.abs() < 1e-5, "orthogonal: cosine ~ 0, got {c}");
}
#[test]
fn cosine_opposite_vectors_is_minus_one() {
let a = vec![1.0, 0.0];
let b = vec![-1.0, 0.0];
let c = cosine(&a, &b);
assert!((c + 1.0).abs() < 1e-5, "opposite: cosine ~ -1, got {c}");
}
#[test]
fn composite_text_joins_upper_name_and_blob() {
let t = composite_text("SemanticRetrievalPipeline", "fn retrieve(query) -> Result<...>");
assert!(t.starts_with("SemanticRetrievalPipeline\n"));
assert!(t.contains("fn retrieve"));
}
#[test]
fn composite_text_empty_blob_is_just_upper_name() {
let t = composite_text("MyClass", "");
assert_eq!(t, "MyClass");
}
#[test]
fn upper_seed_default_name_uses_trailing_segment() {
let s = UpperSeed::new("src/retrieval/pipeline.rs::SemanticRetrievalPipeline", "class");
assert_eq!(s.name, "SemanticRetrievalPipeline");
}
#[test]
fn upper_seed_default_name_extracts_ontology_gid_id_not_version() {
let s = UpperSeed::new("ontology://local:checkout-service:domain_entity:refund:v1", "domain_entity");
assert_eq!(s.name, "refund");
let s2 = UpperSeed::new("ontology://local:ws:workflow:code_indexing_flow:v2", "workflow");
assert_eq!(s2.name, "code_indexing_flow");
}
#[test]
fn upper_seed_with_empty_name_falls_back_to_qn_segment() {
let s = UpperSeed::with_name("src/foo.rs::Bar", "class", "");
assert_eq!(s.name, "Bar");
let s2 = UpperSeed::with_name("src/foo.rs::Bar", "class", "CustomName");
assert_eq!(s2.name, "CustomName");
}
#[test]
fn needs_code_refs_fallback_only_for_ontology_types() {
assert!(needs_code_refs_fallback("domain_entity"));
assert!(needs_code_refs_fallback("workflow"));
assert!(needs_code_refs_fallback("workflow_step"));
assert!(!needs_code_refs_fallback("class"));
assert!(!needs_code_refs_fallback("file"));
assert!(!needs_code_refs_fallback("document"));
}
#[test]
fn record_function_dedups_keeping_shortest_hop() {
let mut discovered: Vec<DiscoveredFunction> = Vec::new();
let mut seen: HashMap<String, usize> = HashMap::new();
let mut total = 0usize;
let make_el = |qn: &str| CodeElement {
qualified_name: qn.to_string(),
element_type: "function".to_string(),
file_path: "src/x.rs".to_string(),
env: "local".to_string(),
..Default::default()
};
record_function(
&mut discovered,
&mut seen,
&mut total,
&make_el("src/x.rs::foo"),
"ontology://w",
"workflow",
"WorkflowA",
"implemented_by",
2,
);
assert_eq!(discovered.len(), 1);
assert_eq!(total, 1);
assert_eq!(discovered[0].hop, 2);
record_function(
&mut discovered,
&mut seen,
&mut total,
&make_el("src/x.rs::foo"),
"src/x.rs::Foo",
"class",
"Foo",
"contains",
1,
);
assert_eq!(discovered.len(), 1, "dedup must not add a second entry");
assert_eq!(total, 1, "total must not double-count");
assert_eq!(discovered[0].hop, 1, "shortest hop should win");
assert_eq!(discovered[0].via_upper_type, "class");
record_function(
&mut discovered,
&mut seen,
&mut total,
&make_el("src/x.rs::bar"),
"src/x.rs::Foo",
"class",
"Foo",
"contains",
1,
);
assert_eq!(discovered.len(), 2);
assert_eq!(total, 2);
}
#[test]
fn record_function_longer_hop_does_not_overwrite_shorter() {
let mut discovered: Vec<DiscoveredFunction> = Vec::new();
let mut seen: HashMap<String, usize> = HashMap::new();
let mut total = 0usize;
let el = CodeElement {
qualified_name: "src/x.rs::foo".to_string(),
element_type: "function".to_string(),
..Default::default()
};
record_function(&mut discovered, &mut seen, &mut total, &el, "C", "class", "C", "contains", 1);
record_function(&mut discovered, &mut seen, &mut total, &el, "W", "workflow", "W", "implemented_by", 2);
assert_eq!(discovered.len(), 1);
assert_eq!(discovered[0].hop, 1);
assert_eq!(discovered[0].via_upper, "C");
}
}