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 { 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, condition: Option, } fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { let mut decode: 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!(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 { 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_decode_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 { 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, ) -> 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 { #decode_fields } } } .into() } fn impl_decode_enum( ident: Ident, generics: Generics, data_enum: DataEnum, context_ty: Option, ) -> 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 = 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); let mut groups: BTreeMap> = 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 = 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 { let id = >::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) }, } }