use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::{HashSet, VecDeque};
use std::path::Path;
use rmcp::model::Tool;
use scryer_db::{CodeGraphEdge, SourceFile, Symbol};
use scryer_engine::EngineService;
use super::admin::{make_tool, read_only};
use super::dependency::{
ResolvedSymbol, SymbolCandidate, lookup_symbol_candidates, narrow_by_file,
};
use crate::context::ProjectContextResolver;
const MAX_CANDIDATES: usize = 10;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum EdgeDirection {
Inbound,
Outbound,
}
impl EdgeDirection {
pub(crate) fn neighbour(self, edge: &CodeGraphEdge) -> u64 {
match self {
Self::Inbound => edge.source_symbol_id,
Self::Outbound => edge.target_symbol_id,
}
}
}
pub(crate) async fn one_hop_edges(
db: &mut toasty::Db,
project_id: u64,
symbol_id: u64,
direction: EdgeDirection,
edge_types: Option<&[&str]>,
) -> anyhow::Result<Vec<CodeGraphEdge>> {
let by_symbol = match direction {
EdgeDirection::Inbound => CodeGraphEdge::fields().target_symbol_id().eq(symbol_id),
EdgeDirection::Outbound => CodeGraphEdge::fields().source_symbol_id().eq(symbol_id),
};
let mut edges = CodeGraphEdge::filter(
CodeGraphEdge::fields()
.project_id()
.eq(project_id)
.and(by_symbol),
)
.exec(db)
.await?;
if let Some(types) = edge_types {
edges.retain(|e| types.contains(&e.edge_type.as_str()));
}
Ok(edges)
}
pub(crate) async fn source_file_path(
db: &mut toasty::Db,
project_id: u64,
file_id: u64,
) -> anyhow::Result<Option<String>> {
let file = SourceFile::filter(
SourceFile::fields()
.project_id()
.eq(project_id)
.and(SourceFile::fields().id().eq(file_id)),
)
.first()
.exec(db)
.await?;
Ok(file.map(|f| f.path))
}
#[derive(Debug, Clone, Deserialize, JsonSchema)]
pub struct TraceCallHierarchyParams {
pub symbol: String,
pub file_path: Option<String>,
pub direction: Option<super::enums::CallDirection>,
pub max_depth: Option<usize>,
pub limit: Option<usize>,
pub file_filter: Option<String>,
pub max_tokens: Option<usize>,
pub no_truncate: Option<bool>,
pub project: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct CallHierarchyNode {
pub symbol_name: String,
pub qualified_name: String,
pub signature: String,
pub file_path: String,
pub depth: usize,
pub edge_type: String,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema)]
pub struct TraceCallHierarchyResult {
pub symbol: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub qualified_name: Option<String>,
pub direction: String,
pub depth: usize,
pub total_nodes: usize,
pub returned_nodes: usize,
pub has_more: bool,
pub nodes: Vec<CallHierarchyNode>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub ambiguous: bool,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub candidates: Vec<SymbolCandidate>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub notes: Vec<String>,
}
#[derive(Debug, Clone, Default, Deserialize, JsonSchema)]
pub struct CalculateBlastRadiusParams {
pub symbol: Option<String>,
pub file_path: Option<String>,
pub limit: Option<usize>,
pub include_symbols: Option<bool>,
pub max_tokens: Option<usize>,
pub no_truncate: Option<bool>,
pub project: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct BlastRadiusResult {
pub target: String,
pub affected_symbols_count: usize,
pub affected_files_count: usize,
pub downstream_files: Vec<String>,
pub affected_symbols: Vec<String>,
pub truncated: bool,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub notes: Vec<String>,
}
pub async fn handle_trace_call_hierarchy(
context: &ProjectContextResolver,
engine: &EngineService,
params: TraceCallHierarchyParams,
) -> anyhow::Result<TraceCallHierarchyResult> {
let file_arg = params.file_path.as_deref().map(Path::new);
let (project, rel_hint) = context
.resolve_project(file_arg, params.project.as_deref())
.await?;
let direction = params
.direction
.unwrap_or(super::enums::CallDirection::Outbound)
.as_str()
.to_string();
let max_depth = params.max_depth.unwrap_or(2).clamp(1, 5);
let mut guard = engine.db().lock().await;
let mut out = TraceCallHierarchyResult {
symbol: params.symbol.clone(),
direction: direction.clone(),
..Default::default()
};
let mut candidates = lookup_symbol_candidates(&mut guard, &project, ¶ms.symbol).await?;
let only_external = !candidates.is_empty() && candidates.iter().all(|c| c.is_external());
candidates.retain(|c| !c.is_external());
if let Some(fp) = ¶ms.file_path {
let hint = rel_hint
.map(|p| p.to_string_lossy().into_owned())
.unwrap_or_else(|| fp.clone());
match narrow_by_file(&candidates, &[&hint, fp]) {
Some(hinted) => candidates = hinted,
None => out.notes.push(format!(
"file_path '{fp}' matched no definition of '{}'; ignoring it",
params.symbol
)),
}
}
if candidates.is_empty() {
out.notes.push(if only_external {
format!(
"'{}' is a dependency symbol; the call graph covers workspace code only",
params.symbol
)
} else {
format!(
"no definition named '{}'; try search_symbols",
params.symbol
)
});
return Ok(out);
}
if candidates.len() > 1 {
out.ambiguous = true;
out.candidates = candidates
.iter()
.take(MAX_CANDIDATES)
.map(ResolvedSymbol::candidate)
.collect();
out.notes.push(format!(
"{} definitions share this name; pass file_path or a qualified name",
candidates.len()
));
return Ok(out);
}
let root = candidates.remove(0).symbol;
out.qualified_name = Some(root.qualified_name.clone());
let symbols = Symbol::filter(Symbol::fields().project_id().eq(project.id))
.exec(&mut *guard)
.await?;
let mut visited = HashSet::new();
let mut queue = VecDeque::new();
let mut result_nodes = Vec::new();
visited.insert(root.id);
queue.push_back((root.id, 0usize));
while let Some((curr_id, curr_depth)) = queue.pop_front() {
if curr_depth >= max_depth {
continue;
}
let edge_direction = if direction == "inbound" {
EdgeDirection::Inbound
} else {
EdgeDirection::Outbound
};
let edges = one_hop_edges(&mut guard, project.id, curr_id, edge_direction, None).await?;
for edge in edges {
let next_id = edge_direction.neighbour(&edge);
if visited.contains(&next_id) {
continue;
}
visited.insert(next_id);
let next_sym = symbols.iter().find(|s| s.id == next_id);
if let Some(sym) = next_sym {
let file_path = source_file_path(&mut guard, project.id, sym.file_id)
.await?
.unwrap_or_else(|| "unknown".to_string());
if let Some(ff) = ¶ms.file_filter
&& !file_path.to_lowercase().contains(&ff.to_lowercase())
{
continue;
}
result_nodes.push(CallHierarchyNode {
symbol_name: sym.name.clone(),
qualified_name: sym.qualified_name.clone(),
signature: sym.signature.clone(),
file_path,
depth: curr_depth + 1,
edge_type: edge.edge_type.clone(),
});
queue.push_back((next_id, curr_depth + 1));
}
}
}
let total_nodes = result_nodes.len();
let (paginated_nodes, has_more) = if let Some(limit) = params.limit {
if result_nodes.len() > limit {
result_nodes.truncate(limit);
(result_nodes, true)
} else {
(result_nodes, false)
}
} else {
(result_nodes, false)
};
let returned_nodes = paginated_nodes.len();
out.depth = max_depth;
out.total_nodes = total_nodes;
out.returned_nodes = returned_nodes;
out.has_more = has_more;
out.nodes = paginated_nodes;
Ok(out)
}
pub async fn handle_calculate_blast_radius(
context: &ProjectContextResolver,
engine: &EngineService,
params: CalculateBlastRadiusParams,
) -> anyhow::Result<BlastRadiusResult> {
anyhow::ensure!(
params.symbol.is_some() || params.file_path.is_some(),
"Pass `symbol` (a symbol name) or `file_path` (a source file) to assess"
);
let file_path = params.file_path.as_deref().map(Path::new);
let (project, rel_path) = context
.resolve_project(file_path, params.project.as_deref())
.await?;
let rel_path = rel_path.or_else(|| {
file_path
.filter(|p| {
p.is_relative()
&& !p
.components()
.any(|c| matches!(c, std::path::Component::ParentDir))
})
.map(Path::to_path_buf)
});
let mut guard = engine.db().lock().await;
let symbols = Symbol::filter(Symbol::fields().project_id().eq(project.id))
.exec(&mut *guard)
.await?;
let mut seed_symbol_ids = Vec::new();
let mut notes = Vec::new();
let target_label = if let Some(sym_name) = ¶ms.symbol {
for s in &symbols {
if s.name == *sym_name || s.qualified_name == *sym_name {
seed_symbol_ids.push(s.id);
}
}
if seed_symbol_ids.is_empty() {
notes.push(format!(
"no symbol named '{sym_name}' in this project; try search_symbols(query: \"{sym_name}\") or check the project"
));
}
sym_name.clone()
} else if let Some(rel) = rel_path {
let rel_str = rel.to_string_lossy().to_string();
let file = SourceFile::filter(
SourceFile::fields()
.project_id()
.eq(project.id)
.and(SourceFile::fields().path().eq(&rel_str)),
)
.first()
.exec(&mut *guard)
.await?;
if let Some(file) = file {
for s in &symbols {
if s.file_id == file.id {
seed_symbol_ids.push(s.id);
}
}
} else {
notes.push(format!(
"'{rel_str}' is not in the index; check the path, or run index_workspace if the file is new"
));
}
rel_str
} else {
notes.push("file_path is outside the resolved project".to_string());
params.file_path.clone().unwrap_or_default()
};
let mut visited_symbols = HashSet::new();
let mut queue = VecDeque::new();
for id in seed_symbol_ids {
visited_symbols.insert(id);
queue.push_back(id);
}
while let Some(curr_id) = queue.pop_front() {
let reverse_edges = one_hop_edges(
&mut guard,
project.id,
curr_id,
EdgeDirection::Inbound,
None,
)
.await?;
for edge in reverse_edges {
if !visited_symbols.contains(&edge.source_symbol_id) {
visited_symbols.insert(edge.source_symbol_id);
queue.push_back(edge.source_symbol_id);
}
}
}
let mut affected_files = HashSet::new();
let mut affected_symbol_names = Vec::new();
for id in &visited_symbols {
if let Some(s) = symbols.iter().find(|sym| sym.id == *id) {
affected_symbol_names.push(s.qualified_name.clone());
if let Some(path) = source_file_path(&mut guard, project.id, s.file_id).await? {
affected_files.insert(path);
}
}
}
let mut downstream_files: Vec<String> = affected_files.into_iter().collect();
downstream_files.sort();
affected_symbol_names.sort();
let total_affected_symbols = visited_symbols.len();
let total_affected_files = downstream_files.len();
let include_symbols = params.include_symbols.unwrap_or(true);
let mut affected_symbols = if include_symbols {
affected_symbol_names
} else {
Vec::new()
};
let mut truncated = false;
if let Some(limit) = params.limit {
if downstream_files.len() > limit {
downstream_files.truncate(limit);
truncated = true;
}
if affected_symbols.len() > limit {
affected_symbols.truncate(limit);
truncated = true;
}
}
Ok(BlastRadiusResult {
target: target_label,
affected_symbols_count: total_affected_symbols,
affected_files_count: total_affected_files,
downstream_files,
affected_symbols,
truncated,
notes,
})
}
pub fn tool_definitions() -> Vec<Tool> {
vec![
make_tool::<TraceCallHierarchyParams>(
"trace_call_hierarchy",
"Use to follow call relationships beyond one hop: direction 'inbound' (callers) or 'outbound' (callees, the default), up to max_depth (default 2, maximum 5) across the workspace call graph. Starts from the definition, never a `use` import; when several definitions share the name it returns candidates, so pass file_path or a qualified name. Paginate with limit and narrow with file_filter.",
read_only(),
),
make_tool::<CalculateBlastRadiusParams>(
"calculate_blast_radius",
"Use before changing or deleting a symbol or file to see which indexed workspace code depends on it. Pass `symbol` or `file_path` (at least one). Returns the reverse dependency graph, affected symbol count and downstream files.",
read_only(),
),
]
}