From 2d330cd45f5c29b855082468e8ea160229e9b2b0 Mon Sep 17 00:00:00 2001 From: Marc Hartmayer Date: Mon, 22 Jun 2026 17:26:31 +0200 Subject: [PATCH] utils_macros: Implement derive_control_flag Add a new derive macro 'derive_control_flag' that is used in the next commit to reimplement how the code deals with Secure Execution control flags. It implements Display, IntoEnumIterator and the ControlFlagTrait for enums using unit variants only. /// Trait for control flags that provide bit position information. pub trait ControlFlagTrait { /// Returns the bit position for this flag. fn bit_position(self) -> u8; } Assisted-by: IBM Bob:1.0.4 Signed-off-by: Marc Hartmayer Reviewed-by: Steffen Eiden Signed-off-by: Steffen Eiden --- rust/utils/src/lib.rs | 2 +- rust/utils_macros/src/lib.rs | 143 ++++++++++- rust/utils_macros/tests/integration_tests.rs | 247 +++++++++++++++++++ rust/utils_macros/tests/test_helpers.rs | 17 ++ 4 files changed, 406 insertions(+), 3 deletions(-) create mode 100644 rust/utils_macros/tests/integration_tests.rs create mode 100644 rust/utils_macros/tests/test_helpers.rs diff --git a/rust/utils/src/lib.rs b/rust/utils/src/lib.rs index 5e98c76e..98e3941e 100644 --- a/rust/utils/src/lib.rs +++ b/rust/utils/src/lib.rs @@ -14,7 +14,7 @@ mod tmpfile; pub use ::log::LevelFilter; // Re-export procedural macros from utils_macros -pub use utils_macros::{ValueEnumDisplay, ValueEnumFromStr}; +pub use utils_macros::{ControlFlag, ValueEnumDisplay, ValueEnumFromStr}; pub use crate::cli::{ combined_path_opt, combined_path_req, get_reader_from_cli_file_arg, diff --git a/rust/utils_macros/src/lib.rs b/rust/utils_macros/src/lib.rs index dca6e530..84f626f9 100644 --- a/rust/utils_macros/src/lib.rs +++ b/rust/utils_macros/src/lib.rs @@ -4,11 +4,150 @@ //! Procedural macros for the utils crate. //! -//! This crate provides derive macros to reduce boilerplate in enum definitions. +//! This crate provides derive macros to reduce boilerplate in enum definitions, +//! particularly for control flags. use proc_macro::TokenStream; use quote::quote; -use syn::{parse_macro_input, DeriveInput}; +use syn::{parse_macro_input, Data, DeriveInput, Fields, Lit, Meta, MetaList}; + +/// Derive macro for control flag enums. +/// +/// This macro generates implementations for `Display`, `IntoEnumIterator`, and `ControlFlagTrait`. +/// It supports the `#[flag(display = "...", value = N)]` attribute to specify custom display +/// strings and discriminant values. +/// +/// # Example +/// +/// ``` +/// use utils_macros::ControlFlag; +/// +/// /// Trait for enums that can be iterated over. +/// pub trait IntoEnumIterator: Sized { +/// /// Returns an iterator over all variants of the enum. +/// fn iter() -> impl Iterator; +/// } +/// +/// /// Trait for control flags that provide bit position information. +/// pub trait ControlFlagTrait { +/// /// Returns the bit position for this flag. +/// fn bit_position(self) -> u8; +/// } +/// +/// #[derive(ControlFlag)] +/// pub enum PcfV1 { +/// #[flag(display = "Confidential dump support", value = 34)] +/// ConfidentialDump, +/// +/// #[flag(display = "V1-specific flag", value = 35)] +/// V1OnlyFlag, +/// } +/// ``` +#[proc_macro_derive(ControlFlag, attributes(flag))] +pub fn derive_control_flag(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); + let name = &input.ident; + + let variants = match &input.data { + Data::Enum(data) => &data.variants, + _ => panic!("ControlFlag can only be derived for enums"), + }; + + // Extract variant information + let mut variant_names = Vec::new(); + let mut variant_displays = Vec::new(); + let mut variant_values = Vec::new(); + + for variant in variants { + if !matches!(variant.fields, Fields::Unit) { + panic!("ControlFlag only supports unit variants"); + } + + let variant_name = &variant.ident; + variant_names.push(variant_name); + + // Parse the #[flag(...)] attribute + let mut display_str = variant_name.to_string(); + let mut value: Option = None; + + for attr in &variant.attrs { + if attr.path().is_ident("flag") { + // Try to parse as MetaList + if let Meta::List(MetaList { tokens, .. }) = &attr.meta { + // Parse the tokens inside the list + let parser = syn::meta::parser(|meta| { + if meta.path.is_ident("display") { + let val = meta.value()?; + let lit: Lit = val.parse()?; + if let Lit::Str(s) = lit { + display_str = s.value(); + } + } else if meta.path.is_ident("value") { + let val = meta.value()?; + let lit: Lit = val.parse()?; + if let Lit::Int(i) = lit { + value = Some(i.base10_parse()?); + } + } + Ok(()) + }); + + let _ = syn::parse::Parser::parse2(parser, tokens.clone()); + } + } + } + + if value.is_none() { + panic!( + "ControlFlag variant {} must have a value attribute", + variant_name + ); + } + + variant_displays.push(display_str); + variant_values.push(value.unwrap()); + } + + let expanded = quote! { + impl #name { + /// Returns the bit position value for this flag. + pub const fn flag_value(&self) -> u8 { + match self { + #(Self::#variant_names => #variant_values),* + } + } + } + + impl ControlFlagTrait for #name { + fn bit_position(self) -> u8 { + self.flag_value() + } + } + + impl std::fmt::Display for #name { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "{}", + match self { + #(Self::#variant_names => #variant_displays),* + } + ) + } + } + + impl IntoEnumIterator for #name { + fn iter() -> impl Iterator { + [ + #(Self::#variant_names),* + ] + .into_iter() + } + } + }; + + TokenStream::from(expanded) +} /// Derive `std::fmt::Display` for enums implementing `clap::ValueEnum`. /// diff --git a/rust/utils_macros/tests/integration_tests.rs b/rust/utils_macros/tests/integration_tests.rs new file mode 100644 index 00000000..f261b025 --- /dev/null +++ b/rust/utils_macros/tests/integration_tests.rs @@ -0,0 +1,247 @@ +// SPDX-License-Identifier: MIT +// +// Copyright IBM Corp. + +//! Integration tests for the ControlFlag derive macro. + +mod test_helpers; + +use test_helpers::{ControlFlagTrait, IntoEnumIterator}; +use utils_macros::ControlFlag; + +// Test enum for basic functionality +#[derive(ControlFlag, Debug, Clone, Copy, PartialEq, Eq)] +enum TestFlag { + #[flag(display = "first flag", value = 1)] + First, + #[flag(display = "second flag", value = 2)] + Second, +} + +// Test enum with multiple variants +#[derive(ControlFlag, Debug, Clone, Copy, PartialEq, Eq)] +enum MultiFlag { + #[flag(display = "flag 1", value = 10)] + Flag1, + #[flag(display = "flag 2", value = 20)] + Flag2, + #[flag(display = "flag 3", value = 30)] + Flag3, + #[flag(display = "flag 4", value = 40)] + Flag4, + #[flag(display = "flag 5", value = 50)] + Flag5, +} + +// Test enum with non-sequential values +#[derive(ControlFlag, Debug, Clone, Copy, PartialEq, Eq)] +enum SparseFlag { + #[flag(display = "low bit", value = 1)] + Low, + #[flag(display = "high bit", value = 63)] + High, + #[flag(display = "middle bit", value = 32)] + Middle, +} + +// Test enum with edge case display strings +#[derive(ControlFlag, Debug, Clone, Copy, PartialEq, Eq)] +enum EdgeCaseFlag { + #[flag(display = "simple", value = 1)] + Simple, + #[flag(display = "with spaces and punctuation!", value = 2)] + WithSpaces, + #[flag(display = "UPPERCASE", value = 3)] + Uppercase, + #[flag(display = "with-dashes-and_underscores", value = 4)] + WithDashes, +} + +// Test enum for Copy/Clone compatibility +#[derive(ControlFlag, Copy, Clone, Debug, PartialEq, Eq)] +enum CopyableFlag { + #[flag(display = "copyable", value = 1)] + Copyable, + #[flag(display = "another", value = 2)] + Another, +} + +#[test] +fn test_basic_derive() { + // Verify the macro successfully derives all required traits + let flag = TestFlag::First; + + // Should compile and be accessible + let _ = flag.flag_value(); + let _ = flag.bit_position(); + let _ = format!("{}", flag); + let _ = TestFlag::iter(); +} + +#[test] +fn test_display_trait() { + // Verify custom display strings are correctly used + assert_eq!(format!("{}", TestFlag::First), "first flag"); + assert_eq!(format!("{}", TestFlag::Second), "second flag"); +} + +#[test] +fn test_enum_iterator() { + // Verify iteration over all enum variants works correctly + let flags: Vec = TestFlag::iter().collect(); + assert_eq!(flags.len(), 2); + assert_eq!(flags[0], TestFlag::First); + assert_eq!(flags[1], TestFlag::Second); +} + +#[test] +fn test_control_flag_trait() { + // Verify bit_position() returns correct values + assert_eq!(TestFlag::First.bit_position(), 1); + assert_eq!(TestFlag::Second.bit_position(), 2); +} + +#[test] +fn test_flag_value_method() { + // Verify flag_value() returns correct bit positions + assert_eq!(TestFlag::First.flag_value(), 1); + assert_eq!(TestFlag::Second.flag_value(), 2); +} + +#[test] +fn test_multiple_variants() { + // Verify the macro handles enums with many variants + + // Test iteration + let flags: Vec = MultiFlag::iter().collect(); + assert_eq!(flags.len(), 5); + assert_eq!(flags[0], MultiFlag::Flag1); + assert_eq!(flags[1], MultiFlag::Flag2); + assert_eq!(flags[2], MultiFlag::Flag3); + assert_eq!(flags[3], MultiFlag::Flag4); + assert_eq!(flags[4], MultiFlag::Flag5); + + // Test bit positions + assert_eq!(MultiFlag::Flag1.bit_position(), 10); + assert_eq!(MultiFlag::Flag2.bit_position(), 20); + assert_eq!(MultiFlag::Flag3.bit_position(), 30); + assert_eq!(MultiFlag::Flag4.bit_position(), 40); + assert_eq!(MultiFlag::Flag5.bit_position(), 50); + + // Test display strings + assert_eq!(format!("{}", MultiFlag::Flag1), "flag 1"); + assert_eq!(format!("{}", MultiFlag::Flag2), "flag 2"); + assert_eq!(format!("{}", MultiFlag::Flag3), "flag 3"); + assert_eq!(format!("{}", MultiFlag::Flag4), "flag 4"); + assert_eq!(format!("{}", MultiFlag::Flag5), "flag 5"); +} + +#[test] +fn test_non_sequential_values() { + // Verify the macro handles non-sequential bit position values + assert_eq!(SparseFlag::Low.bit_position(), 1); + assert_eq!(SparseFlag::High.bit_position(), 63); + assert_eq!(SparseFlag::Middle.bit_position(), 32); + + // Verify iteration order matches declaration order + let flags: Vec = SparseFlag::iter().collect(); + assert_eq!(flags.len(), 3); + assert_eq!(flags[0], SparseFlag::Low); + assert_eq!(flags[1], SparseFlag::High); + assert_eq!(flags[2], SparseFlag::Middle); +} + +#[test] +fn test_display_string_edge_cases() { + // Verify various display string formats work correctly + assert_eq!(format!("{}", EdgeCaseFlag::Simple), "simple"); + assert_eq!( + format!("{}", EdgeCaseFlag::WithSpaces), + "with spaces and punctuation!" + ); + assert_eq!(format!("{}", EdgeCaseFlag::Uppercase), "UPPERCASE"); + assert_eq!( + format!("{}", EdgeCaseFlag::WithDashes), + "with-dashes-and_underscores" + ); +} + +#[test] +fn test_trait_bounds() { + // Verify generated implementations work with common trait bounds + + // Function that requires Display + fn requires_display(flag: T) -> String { + format!("{}", flag) + } + + // Function that requires IntoEnumIterator + fn requires_iterator() -> Vec { + T::iter().collect() + } + + // Test with TestFlag + let result = requires_display(TestFlag::First); + assert_eq!(result, "first flag"); + + let flags: Vec = requires_iterator(); + assert_eq!(flags.len(), 2); +} + +#[test] +fn test_copy_clone_compatibility() { + // Verify the macro works with Copy and Clone derives + let flag1 = CopyableFlag::Copyable; + let flag2 = flag1; // Copy + let flag3 = flag1; // Clone + + assert_eq!(flag1, flag2); + assert_eq!(flag1, flag3); + + // Verify all methods still work + assert_eq!(flag1.bit_position(), 1); + assert_eq!(flag2.flag_value(), 1); + assert_eq!(format!("{}", flag3), "copyable"); +} + +#[test] +fn test_iterator_multiple_calls() { + // Verify iterator can be called multiple times + let iter1: Vec = TestFlag::iter().collect(); + let iter2: Vec = TestFlag::iter().collect(); + + assert_eq!(iter1, iter2); + assert_eq!(iter1.len(), 2); +} + +#[test] +fn test_flag_value_const() { + // Verify flag_value() can be used in const contexts + const FLAG_VALUE: u8 = TestFlag::First.flag_value(); + assert_eq!(FLAG_VALUE, 1); +} + +#[test] +fn test_all_variants_unique_values() { + // Verify all variants have unique bit positions + let flags: Vec = TestFlag::iter().collect(); + let values: Vec = flags.iter().map(|f| f.bit_position()).collect(); + + // Check uniqueness + for i in 0..values.len() { + for j in (i + 1)..values.len() { + assert_ne!(values[i], values[j], "Duplicate bit position found"); + } + } +} + +#[test] +fn test_display_consistency() { + // Verify Display is consistent across multiple calls + let flag = TestFlag::First; + let display1 = format!("{}", flag); + let display2 = format!("{}", flag); + + assert_eq!(display1, display2); + assert_eq!(display1, "first flag"); +} diff --git a/rust/utils_macros/tests/test_helpers.rs b/rust/utils_macros/tests/test_helpers.rs new file mode 100644 index 00000000..566e9aaa --- /dev/null +++ b/rust/utils_macros/tests/test_helpers.rs @@ -0,0 +1,17 @@ +// SPDX-License-Identifier: MIT +// +// Copyright IBM Corp. + +//! Helper traits and utilities for testing the ControlFlag derive macro. + +/// Trait for control flags that provide bit position information. +pub trait ControlFlagTrait { + /// Returns the bit position for this flag. + fn bit_position(self) -> u8; +} + +/// Trait for enums that can be iterated over. +pub trait IntoEnumIterator: Sized { + /// Returns an iterator over all variants of the enum. + fn iter() -> impl Iterator; +}