use nu_protocol::{
Type,
ast::{Block, Expr, Pipeline, PipelineElement, Traverse},
};
use crate::{ast::expression::ExpressionExt, context::LintContext, violation::Detection};
pub struct ConversionSpec<'a> {
pub matches_command: &'a dyn Fn(&str) -> bool,
pub matches_type: &'a dyn Fn(&Type) -> bool,
}
pub fn check_all_pipelines<FixData>(
context: &LintContext,
spec: &ConversionSpec,
create_violation: impl Fn(&Type, &str, &PipelineElement, &PipelineElement) -> (Detection, FixData)
+ Copy,
) -> Vec<(Detection, FixData)> {
let mut violations = Vec::new();
check_block_recursive(
context.ast,
context,
spec,
create_violation,
&mut violations,
);
violations
}
fn check_block_recursive<FixData>(
block: &Block,
context: &LintContext,
spec: &ConversionSpec,
create_violation: impl Fn(&Type, &str, &PipelineElement, &PipelineElement) -> (Detection, FixData)
+ Copy,
violations: &mut Vec<(Detection, FixData)>,
) {
for pipeline in &block.pipelines {
violations.extend(check_pipeline(pipeline, context, spec, create_violation));
}
let mut nested_block_ids = Vec::new();
for pipeline in &block.pipelines {
for element in &pipeline.elements {
element.expr.flat_map(
context.working_set,
&|expr| match &expr.expr {
Expr::Block(id)
| Expr::RowCondition(id)
| Expr::Closure(id)
| Expr::Subexpression(id) => vec![*id],
_ => vec![],
},
&mut nested_block_ids,
);
}
}
for &block_id in &nested_block_ids {
check_block_recursive(
context.working_set.get_block(block_id),
context,
spec,
create_violation,
violations,
);
}
}
fn check_pipeline<FixData>(
pipeline: &Pipeline,
context: &LintContext,
spec: &ConversionSpec,
create_violation: impl Fn(&Type, &str, &PipelineElement, &PipelineElement) -> (Detection, FixData),
) -> Vec<(Detection, FixData)> {
pipeline
.elements
.windows(2)
.filter_map(|window| {
check_pipeline_pair(&window[0], &window[1], context, spec, &create_violation)
})
.collect()
}
fn check_pipeline_pair<FixData>(
left: &PipelineElement,
right: &PipelineElement,
context: &LintContext,
spec: &ConversionSpec,
create_violation: &impl Fn(&Type, &str, &PipelineElement, &PipelineElement) -> (Detection, FixData),
) -> Option<(Detection, FixData)> {
let Expr::ExternalCall(head, _args) = &right.expr.expr else {
return None;
};
let cmd_name = context.expr_text(head);
let clean_name = cmd_name.trim_start_matches('^');
if !(spec.matches_command)(clean_name) {
return None;
}
let input_type = left.expr.infer_output_type(context)?;
(spec.matches_type)(&input_type).then(|| create_violation(&input_type, cmd_name, left, right))
}