luau-syntax 0.732.0

Luau lexer, parser, AST, CST, and source utilities
Documentation
pub(crate) use luau_syntax::allocator::AstArena;
pub(crate) use luau_syntax::ast::{
    AttributeKind, BinaryOp, Block, ClassMember, Expression, ExpressionKind,
    Statement as StatementNode, StatementTag, TableAccess, Type, TypeKind, TypeList, TypeOrPack,
    TypePackKind,
};
pub(crate) use luau_syntax::ast_names::{AstName, AstNameTable};
pub(crate) use luau_syntax::cst::CstNode;
pub(crate) use luau_syntax::location::{Location, Position};
pub(crate) use luau_syntax::parser::{
    Mode, ParseError, ParseErrors, ParseNodeResult, ParseOptions, ParseResult,
};

pub(crate) struct StatementKinds<'ast>(Vec<StatementNode<'ast>>);

impl<'ast> StatementKinds<'ast> {
    pub(crate) fn as_slice(&self) -> &[StatementNode<'ast>] {
        &self.0
    }

    pub(crate) fn exact<const N: usize>(&self) -> [StatementNode<'ast>; N] {
        self.0
            .clone()
            .try_into()
            .unwrap_or_else(|_| panic!("expected {N} statements, got {}", self.0.len()))
    }
}

impl<'ast> std::ops::Deref for StatementKinds<'ast> {
    type Target = [StatementNode<'ast>];

    fn deref(&self) -> &Self::Target {
        self.as_slice()
    }
}

macro_rules! pos {
    ($line:expr, $column:expr) => {
        Position::new($line, $column)
    };
}

macro_rules! loc {
    ($begin:expr, $end:expr) => {
        Location::new($begin, $end)
    };
}

pub(crate) fn with_parse<R>(
    source: &str,
    options: ParseOptions,
    f: impl for<'ast> FnOnce(Result<ParseResult<'ast>, ParseErrors>) -> R,
) -> R {
    let arena = AstArena::new();
    let mut names = AstNameTable::new(&arena);
    f(luau_syntax::parser::parse(
        source, &arena, &mut names, options,
    ))
}

pub(crate) fn with_parse_bytes<R>(
    source: &[u8],
    options: ParseOptions,
    f: impl for<'ast> FnOnce(Result<ParseResult<'ast>, ParseErrors>) -> R,
) -> R {
    let arena = AstArena::new();
    let mut names = AstNameTable::new(&arena);
    f(luau_syntax::parser::parse_bytes(
        source, &arena, &mut names, options,
    ))
}

pub(crate) fn with_parse_type<R>(
    source: &str,
    options: ParseOptions,
    f: impl for<'ast> FnOnce(Result<ParseNodeResult<'ast, Type<'ast>>, ParseErrors>) -> R,
) -> R {
    let arena = AstArena::new();
    let mut names = AstNameTable::new(&arena);
    f(luau_syntax::parser::parse_type(
        source.as_bytes(),
        &arena,
        &mut names,
        options,
    ))
}

pub(crate) fn parse_ok(source: &str) {
    with_parse(source, ParseOptions::default(), |result| {
        result.unwrap();
    })
}

pub(crate) fn with_parse_ok<R>(
    source: &str,
    f: impl for<'ast> FnOnce(ParseResult<'ast>) -> R,
) -> R {
    with_parse(source, ParseOptions::default(), |result| f(result.unwrap()))
}

pub(crate) fn parse_errors(source: &str) -> Vec<ParseError> {
    with_parse(source, ParseOptions::default(), |result| match result {
        Ok(result) => result.metadata.errors,
        Err(errors) => errors.into_errors(),
    })
}

pub(crate) fn parse_errors_with_options(source: &str, options: ParseOptions) -> Vec<ParseError> {
    with_parse(source, options, |result| match result {
        Ok(result) => result.metadata.errors,
        Err(errors) => errors.into_errors(),
    })
}

pub(crate) fn parse_errors_with_declarations(source: &str) -> Vec<ParseError> {
    parse_errors_with_options(
        source,
        ParseOptions::default().with_declaration_syntax(true),
    )
}

pub(crate) fn with_parse_ok_with_declarations<R>(
    source: &str,
    f: impl for<'ast> FnOnce(ParseResult<'ast>) -> R,
) -> R {
    with_parse(
        source,
        ParseOptions::default().with_declaration_syntax(true),
        |result| f(result.unwrap()),
    )
}

pub(crate) trait ParseErrorAssertions {
    fn assert_message(&self, expected: &str);
}

impl ParseErrorAssertions for ParseError {
    fn assert_message(&self, expected: &str) {
        assert_eq!(self.message, expected);
    }
}

pub(crate) trait ParseErrorsAssertions {
    fn first_error(&self) -> &ParseError;
    fn assert_first_message(&self, expected: &str);
    fn assert_first_message_starts_with(&self, expected: &str);
}

impl ParseErrorsAssertions for [ParseError] {
    fn first_error(&self) -> &ParseError {
        self.first().expect("expected parse error")
    }

    fn assert_first_message(&self, expected: &str) {
        self.first_error().assert_message(expected);
    }

    fn assert_first_message_starts_with(&self, expected: &str) {
        let message = &self.first_error().message;
        assert!(message.starts_with(expected), "{message}");
    }
}

pub(crate) fn statement_kinds<'ast>(statements: &[StatementNode<'ast>]) -> StatementKinds<'ast> {
    StatementKinds(statements.to_vec())
}

pub(crate) fn block_statement_kinds<'ast>(block: Block<'ast>) -> StatementKinds<'ast> {
    statement_kinds(block.as_slice())
}

pub(crate) fn with_first_local_annotation<R>(
    source: &str,
    f: impl for<'ast> FnOnce(Type<'ast>) -> R,
) -> R {
    with_parse(source, ParseOptions::default(), |result| {
        let result = result.unwrap();
        let statements = statement_kinds(result.root.as_slice());
        let [statement] = statements.exact();
        let local = statement.as_local().expect("expected local statement");
        f(*local.bindings[0]
            .annotation
            .as_ref()
            .expect("expected annotation"))
    })
}

pub(crate) fn return_strings(source: &str) -> Vec<Vec<u8>> {
    with_parse(source, ParseOptions::default(), |result| {
        let result = result.unwrap();
        let statements = statement_kinds(result.root.as_slice());
        let [statement] = statements.exact();
        let return_statement = statement.as_return().expect("expected return statement");
        return_statement
            .expressions
            .iter()
            .map(|expression| match expression.kind() {
                ExpressionKind::String { value, .. } => value.as_bytes().to_vec(),
                expression => panic!("expected string expression, got {expression:?}"),
            })
            .collect()
    })
}