use std::collections::HashMap;
use badness_parser::parser::conditional::{FlowWord, OpenerScan, Word};
use rowan::{TextRange, TextSize, WalkEvent};
use crate::ast::{AstNode, Command, Group};
use crate::semantic::define::is_definition_command;
use crate::syntax::{SyntaxKind, SyntaxNode};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct Frame {
id: u32,
branch: u32,
}
pub(crate) struct ConditionalIndex {
snapshots: Vec<(TextSize, Vec<Frame>)>,
}
pub(crate) fn guaranteed_before(earlier: &[Frame], later: &[Frame]) -> bool {
later.starts_with(earlier)
}
impl ConditionalIndex {
pub(crate) fn compute(root: &SyntaxNode) -> Self {
let mut stack: Vec<Frame> = Vec::new();
let mut next_id = 0u32;
let mut snapshots: Vec<(TextSize, Vec<Frame>)> = Vec::new();
let mut scan = OpenerScan::new();
let mut suppress_until = TextSize::from(0);
let mut branch_groups = HashMap::new();
let mut scopes: Vec<BranchScope> = Vec::new();
for event in root.preorder() {
let node = match event {
WalkEvent::Enter(node) => node,
WalkEvent::Leave(node) => {
if node.kind() == SyntaxKind::GROUP
&& scopes.last().is_some_and(|s| s.group == node.text_range())
{
let scope = scopes.pop().expect("matching branch scope");
stack = scope.path;
scan = scope.scan;
snapshot(&mut snapshots, node.text_range().end(), &stack);
}
continue;
}
};
let start = node.text_range().start();
if start < suppress_until {
continue;
}
if node.kind() == SyntaxKind::GROUP
&& let Some(frame) = branch_groups.remove(&start)
{
scopes.push(BranchScope {
group: node.text_range(),
path: stack.clone(),
scan: std::mem::take(&mut scan),
});
stack.push(frame);
snapshot(&mut snapshots, start, &stack);
}
let Some(command) = Command::cast(node.clone()) else {
continue;
};
let Some(name) = command.name() else {
continue;
};
if is_definition_command(&name) {
suppress_until = suppress_until.max(definition_span_end(&node));
continue;
}
let branch_depth = scopes.last().map_or(0, |s| s.path.len() + 1);
match scan.visit(&name) {
Word::Flow(FlowWord::Else | FlowWord::Or) => {
if stack.len() > branch_depth {
stack.last_mut().expect("open primitive").branch += 1;
snapshot(&mut snapshots, start, &stack);
}
}
Word::Flow(FlowWord::Fi) => {
if stack.len() > branch_depth {
stack.pop();
snapshot(&mut snapshots, start, &stack);
}
}
Word::Opens => {
stack.push(Frame {
id: next_id,
branch: 0,
});
next_id += 1;
snapshot(&mut snapshots, start, &stack);
}
Word::Inert => {
if let Some(branches) = macro_branches(&command, &name) {
suppress_until = branches[0].start();
for (branch, range) in branches.into_iter().enumerate() {
branch_groups.insert(
range.start(),
Frame {
id: next_id,
branch: branch as u32,
},
);
}
next_id += 1;
}
}
Word::Suppressed => {}
}
}
Self { snapshots }
}
pub(crate) fn path_at(&self, offset: usize) -> &[Frame] {
let offset = TextSize::from(offset as u32);
let i = self.snapshots.partition_point(|(s, _)| *s <= offset);
if i == 0 {
&[]
} else {
&self.snapshots[i - 1].1
}
}
}
struct BranchScope {
group: TextRange,
path: Vec<Frame>,
scan: OpenerScan,
}
fn snapshot(snapshots: &mut Vec<(TextSize, Vec<Frame>)>, at: TextSize, path: &[Frame]) {
if let Some((last_at, last_path)) = snapshots.last_mut()
&& *last_at == at
{
path.clone_into(last_path);
} else {
snapshots.push((at, path.to_vec()));
}
}
fn macro_branches(command: &Command, name: &str) -> Option<[TextRange; 2]> {
let tests = match name {
"ifthenelse" | "iflanguage" | "iftoggle" | "ifbool" | "ifboolexpr" | "ifboolexpe"
| "ifcsdef" | "ifcsundef" | "ifcsmacro" | "ifcsempty" | "ifcsvoid" | "ifstrempty"
| "ifblank" | "ifnumodd" | "IfFileExists" | "IfValueTF" | "IfNoValueTF" | "IfBooleanTF"
| "IfPackageLoadedTF" | "IfClassLoadedTF" | "@ifpackageloaded" | "@ifclassloaded"
| "@ifundefined" => 1,
"ifstrequal" | "ifcsstring" | "ifcsequal" | "ifnumequal" | "ifnumgreater" | "ifnumless"
| "ifdimequal" | "ifdimgreater" | "ifdimless" => 2,
"ifnumcomp" | "ifdimcomp" => 3,
_ => return None,
};
let mut arguments = command
.syntax()
.children_with_tokens()
.skip(1)
.filter(|el| !is_trivia(el.kind()));
let mut branches = [TextRange::default(); 2];
for i in 0..tests + 2 {
let element = arguments.next()?;
let group = Group::cast(element.into_node()?)?;
if group.syntax().last_child_or_token()?.kind() != SyntaxKind::R_BRACE {
return None;
}
if i >= tests {
branches[i - tests] = group.syntax().text_range();
}
}
Some(branches)
}
fn definition_span_end(command: &SyntaxNode) -> TextSize {
let own = command.text_range().end();
if command.children().any(|c| c.kind() == SyntaxKind::GROUP) {
return own;
}
match adjacent_sibling_command(command) {
Some(sibling) => own.max(sibling.text_range().end()),
None => own,
}
}
fn adjacent_sibling_command(command: &SyntaxNode) -> Option<SyntaxNode> {
let mut next = command.next_sibling_or_token();
while let Some(element) = next {
match element {
rowan::NodeOrToken::Token(token) if is_trivia(token.kind()) => {
next = token.next_sibling_or_token();
}
rowan::NodeOrToken::Node(node) if node.kind() == SyntaxKind::COMMAND => {
return Some(node);
}
_ => return None,
}
}
None
}
fn is_trivia(kind: SyntaxKind) -> bool {
matches!(
kind,
SyntaxKind::WHITESPACE | SyntaxKind::NEWLINE | SyntaxKind::COMMENT
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::parse;
fn index(src: &str) -> ConditionalIndex {
let root = SyntaxNode::new_root(parse(src).green);
ConditionalIndex::compute(&root)
}
fn offset(src: &str, needle: &str, n: usize) -> usize {
src.match_indices(needle)
.nth(n)
.map(|(i, _)| i)
.unwrap_or_else(|| panic!("occurrence {n} of {needle:?} in {src:?}"))
}
fn paths_at_loads<'a>(idx: &'a ConditionalIndex, src: &str) -> (&'a [Frame], &'a [Frame]) {
(
idx.path_at(offset(src, "\\usepackage", 0)),
idx.path_at(offset(src, "\\usepackage", 1)),
)
}
#[test]
fn if_else_branches_are_mutually_exclusive() {
let src = "\\iftrue\\usepackage{a}\\else\\usepackage{a}\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
}
#[test]
fn same_branch_guarantees_a_prior_occurrence() {
let src = "\\iftrue\\usepackage{a}\\usepackage{a}\\else x\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(guaranteed_before(a, b));
}
#[test]
fn ifcase_or_branches_are_pairwise_exclusive() {
let src = "\\ifcase 0 \\usepackage{a}\\or\\usepackage{a}\\or\\usepackage{a}\\fi\n";
let idx = index(src);
let a = idx.path_at(offset(src, "\\usepackage", 0));
let b = idx.path_at(offset(src, "\\usepackage", 1));
let c = idx.path_at(offset(src, "\\usepackage", 2));
assert!(!guaranteed_before(a, b));
assert!(!guaranteed_before(b, c));
assert!(!guaranteed_before(a, c));
}
#[test]
fn unconditional_prior_is_guaranteed_but_conditional_prior_is_not() {
let src = "\\iftrue\\usepackage{a}\\fi\n\\usepackage{a}\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(b.is_empty());
assert!(!guaranteed_before(a, b));
assert!(guaranteed_before(b, a));
assert!(guaranteed_before(b, b));
}
#[test]
fn nested_conditionals_compare_by_shared_frame() {
let src = "\\iftrue\\usepackage{a}\\else\\ifodd 1 \\usepackage{a}\\fi\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
}
#[test]
fn unknown_conditional_is_paired_and_trusted() {
let src = "\\ifmyflag\\usepackage{a}\\else\\usepackage{a}\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
}
#[test]
fn unknown_conditional_nested_in_known_resyncs() {
let src = "\\iftrue\\ifmyflag x\\fi\\usepackage{a}\\else\\usepackage{a}\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
}
#[test]
fn unknown_conditionals_else_bumps_its_own_frame() {
let src = "\\iftrue\\usepackage{a}\\ifmyflag\\else\\usepackage{a}\\fi\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(guaranteed_before(a, b));
}
#[test]
fn ifx_operands_open_no_frames() {
let src = "\\ifx\\ifabc\\ifxyz x\\fi\n done";
let idx = index(src);
assert_eq!(idx.path_at(offset(src, "x\\fi", 0)).len(), 1);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn ifdefined_operand_opens_no_frame() {
let src = "\\ifdefined\\iffalse x\\fi\n done";
let idx = index(src);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn textual_operands_do_not_eat_the_else() {
let src = "\\if ab\\usepackage{a}\\else\\usepackage{a}\\fi\n";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
}
#[test]
fn newif_declaration_opens_no_frame() {
let src = "\\newif\\ifmyflag\n done";
let idx = index(src);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn let_aliasing_opens_no_frame() {
let src = "\\let\\ifabc\\iftrue\n done";
let idx = index(src);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn ifcsname_material_is_skipped_and_pairs() {
let src = "\\ifcsname iftex\\endcsname\\usepackage{a}\\else\\usepackage{a}\\fi\n done";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert!(!guaranteed_before(a, b));
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn definition_bodies_change_no_state() {
let src = "\\iftrue x\\newcommand{\\x}{\\else}\\def\\stopit{\\fi} y\\fi\n done";
let idx = index(src);
assert_eq!(idx.path_at(offset(src, " y", 0)).len(), 1);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn macro_branches_do_not_escape_their_groups() {
let src = "\\ifthenelse{\\boolean{x}}{a}{b} $a \\iff b$\n done";
let idx = index(src);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
assert!(!idx.snapshots.is_empty());
}
#[test]
fn braced_macros_track_branch_slots_and_restore_the_outer_path() {
for head in [
"\\ifthenelse{\\boolean{x}}",
"\\iftoggle{x}",
"\\ifbool{x}",
"\\ifboolexpr{bool {x}}",
"\\iflanguage{english}",
"\\ifstrequal{a}{b}",
"\\ifnumcomp{1}{<}{2}",
"\\IfFileExists{example.tex}",
"\\IfNoValueTF{#1}",
] {
let src = format!(
"\\ifouter {head}% comment before the branches\n\
{{\\usepackage{{a}}\\usepackage{{a}}}}% comment between branches\n\
{{\\usepackage{{a}}}}\\usepackage{{a}}\\fi done"
);
let idx = index(&src);
let paths: Vec<_> = (0..4)
.map(|n| idx.path_at(offset(&src, "\\usepackage", n)))
.collect();
assert!(guaranteed_before(paths[0], paths[1]), "{src}");
assert!(!guaranteed_before(paths[0], paths[2]), "{src}");
assert!(!guaranteed_before(paths[0], paths[3]), "{src}");
assert_eq!(paths[0].len(), 2, "{src}");
assert_eq!(paths[3].len(), 1, "{src}");
assert!(idx.path_at(offset(&src, "done", 0)).is_empty(), "{src}");
assert!(idx.snapshots.windows(2).all(|pair| pair[0].0 < pair[1].0));
}
}
#[test]
fn nested_macro_branches_retain_their_enclosing_path() {
let src = "\\ifthenelse{x}{\\usepackage{a}\\iftoggle{y}{\\usepackage{a}}{\\usepackage{a}}}{\\usepackage{a}}";
let idx = index(src);
let path = |n| idx.path_at(offset(src, "\\usepackage", n));
assert!(guaranteed_before(path(0), path(1)));
assert!(guaranteed_before(path(0), path(2)));
assert!(!guaranteed_before(path(1), path(2)));
assert!(!guaranteed_before(path(0), path(3)));
}
#[test]
fn extra_attached_groups_are_not_branches() {
let src = "\\ifthenelse{x}{\\usepackage{a}}{}{\\usepackage{a}}";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert_eq!(a.len(), 1);
assert!(b.is_empty());
}
#[test]
fn malformed_macro_arguments_do_not_get_branch_positions() {
for src in [
"\\ifthenelse[x]{test}{\\usepackage{a}}{\\usepackage{a}}",
"\\ifthenelse{x}{\\usepackage{a}}",
"\\ifthenelse{x}{\\usepackage{a}}{unclosed",
] {
let idx = index(src);
assert!(
idx.path_at(offset(src, "\\usepackage", 0)).is_empty(),
"{src}"
);
}
}
#[test]
fn macro_predicates_and_definition_bodies_do_not_change_outer_state() {
let src = "\\ifouter\\usepackage{a}\\ifthenelse{\\iffoo\\else\\fi}{}{}\\newcommand{\\x}{\\ifthenelse{x}{\\fi}{\\else}}\\usepackage{a}\\fi done";
let idx = index(src);
let (a, b) = paths_at_loads(&idx, src);
assert_eq!(a, b);
assert_eq!(a.len(), 1);
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
#[test]
fn primitive_state_cannot_escape_a_macro_branch() {
for body in ["\\iffoo", "\\else\\or\\fi", "\\ifdefined", "\\ifcsname"] {
let src = format!(
"\\ifouter\\ifthenelse{{x}}{{{body}}}{{\\usepackage{{a}}\\ifinner\\usepackage{{a}}\\fi}}\\usepackage{{a}}\\fi done"
);
let idx = index(&src);
let a = idx.path_at(offset(&src, "\\usepackage", 0));
let b = idx.path_at(offset(&src, "\\usepackage", 1));
let c = idx.path_at(offset(&src, "\\usepackage", 2));
assert_eq!(a.len(), 2, "{src}");
assert_eq!(b.len(), 3, "{src}");
assert_eq!(c.len(), 1, "{src}");
assert!(guaranteed_before(a, b), "{src}");
assert!(!guaranteed_before(a, c), "{src}");
assert!(idx.path_at(offset(&src, "done", 0)).is_empty(), "{src}");
}
}
#[test]
fn stray_flow_words_are_no_ops() {
let src = "\\else\\or\\fi\n done";
let idx = index(src);
assert!(idx.snapshots.is_empty());
assert!(idx.path_at(offset(src, "done", 0)).is_empty());
}
}