use brink_syntax_native::SyntaxKind as N;
use brink_syntax_native::ast::{self, AstNode as _};
use brink_syntax_native::{SyntaxNode, SyntaxToken};
use crate::hir::FileId;
use crate::provenance::NodeClass;
use crate::{Diagnostic, DiagnosticCode, Expr, LambdaBody, LambdaExpr, Name, Param};
use super::provenance::native_provenance;
fn diag(file: FileId, range: rowan::TextRange, code: DiagnosticCode) -> Diagnostic {
Diagnostic {
file,
range,
message: code.title().to_string(),
code,
}
}
fn name_from(tok: Option<SyntaxToken>) -> Option<Name> {
tok.map(|t| Name {
text: t.text().to_string(),
range: t.text_range(),
})
}
pub(super) fn lower_lambda(
file_id: FileId,
node: &SyntaxNode,
diags: &mut Vec<Diagnostic>,
) -> Expr {
let Some(lambda) = ast::LambdaExpr::cast(node.clone()) else {
diags.push(diag(file_id, node.text_range(), DiagnosticCode::E015));
return Expr::Null;
};
let Some(body_node) = lambda.body() else {
diags.push(diag(file_id, node.text_range(), DiagnosticCode::E015));
return Expr::Null;
};
let params = lower_lambda_params(lambda.params().as_ref());
let return_type = lambda
.return_annotation()
.as_ref()
.and_then(super::types::lower_type_annotation);
let body = if let Some(block) = ast::StmtBlock::cast(body_node.clone()) {
let tail = block.tail();
let stmts = super::control_flow::lower_stmt_block_stmts(file_id, &block, diags);
LambdaBody::Block {
stmts,
tail: tail.map(|t| Box::new(super::expr::lower_expr(file_id, &t, diags))),
}
} else {
LambdaBody::Expr(Box::new(super::expr::lower_expr(
file_id, &body_node, diags,
)))
};
check_capture_writes(file_id, &lambda, diags);
Expr::Lambda(Box::new(LambdaExpr {
ptr: native_provenance(file_id, NodeClass::Lambda, node),
params,
return_type,
body,
container_id: None,
}))
}
fn lower_lambda_params(params: Option<&ast::LambdaParams>) -> Vec<Param> {
params
.into_iter()
.flat_map(|row| row.params().collect::<Vec<_>>())
.filter_map(|p| {
name_from(p.name_token()).map(|name| Param {
name,
is_ref: false,
is_divert: false,
annotation: p
.type_annotation()
.as_ref()
.and_then(super::types::lower_type_annotation),
})
})
.collect()
}
fn check_capture_writes(file_id: FileId, lambda: &ast::LambdaExpr, diags: &mut Vec<Diagnostic>) {
let node = lambda.syntax();
let inner = inner_binders(lambda);
let outer = outer_binders(node);
for assign in node
.descendants()
.filter(|n| n.kind() == N::ASSIGN_STMT)
.filter(|n| nearest_lambda(n).as_ref() == Some(node))
{
let Some(root) = ast::AssignStmt::cast(assign.clone())
.and_then(|a| a.place())
.and_then(|p| p.segments().next())
else {
continue;
};
let text = root.text().to_string();
if inner.contains(&text) {
continue;
}
if outer.contains(&text) {
diags.push(diag(file_id, assign.text_range(), DiagnosticCode::E156));
}
}
}
fn nearest_lambda(node: &SyntaxNode) -> Option<SyntaxNode> {
node.ancestors()
.skip(1)
.find(|a| a.kind() == N::LAMBDA_EXPR)
}
fn inner_binders(lambda: &ast::LambdaExpr) -> Vec<String> {
let mut names: Vec<String> = lambda
.params()
.into_iter()
.flat_map(|row| row.params().collect::<Vec<_>>())
.filter_map(|p| p.name_token().map(|t| t.text().to_string()))
.collect();
for node in lambda.syntax().descendants() {
collect_binder_names(&node, &mut names);
}
names
}
fn outer_binders(lambda_node: &SyntaxNode) -> Vec<String> {
let mut names = Vec::new();
for ancestor in lambda_node.ancestors().skip(1) {
match ancestor.kind() {
N::FN_DECL | N::FLOW_DECL => {
names.extend(
ancestor
.children()
.filter(|c| c.kind() == N::PARAM_LIST)
.flat_map(|pl| pl.children().collect::<Vec<_>>())
.filter_map(ast::Param::cast)
.filter_map(|p| p.name_token().map(|t| t.text().to_string())),
);
}
N::LAMBDA_EXPR => {
names.extend(
ast::LambdaExpr::cast(ancestor.clone())
.and_then(|l| l.params())
.into_iter()
.flat_map(|row| row.params().collect::<Vec<_>>())
.filter_map(|p| p.name_token().map(|t| t.text().to_string())),
);
}
N::STMT_BLOCK
| N::IF_STMT
| N::WHILE_STMT
| N::UNTIL_STMT
| N::CONDITIONAL_BLOCK
| N::CHOICE_GUARD => {
for child in ancestor.children() {
collect_binder_names(&child, &mut names);
}
}
_ => collect_binder_names(&ancestor, &mut names),
}
}
names
}
fn collect_binder_names(node: &SyntaxNode, out: &mut Vec<String>) {
match node.kind() {
N::LET_STMT => out.extend(
ast::LetStmt::cast(node.clone())
.and_then(|s| s.name_token())
.map(|t| t.text().to_string()),
),
N::FOR_STMT => {
if let Some(f) = ast::ForStmt::cast(node.clone()) {
out.extend(f.name_token().map(|t| t.text().to_string()));
out.extend(f.val_name_token().map(|t| t.text().to_string()));
}
}
N::AS_BINDING => out.extend(
ast::AsBinding::cast(node.clone())
.and_then(|b| b.name_token())
.map(|t| t.text().to_string()),
),
_ => {}
}
}