gproxy_protocol_macros/
lib.rs1use 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#[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}