use ryo_analysis::context::AnalysisContext;
use ryo_analysis::query::CodeGraphV2;
use ryo_analysis::ASTRegistry;
use ryo_mutations::basic::trait_ops::cross_crate_caller_pattern::{classify_fn, CallerPattern};
use ryo_source::pure::{PureItem, PureTraitItem, PureVis};
use ryo_symbol::{SymbolId, SymbolKind, SymbolRegistry};
use std::collections::{HashSet, VecDeque};
#[derive(Debug, Clone)]
pub struct CallerRewrite {
pub caller_id: SymbolId,
pub patterns: Vec<CallerPattern>,
pub is_public: bool,
}
pub fn scan_cross_crate_callers_raw(
ast_registry: &ASTRegistry,
symbol_registry: &SymbolRegistry,
trait_id: SymbolId,
) -> Vec<CallerRewrite> {
let trait_crate = match symbol_registry.path(trait_id) {
Some(p) => p.crate_name().to_string(),
None => return Vec::new(),
};
let trait_name = match symbol_registry.path(trait_id) {
Some(p) => p.name().to_string(),
None => return Vec::new(),
};
if trait_crate.is_empty() || trait_name.is_empty() {
return Vec::new();
}
let trait_methods: Vec<String> = match ast_registry.get(trait_id) {
Some(PureItem::Trait(t)) => t
.items
.iter()
.filter_map(|it| match it {
PureTraitItem::Fn(f) => Some(f.name.clone()),
_ => None,
})
.collect(),
_ => Vec::new(),
};
let mut out = Vec::new();
let fn_ids: Vec<SymbolId> = symbol_registry
.iter()
.filter(|(id, path)| {
matches!(symbol_registry.kind(*id), Some(SymbolKind::Function))
&& path.crate_name() != trait_crate
})
.map(|(id, _)| id)
.collect();
for caller_id in fn_ids {
if let Some(PureItem::Fn(f)) = ast_registry.get(caller_id) {
let patterns = classify_fn(f, &trait_name, &trait_methods);
if !patterns.is_empty() {
let is_public = matches!(f.vis, PureVis::Public);
out.push(CallerRewrite {
caller_id,
patterns,
is_public,
});
}
}
}
out
}
pub fn scan_cross_crate_callers(ctx: &AnalysisContext, trait_id: SymbolId) -> Vec<CallerRewrite> {
scan_cross_crate_callers_raw(&ctx.ast_registry, &ctx.registry, trait_id)
}
pub fn scan_same_crate_external_callers_raw(
ast_registry: &ASTRegistry,
symbol_registry: &SymbolRegistry,
trait_id: SymbolId,
internal_ids: &HashSet<SymbolId>,
) -> Vec<CallerRewrite> {
let trait_crate = match symbol_registry.path(trait_id) {
Some(p) => p.crate_name().to_string(),
None => return Vec::new(),
};
let trait_name = match symbol_registry.path(trait_id) {
Some(p) => p.name().to_string(),
None => return Vec::new(),
};
if trait_crate.is_empty() || trait_name.is_empty() {
return Vec::new();
}
let trait_methods: Vec<String> = match ast_registry.get(trait_id) {
Some(PureItem::Trait(t)) => t
.items
.iter()
.filter_map(|it| match it {
PureTraitItem::Fn(f) => Some(f.name.clone()),
_ => None,
})
.collect(),
_ => Vec::new(),
};
let mut out = Vec::new();
let fn_ids: Vec<SymbolId> = symbol_registry
.iter()
.filter(|(id, path)| {
matches!(symbol_registry.kind(*id), Some(SymbolKind::Function))
&& path.crate_name() == trait_crate
&& !internal_ids.contains(id)
})
.map(|(id, _)| id)
.collect();
for caller_id in fn_ids {
if let Some(PureItem::Fn(f)) = ast_registry.get(caller_id) {
let patterns = classify_fn(f, &trait_name, &trait_methods);
if !patterns.is_empty() {
let is_public = matches!(f.vis, PureVis::Public);
out.push(CallerRewrite {
caller_id,
patterns,
is_public,
});
}
}
}
out
}
pub fn scan_same_crate_external_callers(
ctx: &AnalysisContext,
trait_id: SymbolId,
internal_ids: &HashSet<SymbolId>,
) -> Vec<CallerRewrite> {
scan_same_crate_external_callers_raw(&ctx.ast_registry, &ctx.registry, trait_id, internal_ids)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CascadeVerdict {
Safe,
Escapes,
}
pub fn walk_transitive_callers(
code_graph: &CodeGraphV2,
ast_registry: &ASTRegistry,
root_caller_id: SymbolId,
) -> CascadeVerdict {
let mut visited: HashSet<SymbolId> = HashSet::new();
let mut queue: VecDeque<SymbolId> = VecDeque::new();
queue.push_back(root_caller_id);
visited.insert(root_caller_id);
while let Some(id) = queue.pop_front() {
let callers: Vec<SymbolId> = code_graph.callers_of(id).collect();
if callers.is_empty() {
if let Some(PureItem::Fn(f)) = ast_registry.get(id) {
if matches!(f.vis, PureVis::Public) {
return CascadeVerdict::Escapes;
}
}
continue;
}
for caller in callers {
if visited.insert(caller) {
queue.push_back(caller);
}
}
}
CascadeVerdict::Safe
}