use std::collections::HashMap;
use tree_sitter::Node;
use crate::cop::shared::call_method_name;
use crate::cop::{Cop, CopConfig};
use crate::correction::Correction;
use crate::diagnostic::Diagnostic;
use crate::parse::source::SourceFile;
fn is_call_kind(kind: &str) -> bool {
matches!(kind, "call" | "command" | "command_call")
}
pub struct BatchedWalker<'a> {
cops: Vec<&'a dyn Cop>,
configs: Vec<&'a CopConfig>,
kinds: HashMap<&'static str, Vec<usize>>,
call_all: Vec<usize>,
call_by_name: HashMap<&'static [u8], Vec<usize>>,
}
impl<'a> BatchedWalker<'a> {
pub fn new(cops: Vec<&'a dyn Cop>, configs: Vec<&'a CopConfig>) -> Self {
let mut kinds: HashMap<&'static str, Vec<usize>> = HashMap::new();
let mut call_all = Vec::new();
let mut call_by_name: HashMap<&'static [u8], Vec<usize>> = HashMap::new();
for (i, cop) in cops.iter().enumerate() {
register_cop(i, *cop, &mut kinds, &mut call_all, &mut call_by_name);
}
Self {
cops,
configs,
kinds,
call_all,
call_by_name,
}
}
pub fn walk(
&self,
source: &SourceFile,
root: Node<'_>,
diagnostics: &mut Vec<Diagnostic>,
mut corrections: Option<&mut Vec<Correction>>,
) {
self.visit(source, root, diagnostics, &mut corrections);
}
fn visit(
&self,
source: &SourceFile,
node: Node<'_>,
diagnostics: &mut Vec<Diagnostic>,
corrections: &mut Option<&mut Vec<Correction>>,
) {
if node.is_named() {
self.dispatch(source, node, diagnostics, corrections);
}
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
self.visit(source, child, diagnostics, corrections);
}
}
fn dispatch(
&self,
source: &SourceFile,
node: Node<'_>,
diagnostics: &mut Vec<Diagnostic>,
corrections: &mut Option<&mut Vec<Correction>>,
) {
let kind = node.kind();
if is_call_kind(kind) {
self.dispatch_call(source, node, diagnostics, corrections);
return;
}
if let Some(idxs) = self.kinds.get(kind) {
for &i in idxs {
self.run_cop(i, source, node, diagnostics, corrections);
}
}
}
fn dispatch_call(
&self,
source: &SourceFile,
node: Node<'_>,
diagnostics: &mut Vec<Diagnostic>,
corrections: &mut Option<&mut Vec<Correction>>,
) {
for &i in &self.call_all {
self.run_cop(i, source, node, diagnostics, corrections);
}
if let Some(name) = call_method_name(source, node) {
if let Some(idxs) = self.call_by_name.get(name) {
for &i in idxs {
self.run_cop(i, source, node, diagnostics, corrections);
}
}
}
}
fn run_cop(
&self,
i: usize,
source: &SourceFile,
node: Node<'_>,
diagnostics: &mut Vec<Diagnostic>,
corrections: &mut Option<&mut Vec<Correction>>,
) {
self.cops[i].check_node(
source,
node,
self.configs[i],
diagnostics,
corrections.as_deref_mut(),
);
}
}
fn register_cop(
i: usize,
cop: &dyn Cop,
kinds: &mut HashMap<&'static str, Vec<usize>>,
call_all: &mut Vec<usize>,
call_by_name: &mut HashMap<&'static [u8], Vec<usize>>,
) {
let mut saw_call = false;
for &k in cop.interested_node_kinds() {
if is_call_kind(k) {
if saw_call {
continue;
}
saw_call = true;
let names = cop.interested_call_names();
if names.is_empty() {
call_all.push(i);
} else {
for &n in names {
call_by_name.entry(n).or_default().push(i);
}
}
} else {
kinds.entry(k).or_default().push(i);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CountingCop {
kinds: &'static [&'static str],
names: &'static [&'static [u8]],
hits: AtomicUsize,
}
impl Cop for CountingCop {
fn name(&self) -> &'static str {
"Test/Counting"
}
fn interested_node_kinds(&self) -> &'static [&'static str] {
self.kinds
}
fn interested_call_names(&self) -> &'static [&'static [u8]] {
self.names
}
fn check_node(
&self,
_source: &SourceFile,
_node: Node<'_>,
_config: &CopConfig,
_diagnostics: &mut Vec<Diagnostic>,
_corrections: Option<&mut Vec<Correction>>,
) {
self.hits.fetch_add(1, Ordering::Relaxed);
}
}
fn counting(names: &'static [&'static [u8]]) -> CountingCop {
CountingCop {
kinds: &["call", "command"],
names,
hits: AtomicUsize::new(0),
}
}
fn hit_count(cop: &CountingCop) -> usize {
cop.hits.load(Ordering::Relaxed)
}
fn walk_sample(gated: &CountingCop, ungated: &CountingCop) {
let cfg = CopConfig::default();
let walker = BatchedWalker::new(vec![gated, ungated], vec![&cfg, &cfg]);
let sf = SourceFile::from_bytes(PathBuf::from("t.rb"), b"foo.each {}; bar.map {}".to_vec());
let tree = crate::parse::parse_ruby(&sf).unwrap();
walker.walk(&sf, tree.root_node(), &mut Vec::new(), None);
}
#[test]
fn call_name_gate_skips_non_matching_methods() {
let gated = counting(&[b"map"]);
let ungated = counting(&[]);
walk_sample(&gated, &ungated);
assert_eq!(hit_count(&gated), 1);
assert_eq!(hit_count(&ungated), 2);
}
}