Skip to main content

gproxy_protocol_macros/
lib.rs

1//! Derives and expression helpers used by `gproxy-protocol`.
2
3use proc_macro::TokenStream;
4use quote::{format_ident, quote};
5use syn::visit_mut::{self, VisitMut};
6use syn::{Data, DeriveInput, Expr, ExprStruct, Fields, Meta, Type, parse_macro_input};
7
8#[proc_macro_derive(WireBuilder, attributes(serde))]
9pub fn derive_wire_builder(input: TokenStream) -> TokenStream {
10    let input = parse_macro_input!(input as DeriveInput);
11    let name = input.ident;
12    let builder = format_ident!("Wire{name}Builder");
13    let generics = input.generics;
14    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
15    let Data::Struct(data) = input.data else {
16        return syn::Error::new_spanned(name, "WireBuilder only supports structs")
17            .to_compile_error()
18            .into();
19    };
20    let Fields::Named(fields) = data.fields else {
21        return syn::Error::new_spanned(name, "WireBuilder requires named fields")
22            .to_compile_error()
23            .into();
24    };
25
26    let fields: Vec<_> = fields.named.into_iter().collect();
27    let names: Vec<_> = fields
28        .iter()
29        .map(|field| field.ident.as_ref().expect("named field"))
30        .collect();
31    let types: Vec<_> = fields.iter().map(|field| &field.ty).collect();
32    let values = fields.iter().map(|field| {
33        let field_name = field.ident.as_ref().expect("named field");
34        if field_has_default(field) {
35            quote! { self.#field_name.unwrap_or_default() }
36        } else {
37            quote! {
38                self.#field_name.ok_or_else(|| crate::WireBuildError::missing(
39                    stringify!(#name),
40                    stringify!(#field_name),
41                ))?
42            }
43        }
44    });
45
46    quote! {
47        #[doc(hidden)]
48        pub struct #builder #generics {
49            #(#names: ::core::option::Option<#types>,)*
50        }
51
52        impl #impl_generics #builder #type_generics #where_clause {
53            #(
54                pub fn #names(mut self, value: #types) -> Self {
55                    self.#names = ::core::option::Option::Some(value);
56                    self
57                }
58            )*
59
60            pub fn build(self) -> ::core::result::Result<#name #type_generics, crate::WireBuildError> {
61                ::core::result::Result::Ok(#name {
62                    #(#names: #values,)*
63                })
64            }
65        }
66
67        impl #impl_generics #name #type_generics #where_clause {
68            pub fn builder() -> #builder #type_generics {
69                #builder {
70                    #(#names: ::core::option::Option::None,)*
71                }
72            }
73
74            #[doc(hidden)]
75            pub fn builder_from(value: Self) -> #builder #type_generics {
76                #builder {
77                    #(#names: ::core::option::Option::Some(value.#names),)*
78                }
79            }
80        }
81    }
82    .into()
83}
84
85fn field_has_default(field: &syn::Field) -> bool {
86    if matches!(
87        &field.ty,
88        Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "Option")
89    ) {
90        return true;
91    }
92    field.attrs.iter().any(|attribute| {
93        if !attribute.path().is_ident("serde") {
94            return false;
95        }
96        let Meta::List(list) = &attribute.meta else {
97            return false;
98        };
99        let tokens = list.tokens.to_string();
100        tokens.split(',').any(|part| {
101            let part = part.trim();
102            part == "default"
103                || part.starts_with("default =")
104                || part == "skip"
105                || part == "skip_deserializing"
106        })
107    })
108}
109
110/// Construct an extensible wire struct with familiar named-field syntax.
111#[proc_macro]
112pub fn wire(input: TokenStream) -> TokenStream {
113    let expression = parse_macro_input!(input as ExprStruct);
114    expand_struct(expression).into()
115}
116
117struct NestedWireRewriter;
118
119impl VisitMut for NestedWireRewriter {
120    fn visit_expr_mut(&mut self, expression: &mut Expr) {
121        visit_mut::visit_expr_mut(self, expression);
122        let Expr::Struct(struct_expression) = expression else {
123            return;
124        };
125        if !looks_like_enum_variant(&struct_expression.path) {
126            *expression = Expr::Verbatim(expand_struct(struct_expression.clone()));
127        }
128    }
129}
130
131fn looks_like_enum_variant(path: &syn::Path) -> bool {
132    path.segments
133        .iter()
134        .rev()
135        .nth(1)
136        .and_then(|segment| segment.ident.to_string().chars().next())
137        .is_some_and(char::is_uppercase)
138}
139
140fn expand_struct(mut expression: ExprStruct) -> proc_macro2::TokenStream {
141    let mut rewriter = NestedWireRewriter;
142    for field in &mut expression.fields {
143        rewriter.visit_expr_mut(&mut field.expr);
144    }
145    if let Some(rest) = expression.rest.as_mut() {
146        rewriter.visit_expr_mut(rest);
147    }
148    if looks_like_enum_variant(&expression.path) {
149        return quote! { #expression };
150    }
151    let path = expression.path;
152    let fields = expression.fields;
153
154    if let Some(rest) = expression.rest {
155        let setters = fields.iter().map(|field| {
156            let member = &field.member;
157            let value = &field.expr;
158            quote! { .#member(#value) }
159        });
160        return quote! {{
161            #path::builder_from(#rest)
162                #(#setters)*
163                .build()
164                .expect(concat!("complete ", stringify!(#path), " wire construction"))
165        }};
166    }
167
168    let setters = fields.iter().map(|field| {
169        let member = &field.member;
170        let value = &field.expr;
171        quote! { .#member(#value) }
172    });
173    quote! {{
174        #path::builder()
175            #(#setters)*
176            .build()
177            .expect(concat!("complete ", stringify!(#path), " wire construction"))
178    }}
179}