use std::collections::{BTreeMap, BTreeSet};
use harn_lexer::{FixEdit, Span};
use harn_parser::{visit, Node, SNode, TypeExpr};
use super::signature_threading::{add_call_argument_edit, prepend_list_item};
#[derive(Debug, Default)]
pub(super) struct AliasWidening {
by_callable: BTreeMap<String, String>,
edits: BTreeMap<String, Vec<FixEdit>>,
}
impl AliasWidening {
pub(super) fn covers(&self, callable: &str) -> bool {
self.by_callable.contains_key(callable)
}
pub(super) fn edits_for(&self, callable: &str) -> &[FixEdit] {
self.by_callable
.get(callable)
.and_then(|alias| self.edits.get(alias))
.map_or(&[], Vec::as_slice)
}
pub(super) fn analyze(
program: &[SNode],
source: &str,
referenced_by_value: &BTreeSet<String>,
) -> Self {
let aliases = local_fn_aliases(program);
if aliases.is_empty() {
return Self::default();
}
let defaults = parameter_default_sites(program);
let identifier_spans = identifier_spans(program);
let mut by_callable: BTreeMap<String, String> = BTreeMap::new();
for callable in referenced_by_value {
let sites: Vec<&DefaultSite> = defaults
.iter()
.filter(|site| &site.callable == callable)
.collect();
if sites.is_empty() {
continue;
}
let site_spans: BTreeSet<usize> =
sites.iter().map(|site| site.value_span.start).collect();
let all_reads_are_defaults = identifier_spans
.get(callable)
.is_some_and(|spans| spans.iter().all(|span| site_spans.contains(&span.start)))
&& identifier_spans.get(callable).map_or(0, Vec::len) == sites.len();
if !all_reads_are_defaults {
continue;
}
let mut governing = None;
let consistent = sites.iter().all(|site| match &site.declared {
Some(TypeExpr::Named(alias)) if aliases.contains_key(alias) => {
let first = governing.get_or_insert(alias.clone());
first == alias
}
_ => false,
});
if let (true, Some(alias)) = (consistent, governing) {
by_callable.insert(callable.clone(), alias);
}
}
let candidates = by_callable.clone();
by_callable.retain(|_, alias| {
let Some(decl) = aliases.get(alias) else {
return false;
};
if decl.is_pub {
return false;
}
let governed: Vec<&DefaultSite> = defaults
.iter()
.filter(
|site| matches!(&site.declared, Some(TypeExpr::Named(name)) if name == alias),
)
.collect();
governed
.iter()
.all(|site| candidates.contains_key(&site.callable))
&& word_occurrences(source, alias) == governed.len() + 1
});
let edits = by_callable
.values()
.collect::<BTreeSet<_>>()
.into_iter()
.filter_map(|alias| {
let decl = aliases.get(alias)?;
let mut alias_edits = vec![alias_widening_edit(source, decl)?];
alias_edits.extend(dispatch_argument_edits(program, source, alias)?);
Some((alias.clone(), alias_edits))
})
.collect::<BTreeMap<_, _>>();
by_callable.retain(|_, alias| edits.contains_key(alias));
Self { by_callable, edits }
}
}
#[derive(Debug)]
struct AliasDecl {
span: Span,
is_pub: bool,
}
fn local_fn_aliases(program: &[SNode]) -> BTreeMap<String, AliasDecl> {
let mut aliases = BTreeMap::new();
for node in program {
let inner = match &node.node {
Node::AttributedDecl { inner, .. } => inner.as_ref(),
_ => node,
};
if let Node::TypeDecl {
name,
type_params,
type_expr: TypeExpr::FnType { .. },
is_pub,
} = &inner.node
{
if type_params.is_empty() {
aliases.insert(
name.clone(),
AliasDecl {
span: inner.span,
is_pub: *is_pub,
},
);
}
}
}
aliases
}
#[derive(Debug)]
struct DefaultSite {
callable: String,
declared: Option<TypeExpr>,
value_span: Span,
}
fn parameter_default_sites(program: &[SNode]) -> Vec<DefaultSite> {
let mut sites = Vec::new();
visit::walk_program(program, &mut |node| {
let params = match &node.node {
Node::FnDecl { params, .. }
| Node::ToolDecl { params, .. }
| Node::Pipeline { params, .. }
| Node::Closure { params, .. } => params,
_ => return,
};
for param in params {
let Some(default) = ¶m.default_value else {
continue;
};
let Node::Identifier(callable) = &default.node else {
continue;
};
sites.push(DefaultSite {
callable: callable.clone(),
declared: param.type_expr.clone(),
value_span: default.span,
});
}
});
sites
}
fn identifier_spans(program: &[SNode]) -> BTreeMap<String, Vec<Span>> {
let mut spans: BTreeMap<String, Vec<Span>> = BTreeMap::new();
visit::walk_program(program, &mut |node| {
if let Node::Identifier(name) = &node.node {
spans.entry(name.clone()).or_default().push(node.span);
}
});
spans
}
fn word_occurrences(source: &str, name: &str) -> usize {
let is_word = |byte: u8| byte.is_ascii_alphanumeric() || byte == b'_';
let bytes = source.as_bytes();
source
.match_indices(name)
.filter(|(index, _)| {
let before_ok = *index == 0 || !is_word(bytes[index - 1]);
let after = index + name.len();
let after_ok = after >= bytes.len() || !is_word(bytes[after]);
before_ok && after_ok
})
.count()
}
fn alias_widening_edit(source: &str, decl: &AliasDecl) -> Option<FixEdit> {
let region = source.get(decl.span.start..decl.span.end)?;
let fn_at = region.find("fn(")?;
let open_paren = decl.span.start + fn_at + 3;
let close_paren = region[fn_at + 3..].find(')')? + fn_at + 3 + decl.span.start;
let has_params = !source.get(open_paren..close_paren)?.trim().is_empty();
Some(FixEdit {
span: Span::with_offsets(open_paren, open_paren, decl.span.line, decl.span.column),
replacement: prepend_list_item(source, open_paren, "Harness", has_params),
})
}
fn dispatch_argument_edits(program: &[SNode], source: &str, alias: &str) -> Option<Vec<FixEdit>> {
let mut edits = Vec::new();
let mut refused = false;
visit::walk_program(program, &mut |node| {
if refused {
return;
}
let (params, body) = match &node.node {
Node::FnDecl { params, body, .. }
| Node::ToolDecl { params, body, .. }
| Node::Pipeline { params, body, .. }
| Node::Closure { params, body, .. } => (params, body),
_ => return,
};
let bindings: Vec<&str> = params
.iter()
.filter(
|param| matches!(¶m.type_expr, Some(TypeExpr::Named(name)) if name == alias),
)
.map(|param| param.name.as_str())
.collect();
if bindings.is_empty() {
return;
}
let Some(harness) = params
.iter()
.find(|param| {
matches!(¶m.type_expr, Some(TypeExpr::Named(name)) if name == "Harness")
})
.map(|param| param.name.clone())
else {
refused = true;
return;
};
for binding in bindings {
let mut call_spans = BTreeSet::new();
let mut reads = Vec::new();
visit::walk_program(body, &mut |child| {
match &child.node {
Node::FunctionCall { name, .. } if name == binding => {
edits.push((child.span, harness.clone()));
}
Node::ValueCall { callee, .. } => {
if matches!(&callee.node, Node::Identifier(name) if name == binding) {
call_spans.insert(callee.span.start);
edits.push((child.span, harness.clone()));
}
}
Node::Identifier(name) if name == binding => reads.push(child.span),
_ => {}
}
});
if reads.iter().any(|span| !call_spans.contains(&span.start)) {
refused = true;
return;
}
}
});
if refused {
return None;
}
edits
.into_iter()
.map(|(span, harness)| add_call_argument_edit(source, &span, &harness))
.collect()
}