extern crate proc_macro;
use std::collections::BTreeMap;
use convert_case::{Case, Casing};
use proc_macro::TokenStream;
use proc_macro2::{Ident, Literal, TokenStream as TokenStream2, TokenTree};
use quote::quote;
use syn::{parse2, Attribute, Data, DeriveInput, Meta, Type};
#[proc_macro_derive(FromRecord, attributes(from_record))]
pub fn derive_fromrecord(tokens: TokenStream) -> TokenStream {
inner(tokens.into())
.unwrap_or_else(syn::Error::into_compile_error)
.into()
}
fn inner(tokens: TokenStream2) -> syn::Result<TokenStream2> {
let parsed: DeriveInput = parse2(tokens)?;
let name = parsed.ident;
let Data::Struct(data) = parsed.data else {
panic!("only usable on structs");
};
let mut attr_map = parse_attrs(parsed.attrs).expect("attribute with `id = n` required");
let TokenTree::Literal(id) = attr_map.remove("id").expect("record ID required") else {
panic!("record id should be a literal");
};
let use_box = match attr_map.remove("use_box") {
Some(TokenTree::Ident(val)) if val == "true" => true,
Some(TokenTree::Ident(val)) if val == "false" => true,
Some(v) => panic!("Expected ident but got {v:?}"),
None => false,
};
error_if_map_not_empty(&attr_map);
let mut match_stmts = Vec::new();
for field in data.fields {
let Type::Path(path) = field.ty else {
panic!("invalid type")
};
let field_name = field.ident.unwrap();
let mut field_map = parse_attrs(field.attrs).unwrap_or_default();
let match_pat = match field_map.remove("rename") {
Some(TokenTree::Literal(v)) => v,
Some(v) => panic!("expected literal, got {v:?}"),
None => create_key_name(&field_name),
};
error_if_map_not_empty(&field_map);
let match_lit = match_pat.to_string();
let ret_stmt = if path.path.segments.first().unwrap().ident == "Option" {
quote! { ret.#field_name = Some(parsed); }
} else {
quote! { ret.#field_name = parsed; }
};
let quoted = quote! {
#match_pat => {
let parsed = val.parse_as_utf8()
.context(concat!("while parsing `", #match_lit, "`"))?;
#ret_stmt
},
};
match_stmts.push(quoted);
}
let ret_val = if use_box {
quote! { Ok(SchRecord::#name(Box::new(ret))) }
} else {
quote! { Ok(SchRecord::#name(ret)) }
};
let ret = quote! {
impl FromRecord for #name {
const RECORD_ID: u32 = #id;
fn from_record<'a, I: Iterator<Item = (&'a [u8], &'a [u8])>>(
records: I,
) -> Result<SchRecord, crate::Error> {
let mut ret = Self::default();
for (key, val) in records {
match key {
#(#match_stmts)*
_ => crate::__private::macro_unsupported_key(stringify!(#name), key, val)
}
}
#ret_val
}
}
};
Ok(ret)
}
#[derive(Clone, Debug, PartialEq)]
enum AttrState {
Key,
Eq(String),
Val(String),
Comma,
}
fn parse_attrs(attrs: Vec<Attribute>) -> Option<BTreeMap<String, TokenTree>> {
let attr = attrs
.into_iter()
.find(|attr| attr.path().is_ident("from_record"))?;
let Meta::List(list) = attr.meta else {
panic!("invalid usage; use `#[from_record(...=..., ...)]`");
};
let mut state = AttrState::Key;
let mut map = BTreeMap::new();
for token in list.tokens {
match state {
AttrState::Key => {
let TokenTree::Ident(idtoken) = token else {
panic!("expected an identifier at {token}");
};
state = AttrState::Eq(idtoken.to_string());
}
AttrState::Eq(key) => {
match token {
TokenTree::Punct(v) if v.as_char() == '=' => (),
_ => panic!("expected `=` at {token}"),
}
state = AttrState::Val(key);
}
AttrState::Val(key) => {
map.insert(key, token);
state = AttrState::Comma;
}
AttrState::Comma => {
match token {
TokenTree::Punct(v) if v.as_char() == ',' => (),
_ => panic!("expected `,` at {token}"),
};
state = AttrState::Key;
}
}
}
Some(map)
}
fn error_if_map_not_empty(map: &BTreeMap<String, TokenTree>) {
assert!(map.is_empty(), "unexpected pairs {map:?}");
}
fn create_key_name(id: &Ident) -> Literal {
const REPLACE: &[(&str, &str)] = &[
("LocationX", "Location.X"),
("LocationY", "Location.Y"),
("CornerX", "Corner.X"),
("CornerY", "Corner.Y"),
("UniqueId", "UniqueID"),
("FontId", "FontID"),
("PartIdLocked", "PartIDLocked"),
("Accessible", "Accesible"),
("Frac", "_Frac"),
];
let mut id_str = id.to_string().to_case(Case::Pascal);
for (from, to) in REPLACE {
id_str = id_str.replace(from, to);
}
Literal::byte_string(id_str.as_bytes())
}