summaryrefslogtreecommitdiff
path: root/macros
diff options
context:
space:
mode:
authorzirkonya <zirkonya@iridium.lan>2026-09-01 09:51:18 +0200
committerzirkonya <zirkonya@iridium.lan>2026-09-01 09:51:18 +0200
commitbe62f065a62048b798e8bb2b8e3699bdd5b6d517 (patch)
tree8b3f4b05838b5fa34c3ae24e35004343baadf2af /macros
parentfdc02f07cbd1994c1efb057f24a37a96faaa51fa (diff)
add proc macros ; benchmark ; example
Diffstat (limited to 'macros')
-rw-r--r--macros/Cargo.toml3
-rw-r--r--macros/src/decode.rs508
-rw-r--r--macros/src/encode.rs516
-rw-r--r--macros/src/lib.rs33
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()
+}