diff --git a/Cargo.toml b/Cargo.toml index b9fef2d..c66c04e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,9 +45,11 @@ members = [ "src/std/ndslice", "src/std/range_rotation", "src/utils/librstest", + "src/utils/loop_xform", "src/utils/rstest", "tests/bitstream/bitstreams", "tests/metadata/camerasxml_parser", + "tests/utils/loop_xform", ] [workspace.dependencies] diff --git a/src/utils/loop_xform/Cargo.toml b/src/utils/loop_xform/Cargo.toml new file mode 100644 index 0000000..16093f3 --- /dev/null +++ b/src/utils/loop_xform/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "rawspeed-utils-loop_xform" +version.workspace = true +authors.workspace = true +edition.workspace = true +rust-version.workspace = true +documentation.workspace = true +homepage.workspace = true +repository.workspace = true +license.workspace = true + +[lints] +workspace = true + +[dependencies] +proc-macro2 = { version = "1.0", default-features = false, features = [] } +syn = { version = "2.0", default-features = false, features = ["proc-macro", "parsing", "full", "visit-mut", "printing"] } +quote = { version = "1.0", default-features = false, features = [] } + +[lib] +proc-macro = true +path = "mod.rs" diff --git a/src/utils/loop_xform/mod.rs b/src/utils/loop_xform/mod.rs new file mode 100644 index 0000000..de685d9 --- /dev/null +++ b/src/utils/loop_xform/mod.rs @@ -0,0 +1,135 @@ +use quote::ToTokens; + +mod kw { + syn::custom_keyword!(runtime); + syn::custom_keyword!(with_remainder); +} + +enum Loop { + ExprForLoop(syn::ExprForLoop), + ExprWhile(syn::ExprWhile), + ExprLoop(syn::ExprLoop), +} + +impl From for Loop { + fn from(v: syn::ExprForLoop) -> Self { + Self::ExprForLoop(v) + } +} + +impl From for Loop { + fn from(v: syn::ExprWhile) -> Self { + Self::ExprWhile(v) + } +} + +impl From for Loop { + fn from(v: syn::ExprLoop) -> Self { + Self::ExprLoop(v) + } +} + +impl ToTokens for Loop { + fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) { + match self { + Self::ExprForLoop(expr_for_loop) => expr_for_loop.to_tokens(tokens), + Self::ExprWhile(expr_while) => expr_while.to_tokens(tokens), + Self::ExprLoop(expr_loop) => expr_loop.to_tokens(tokens), + } + } +} + +#[derive(PartialEq, Eq, Debug)] +enum UnrollMethod { + Runtime, + WithRemainder, +} + +struct LoopUnrollParams { + pub unroll_method: UnrollMethod, + pub unroll_factor: usize, +} + +struct LoopPeelParams { + pub peel_count: usize, +} + +enum LoopXFormParams { + LoopUnroll(LoopUnrollParams), + LoopPeel(LoopPeelParams), +} + +impl From for LoopXFormParams { + fn from(value: LoopUnrollParams) -> Self { + Self::LoopUnroll(value) + } +} + +impl From for LoopXFormParams { + fn from(value: LoopPeelParams) -> Self { + Self::LoopPeel(value) + } +} + +struct LoopXFormConf { + pub params: Params, + pub the_loop: Loop, + pub rest_of_tokenstream: proc_macro2::TokenStream, +} + +enum Item { + LoopUnrollAttr(LoopXFormConf), + LoopPeelAttr(LoopXFormConf), +} + +impl Item { + const fn new( + params: LoopXFormParams, + the_loop: Loop, + rest_of_tokenstream: proc_macro2::TokenStream, + ) -> Self { + match params { + LoopXFormParams::LoopUnroll(loop_unroll_params) => { + Item::LoopUnrollAttr(LoopXFormConf { + params: loop_unroll_params, + the_loop, + rest_of_tokenstream, + }) + } + LoopXFormParams::LoopPeel(loop_peel_params) => { + Item::LoopPeelAttr(LoopXFormConf { + params: loop_peel_params, + the_loop, + rest_of_tokenstream, + }) + } + } + } + + fn the_loop(&self) -> &Loop { + match self { + Self::LoopUnrollAttr(loop_xform_conf) => &loop_xform_conf.the_loop, + Self::LoopPeelAttr(loop_xform_conf) => &loop_xform_conf.the_loop, + } + } +} + +#[proc_macro] +#[inline(never)] +pub fn enable_loop_xforms( + tokens: proc_macro::TokenStream, +) -> proc_macro::TokenStream { + let input = syn::parse_macro_input!(tokens as Item); + if cfg!(clippy) { + use quote::ToTokens as _; + return input.the_loop().to_token_stream().into(); + } + match input { + Item::LoopUnrollAttr(c) => transform::perform_loop_unroll(c).into(), + Item::LoopPeelAttr(c) => transform::perform_loop_peel(c).into(), + _ => unreachable!(), + } +} + +mod parse; +mod transform; diff --git a/src/utils/loop_xform/parse/mod.rs b/src/utils/loop_xform/parse/mod.rs new file mode 100644 index 0000000..3af7149 --- /dev/null +++ b/src/utils/loop_xform/parse/mod.rs @@ -0,0 +1,230 @@ +use crate::{Loop, LoopPeelParams, LoopUnrollParams, LoopXFormParams}; + +use super::Item; +use super::UnrollMethod; +use super::kw; +use syn::LitInt; +use syn::Result; +use syn::parenthesized; +use syn::parse::Parse; +use syn::parse::ParseStream; +use syn::{Attribute, ExprForLoop}; + +impl Parse for UnrollMethod { + fn parse(input: ParseStream<'_>) -> Result { + let lookahead = input.lookahead1(); + if lookahead.peek(kw::runtime) { + input.parse::()?; + Ok(UnrollMethod::Runtime) + } else if lookahead.peek(kw::with_remainder) { + input.parse::()?; + Ok(UnrollMethod::WithRemainder) + } else { + Err(lookahead.error()) + } + } +} + +fn parse_method( + meta: &syn::meta::ParseNestedMeta<'_>, + unroll_method: &mut Option, +) -> Result<()> { + assert!(meta.path.is_ident("method")); + + if unroll_method.is_some() { + return Err( + meta.error("only a single unroll method shall be specified") + ); + } + + let content; + parenthesized!(content in meta.input); + let head = content.fork(); + match content.parse::() { + Ok(m) => *unroll_method = Some(m), + Err(_) => { + return Err(head.error("expected valid unroll method")); + } + } + if !content.is_empty() { + return Err(syn::Error::new_spanned( + content.parse::()?, + "unexpected garbage in unroll method argument", + )); + } + Ok(()) +} + +fn parse_factor( + meta: &syn::meta::ParseNestedMeta<'_>, + unroll_factor: &mut Option, +) -> Result<()> { + assert!(meta.path.is_ident("factor")); + + if unroll_factor.is_some() { + return Err( + meta.error("only a single unroll factor shall be specified") + ); + } + let content; + parenthesized!(content in meta.input); + let lit: LitInt = content.parse()?; + if !lit.suffix().is_empty() { + return Err(syn::Error::new_spanned( + lit, + "unroll factor should not have any suffix", + )); + } + if !content.is_empty() { + return Err(syn::Error::new_spanned( + content.parse::()?, + "unexpected garbage in unroll factor argument", + )); + } + let n: usize = lit.base10_parse()?; + if n < 1 { + return Err(meta.error("Unroll factor can not be zero")); + } + *unroll_factor = Some(n); + Ok(()) +} + +fn parse_unroll_attr(attr: &Attribute) -> Option> { + if !attr.path().is_ident("loop_unroll") { + return None; + } + + let mut unroll_method: Option = None; + let mut unroll_factor: Option = None; + if let Err(err) = attr.parse_nested_meta(|meta| { + if meta.path.is_ident("method") { + return parse_method(&meta, &mut unroll_method); + } + if meta.path.is_ident("factor") { + return parse_factor(&meta, &mut unroll_factor); + } + Err(meta.error( + "unrecognized parameter, expected `method(...)` and `factor(..)`", + )) + }) { + return Some(Err(err)); + }; + + let Some(unroll_method) = unroll_method else { + return Some(Err(syn::Error::new_spanned( + attr, + "The attribute must specify unroll `method`", + ))); + }; + + let Some(unroll_factor) = unroll_factor else { + return Some(Err(syn::Error::new_spanned( + attr, + "The attribute must specify `factor`", + ))); + }; + + Some(Ok(LoopUnrollParams { + unroll_method, + unroll_factor, + })) +} + +struct PeelCount(usize); +impl syn::parse::Parse for PeelCount { + fn parse(input: ParseStream) -> Result { + let Ok(lit) = input.parse::() else { + return Err(syn::Error::new_spanned( + input.parse::()?, + "The attribute must specify peel count", + )); + }; + if !lit.suffix().is_empty() { + return Err(syn::Error::new_spanned( + lit, + "Peel count should not have any suffix", + )); + } + if !input.is_empty() { + return Err(syn::Error::new_spanned( + input.parse::()?, + "unexpected garbage in peel count argument", + )); + } + let n: usize = lit.base10_parse()?; + if n < 1 { + return Err(input.error("Peel count can not be zero")); + } + Ok(Self(n)) + } +} + +fn parse_peel_attr(attr: &Attribute) -> Option> { + if !attr.path().is_ident("loop_peel") { + return None; + } + + match attr.parse_args::() { + Ok(PeelCount(peel_count)) => Some(Ok(LoopPeelParams { peel_count })), + Err(err) => Some(Err(err)), + } +} + +fn parse_attr(attr: &Attribute) -> Result { + if let Some(p) = parse_unroll_attr(attr) { + return Ok(p?.into()); + } + + if let Some(p) = parse_peel_attr(attr) { + return Ok(p?.into()); + } + + Err(syn::Error::new_spanned( + attr, + "`loop_unroll` or `loop_peel` attribute expected", + )) +} + +impl Parse for Loop { + fn parse(input: ParseStream<'_>) -> Result { + let origin = input.fork(); + match input.parse::() { + Ok(syn::Expr::ForLoop(expr_for_loop)) => Ok(expr_for_loop.into()), + Ok(syn::Expr::Loop(expr_loop)) => Ok(expr_loop.into()), + Ok(syn::Expr::While(expr_while)) => Ok(expr_while.into()), + _ => Err(origin.error("expected some kind of loop")), + } + } +} + +impl Parse for Item { + fn parse(input: ParseStream<'_>) -> Result { + let attrs = input.call(Attribute::parse_outer)?; + + let params = if let Some(attr) = attrs.first() { + if let Some(ea) = attrs.get(1) { + return Err(syn::Error::new_spanned( + ea, + "There should only be a single attribute", + )); + } + + parse_attr(attr)? + } else { + return Err(syn::Error::new_spanned( + input.parse::()?, + "There must be an attribute", + )); + }; + + let the_loop = input.parse()?; + let remainder = input.parse()?; + assert!(input.is_empty()); + + Ok(Self::new(params, the_loop, remainder)) + } +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod tests; diff --git a/src/utils/loop_xform/parse/tests/mod.rs b/src/utils/loop_xform/parse/tests/mod.rs new file mode 100644 index 0000000..b0df5f7 --- /dev/null +++ b/src/utils/loop_xform/parse/tests/mod.rs @@ -0,0 +1,11 @@ +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod unroll_runtime; + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod peel; + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod unroll_with_remainder; diff --git a/src/utils/loop_xform/parse/tests/peel.rs b/src/utils/loop_xform/parse/tests/peel.rs new file mode 100644 index 0000000..63eae16 --- /dev/null +++ b/src/utils/loop_xform/parse/tests/peel.rs @@ -0,0 +1,163 @@ +use crate::Item; +use crate::UnrollMethod; +use quote::ToTokens as _; +use quote::quote; + +#[test] +#[should_panic(expected = "There must be an attribute")] +fn t0_test() { + let tokens = quote! {}; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "attribute expected")] +fn t1_test() { + let tokens = quote! { + #[attr] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "expected attribute arguments in parentheses: #[loop_peel(...)]" +)] +fn t2_test() { + let tokens = quote! { + #[loop_peel] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "The attribute must specify peel count")] +fn t3_test() { + let tokens = quote! { + #[loop_peel()] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "unexpected end of input, expected some kind of loop" +)] +fn t4_test() { + let tokens = quote! { + #[loop_peel(42)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "Peel count should not have any suffix")] +fn t5_test() { + let tokens = quote! { + #[loop_peel(42u16)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "unexpected end of input, Peel count can not be zero" +)] +fn t6_test() { + let tokens = quote! { + #[loop_peel(0)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "unexpected garbage in peel count argument")] +fn t7_test() { + let tokens = quote! { + #[loop_peel(1,)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "There should only be a single attribute")] +fn t8_test() { + let tokens = quote! { + #[loop_peel(1)] + #[loop_peel(1)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +fn good_test() { + let tokens = quote! { + #[loop_peel(42)] + 'loop_label: for elt in iter { body } rest + }; + match syn::parse2::(tokens).unwrap() { + Item::LoopPeelAttr(attr) => { + assert_eq!(attr.params.peel_count, 42); + let crate::Loop::ExprForLoop(for_loop) = attr.the_loop else { + unreachable!() + }; + assert_eq!( + for_loop.label.to_token_stream().to_string(), + ("'loop_label :") + ); + assert_eq!(for_loop.pat.to_token_stream().to_string(), ("elt")); + assert_eq!(for_loop.expr.to_token_stream().to_string(), ("iter")); + assert_eq!( + for_loop.body.to_token_stream().to_string(), + ("{ body }") + ); + assert_eq!(attr.rest_of_tokenstream.to_string(), "rest"); + } + Item::LoopUnrollAttr(_) => unreachable!(), + } +} diff --git a/src/utils/loop_xform/parse/tests/unroll_runtime.rs b/src/utils/loop_xform/parse/tests/unroll_runtime.rs new file mode 100644 index 0000000..074f80f --- /dev/null +++ b/src/utils/loop_xform/parse/tests/unroll_runtime.rs @@ -0,0 +1,264 @@ +use crate::Item; +use crate::UnrollMethod; +use quote::ToTokens as _; +use quote::quote; + +#[test] +#[should_panic(expected = "There must be an attribute")] +fn t0_test() { + let tokens = quote! {}; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "attribute expected")] +fn t1_test() { + let tokens = quote! { + #[attr] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "expected attribute arguments in parentheses: #[loop_unroll(...)" +)] +fn t2_test() { + let tokens = quote! { + #[loop_unroll] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "The attribute must specify unroll `method`")] +fn t3_test() { + let tokens = quote! { + #[loop_unroll()] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "unexpected end of input, expected parentheses")] +fn t4_test() { + let tokens = quote! { + #[loop_unroll(method)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "unexpected end of input, expected valid unroll method" +)] +fn t5_test() { + let tokens = quote! { + #[loop_unroll(method())] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "expected valid unroll method")] +fn t6_test() { + let tokens = quote! { + #[loop_unroll(method(run1time))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "The attribute must specify `factor`")] +fn t7_test() { + let tokens = quote! { + #[loop_unroll(method(runtime))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "expected `,")] +fn t8_test() { + let tokens = quote! { + #[loop_unroll(method(runtime) factor)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "unexpected end of input, expected parentheses")] +fn t9_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor)] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "unexpected end of input, expected integer literal")] +fn t10_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor())] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "unexpected end of input, expected some kind of loop" +)] +fn t11_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(42))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "unroll factor should not have any suffix")] +fn t12_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(42u16))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "Unroll factor can not be zero")] +fn t13_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(0))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic( + expected = "unexpected end of input, expected some kind of loop" +)] +fn t14_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(1))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +#[should_panic(expected = "There should only be a single attribute")] +fn t15_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(1))] + #[loop_unroll(method(runtime), factor(1))] + }; + match syn::parse2::(tokens) { + Ok(_) => (), + Err(e) => { + panic!("{}", e) + } + } +} + +#[test] +fn good_test() { + let tokens = quote! { + #[loop_unroll(method(runtime), factor(42))] + 'loop_label: for elt in iter { body } rest + }; + match syn::parse2::(tokens).unwrap() { + Item::LoopUnrollAttr(attr) => { + assert_eq!(attr.params.unroll_method, UnrollMethod::Runtime); + assert_eq!(attr.params.unroll_factor, 42); + let crate::Loop::ExprForLoop(for_loop) = attr.the_loop else { + unreachable!() + }; + assert_eq!( + for_loop.label.to_token_stream().to_string(), + ("'loop_label :") + ); + assert_eq!(for_loop.pat.to_token_stream().to_string(), ("elt")); + assert_eq!(for_loop.expr.to_token_stream().to_string(), ("iter")); + assert_eq!( + for_loop.body.to_token_stream().to_string(), + ("{ body }") + ); + assert_eq!(attr.rest_of_tokenstream.to_string(), "rest"); + } + Item::LoopPeelAttr(loop_xform_conf) => unreachable!(), + } +} diff --git a/src/utils/loop_xform/parse/tests/unroll_with_remainder.rs b/src/utils/loop_xform/parse/tests/unroll_with_remainder.rs new file mode 100644 index 0000000..0960526 --- /dev/null +++ b/src/utils/loop_xform/parse/tests/unroll_with_remainder.rs @@ -0,0 +1,33 @@ +use crate::Item; +use crate::UnrollMethod; +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn good_test() { + let tokens = quote! { + #[loop_unroll(method(with_remainder), factor(42))] + 'loop_label: for elt in iter { body } rest + }; + match syn::parse2::(tokens).unwrap() { + Item::LoopUnrollAttr(attr) => { + assert_eq!(attr.params.unroll_method, UnrollMethod::WithRemainder); + assert_eq!(attr.params.unroll_factor, 42); + let crate::Loop::ExprForLoop(for_loop) = attr.the_loop else { + unreachable!() + }; + assert_eq!( + for_loop.label.to_token_stream().to_string(), + ("'loop_label :") + ); + assert_eq!(for_loop.pat.to_token_stream().to_string(), ("elt")); + assert_eq!(for_loop.expr.to_token_stream().to_string(), ("iter")); + assert_eq!( + for_loop.body.to_token_stream().to_string(), + ("{ body }") + ); + assert_eq!(attr.rest_of_tokenstream.to_string(), "rest"); + } + Item::LoopPeelAttr(loop_xform_conf) => unreachable!(), + } +} diff --git a/src/utils/loop_xform/transform/mod.rs b/src/utils/loop_xform/transform/mod.rs new file mode 100644 index 0000000..83d92bc --- /dev/null +++ b/src/utils/loop_xform/transform/mod.rs @@ -0,0 +1,25 @@ +use crate::{LoopPeelParams, LoopUnrollParams, LoopXFormConf}; + +use super::UnrollMethod; + +pub fn perform_loop_unroll( + c: LoopXFormConf, +) -> proc_macro2::TokenStream { + match c.params.unroll_method { + UnrollMethod::Runtime => unroll_runtime::transform(&c), + UnrollMethod::WithRemainder => unroll_with_remainder::transform(c), + } +} +pub fn perform_loop_peel( + c: LoopXFormConf, +) -> proc_macro2::TokenStream { + peel::transform(c) +} + +mod utils { + pub mod loop_break_labeller; +} + +mod peel; +mod unroll_runtime; +mod unroll_with_remainder; diff --git a/src/utils/loop_xform/transform/peel/mod.rs b/src/utils/loop_xform/transform/peel/mod.rs new file mode 100644 index 0000000..b95bc20 --- /dev/null +++ b/src/utils/loop_xform/transform/peel/mod.rs @@ -0,0 +1,179 @@ +use quote::quote; +use syn::{Label, Lifetime}; + +use crate::{LoopPeelParams, LoopXFormConf}; +// use syn::{ExprForLoop, token::Loop}; + +fn transform_ExprForLoop( + params: LoopPeelParams, + mut expr_for_loop: syn::ExprForLoop, + rest_of_tokenstream: proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + let label = match expr_for_loop.label { + Some(l) => l.name, + None => Lifetime::new("'loop_label", proc_macro2::Span::mixed_site()), + }; + + expr_for_loop.label = Some(Label { + name: label.clone(), + colon_token: syn::token::Colon { + spans: [proc_macro2::Span::mixed_site()], + }, + }); + + super::utils::loop_break_labeller::LabelUnlabelledBreaks::visit_expr_for_loop( + &mut expr_for_loop, + ); + + expr_for_loop.label = None; + + let pat = expr_for_loop.pat.as_ref(); + let expr = expr_for_loop.expr.as_ref(); + let body_stmts = expr_for_loop.body.stmts.as_slice(); + + let iter = syn::Ident::new_raw("iter", proc_macro2::Span::mixed_site()); + + let prelude = core::iter::repeat_n( + quote! { + if let Some(#pat) = #iter.next() { + #(#body_stmts)* + } else { + break #label; + } + }, + params.peel_count, + ); + + quote! { + #label: while true { + let mut #iter = #expr; + #(#prelude)* + for #pat in #iter { + #(#body_stmts)* + } + break #label; + } + #rest_of_tokenstream + } +} + +fn transform_ExprWhile( + params: LoopPeelParams, + mut expr_while: syn::ExprWhile, + rest_of_tokenstream: proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + let label = match expr_while.label { + Some(l) => l.name, + None => Lifetime::new("'loop_label", proc_macro2::Span::mixed_site()), + }; + + expr_while.label = Some(Label { + name: label.clone(), + colon_token: syn::token::Colon { + spans: [proc_macro2::Span::mixed_site()], + }, + }); + + super::utils::loop_break_labeller::LabelUnlabelledBreaks::visit_expr_while( + &mut expr_while, + ); + + expr_while.label = None; + + let cond = &*expr_while.cond; + let body_stmts = expr_while.body.stmts.as_slice(); + + let prelude = core::iter::repeat_n( + quote! { + if #cond { + #(#body_stmts)* + } else { + break #label; + } + }, + params.peel_count, + ); + + quote! { + #label: while true { + #(#prelude)* + while #cond { + #(#body_stmts)* + } + break #label; + } + #rest_of_tokenstream + } +} + +fn transform_ExprLoop( + params: LoopPeelParams, + mut expr_loop: syn::ExprLoop, + rest_of_tokenstream: proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + let label = match expr_loop.label { + Some(l) => l.name, + None => Lifetime::new("'loop_label", proc_macro2::Span::mixed_site()), + }; + + expr_loop.label = Some(Label { + name: label.clone(), + colon_token: syn::token::Colon { + spans: [proc_macro2::Span::mixed_site()], + }, + }); + + super::utils::loop_break_labeller::LabelUnlabelledBreaks::visit_expr_loop( + &mut expr_loop, + ); + + expr_loop.label = None; + + let body_stmts = expr_loop.body.stmts.as_slice(); + + let prelude = core::iter::repeat_n( + quote! { + { + #(#body_stmts)* + } + }, + params.peel_count, + ); + + quote! { + #label: while true { + #(#prelude)* + loop { + #(#body_stmts)* + } + break #label; + } + #rest_of_tokenstream + } +} + +pub fn transform( + ast: LoopXFormConf, +) -> proc_macro2::TokenStream { + let LoopXFormConf { + params, + the_loop, + rest_of_tokenstream, + } = ast; + match the_loop { + crate::Loop::ExprForLoop(expr_for_loop) => { + transform_ExprForLoop(params, expr_for_loop, rest_of_tokenstream) + } + crate::Loop::ExprWhile(expr_while) => { + transform_ExprWhile(params, expr_while, rest_of_tokenstream) + } + crate::Loop::ExprLoop(expr_loop) => { + transform_ExprLoop(params, expr_loop, rest_of_tokenstream) + } + _ => unreachable!(), + } +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod tests; diff --git a/src/utils/loop_xform/transform/peel/tests/expr_for_loop.rs b/src/utils/loop_xform/transform/peel/tests/expr_for_loop.rs new file mode 100644 index 0000000..aa4de01 --- /dev/null +++ b/src/utils/loop_xform/transform/peel/tests/expr_for_loop.rs @@ -0,0 +1,118 @@ +use crate::{Loop, LoopPeelParams, LoopXFormConf, transform::peel::transform}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn peel1_test() { + for src in [ + quote! { for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + if let Some(elt) = r#iter.next() { body; break 'loop_label; } else { break 'loop_label; } + for elt in r#iter { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel2_test() { + for src in [ + quote! { for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 2 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + if let Some(elt) = r#iter.next() { body; break 'loop_label; } else { break 'loop_label; } + if let Some(elt) = r#iter.next() { body; break 'loop_label; } else { break 'loop_label; } + for elt in r#iter { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel1_with_nested_loop_test() { + let body = quote! { + for elt in other_iter { body; break; }; + while other_iter { body; break; }; + loop { body; break; }; + }; + for src in [ + quote! { + for elt in iter { + #body + break; + } + }, + quote! { + 'loop_label: for elt in iter { + #body + break; + } + }, + quote! { + 'loop_label: for elt in iter { + #body + break 'loop_label; + } + }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + if let Some(elt) = r#iter.next() { #body break 'loop_label; } else { break 'loop_label; } + for elt in r#iter { #body break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} diff --git a/src/utils/loop_xform/transform/peel/tests/expr_loop.rs b/src/utils/loop_xform/transform/peel/tests/expr_loop.rs new file mode 100644 index 0000000..77498bd --- /dev/null +++ b/src/utils/loop_xform/transform/peel/tests/expr_loop.rs @@ -0,0 +1,115 @@ +use crate::{Loop, LoopPeelParams, LoopXFormConf, transform::peel::transform}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn peel1_test() { + for src in [ + quote! { loop { body; break; } }, + quote! { 'loop_label: loop { body; break; } }, + quote! { 'loop_label: loop { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + { body; break 'loop_label; } + loop { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel2_test() { + for src in [ + quote! { loop { body; break; } }, + quote! { 'loop_label: loop { body; break; } }, + quote! { 'loop_label: loop { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 2 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + { body; break 'loop_label; } + { body; break 'loop_label; } + loop { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel1_with_nested_loop_test() { + let body = quote! { + for elt in other_iter { body; break; }; + while other_iter { body; break; }; + loop { body; break; }; + }; + for src in [ + quote! { + loop { + #body + break; + } + }, + quote! { + 'loop_label: loop { + #body + break; + } + }, + quote! { + 'loop_label: loop { + #body + break 'loop_label; + } + }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + { #body break 'loop_label; } + loop { #body break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} diff --git a/src/utils/loop_xform/transform/peel/tests/expr_while.rs b/src/utils/loop_xform/transform/peel/tests/expr_while.rs new file mode 100644 index 0000000..484743c --- /dev/null +++ b/src/utils/loop_xform/transform/peel/tests/expr_while.rs @@ -0,0 +1,115 @@ +use crate::{Loop, LoopPeelParams, LoopXFormConf, transform::peel::transform}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn peel1_test() { + for src in [ + quote! { while cond { body; break; } }, + quote! { 'loop_label: while cond { body; break; } }, + quote! { 'loop_label: while cond { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + if cond { body; break 'loop_label; } else { break 'loop_label; } + while cond { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel2_test() { + for src in [ + quote! { while cond { body; break; } }, + quote! { 'loop_label: while cond { body; break; } }, + quote! { 'loop_label: while cond { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 2 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + if cond { body; break 'loop_label; } else { break 'loop_label; } + if cond { body; break 'loop_label; } else { break 'loop_label; } + while cond { body; break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn peel1_with_nested_loop_test() { + let body = quote! { + for elt in other_iter { body; break; }; + while other_iter { body; break; }; + loop { body; break; }; + }; + for src in [ + quote! { + while cond { + #body + break; + } + }, + quote! { + 'loop_label: while cond { + #body + break; + } + }, + quote! { + 'loop_label: while cond { + #body + break 'loop_label; + } + }, + ] { + let conf = LoopXFormConf { + params: LoopPeelParams { peel_count: 1 }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + if cond { #body break 'loop_label; } else { break 'loop_label; } + while cond { #body break 'loop_label; } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} diff --git a/src/utils/loop_xform/transform/peel/tests/mod.rs b/src/utils/loop_xform/transform/peel/tests/mod.rs new file mode 100644 index 0000000..3d31069 --- /dev/null +++ b/src/utils/loop_xform/transform/peel/tests/mod.rs @@ -0,0 +1,3 @@ +mod expr_for_loop; +mod expr_loop; +mod expr_while; diff --git a/src/utils/loop_xform/transform/unroll_runtime/mod.rs b/src/utils/loop_xform/transform/unroll_runtime/mod.rs new file mode 100644 index 0000000..86a7b57 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_runtime/mod.rs @@ -0,0 +1,109 @@ +use quote::quote; + +use crate::{LoopUnrollParams, LoopXFormConf}; +// use syn::{ExprForLoop, token::Loop}; + +fn transform_ExprForLoop( + ast: &LoopXFormConf, + expr_for_loop: &syn::ExprForLoop, +) -> proc_macro2::TokenStream { + let label = &expr_for_loop.label; + let pat = expr_for_loop.pat.as_ref(); + let expr = expr_for_loop.expr.as_ref(); + let body_stmts = expr_for_loop.body.stmts.as_slice(); + let remainder = &ast.rest_of_tokenstream; + + let iter = syn::Ident::new_raw("iter", proc_macro2::Span::mixed_site()); + + let new_body = core::iter::repeat_n( + quote! { + if let Some(#pat) = #iter.next() { + #(#body_stmts)* + } else { + break; + } + }, + ast.params.unroll_factor, + ); + + quote! { + { + let mut #iter = #expr; + #label while true { + #(#new_body)* + } + } + #remainder + } +} + +fn transform_ExprWhile( + ast: &LoopXFormConf, + expr_while: &syn::ExprWhile, +) -> proc_macro2::TokenStream { + let label = &expr_while.label; + let cond = expr_while.cond.as_ref(); + let body_stmts = expr_while.body.stmts.as_slice(); + let remainder = &ast.rest_of_tokenstream; + + let new_body = core::iter::repeat_n( + quote! { + if #cond { + #(#body_stmts)* + } else { + break; + } + }, + ast.params.unroll_factor, + ); + + quote! { + #label while true { + #(#new_body)* + } + #remainder + } +} + +fn transform_ExprLoop( + ast: &LoopXFormConf, + expr_loop: &syn::ExprLoop, +) -> proc_macro2::TokenStream { + let label = &expr_loop.label; + let body_stmts = expr_loop.body.stmts.as_slice(); + let remainder = &ast.rest_of_tokenstream; + + let new_body = core::iter::repeat_n( + quote! { + { + #(#body_stmts)* + } + }, + ast.params.unroll_factor, + ); + + quote! { + #label loop { + #(#new_body)* + } + #remainder + } +} + +pub fn transform( + ast: &LoopXFormConf, +) -> proc_macro2::TokenStream { + match &ast.the_loop { + crate::Loop::ExprForLoop(expr_for_loop) => { + transform_ExprForLoop(ast, expr_for_loop) + } + crate::Loop::ExprWhile(expr_while) => { + transform_ExprWhile(ast, expr_while) + } + crate::Loop::ExprLoop(expr_loop) => transform_ExprLoop(ast, expr_loop), + } +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod tests; diff --git a/src/utils/loop_xform/transform/unroll_runtime/tests/expr_for_loop.rs b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_for_loop.rs new file mode 100644 index 0000000..bf2ca81 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_for_loop.rs @@ -0,0 +1,70 @@ +use crate::{ + Loop, LoopUnrollParams, LoopXFormConf, UnrollMethod, + transform::unroll_runtime::transform, +}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn unroll1_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 1, + }, + the_loop: syn::parse2::( + quote! { 'loop_label: for elt in iter { body } }, + ) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + { + let mut r#iter = iter; + 'loop_label : while true { + if let Some(elt) = r#iter.next() { body } else { break; } + } + } + rest + } + .to_token_stream() + .to_string() + ); +} + +#[test] +fn unroll2_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 2, + }, + the_loop: syn::parse2::( + quote! { 'loop_label: for elt in iter { body } }, + ) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + { + let mut r#iter = iter; + 'loop_label : while true { + if let Some(elt) = r#iter.next() { body } else { break; } + if let Some(elt) = r#iter.next() { body } else { break; } + } + } + rest + } + .to_token_stream() + .to_string() + ); +} diff --git a/src/utils/loop_xform/transform/unroll_runtime/tests/expr_loop.rs b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_loop.rs new file mode 100644 index 0000000..35af83d --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_loop.rs @@ -0,0 +1,60 @@ +use crate::{ + Loop, LoopUnrollParams, LoopXFormConf, UnrollMethod, + transform::unroll_runtime::transform, +}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn unroll1_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 1, + }, + the_loop: syn::parse2::(quote! { 'loop_label: loop { body } }) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: loop { + { body } + } + rest + } + .to_token_stream() + .to_string() + ); +} + +#[test] +fn unroll2_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 2, + }, + the_loop: syn::parse2::(quote! { 'loop_label: loop { body } }) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: loop { + { body } + { body } + } + rest + } + .to_token_stream() + .to_string() + ); +} diff --git a/src/utils/loop_xform/transform/unroll_runtime/tests/expr_while.rs b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_while.rs new file mode 100644 index 0000000..88fd9c0 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_runtime/tests/expr_while.rs @@ -0,0 +1,64 @@ +use crate::{ + Loop, LoopUnrollParams, LoopXFormConf, UnrollMethod, + transform::unroll_runtime::transform, +}; + +use quote::ToTokens as _; +use quote::quote; + +#[test] +fn unroll1_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 1, + }, + the_loop: syn::parse2::( + quote! { 'loop_label: while cond { body } }, + ) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + if cond { body } else { break; } + } + rest + } + .to_token_stream() + .to_string() + ); +} + +#[test] +fn unroll2_test() { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::Runtime, + unroll_factor: 2, + }, + the_loop: syn::parse2::( + quote! { 'loop_label: while cond { body } }, + ) + .unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = transform(&conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + if cond { body } else { break; } + if cond { body } else { break; } + } + rest + } + .to_token_stream() + .to_string() + ); +} diff --git a/src/utils/loop_xform/transform/unroll_runtime/tests/mod.rs b/src/utils/loop_xform/transform/unroll_runtime/tests/mod.rs new file mode 100644 index 0000000..3d31069 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_runtime/tests/mod.rs @@ -0,0 +1,3 @@ +mod expr_for_loop; +mod expr_loop; +mod expr_while; diff --git a/src/utils/loop_xform/transform/unroll_with_remainder/mod.rs b/src/utils/loop_xform/transform/unroll_with_remainder/mod.rs new file mode 100644 index 0000000..eac66a2 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_with_remainder/mod.rs @@ -0,0 +1,161 @@ +use crate::{LoopUnrollParams, LoopXFormConf}; +use quote::{ToTokens as _, quote}; +use syn::{Expr, Label, Lifetime}; + +pub fn transform( + ast: LoopXFormConf, +) -> proc_macro2::TokenStream { + let crate::Loop::ExprForLoop(mut for_loop) = ast.the_loop else { + return syn::Error::new_spanned( + ast.the_loop.to_token_stream(), + "Loop unroll with remainder only supported expr-for-loops", + ) + .to_compile_error(); + }; + + let label_outer = match for_loop.label { + Some(l) => l.name.clone(), + None => Lifetime::new("'loop_label", proc_macro2::Span::mixed_site()), + }; + + for_loop.label = Some(Label { + name: label_outer.clone(), + colon_token: syn::token::Colon { + spans: [proc_macro2::Span::mixed_site()], + }, + }); + + super::utils::loop_break_labeller::LabelUnlabelledBreaks::visit_expr_for_loop( + &mut for_loop, + ); + + let label_inner = + Lifetime::new("'label_inner", proc_macro2::Span::mixed_site()); + + let iter = syn::Ident::new_raw("iter", proc_macro2::Span::mixed_site()); + + let mut iter_evals = vec![]; + let mut iter_elts = vec![]; + for i in 0..ast.params.unroll_factor { + let suffix = format!("{}_of_{}", i + 1, ast.params.unroll_factor); + iter_evals.push(syn::Ident::new_raw( + &format!("iter_{suffix}"), + proc_macro2::Span::mixed_site(), + )); + iter_elts.push(syn::Ident::new_raw( + &format!("elt_{suffix}"), + proc_macro2::Span::mixed_site(), + )); + } + + let p = Pieces { + label_outer, + pat: for_loop.pat, + expr: for_loop.expr, + body_stmts: for_loop.body.stmts, + remainder: ast.rest_of_tokenstream, + label_inner, + iter, + iter_evals, + iter_elts, + }; + builder(&p) +} + +struct Pieces { + label_outer: Lifetime, + pat: Box, + expr: Box, + body_stmts: Vec, + remainder: proc_macro2::TokenStream, + label_inner: Lifetime, + iter: syn::Ident, + iter_evals: Vec, + iter_elts: Vec, +} + +fn builder(s: &Pieces) -> proc_macro2::TokenStream { + let label_outer = &s.label_outer; + let pat = s.pat.as_ref(); + let expr = s.expr.as_ref(); + let body_stmts = s.body_stmts.as_slice(); + let remainder = &s.remainder; + let label_inner = &s.label_inner; + let iter = &s.iter; + + let prologue = s.iter_evals.iter().rev().map(|curr_pat| { + quote! { + let mut #curr_pat = None; + } + }); + + let unrolled_iter_init = s.iter_evals.iter().map(Some).scan(None, |prev, curr| { + let out = (*prev, curr); + *prev = curr; + Some(out) + }).map(|(prev_iter, cur_iter)| { + match (prev_iter, cur_iter) { + (None, Some(cur_iter)) => { + quote! { + #cur_iter = #iter.next(); + } + } + (Some(prev_iter), Some(cur_iter)) => { + quote! { + #cur_iter = if #prev_iter.is_some() { #iter.next() } else { None }; + } + } + _ => unreachable!(), + } + }); + + let unrolled_iter_cond = + s.iter_evals.iter().zip(s.iter_elts.iter()).rev().map( + |(curr_iter, curr_elt)| { + quote! { + let Some(#curr_elt) = #curr_iter.take() + } + }, + ); + + let unrolled_body = s.iter_elts.iter().map(|curr_elt| { + quote! { + { + let #pat = #curr_elt; + #(#body_stmts)* + } + } + }); + + let epilogue = s.iter_evals.iter().rev().skip(1).rev().map(|curr_pat| { + quote! { + if let Some(#pat) = #curr_pat.take() { + #(#body_stmts)* + } else { + break #label_outer; + } + } + }); + + quote! { + #label_outer: while true { + let mut #iter = #expr; + #(#prologue)* + #label_inner: while true { + #(#unrolled_iter_init)* + if #(#unrolled_iter_cond)&&* { + #(#unrolled_body)* + } else { + break #label_inner; + } + } + #(#epilogue)* + break #label_outer; + } + #remainder + } +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod tests; diff --git a/src/utils/loop_xform/transform/unroll_with_remainder/tests.rs b/src/utils/loop_xform/transform/unroll_with_remainder/tests.rs new file mode 100644 index 0000000..7edea52 --- /dev/null +++ b/src/utils/loop_xform/transform/unroll_with_remainder/tests.rs @@ -0,0 +1,243 @@ +use crate::{Loop, LoopUnrollParams, LoopXFormConf, UnrollMethod}; +use quote::ToTokens as _; +use quote::quote; +use syn::ExprForLoop; + +#[test] +fn unroll1_test() { + for src in [ + quote! { for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::WithRemainder, + unroll_factor: 1, + }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = super::transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + let mut r#iter_1_of_1 = None; + 'label_inner: while true { + r#iter_1_of_1 = r#iter.next(); + if let Some(r#elt_1_of_1) = r#iter_1_of_1.take() { + { + let elt = r#elt_1_of_1; + body; + break 'loop_label; + } + } else { + break 'label_inner; + } + } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn unroll2_test() { + for src in [ + quote! { for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::WithRemainder, + unroll_factor: 2, + }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = super::transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + let mut r#iter_2_of_2 = None; + let mut r#iter_1_of_2 = None; + 'label_inner: while true { + r#iter_1_of_2 = r#iter.next(); + r#iter_2_of_2 = if r#iter_1_of_2.is_some() { r#iter.next() } else { None }; + if let Some(r#elt_2_of_2) = r#iter_2_of_2.take() && + let Some(r#elt_1_of_2) = r#iter_1_of_2.take() { + { + let elt = r#elt_1_of_2; + body; + break 'loop_label; + } + { + let elt = r#elt_2_of_2; + body; + break 'loop_label; + } + } else { + break 'label_inner; + } + } + if let Some(elt) = r#iter_1_of_2.take() { + body; + break 'loop_label; + } else { + break 'loop_label; + } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn unroll3_test() { + for src in [ + quote! { for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break; } }, + quote! { 'loop_label: for elt in iter { body; break 'loop_label; } }, + ] { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::WithRemainder, + unroll_factor: 3, + }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = super::transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + let mut r#iter_3_of_3 = None; + let mut r#iter_2_of_3 = None; + let mut r#iter_1_of_3 = None; + 'label_inner: while true { + r#iter_1_of_3 = r#iter.next(); + r#iter_2_of_3 = if r#iter_1_of_3.is_some() { r#iter.next() } else { None }; + r#iter_3_of_3 = if r#iter_2_of_3.is_some() { r#iter.next() } else { None }; + if let Some(r#elt_3_of_3) = r#iter_3_of_3.take() && + let Some(r#elt_2_of_3) = r#iter_2_of_3.take() && + let Some(r#elt_1_of_3) = r#iter_1_of_3.take() { + { + let elt = r#elt_1_of_3; + body; + break 'loop_label; + } + { + let elt = r#elt_2_of_3; + body; + break 'loop_label; + } + { + let elt = r#elt_3_of_3; + body; + break 'loop_label; + } + } else { + break 'label_inner; + } + } + if let Some(elt) = r#iter_1_of_3.take() { + body; + break 'loop_label; + } else { + break 'loop_label; + } + if let Some(elt) = r#iter_2_of_3.take() { + body; break 'loop_label; + } else { + break 'loop_label; + } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} + +#[test] +fn unroll1_with_nested_loop_test() { + let body = quote! { + for elt in other_iter { body; break; }; + while other_iter { body; break; }; + loop { body; break; }; + }; + for src in [ + quote! { + for elt in iter { + #body + break; + } + }, + quote! { + 'loop_label: for elt in iter { + #body + break; + } + }, + quote! { + 'loop_label: for elt in iter { + #body + break 'loop_label; + } + }, + ] { + let conf = LoopXFormConf { + params: LoopUnrollParams { + unroll_method: UnrollMethod::WithRemainder, + unroll_factor: 1, + }, + the_loop: syn::parse2::(src).unwrap(), + rest_of_tokenstream: quote! { rest }, + }; + + let res = super::transform(conf); + assert_eq!( + res.to_string(), + quote! { + 'loop_label: while true { + let mut r#iter = iter; + let mut r#iter_1_of_1 = None; + 'label_inner: while true { + r#iter_1_of_1 = r#iter.next(); + if let Some(r#elt_1_of_1) = r#iter_1_of_1.take() { + { + let elt = r#elt_1_of_1; #body break 'loop_label; + } + } else { + break 'label_inner; + } + } + break 'loop_label; + } + rest + } + .to_token_stream() + .to_string() + ); + } +} diff --git a/src/utils/loop_xform/transform/utils/loop_break_labeller/mod.rs b/src/utils/loop_xform/transform/utils/loop_break_labeller/mod.rs new file mode 100644 index 0000000..4caba72 --- /dev/null +++ b/src/utils/loop_xform/transform/utils/loop_break_labeller/mod.rs @@ -0,0 +1,49 @@ +use syn::visit_mut::VisitMut; + +pub struct LabelUnlabelledBreaks { + loop_label: syn::Lifetime, +} + +impl LabelUnlabelledBreaks { + pub fn visit_expr_for_loop(i: &mut syn::ExprForLoop) { + let mut this = Self { + loop_label: i.label.as_ref().unwrap().name.clone(), + }; + for stmt in &mut i.body.stmts { + this.visit_stmt_mut(stmt); + } + } + pub fn visit_expr_while(i: &mut syn::ExprWhile) { + let mut this = Self { + loop_label: i.label.as_ref().unwrap().name.clone(), + }; + for stmt in &mut i.body.stmts { + this.visit_stmt_mut(stmt); + } + } + pub fn visit_expr_loop(i: &mut syn::ExprLoop) { + let mut this = Self { + loop_label: i.label.as_ref().unwrap().name.clone(), + }; + for stmt in &mut i.body.stmts { + this.visit_stmt_mut(stmt); + } + } +} + +#[allow(clippy::missing_trait_methods)] +impl VisitMut for LabelUnlabelledBreaks { + fn visit_expr_loop_mut(&mut self, _i: &mut syn::ExprLoop) {} + fn visit_expr_while_mut(&mut self, _i: &mut syn::ExprWhile) {} + fn visit_expr_for_loop_mut(&mut self, _i: &mut syn::ExprForLoop) {} + + fn visit_expr_break_mut(&mut self, i: &mut syn::ExprBreak) { + if i.label.is_none() { + i.label = Some(self.loop_label.clone()); + } + } +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod tests; diff --git a/src/utils/loop_xform/transform/utils/loop_break_labeller/tests.rs b/src/utils/loop_xform/transform/utils/loop_break_labeller/tests.rs new file mode 100644 index 0000000..e500bd2 --- /dev/null +++ b/src/utils/loop_xform/transform/utils/loop_break_labeller/tests.rs @@ -0,0 +1,69 @@ +use quote::ToTokens as _; +use quote::quote; + +#[test] +#[should_panic(expected = "called `Option::unwrap()` on a `None` value")] +fn t0_test() { + let tokens = quote! { + for i in e {} + }; + let mut for_loop = syn::parse2::(tokens).unwrap(); + super::LabelUnlabelledBreaks::visit_expr_for_loop(&mut for_loop); +} + +#[test] +fn test() { + let body_verbatim = quote! { + break 'loop_label; + for a in b { + break; + break 'loop_label; + break 'other_loop_label; + } + 'other_loop_label: for a in b { + break; + break 'loop_label; + break 'other_loop_label; + } + while c { + break; + break 'loop_label; + break 'other_loop_label; + } + 'other_loop_label: while c { + break; + break 'loop_label; + break 'other_loop_label; + } + loop { + break; + break 'loop_label; + break 'other_loop_label; + } + 'other_loop_label: loop { + break; + break 'loop_label; + break 'other_loop_label; + } + }; + let tokens = quote! { + 'loop_label: for i in e { + break; + #body_verbatim + } + }; + let mut for_loop = syn::parse2::(tokens).unwrap(); + super::LabelUnlabelledBreaks::visit_expr_for_loop(&mut for_loop); + + assert_eq!( + quote! { #for_loop }.to_string(), + quote! { + 'loop_label: for i in e { + break 'loop_label; + #body_verbatim + } + } + .to_token_stream() + .to_string() + ); +} diff --git a/tests/utils/loop_xform/Cargo.toml b/tests/utils/loop_xform/Cargo.toml new file mode 100644 index 0000000..0c43379 --- /dev/null +++ b/tests/utils/loop_xform/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "loop_xform-test" +version.workspace = true +authors.workspace = true +edition.workspace = true +rust-version.workspace = true +documentation.workspace = true +homepage.workspace = true +repository.workspace = true +license.workspace = true + +[lints] +workspace = true + +[dependencies] +rawspeed-utils-loop_xform = { path = "../../../src/utils/loop_xform" } + +[[test]] +name = "loop_xform-test" +path = "mod.rs" diff --git a/tests/utils/loop_xform/mod.rs b/tests/utils/loop_xform/mod.rs new file mode 100644 index 0000000..75547fc --- /dev/null +++ b/tests/utils/loop_xform/mod.rs @@ -0,0 +1,146 @@ +use core::cell::RefCell; +use std::rc::Rc; + +struct LoggingIter { + log: Rc>>, + pos: usize, + end: usize, +} + +impl LoggingIter { + fn new( + log: Rc>>, + core::ops::Range { start, end }: core::ops::Range, + ) -> Self { + log.borrow_mut().push(format!("Iter created, at {start}")); + + Self { + log, + pos: start, + end, + } + } +} + +#[allow(clippy::missing_trait_methods)] +impl Iterator for LoggingIter { + type Item = IterVal; + + fn next(&mut self) -> Option { + self.log + .borrow_mut() + .push(format!("Iter next() called at pos = {}", self.pos)); + + if self.pos >= self.end { + self.log.borrow_mut().push(format!( + "Iter next() called at pos = {}, returning None", + self.pos + )); + return None; + } + + let current = self.pos; + let next = self.pos + 1; + self.log.borrow_mut().push(format!( + "Iter next() called at pos = {}, returning {}, next is {}", + self.pos, current, next + )); + self.pos = next; + Some(IterVal::new(Rc::clone(&self.log), current)) + } +} + +impl Drop for LoggingIter { + fn drop(&mut self) { + self.log + .borrow_mut() + .push(format!("Iter dropped, was at {}", self.pos)); + } +} + +struct IterVal { + log: Rc>>, + val: usize, +} + +impl IterVal { + fn new(log: Rc>>, val: usize) -> Self { + log.borrow_mut().push(format!("IterVal({val}) created")); + Self { log, val } + } +} + +impl core::ops::Deref for IterVal { + type Target = usize; + + fn deref(&self) -> &Self::Target { + self.log + .borrow_mut() + .push(format!("IterVal({}) deref", self.val)); + &self.val + } +} + +impl Drop for IterVal { + fn drop(&mut self) { + self.log + .borrow_mut() + .push(format!("IterVal({}) dropped", self.val)); + } +} + +fn gen_native_output( + r: core::ops::Range, + break_after: Option, +) -> Vec { + let mut vec = vec![]; + vec.push("Before macro".to_owned()); + vec.push(format!("Iter created, at {}", r.start).to_owned()); + for i in r.clone() { + vec.push(format!("Iter next() called at pos = {i}")); + vec.push(format!( + "Iter next() called at pos = {i}, returning {i}, next is {}", + i + 1 + )); + vec.push(format!("IterVal({i}) created")); + vec.push(format!("IterVal({i}) deref")); + vec.push(format!("Loop body at i = {i}")); + vec.push(format!("IterVal({i}) dropped")); + if Some(i) == break_after { + break; + } + } + if let Some(break_after) = break_after + && Some(break_after) <= r.clone().last() + { + vec.push( + format!("Iter dropped, was at {}", 1 + break_after).to_owned(), + ); + } else { + vec.push(format!("Iter next() called at pos = {}", r.end)); + vec.push(format!( + "Iter next() called at pos = {}, returning None", + r.end, + )); + vec.push(format!("Iter dropped, was at {}", r.end).to_owned()); + } + vec.push("After loop".to_owned()); + vec.push("After macro".to_owned()); + vec +} + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod naive; + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod unroll_runtime; + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod unroll_with_remainder; + +#[cfg(test)] +#[allow(clippy::large_stack_frames)] +mod peel; diff --git a/tests/utils/loop_xform/naive.rs b/tests/utils/loop_xform/naive.rs new file mode 100644 index 0000000..9a56de0 --- /dev/null +++ b/tests/utils/loop_xform/naive.rs @@ -0,0 +1,96 @@ +use crate::LoggingIter; +use crate::gen_native_output; +use core::cell::RefCell; +use std::rc::Rc; + +macro_rules! gen_native_test { + ($name:ident, $len:expr) => { + #[test] + fn $name() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + for i in LoggingIter::new(Rc::clone(&log), 0..$len) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + } + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], gen_native_output(0..$len, None)); + } + }; +} + +gen_native_test!(baseline_len0, 0); +gen_native_test!(baseline_len1, 1); +gen_native_test!(baseline_len2, 2); +gen_native_test!(baseline_len3, 3); +gen_native_test!(baseline_len4, 4); +gen_native_test!(baseline_len5, 5); + +#[test] +fn break0_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + for i in LoggingIter::new(Rc::clone(&log), 0..16) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "Iter dropped, was at 1", + "After loop", + "After macro" + ] + ); +} + +#[test] +fn break_label0_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + 'my_loop: for i in LoggingIter::new(Rc::clone(&log), 0..16) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break 'my_loop; + } + } + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "Iter dropped, was at 1", + "After loop", + "After macro" + ] + ); +} diff --git a/tests/utils/loop_xform/peel.rs b/tests/utils/loop_xform/peel.rs new file mode 100644 index 0000000..c10e29f --- /dev/null +++ b/tests/utils/loop_xform/peel.rs @@ -0,0 +1,88 @@ +macro_rules! gen_test { + ($name:ident, $peel_count:expr) => { + mod $name { + #[test] + fn expr_for_loop() { + for len in 0..5 { + for break_after in [None].iter().copied().chain((0..=len).into_iter().map(Some)) { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_peel($peel_count)] + for i in crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..len) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if Some(i) == break_after { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..len, break_after)); + } + } + } + #[test] + fn expr_while() { + for len in 0..5 { + for break_after in [None].iter().copied().chain((0..=len).into_iter().map(Some)) { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + let mut iter = crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..len); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_peel($peel_count)] + while let Some(i) = iter.next() { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if Some(i) == break_after { + break; + } + }; + ); + drop(iter); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..len, break_after)); + } + } + } + #[test] + fn expr_loop() { + for len in 0..5 { + for break_after in [None].iter().copied().chain((0..=len).into_iter().map(Some)) { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + let mut iter = crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..len); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_peel($peel_count)] + loop { + let Some(i) = iter.next() else { break; }; + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if Some(i) == break_after { + break; + } + }; + ); + drop(iter); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..len, break_after)); + } + } + } + } + } +} + +gen_test!(peel1, 1); +gen_test!(peel2, 2); +gen_test!(peel3, 3); +gen_test!(peel4, 4); diff --git a/tests/utils/loop_xform/unroll_runtime.rs b/tests/utils/loop_xform/unroll_runtime.rs new file mode 100644 index 0000000..fe8f332 --- /dev/null +++ b/tests/utils/loop_xform/unroll_runtime.rs @@ -0,0 +1,127 @@ +macro_rules! gen_test { + ($name:ident, $uf:expr, $len:expr) => { + mod $name { + #[test] + fn expr_for_loop() { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_unroll(method(runtime), factor($uf))] + for i in crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..$len) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..$len, None)); + } + #[test] + fn expr_while() { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + let mut iter = crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..$len); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_unroll(method(runtime), factor($uf))] + while let Some(i) = iter.next() { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + }; + ); + drop(iter); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..$len, None)); + } + #[test] + fn expr_loop() { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + let mut iter = crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..$len); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_unroll(method(runtime), factor($uf))] + loop { + let Some(i) = iter.next() else { break; }; + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + }; + ); + drop(iter); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..$len, None)); + } + } + }; +} + +gen_test!(unroll1_len0, 1, 0); +gen_test!(unroll1_len1, 1, 1); +gen_test!(unroll1_len2, 1, 2); +gen_test!(unroll1_len3, 1, 3); +gen_test!(unroll1_len4, 1, 4); +gen_test!(unroll1_len5, 1, 5); + +gen_test!(unroll2_len0, 2, 0); +gen_test!(unroll2_len1, 2, 1); +gen_test!(unroll2_len2, 2, 2); +gen_test!(unroll2_len3, 2, 3); +gen_test!(unroll2_len4, 2, 4); +gen_test!(unroll2_len5, 2, 5); + +gen_test!(unroll3_len0, 3, 0); +gen_test!(unroll3_len1, 3, 1); +gen_test!(unroll3_len2, 3, 2); +gen_test!(unroll3_len3, 3, 3); +gen_test!(unroll3_len4, 3, 4); +gen_test!(unroll3_len5, 3, 5); + +#[test] +fn break0_test() { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_unroll(method(runtime), factor(16))] + for i in crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..16) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..16, Some(0))); +} + +#[test] +fn break_label0_test() { + let log = std::rc::Rc::new(core::cell::RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + rawspeed_utils_loop_xform::enable_loop_xforms!( + #[loop_unroll(method(runtime), factor(16))] + 'my_loop: for i in + crate::LoggingIter::new(std::rc::Rc::clone(&log), 0..16) + { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break 'my_loop; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], crate::gen_native_output(0..16, Some(0))); +} diff --git a/tests/utils/loop_xform/unroll_with_remainder.rs b/tests/utils/loop_xform/unroll_with_remainder.rs new file mode 100644 index 0000000..b1f776d --- /dev/null +++ b/tests/utils/loop_xform/unroll_with_remainder.rs @@ -0,0 +1,258 @@ +use crate::LoggingIter; +use core::cell::RefCell; +use rawspeed_utils_loop_xform::enable_loop_xforms; +use std::rc::Rc; + +fn gen_unroll_output(uf: usize, r: core::ops::Range) -> Vec { + let mut vec = vec![]; + vec.push("Before macro".to_owned()); + vec.push(format!("Iter created, at {}", r.start).to_owned()); + let iterspace = r.clone().collect::>(); + let mut chunks = iterspace[..].chunks_exact(uf); + for chunk in chunks.by_ref() { + for i in chunk { + vec.push(format!("Iter next() called at pos = {i}")); + vec.push(format!( + "Iter next() called at pos = {i}, returning {i}, next is {}", + i + 1 + )); + vec.push(format!("IterVal({i}) created")); + } + for i in chunk { + vec.push(format!("IterVal({i}) deref")); + vec.push(format!("Loop body at i = {i}")); + vec.push(format!("IterVal({i}) dropped")); + } + } + for i in chunks.remainder() { + vec.push(format!("Iter next() called at pos = {i}")); + vec.push(format!( + "Iter next() called at pos = {i}, returning {i}, next is {}", + i + 1 + )); + vec.push(format!("IterVal({i}) created")); + } + vec.push(format!("Iter next() called at pos = {}", r.end)); + vec.push(format!( + "Iter next() called at pos = {}, returning None", + r.end, + )); + for i in chunks.remainder() { + vec.push(format!("IterVal({i}) deref")); + vec.push(format!("Loop body at i = {i}")); + vec.push(format!("IterVal({i}) dropped")); + } + vec.push(format!("Iter dropped, was at {}", r.end,)); + vec.push("After loop".to_owned()); + vec.push("After macro".to_owned()); + vec +} + +macro_rules! gen_test { + ($name:ident, $uf:expr, $len:expr) => { + #[test] + fn $name() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + enable_loop_xforms!( + #[loop_unroll(method(with_remainder), factor($uf))] + for i in LoggingIter::new(Rc::clone(&log), 0..$len) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!(log.borrow()[..], gen_unroll_output($uf, 0..$len)); + } + }; +} + +gen_test!(unroll1_len0, 1, 0); +gen_test!(unroll1_len1, 1, 1); +gen_test!(unroll1_len2, 1, 2); +gen_test!(unroll1_len3, 1, 3); +gen_test!(unroll1_len4, 1, 4); +gen_test!(unroll1_len5, 1, 5); + +gen_test!(unroll2_len0, 2, 0); +gen_test!(unroll2_len1, 2, 1); +gen_test!(unroll2_len2, 2, 2); +gen_test!(unroll2_len3, 2, 3); +gen_test!(unroll2_len4, 2, 4); +gen_test!(unroll2_len5, 2, 5); + +gen_test!(unroll3_len0, 3, 0); +gen_test!(unroll3_len1, 3, 1); +gen_test!(unroll3_len2, 3, 2); +gen_test!(unroll3_len3, 3, 3); +gen_test!(unroll3_len4, 3, 4); +gen_test!(unroll3_len5, 3, 5); + +#[test] +fn break0_unroll2_len2_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + enable_loop_xforms!( + #[loop_unroll(method(with_remainder), factor(2))] + for i in LoggingIter::new(Rc::clone(&log), 0..2) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "Iter next() called at pos = 1", + "Iter next() called at pos = 1, returning 1, next is 2", + "IterVal(1) created", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "IterVal(1) dropped", + "Iter dropped, was at 2", + "After loop", + "After macro" + ] + ); +} + +#[test] +fn break0_unroll3_len2_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + enable_loop_xforms!( + #[loop_unroll(method(with_remainder), factor(3))] + for i in LoggingIter::new(Rc::clone(&log), 0..2) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "Iter next() called at pos = 1", + "Iter next() called at pos = 1, returning 1, next is 2", + "IterVal(1) created", + "Iter next() called at pos = 2", + "Iter next() called at pos = 2, returning None", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "IterVal(1) dropped", + "Iter dropped, was at 2", + "After loop", + "After macro" + ] + ); +} + +#[test] +fn break0_unroll2_len3_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + enable_loop_xforms!( + #[loop_unroll(method(with_remainder), factor(2))] + for i in LoggingIter::new(Rc::clone(&log), 0..3) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "Iter next() called at pos = 1", + "Iter next() called at pos = 1, returning 1, next is 2", + "IterVal(1) created", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "IterVal(1) dropped", + "Iter dropped, was at 2", + "After loop", + "After macro" + ] + ); +} + +#[test] +fn break0_unroll3_len3_test() { + let log = Rc::new(RefCell::new(Vec::::new())); + + log.borrow_mut().push("Before macro".to_owned()); + enable_loop_xforms!( + #[loop_unroll(method(with_remainder), factor(3))] + for i in LoggingIter::new(Rc::clone(&log), 0..3) { + let i = *i; + log.borrow_mut().push(format!("Loop body at i = {i}")); + if i == 0 { + break; + } + } + ); + log.borrow_mut().push("After loop".to_owned()); + log.borrow_mut().push("After macro".to_owned()); + + assert_eq!( + log.borrow()[..], + [ + "Before macro", + "Iter created, at 0", + "Iter next() called at pos = 0", + "Iter next() called at pos = 0, returning 0, next is 1", + "IterVal(0) created", + "Iter next() called at pos = 1", + "Iter next() called at pos = 1, returning 1, next is 2", + "IterVal(1) created", + "Iter next() called at pos = 2", + "Iter next() called at pos = 2, returning 2, next is 3", + "IterVal(2) created", + "IterVal(0) deref", + "Loop body at i = 0", + "IterVal(0) dropped", + "IterVal(1) dropped", + "IterVal(2) dropped", + "Iter dropped, was at 3", + "After loop", + "After macro" + ] + ); +}