use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::quote;
use syn::parse::{Parse, ParseStream, Result as ParseResult};
use syn::{Data, DeriveInput, Field, Ident, LitInt, LitStr, Token, Type};
#[derive(Debug)]
struct BitField<'a> {
parent: Option<&'a Field>,
getter: Option<Ident>,
setter: Option<Ident>,
msb: LitInt,
lsb: LitInt,
as_type: Option<Type>,
doc: Option<LitStr>,
}
impl<'a> BitField<'a> {
fn codegen(&self) -> TokenStream2 {
let empty_str = LitStr::new("", Span::call_site());
let getter = &self.getter;
let setter = &self.setter;
let msb = &self.msb;
let lsb = &self.lsb;
let as_type = &self.as_type;
let doc = self.doc.as_ref().unwrap_or(&empty_str);
let field = self.parent.unwrap();
let field_name = &field.ident;
let value_type = &field.ty;
let vis = &field.vis;
let getter_tokens = if getter.is_some() {
if self.is_as_bool() {
quote! {
#[doc = #doc]
#[inline]
#vis fn #getter(&self) -> bool {
let mask = ((1 << (#msb - #lsb + 1)) - 1) << #lsb;
((self.#field_name & mask) >> #lsb) != 0
}
}
} else if as_type.is_some() {
quote! {
#[doc = #doc]
#[inline]
#vis fn #getter(&self) -> #as_type {
let one: #value_type = 1;
let (mask, over) = one.overflowing_shl(#msb - #lsb + 1);
let mask = if over {
#value_type::MAX
} else {
(mask - 1) << #lsb
};
((self.#field_name & mask) >> #lsb) as #as_type
}
}
} else {
quote! {
#[doc = #doc]
#[inline]
#vis fn #getter(&self) -> #value_type {
let one: #value_type = 1;
let (mask, over) = one.overflowing_shl(#msb - #lsb + 1);
let mask = if over {
#value_type::MAX
} else {
(mask - 1) << #lsb
};
(self.#field_name & mask) >> #lsb
}
}
}
} else {
quote! {}
};
let setter_tokens = if setter.is_some() {
if self.is_as_bool() {
quote! {
#[doc = #doc]
#[inline]
#vis fn #setter(&mut self, value: bool) {
if value {
self.#field_name |= 1 << #lsb;
} else {
self.#field_name &= !(1 << #lsb);
}
}
}
} else if as_type.is_some() {
quote! {
#[doc = #doc]
#[inline]
#vis fn #setter(&mut self, value: #as_type) {
let one: #value_type = 1;
let (mask, over) = one.overflowing_shl(#msb - #lsb + 1);
let mask = if over {
#value_type::MAX
} else {
(mask - 1) << #lsb
};
self.#field_name &= !mask;
self.#field_name |= ((value as #value_type) << #lsb) & mask;
}
}
} else {
quote! {
#[doc = #doc]
#[inline]
#vis fn #setter(&mut self, value: #value_type) {
let one: #value_type = 1;
let (mask, over) = one.overflowing_shl(#msb - #lsb + 1);
let mask = if over {
#value_type::MAX
} else {
(mask - 1) << #lsb
};
self.#field_name &= !mask;
self.#field_name |= (value << #lsb) & mask;
}
}
}
} else {
quote! {}
};
quote! {
#getter_tokens
#setter_tokens
}
}
fn is_as_bool(&self) -> bool {
self.as_type
.as_ref()
.map(|x| match x {
Type::Path(p) => p.path.is_ident("bool"),
_ => false,
})
.unwrap_or(false)
}
}
impl<'a> Parse for BitField<'a> {
fn parse(input: ParseStream<'_>) -> ParseResult<Self> {
let getter = if input.peek(Token![_]) {
let _ = input.parse::<Token![_]>()?;
None
} else {
let getter: Ident = input.parse()?;
Some(getter)
};
let comma_token: Option<Token![,]> = input.parse()?;
let setter = if comma_token.is_some() {
if input.peek(Token![_]) {
input.parse::<Token![_]>()?;
None
} else {
let setter: Ident = input.parse()?;
Some(setter)
}
} else {
let setter = Ident::new(
&format!("set_{}", getter.as_ref().unwrap()),
Span::call_site(),
);
Some(setter)
};
let _colon_token: Token![:] = input.parse()?;
let content;
let _bracket_token = syn::bracketed!(content in input);
let msb = content.parse()?;
let colon_token: Option<Token![:]> = content.parse()?;
let lsb = if colon_token.is_some() {
content.parse()?
} else {
Clone::clone(&msb)
};
let as_token: Option<Token![as]> = input.parse()?;
let as_type = if as_token.is_some() {
Some(input.parse()?)
} else {
None
};
let doc: Option<LitStr> = input.parse()?;
Ok(Self {
parent: None,
getter,
setter,
msb,
lsb,
as_type,
doc,
})
}
}
#[proc_macro_derive(BitFields, attributes(bitfield))]
pub fn derive_bitfields(input: TokenStream) -> TokenStream {
let ast: DeriveInput = syn::parse(input).unwrap();
let bitfields = parse_bitfields(&ast);
let bitfields_impls: Vec<TokenStream2> = bitfields.iter().map(|x| x.codegen()).collect();
let ident = &ast.ident;
let generics = &ast.generics;
let tokens = quote! {
impl #generics #ident #generics {
#(#bitfields_impls)*
}
};
tokens.into()
}
fn parse_bitfields(ast: &DeriveInput) -> Vec<BitField> {
let mut bitfields: Vec<BitField> = vec![];
if let Data::Struct(s) = &ast.data {
for f in &s.fields {
for a in &f.attrs {
if a.path.is_ident("bitfield") {
let mut bitfield: BitField = a.parse_args().unwrap();
bitfield.parent = Some(f);
bitfields.push(bitfield);
}
}
}
}
bitfields
}