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 { 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 { 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, } fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { let mut encode: Option = None; let mut condition: Option = 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, pub condition: Option, pub skip: bool, } pub(crate) fn parse_variant_attrs(variant: &syn::Variant) -> VariantAttrs { let mut id: Option = None; let mut condition: Option = 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 { attrs.iter().find_map(|attr| { if attr.path().is_ident("context") { let syn::Meta::List(list) = &attr.meta else { return None; }; syn::parse2::(list.tokens.clone()).ok() } else { None } }) } // ── Build impl generics ───────────────────────────────────────────── fn build_encode_impl_generics( params: &syn::punctuated::Punctuated, where_clause: &Option, context_ty: &Option, ) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { let params = params.iter().cloned().collect::>(); let has_params = !params.is_empty(); if let Some(ctx_ty) = context_ty { // With context type: impl<#params> Encode 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 Encode 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, ) -> 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, ) -> 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, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result { let mut len = 0; #encode_fields Ok(len) } } } .into() } fn impl_encode_enum( ident: Ident, generics: Generics, data_enum: DataEnum, context_ty: Option, ) -> 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, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result { let mut len = 0; match self { #(#match_arms),* } Ok(len) } } } .into() }