use item::{parse_items, Item};
use proc_macro::TokenStream;
use quote::quote;
use syn::{DataEnum, DeriveInput, Fields, Ident};
mod item;
#[proc_macro_derive(KeyMap, attributes(key))]
pub fn keymap(input: TokenStream) -> TokenStream {
let ast: DeriveInput = syn::parse(input).unwrap();
let syn::Data::Enum(DataEnum { variants, .. }) = ast.data else {
return syn::Error::new_spanned(
ast.ident,
"#[derive(KeyMap)] can only be derived for enums",
)
.to_compile_error()
.into();
};
match parse_items(&variants) {
Ok(items) => {
let config = impl_keymap_config(&ast.ident, &items);
quote! {
#config
}
.into()
}
Err(err) => err.to_compile_error().into(),
}
}
fn impl_keymap_config(name: &Ident, items: &Vec<Item>) -> proc_macro2::TokenStream {
let mut entries = Vec::new();
let mut match_arms = Vec::new();
let mut match_arms_serialize = Vec::new();
let mut match_arms_deserialize = Vec::new();
let mut match_arms_bind = Vec::new();
for item in items {
let ident = &item.variant.ident;
let keys = &item
.keys
.iter()
.map(|key| quote! { #key.to_string() })
.collect::<Vec<_>>();
let doc = &item.description;
let mut char_idx: Option<usize> = None;
if let Some(first_node_seq) = item.nodes.first() {
for (idx, node) in first_node_seq.iter().enumerate() {
if let keymap_parser::node::Key::Group(_) = node.key {
char_idx = Some(idx);
}
}
}
let extract_via_trait = |ty: &syn::Type| -> proc_macro2::TokenStream {
if let Some(idx) = char_idx {
quote! {
match keys.get(#idx) {
Some(node) => <#ty as ::keymap::KeyGroupValue>::from_keymap_node(node),
None => Default::default(),
}
}
} else {
quote! { Default::default() }
}
};
let variant_expr = match &item.variant.fields {
Fields::Unit => quote! { #name::#ident },
Fields::Unnamed(fields) => {
let defaults = fields.unnamed.iter().map(|f| {
if char_idx.is_some() {
extract_via_trait(&f.ty)
} else {
quote! { Default::default() }
}
});
quote! { #name::#ident(#(#defaults),*) }
}
Fields::Named(fields) => {
let defaults = fields.named.iter().map(|f| {
let field_name = f.ident.as_ref().unwrap();
if char_idx.is_some() {
let expr = extract_via_trait(&f.ty);
quote! { #field_name: #expr }
} else {
quote! { #field_name: Default::default() }
}
});
quote! { #name::#ident { #(#defaults),* } }
}
};
let variant_pat = match &item.variant.fields {
Fields::Unit => quote! { #name::#ident },
Fields::Unnamed(_) => quote! { #name::#ident(..) },
Fields::Named(_) => quote! { #name::#ident { .. } },
};
let variant_name_str = ident.to_string();
match_arms_serialize.push(quote! {
#variant_pat => #variant_name_str,
});
if !item.ignore {
match_arms_bind.push(quote! {
#variant_pat => #variant_expr,
});
let variant_expr_default = match &item.variant.fields {
Fields::Unit => quote! { #name::#ident },
Fields::Unnamed(fields) => {
let defaults = fields.unnamed.iter().map(|_| quote! { Default::default() });
quote! { #name::#ident(#(#defaults),*) }
}
Fields::Named(fields) => {
let defaults = fields.named.iter().map(|f| {
let name = &f.ident;
quote! { #name: Default::default() }
});
quote! { #name::#ident { #(#defaults),* } }
}
};
let symbol_opt = match &item.symbol {
Some(sym) => quote! { .with_symbol(Some(#sym)) },
None => quote! {},
};
let help_opt = match &item.help {
Some(h) => quote! { .with_help(Some(#h)) },
None => quote! {},
};
match_arms_deserialize.push(quote! {
#variant_name_str => Ok(#variant_expr_default),
});
match_arms.push(quote! {
#variant_pat => ::keymap::Item::new(
vec![#(#keys),*],
#doc.to_string()
) #symbol_opt #help_opt,
});
entries.push(quote! {
(
#variant_expr_default,
::keymap::Item::new(
vec![#(#keys),*],
#doc.to_string()
) #symbol_opt #help_opt
),
});
}
}
let serde_impls = quote! {
impl ::serde::Serialize for #name {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: ::serde::Serializer,
{
let variant_name = match self {
#(#match_arms_serialize)*
};
serializer.serialize_str(variant_name)
}
}
impl<'de> ::serde::Deserialize<'de> for #name {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: ::serde::Deserializer<'de>,
{
struct EnumVisitor;
impl<'de> ::serde::de::Visitor<'de> for EnumVisitor {
type Value = #name;
fn expecting(&self, formatter: &mut ::std::fmt::Formatter) -> ::std::fmt::Result {
formatter.write_str("a valid variant name for #name")
}
fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
where
E: ::serde::de::Error,
{
match value {
#(#match_arms_deserialize)*
_ => Err(E::unknown_variant(value, &[])),
}
}
}
deserializer.deserialize_str(EnumVisitor)
}
}
};
quote! {
impl ::keymap::KeyMapConfig<#name> for #name {
fn keymap_config() -> ::keymap::Config<#name> {
::keymap::Config::new(vec![#(#entries)*])
}
fn keymap_item(&self) -> ::keymap::Item {
match self {
#(#match_arms)*
_ => ::core::unreachable!("ignored variant has no keymap"),
}
}
fn bind(&self, keys: &[::keymap::KeyMap]) -> Self
where
Self: Clone,
{
match self {
#(#match_arms_bind)*
_ => ::core::unreachable!("ignored variant cannot be bound"),
}
}
}
#serde_impls
}
}