use proc_macro2::TokenStream;
use quote::{ToTokens, quote};
use syn::parse2;
use crate::{
character_classes::CharacterClasses, dfa::Dfa, nfa::Nfa, scanner_data::ScannerData,
scanner_mode::ScannerMode,
};
pub fn generate(input: TokenStream) -> TokenStream {
let scanner_data: ScannerData = parse2(input).expect("Failed to parse input");
let scanner_modes: Vec<ScannerMode> = scanner_data
.build_scanner_modes()
.expect("Failed to build scanner modes");
let mut nfas = scanner_modes
.iter()
.map(|mode| {
Nfa::build_from_patterns(&mode.patterns).expect("Failed to build NFA for pattern")
})
.collect::<Vec<_>>();
let mut character_classes = CharacterClasses::new();
for nfa in &nfas {
nfa.collect_character_classes(&mut character_classes)
}
character_classes.create_disjoint_character_classes();
for nfa in &mut nfas {
nfa.convert_to_disjoint_character_classes(&character_classes);
}
let dfas = nfas
.into_iter()
.try_fold(Vec::new(), |mut acc, nfa| -> Result<Vec<Dfa>, syn::Error> {
let dfa = Dfa::try_from(&nfa).map_err(|e| {
syn::Error::new(
proc_macro2::Span::call_site(),
format!("Failed to convert NFA to DFA: {}", e),
)
})?;
acc.push(dfa);
Ok(acc)
})
.expect("Failed to convert NFAs to DFAs");
let module_name = to_snake_case(&scanner_data.name);
let module_name_ident = syn::Ident::new(&module_name, proc_macro2::Span::call_site());
let scanner_name = syn::Ident::new(&scanner_data.name, proc_macro2::Span::call_site());
let match_function_code = character_classes.generate("match_function");
let modes = scanner_modes.into_iter().enumerate().map(|(index, mode)| {
let transitions = mode.transitions.iter().map(|(token_type, new_mode_index)| {
quote! {
(#token_type, #new_mode_index)
}
});
let states = dfas[index]
.states
.iter()
.map(|state| state.to_token_stream());
let mode_name = mode.name;
quote! {
ScannerMode {
name: #mode_name,
transitions: &[#(#transitions),*],
dfa: Dfa { states: &[#(#states),*] }
}
}
});
let output = quote! {
pub mod #module_name_ident {
use scnr2::{AcceptData, Dfa, DfaState, DfaTransition, Lookahead, ScannerMode, ScannerImpl};
pub const MODES: &[ScannerMode] = &[
#(
#modes
),*
];
pub struct #scanner_name {
scanner_impl: ScannerImpl,
}
impl #scanner_name {
pub fn new() -> Self {
#scanner_name {
scanner_impl: ScannerImpl::new(MODES),
}
}
#match_function_code
pub fn find_matches<'a, F>(
&'a self,
haystack: &'a str,
offset: usize,
match_function: &'static F,
) -> scnr2::internals::find_matches::FindMatches<'a, F>
where
F: Fn(char) -> Option<usize> + 'static,
{
self.scanner_impl.find_matches(haystack, offset, match_function)
}
pub fn find_matches_with_position<'a, F>(
&'a self,
haystack: &'a str,
offset: usize,
match_function: &'static F,
) -> scnr2::internals::find_matches::FindMatchesWithPosition<'a, F>
where
F: Fn(char) -> Option<usize> + 'static,
{
self.scanner_impl.find_matches_with_position(haystack, offset, match_function)
}
}
}
};
output
}
fn to_snake_case(s: &str) -> String {
let mut result = String::new();
let chars = s.chars().peekable();
for c in chars {
if c.is_uppercase() {
if !result.is_empty() && !result.ends_with('_') {
result.push('_');
}
result.push(c.to_lowercase().next().unwrap());
} else {
result.push(c);
}
}
result
}
#[cfg(test)]
mod tests {
use std::io::Write;
use super::*;
use crate::Result;
use std::path::Path;
use std::process::Command;
fn try_format(path_to_file: &Path) -> Result<()> {
Command::new("rustfmt")
.args([path_to_file])
.status()
.map(|_| ())
.map_err(|e| {
std::io::Error::new(e.kind(), format!("Failed to format file: {}", e)).into()
})
}
#[test]
fn test_generate() {
let input = quote::quote! {
TestScanner {
mode INITIAL {
token r"\r\n|\r|\n" => 1;
token r"[\s--\r\n]+" => 2;
token r"//.*(\r\n|\r|\n)?" => 3;
token r"/\*([^*]|\*[^/])*\*/" => 4;
token r#"""# => 8;
token r"Hello" => 9;
token r"World" => 10;
token r"World" => 11 followed by r"!";
token r"!" => 12 not followed by r"!";
token r"[a-zA-Z_]\w*" => 13;
token r"." => 14;
transition 8 => STRING;
}
mode STRING {
token r#"\\[\"\\bfnt]"# => 5;
token r"\\[\s--\r\n]*\r?\n" => 6;
token r#"[^\"\\]+"# => 7;
token r#"""# => 8;
token r"." => 14;
transition 8 => INITIAL;
}
}
};
let code = generate(input).to_string();
let mut temp_file =
tempfile::NamedTempFile::new().expect("Failed to create temporary file");
temp_file
.write_all(code.as_bytes())
.expect("Failed to write to temporary file");
println!("Temporary file created at: {:?}", temp_file.path());
try_format(temp_file.path()).expect("Failed to format the temporary file");
let formatted_code = std::fs::read_to_string(temp_file.path())
.expect("Failed to read the formatted temporary file")
.replace("\r\n", "\n");
let expected_code = std::fs::read_to_string("data/expected_generated_code.rs")
.expect("Failed to read the expected code file")
.replace("\r\n", "\n");
assert_eq!(formatted_code, expected_code);
}
}