summaryrefslogtreecommitdiff
path: root/macros/src/encode.rs
diff options
context:
space:
mode:
Diffstat (limited to 'macros/src/encode.rs')
-rw-r--r--macros/src/encode.rs516
1 files changed, 516 insertions, 0 deletions
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()
+}