diff options
| author | zirkonya <zirkonya@iridium.lan> | 2026-09-01 09:51:18 +0200 |
|---|---|---|
| committer | zirkonya <zirkonya@iridium.lan> | 2026-09-01 09:51:18 +0200 |
| commit | be62f065a62048b798e8bb2b8e3699bdd5b6d517 (patch) | |
| tree | 8b3f4b05838b5fa34c3ae24e35004343baadf2af /macros | |
| parent | fdc02f07cbd1994c1efb057f24a37a96faaa51fa (diff) | |
add proc macros ; benchmark ; example
Diffstat (limited to 'macros')
| -rw-r--r-- | macros/Cargo.toml | 3 | ||||
| -rw-r--r-- | macros/src/decode.rs | 508 | ||||
| -rw-r--r-- | macros/src/encode.rs | 516 | ||||
| -rw-r--r-- | macros/src/lib.rs | 33 |
4 files changed, 1060 insertions, 0 deletions
diff --git a/macros/Cargo.toml b/macros/Cargo.toml index 3820f50..9a3fcaf 100644 --- a/macros/Cargo.toml +++ b/macros/Cargo.toml @@ -7,3 +7,6 @@ edition = "2024" proc-macro = true [dependencies] +proc-macro2 = "1.0.107" +quote = "1.0.47" +syn = "3.0.3" diff --git a/macros/src/decode.rs b/macros/src/decode.rs new file mode 100644 index 0000000..c289cce --- /dev/null +++ b/macros/src/decode.rs @@ -0,0 +1,508 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use quote::quote; +use syn::{Data, DataEnum, DataStruct, DeriveInput, Generics, Ident, PathArguments, Type}; + +use crate::encode::{next_group, next_ident, parse_variant_attrs, skip_commas}; + +// ── Helper: extract element type from collection ──────────────────── + +fn extract_element_type(ty: &Type) -> Option<Type> { + if let Type::Path(type_path) = ty + && let Some(segment) = type_path.path.segments.last() + { + match segment.ident.to_string().as_str() { + "Vec" | "Option" | "HashSet" | "BTreeSet" | "BinaryHeap" | "LinkedList" + | "VecDeque" => { + if let PathArguments::AngleBracketed(args) = &segment.arguments + && let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() + { + return Some(inner_ty.clone()); + } + } + "HashMap" | "BTreeMap" => { + if let PathArguments::AngleBracketed(args) = &segment.arguments { + let mut types = args.args.iter().filter_map(|arg| { + if let syn::GenericArgument::Type(t) = arg { + Some(t.clone()) + } else { + None + } + }); + if let (Some(k), Some(v)) = (types.next(), types.next()) { + return Some(syn::parse_quote! { (#k, #v) }); + } + } + } + _ => {} + } + } + None +} + +// ── Field attribute parsing ───────────────────────────────────────── + +enum DecodeAttr { + Normal, + Count(syn::Type), + Len(syn::Type), + Custom(syn::Path), +} + +struct FieldAttrs { + decode: DecodeAttr, + element_type: Option<Type>, + condition: Option<syn::Expr>, +} + +fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { + let mut decode: Option<DecodeAttr> = None; + let mut condition: Option<syn::Expr> = None; + + for attr in &field.attrs { + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + decode = Some(DecodeAttr::Normal); + } + "count" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected count(Type)"); + let len_ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in count(...)"); + let element_type = extract_element_type(&field.ty); + decode = Some(DecodeAttr::Count(len_ty)); + return FieldAttrs { + decode: decode.unwrap(), + element_type, + condition, + }; + } + "len" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected len(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in len(...)"); + decode = Some(DecodeAttr::Len(ty)); + } + "with" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected with(path)"); + let path: syn::Path = + syn::parse2(group.stream()).expect("expected path in with(...)"); + decode = Some(DecodeAttr::Custom(path)); + } + "if" => { + assert!(condition.is_none(), "multiple codec(if(...)) on one field"); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper on decode: `{other}`"), + } + } + } + + let element_type = if matches!(decode, Some(DecodeAttr::Count(_))) { + extract_element_type(&field.ty) + } else { + None + }; + + FieldAttrs { + decode: decode.unwrap_or(DecodeAttr::Normal), + element_type, + condition, + } +} + +// ── Context type extraction ───────────────────────────────────────── + +fn extract_context_ty(attrs: &[syn::Attribute]) -> Option<Type> { + attrs.iter().find_map(|attr| { + if attr.path().is_ident("context") { + let syn::Meta::List(list) = &attr.meta else { + return None; + }; + syn::parse2::<Type>(list.tokens.clone()).ok() + } else { + None + } + }) +} + +// ── Build impl generics ───────────────────────────────────────────── + +fn build_decode_impl_generics( + params: &syn::punctuated::Punctuated<syn::GenericParam, syn::Token![,]>, + where_clause: &Option<syn::WhereClause>, + context_ty: &Option<Type>, +) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { + let params = params.iter().cloned().collect::<Vec<_>>(); + let has_params = !params.is_empty(); + + if let Some(ctx_ty) = context_ty { + let impl_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { #ctx_ty }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } else { + let mut all_params = vec![syn::parse_quote! { Data }]; + all_params.extend(params.clone()); + let impl_generics = quote! { <#(#all_params),*> }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { Data }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } +} + +// ── Decode code generation ────────────────────────────────────────── + +pub fn derive_decode( + DeriveInput { + ident, + generics, + data, + attrs, + .. + }: DeriveInput, +) -> TokenStream { + let context_ty = extract_context_ty(&attrs); + match data { + Data::Struct(data_struct) => impl_decode_struct(ident, generics, data_struct, context_ty), + Data::Enum(data_enum) => impl_decode_enum(ident, generics, data_enum, context_ty), + Data::Union(_) => panic!("Not implemented for Union"), + } +} + +fn gen_decode_field( + field_type: &Type, + attrs: &FieldAttrs, + data_param: &TokenStream2, +) -> TokenStream2 { + let inner = match &attrs.decode { + DecodeAttr::Normal => quote! { + <#field_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)? + }, + DecodeAttr::Count(len_type) => { + let element_type = attrs + .element_type + .clone() + .unwrap_or_else(|| field_type.clone()); + quote! { + { + let count: #len_type = <#len_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + let mut items = Vec::with_capacity(count as usize); + for _ in 0..count { + items.push( + <#element_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)? + ); + } + items + } + } + } + DecodeAttr::Len(len_type) => quote! { + { + let byte_len: #len_type = <#len_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + let byte_len = byte_len as usize; + if buf.len() < byte_len { + return Err(zr_protocol::codec::error::CodecError::IoError(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "insufficient bytes for len-prefixed field", + ))); + } + let (mut data, rest) = buf.split_at(byte_len); + *buf = rest; + <#field_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(&mut data, ctx)? + } + }, + DecodeAttr::Custom(fn_path) => quote! { + #fn_path(buf, ctx)? + }, + }; + + match &attrs.condition { + Some(expr) => quote! { + if #expr { + #inner + } else { + Default::default() + } + }, + None => inner, + } +} + +fn impl_decode_struct( + ident: Ident, + generics: Generics, + data_struct: DataStruct, + context_ty: Option<Type>, +) -> TokenStream { + let (impl_generics, ty_generics, where_clause, data_param) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + let decode_fields = match &data_struct.fields { + syn::Fields::Named(fields) => { + let field_decodes: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let ty = &field.ty; + let attrs = parse_field_attrs(field); + let decoded = gen_decode_field(ty, &attrs, &data_param); + quote! { #name: #decoded } + }) + .collect(); + quote! { Ok(Self { #(#field_decodes),* }) } + } + syn::Fields::Unnamed(fields) => { + let field_decodes: Vec<_> = fields + .unnamed + .iter() + .map(|field| { + let ty = &field.ty; + let attrs = parse_field_attrs(field); + gen_decode_field(ty, &attrs, &data_param) + }) + .collect(); + quote! { Ok(Self(#(#field_decodes),*)) } + } + syn::Fields::Unit => quote! { Ok(Self) }, + }; + + quote! { + impl #impl_generics zr_protocol::codec::decode::Decode<#data_param> for #ident #ty_generics #where_clause { + fn decode(buf: &mut &[u8], ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<Self> { + #decode_fields + } + } + } + .into() +} + +fn impl_decode_enum( + ident: Ident, + generics: Generics, + data_enum: DataEnum, + context_ty: Option<Type>, +) -> TokenStream { + let (_, _, _, data_param) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + // Parse all variant attributes + let parsed: Vec<_> = data_enum.variants.iter().map(parse_variant_attrs).collect(); + + // Build explicit ID set and auto-ID counter + let mut explicit_ids = std::collections::HashSet::new(); + for attrs in &parsed { + if let Some(id) = attrs.id { + explicit_ids.insert(id); + } + } + + // Assign discriminants (same logic as encode) + let mut auto_counter: u8 = 0; + let mut discriminants: Vec<u8> = Vec::with_capacity(data_enum.variants.len()); + for attrs in &parsed { + if let Some(id) = attrs.id { + discriminants.push(id); + } else { + while explicit_ids.contains(&auto_counter) { + auto_counter = auto_counter.wrapping_add(1); + } + let d = auto_counter; + auto_counter = auto_counter.wrapping_add(1); + discriminants.push(d); + } + } + + // Group variants by discriminant ID + use std::collections::BTreeMap; + type VariantEntry<'a> = (usize, &'a Ident, &'a syn::Fields, &'a Option<syn::Expr>); + let mut groups: BTreeMap<u8, Vec<VariantEntry>> = BTreeMap::new(); + + for (i, (variant, attrs)) in data_enum.variants.iter().zip(parsed.iter()).enumerate() { + let disc = discriminants[i]; + if attrs.skip { + continue; + } + groups.entry(disc).or_default().push(( + i, + &variant.ident, + &variant.fields, + &attrs.condition, + )); + } + + // Generate decode match arms + let match_arms: Vec<_> = groups + .iter() + .map(|(disc, entries)| { + if entries.len() == 1 { + let (_, variant_ident, fields, condition) = &entries[0]; + let field_decode = gen_variant_decode(&ident, variant_ident, fields, &data_param); + + match condition { + Some(expr) => quote! { + #disc => { + if #expr { + #field_decode + } else { + Err(zr_protocol::codec::error::CodecError::Custom( + concat!("variant ", stringify!(#variant_ident), " not valid in current context").into() + )) + } + } + }, + None => quote! { + #disc => { #field_decode } + }, + } + } else { + let mut arms: Vec<TokenStream2> = Vec::new(); + let mut has_unconditional = false; + + for (_, variant_ident, fields, condition) in entries { + let field_decode = gen_variant_decode(&ident, variant_ident, fields, &data_param); + + match condition { + Some(expr) => { + arms.push(quote! { + if #expr { + #field_decode + } + }); + } + None => { + has_unconditional = true; + arms.push(field_decode); + } + } + } + + if has_unconditional { + let last = arms.pop().unwrap(); + let chain = arms.into_iter().rev().fold(last, |acc, arm| { + quote! { #arm else { #acc } } + }); + quote! { #disc => { #chain } } + } else { + let chain = arms.into_iter().rev().fold( + quote! { + Err(zr_protocol::codec::error::CodecError::Custom( + format!("no variant for ID {:#04x} in current context", #disc).into() + )) + }, + |acc, arm| { + quote! { #arm else { #acc } } + }, + ); + quote! { #disc => { #chain } } + } + } + }) + .collect(); + + let (impl_generics, ty_generics, where_clause, _) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + quote! { + impl #impl_generics zr_protocol::codec::decode::Decode<#data_param> for #ident #ty_generics #where_clause { + fn decode(buf: &mut &[u8], ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<Self> { + let id = <u8 as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + match id { + #(#match_arms),* + other => Err(zr_protocol::codec::error::CodecError::Custom( + format!("unknown discriminant: {other:#04x}").into() + )), + } + } + } + } + .into() +} + +fn gen_variant_decode( + enum_ident: &Ident, + variant_ident: &Ident, + fields: &syn::Fields, + data_param: &TokenStream2, +) -> TokenStream2 { + match fields { + syn::Fields::Named(fields) => { + let field_decodes: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let ty = &field.ty; + let attrs = parse_field_attrs(field); + let decoded = gen_decode_field(ty, &attrs, data_param); + quote! { #name: #decoded } + }) + .collect(); + quote! { + Ok(#enum_ident::#variant_ident { #(#field_decodes),* }) + } + } + syn::Fields::Unnamed(fields) => { + let field_decodes: Vec<_> = fields + .unnamed + .iter() + .map(|field| { + let ty = &field.ty; + let attrs = parse_field_attrs(field); + gen_decode_field(ty, &attrs, data_param) + }) + .collect(); + quote! { + Ok(#enum_ident::#variant_ident(#(#field_decodes),*)) + } + } + syn::Fields::Unit => quote! { + Ok(#enum_ident::#variant_ident) + }, + } +} diff --git a/macros/src/encode.rs b/macros/src/encode.rs new file mode 100644 index 0000000..d8df604 --- /dev/null +++ b/macros/src/encode.rs @@ -0,0 +1,516 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use quote::quote; +use syn::{Data, DataEnum, DataStruct, DeriveInput, Generics, Ident, Type}; + +// ── Shared token helpers ──────────────────────────────────────────── + +pub(crate) fn next_ident(iter: &mut proc_macro2::token_stream::IntoIter) -> Option<Ident> { + if let Some(proc_macro2::TokenTree::Ident(ident)) = iter.clone().next() { + iter.next(); + return Some(ident); + } + None +} + +pub(crate) fn next_group( + iter: &mut proc_macro2::token_stream::IntoIter, +) -> Option<proc_macro2::Group> { + if let Some(proc_macro2::TokenTree::Group(group)) = iter.clone().next() { + iter.next(); + return Some(group); + } + None +} + +pub(crate) fn skip_commas(iter: &mut proc_macro2::token_stream::IntoIter) { + loop { + let peek = iter.clone().next(); + match peek { + Some(proc_macro2::TokenTree::Punct(p)) if p.as_char() == ',' => { + iter.next(); + } + _ => break, + } + } +} + +// ── Field attribute parsing ───────────────────────────────────────── + +enum EncodeAttr { + Normal, + Skip, + Count(syn::Type), + Len(syn::Type), + Custom(syn::Path), +} + +struct FieldAttrs { + encode: EncodeAttr, + condition: Option<syn::Expr>, +} + +fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { + let mut encode: Option<EncodeAttr> = None; + let mut condition: Option<syn::Expr> = None; + + for attr in &field.attrs { + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + encode = Some(EncodeAttr::Skip); + } + "count" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected count(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in count(...)"); + encode = Some(EncodeAttr::Count(ty)); + } + "len" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected len(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in len(...)"); + encode = Some(EncodeAttr::Len(ty)); + } + "with" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected with(path)"); + let path: syn::Path = + syn::parse2(group.stream()).expect("expected path in with(...)"); + encode = Some(EncodeAttr::Custom(path)); + } + "if" => { + assert!(condition.is_none(), "multiple codec(if(...)) on one field"); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper: `{other}`"), + } + } + } + + FieldAttrs { + encode: encode.unwrap_or(EncodeAttr::Normal), + condition, + } +} + +// ── Variant attribute parsing ─────────────────────────────────────── + +pub(crate) struct VariantAttrs { + pub id: Option<u8>, + pub condition: Option<syn::Expr>, + pub skip: bool, +} + +pub(crate) fn parse_variant_attrs(variant: &syn::Variant) -> VariantAttrs { + let mut id: Option<u8> = None; + let mut condition: Option<syn::Expr> = None; + let mut skip = false; + + for attr in &variant.attrs { + // #[id(0x00)] + if attr.path().is_ident("id") { + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + let lit: syn::LitInt = + syn::parse2(list.tokens.clone()).expect("expected u8 literal in #[id(...)]"); + let val: u8 = lit.base10_parse().expect("id must be a u8 value"); + assert!(id.is_none(), "multiple #[id(...)] on one variant"); + id = Some(val); + continue; + } + + // #[codec(if(...)), codec(skip)] + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + skip = true; + } + "if" => { + assert!( + condition.is_none(), + "multiple codec(if(...)) on one variant" + ); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper on variant: `{other}`"), + } + } + } + + VariantAttrs { + id, + condition, + skip, + } +} + +// ── Context type extraction ───────────────────────────────────────── + +fn extract_context_ty(attrs: &[syn::Attribute]) -> Option<Type> { + attrs.iter().find_map(|attr| { + if attr.path().is_ident("context") { + let syn::Meta::List(list) = &attr.meta else { + return None; + }; + syn::parse2::<Type>(list.tokens.clone()).ok() + } else { + None + } + }) +} + +// ── Build impl generics ───────────────────────────────────────────── + +fn build_encode_impl_generics( + params: &syn::punctuated::Punctuated<syn::GenericParam, syn::Token![,]>, + where_clause: &Option<syn::WhereClause>, + context_ty: &Option<Type>, +) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { + let params = params.iter().cloned().collect::<Vec<_>>(); + let has_params = !params.is_empty(); + + if let Some(ctx_ty) = context_ty { + // With context type: impl<#params> Encode<Ctx> for Ident + let impl_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { #ctx_ty }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } else { + // Without context: impl<Data, #params> Encode<Data> for Ident + let mut all_params = vec![syn::parse_quote! { Data }]; + all_params.extend(params.iter().cloned()); + let impl_generics = quote! { <#(#all_params),*> }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { Data }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } +} + +// ── Encode code generation ────────────────────────────────────────── + +pub fn derive_encode( + DeriveInput { + ident, + generics, + data, + attrs, + .. + }: DeriveInput, +) -> TokenStream { + let context_ty = extract_context_ty(&attrs); + match data { + Data::Struct(data_struct) => impl_encode_struct(ident, generics, data_struct, context_ty), + Data::Enum(data_enum) => impl_encode_enum(ident, generics, data_enum, context_ty), + Data::Union(_) => panic!("Not implemented for Union"), + } +} + +fn gen_encode_field( + field_access: TokenStream2, + attrs: &FieldAttrs, + _data_param: &TokenStream2, +) -> TokenStream2 { + let inner = match &attrs.encode { + EncodeAttr::Normal => quote! { + len += #field_access.encode(buffer, ctx)?; + }, + EncodeAttr::Skip => quote! {}, + EncodeAttr::Count(len_type) => quote! { + let count: #len_type = #field_access.len() as #len_type; + len += count.encode(buffer, ctx)?; + for item in &#field_access { + len += item.encode(buffer, ctx)?; + } + }, + EncodeAttr::Len(len_type) => quote! { + let mut __buf = Vec::with_capacity(2048); + #field_access.encode(&mut __buf, ctx)?; + let __len: #len_type = zr_protocol::types::size::Size::from_size(__buf.len()); + len += __len.encode(buffer, ctx)?; + buffer.extend_from_slice(&__buf); + len += __buf.len(); + }, + EncodeAttr::Custom(fn_path) => quote! { + len += #fn_path(&#field_access, buffer, ctx)?; + }, + }; + + match &attrs.condition { + Some(expr) => quote! { + if #expr { + #inner + } + }, + None => inner, + } +} + +fn gen_encode_body( + enum_ident: &Ident, + variant_ident: &Ident, + pattern: &TokenStream2, + encode_fields: &TokenStream2, + discriminant: u8, + condition: &Option<syn::Expr>, +) -> TokenStream2 { + let encode_disc_and_fields = quote! { + len += #discriminant.encode(buffer, ctx)?; + #encode_fields + }; + + match condition { + Some(expr) => quote! { + #enum_ident::#variant_ident #pattern => { + if #expr { + #encode_disc_and_fields + } else { + return Err(zr_protocol::codec::error::CodecError::Custom( + concat!("variant ", stringify!(#variant_ident), " not valid in current context").into() + )); + } + } + }, + None => quote! { + #enum_ident::#variant_ident #pattern => { + #encode_disc_and_fields + } + }, + } +} + +fn impl_encode_struct( + ident: Ident, + generics: Generics, + data_struct: DataStruct, + context_ty: Option<Type>, +) -> TokenStream { + let (impl_generics, ty_generics, where_clause, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + let encode_fields = match &data_struct.fields { + syn::Fields::Named(fields) => { + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#name }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unnamed(fields) => { + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let idx = syn::Index::from(i); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#idx }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unit => quote! {}, + }; + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + #encode_fields + Ok(len) + } + } + } + .into() +} + +fn impl_encode_enum( + ident: Ident, + generics: Generics, + data_enum: DataEnum, + context_ty: Option<Type>, +) -> TokenStream { + let (_, _, _, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + // First pass: collect explicit IDs and count auto-IDs needed + let mut explicit_ids = std::collections::HashSet::new(); + + // Pre-parse all variant attributes + let parsed: Vec<_> = data_enum + .variants + .iter() + .map(|v| { + let attrs = parse_variant_attrs(v); + if let Some(id) = attrs.id { + explicit_ids.insert(id); + } + attrs + }) + .collect(); + + let match_arms: Vec<_> = { + let mut auto_counter: u8 = 0; + data_enum + .variants + .iter() + .zip(parsed.iter()) + .map(|(variant, attrs)| { + let variant_ident = &variant.ident; + + if attrs.skip { + return quote! { + #ident::#variant_ident { .. } => { + unreachable!("skipped variant") + } + }; + } + + // Resolve discriminant: explicit #[id(...)] or auto-increment + let discriminant = if let Some(id) = attrs.id { + id + } else { + while explicit_ids.contains(&auto_counter) { + auto_counter = auto_counter.wrapping_add(1); + } + let d = auto_counter; + auto_counter = auto_counter.wrapping_add(1); + d + }; + + let (pattern, encode_fields) = match &variant.fields { + syn::Fields::Named(fields) => { + let names: Vec<_> = fields + .named + .iter() + .map(|f| f.ident.as_ref().unwrap()) + .collect(); + let pattern = quote! { { #(#names),* } }; + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unnamed(fields) => { + let names: Vec<_> = (0..fields.unnamed.len()) + .map(|i| Ident::new(&format!("_{}", i), proc_macro2::Span::call_site())) + .collect(); + let pattern = quote! { ( #(#names),* ) }; + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let name = &names[i]; + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unit => (quote! {}, quote! {}), + }; + + gen_encode_body( + &ident, + variant_ident, + &pattern, + &encode_fields, + discriminant, + &attrs.condition, + ) + }) + .collect() + }; + + let (impl_generics, ty_generics, where_clause, _) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + match self { + #(#match_arms),* + } + Ok(len) + } + } + } + .into() +} diff --git a/macros/src/lib.rs b/macros/src/lib.rs index 8b13789..f20edae 100644 --- a/macros/src/lib.rs +++ b/macros/src/lib.rs @@ -1 +1,34 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use syn::DeriveInput; +mod decode; +mod encode; + +#[proc_macro_derive(Encode, attributes(codec, id, context))] +pub fn encode(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + encode::derive_encode(input) +} + +#[proc_macro_derive(Decode, attributes(codec, id, context))] +pub fn decode(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + decode::derive_decode(input) +} + +#[proc_macro_derive(Codec, attributes(codec, id, context))] +pub fn codec(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + let encode_impl = encode::derive_encode(input.clone()); + let decode_impl = decode::derive_decode(input.clone()); + + let encode_ts: TokenStream2 = encode_impl.into(); + let decode_ts: TokenStream2 = decode_impl.into(); + + quote::quote! { + #encode_ts + #decode_ts + } + .into() +} |
