Skip to main content

libmaxminddb_rs_derive/
lib.rs

1#![warn(missing_docs)]
2//! Derive macros for `libmaxminddb-rs` records.
3//!
4//! This package is normally consumed through the main crate's `derive` feature.
5//! The macros generate implementations for the public traits re-exported there.
6
7use proc_macro::TokenStream;
8use quote::{format_ident, quote};
9use syn::ext::IdentExt;
10use syn::{Data, DeriveInput, Field, Fields, GenericParam, LitByteStr, parse_macro_input};
11
12fn is_network_field(field: &Field) -> syn::Result<bool> {
13    let mut network = false;
14    for attr in &field.attrs {
15        if attr.path().is_ident("mmdb") {
16            attr.parse_nested_meta(|meta| {
17                if meta.path.is_ident("network") {
18                    network = true;
19                    Ok(())
20                } else {
21                    Err(meta.error("unsupported mmdb attribute; expected `network`"))
22                }
23            })?;
24        }
25    }
26    Ok(network)
27}
28
29#[proc_macro_derive(MmdbDecode)]
30/// Derives borrowed MMDB map decoding for a struct.
31///
32/// Fields such as `&str` and `&[u8]` borrow the source database bytes.
33/// Unsupported or missing required fields produce a decoding error.
34///
35/// # Examples
36///
37/// ```ignore
38/// use libmaxminddb_rs::MmdbDecode;
39/// #[derive(MmdbDecode)]
40/// struct Record<'a> { country: &'a str, asn: u32 }
41/// ```
42pub fn derive_decode(input: TokenStream) -> TokenStream {
43    let input = parse_macro_input!(input as DeriveInput);
44    derive_decode_impl(input).into()
45}
46
47fn derive_decode_impl(input: DeriveInput) -> proc_macro2::TokenStream {
48    let name = input.ident;
49    let Data::Struct(data) = input.data else {
50        return syn::Error::new_spanned(name, "MmdbDecode only supports structs")
51            .to_compile_error();
52    };
53    let Fields::Named(fields) = data.fields else {
54        return syn::Error::new_spanned(name, "MmdbDecode requires named fields")
55            .to_compile_error();
56    };
57
58    let lifetimes: Vec<_> = input
59        .generics
60        .params
61        .iter()
62        .filter_map(|p| match p {
63            GenericParam::Lifetime(l) => Some(l.lifetime.clone()),
64            _ => None,
65        })
66        .collect();
67    if input
68        .generics
69        .params
70        .iter()
71        .any(|p| !matches!(p, GenericParam::Lifetime(_)))
72    {
73        return syn::Error::new_spanned(
74            name,
75            "MmdbDecode currently supports lifetime generics only",
76        )
77        .to_compile_error();
78    }
79
80    let decode_lt: syn::Lifetime = lifetimes
81        .first()
82        .cloned()
83        .unwrap_or_else(|| syn::parse_quote!('__mmdb));
84    let impl_generics = if lifetimes.is_empty() {
85        quote!(<'__mmdb>)
86    } else {
87        let ls = &lifetimes;
88        quote!(<#(#ls),*>)
89    };
90    let ty_generics = if lifetimes.is_empty() {
91        quote!()
92    } else {
93        let ls = &lifetimes;
94        quote!(<#(#ls),*>)
95    };
96
97    // Decode generated structs with one linear map pass instead of calling
98    // `ValueRef::get` once per field (O(fields * entries)).  Each slot only accepts
99    // its first match, preserving the previous first-duplicate-wins behaviour.
100    let field_count = fields.named.len();
101    let field_info: Vec<_> = fields
102        .named
103        .iter()
104        .map(|f| {
105            let ident = f.ident.as_ref().expect("named field").clone();
106            // `r#type` must match the MMDB key `type`, not `r#type`.
107            let key = ident.unraw().to_string();
108            let slot = format_ident!("__mmdb_field_{}", key);
109            let ty = f.ty.clone();
110            (ident, slot, key, ty)
111        })
112        .collect();
113    let declarations = field_info.iter().map(|(_, slot, _, _)| {
114        quote! { let mut #slot = None; }
115    });
116    let match_arms = field_info.iter().map(|(_, slot, key, _)| {
117        quote! {
118            #key if #slot.is_none() => {
119                #slot = Some(__mmdb_value);
120                __mmdb_matched += 1;
121            }
122        }
123    });
124    let initializers = field_info.iter().map(|(ident, slot, _, ty)| {
125        quote! {
126            #ident: <#ty as ::libmaxminddb_rs::DecodeField<#decode_lt>>::decode_field(#slot)?
127        }
128    });
129
130    // Single-pass decoder over the encoded map: each key is compared as raw
131    // bytes with the field names, matching values are decoded in place, and
132    // unknown or duplicate entries are skipped without being decoded. The loop
133    // stops as soon as every field is filled; `finish_map` then skips the tail
134    // only when the parent still has to read past it.
135    let raw_declarations = field_info.iter().map(|(_, slot, _, ty)| {
136        quote! { let mut #slot: ::core::option::Option<#ty> = ::core::option::Option::None; }
137    });
138    let raw_arms = field_info.iter().map(|(_, slot, key, ty)| {
139        let key_bytes = LitByteStr::new(key.as_bytes(), proc_macro2::Span::call_site());
140        quote! {
141            #key_bytes if #slot.is_none() => {
142                #slot = ::core::option::Option::Some(
143                    <#ty as ::libmaxminddb_rs::DecodeField<#decode_lt>>::decode_raw(__mmdb_decoder)?,
144                );
145                __mmdb_matched += 1;
146                if __mmdb_matched == #field_count {
147                    break;
148                }
149            }
150        }
151    });
152    let raw_initializers = field_info.iter().map(|(ident, slot, _, ty)| {
153        quote! {
154            #ident: match #slot {
155                ::core::option::Option::Some(value) => value,
156                ::core::option::Option::None => {
157                    <#ty as ::libmaxminddb_rs::DecodeField<#decode_lt>>::decode_missing()?
158                }
159            }
160        }
161    });
162
163    quote! {
164        impl #impl_generics ::libmaxminddb_rs::MmdbDecode<#decode_lt> for #name #ty_generics {
165            fn decode(value: &::libmaxminddb_rs::ValueRef<#decode_lt>) -> ::libmaxminddb_rs::Result<Self> {
166                let __mmdb_entries = match value {
167                    ::libmaxminddb_rs::ValueRef::Map(entries) => entries,
168                    _ => return Err(::libmaxminddb_rs::Error::DecodingError("MmdbDecode expected a map".into())),
169                };
170                #(#declarations)*
171                let mut __mmdb_matched = 0usize;
172                for (__mmdb_key, __mmdb_value) in __mmdb_entries {
173                    match *__mmdb_key {
174                        #(#match_arms,)*
175                        _ => {}
176                    }
177                    if __mmdb_matched == #field_count {
178                        break;
179                    }
180                }
181                Ok(Self { #(#initializers),* })
182            }
183
184            #[inline]
185            #[allow(unused_mut, unused_variables)]
186            fn decode_raw(
187                __mmdb_decoder: &mut ::libmaxminddb_rs::__private::RawDecoder<#decode_lt>,
188            ) -> ::libmaxminddb_rs::Result<Self> {
189                let __mmdb_map = __mmdb_decoder.enter_map("MmdbDecode expected a map")?;
190                #(#raw_declarations)*
191                let mut __mmdb_matched = 0usize;
192                let mut __mmdb_remaining = __mmdb_map.len();
193                while __mmdb_remaining != 0 {
194                    __mmdb_remaining -= 1;
195                    match __mmdb_decoder.read_key()? {
196                        #(#raw_arms)*
197                        _ => __mmdb_decoder.skip_value()?,
198                    }
199                }
200                __mmdb_decoder.finish_map(__mmdb_map, __mmdb_remaining)?;
201                Ok(Self { #(#raw_initializers),* })
202            }
203        }
204
205        impl #impl_generics ::libmaxminddb_rs::DecodeField<#decode_lt> for #name #ty_generics {
206            fn decode_field(
207                value: Option<&::libmaxminddb_rs::ValueRef<#decode_lt>>,
208            ) -> ::libmaxminddb_rs::Result<Self> {
209                let value = value.ok_or_else(|| {
210                    ::libmaxminddb_rs::Error::DecodingError("missing nested MMDB struct".into())
211                })?;
212                <Self as ::libmaxminddb_rs::MmdbDecode<#decode_lt>>::decode(value)
213            }
214
215            #[inline]
216            fn decode_raw(
217                decoder: &mut ::libmaxminddb_rs::__private::RawDecoder<#decode_lt>,
218            ) -> ::libmaxminddb_rs::Result<Self> {
219                <Self as ::libmaxminddb_rs::MmdbDecode<#decode_lt>>::decode_raw(decoder)
220            }
221        }
222    }
223}
224
225#[proc_macro_derive(MmdbEncode, attributes(mmdb))]
226/// Derives encoding of struct fields to an owned MMDB value.
227///
228/// `Option::None` fields are omitted because MMDB has no null data type.
229///
230/// # Examples
231///
232/// ```ignore
233/// use libmaxminddb_rs::MmdbEncode;
234/// #[derive(MmdbEncode)]
235/// struct Record<'a> { country: &'a str, asn: u32 }
236/// ```
237pub fn derive_encode(input: TokenStream) -> TokenStream {
238    let input = parse_macro_input!(input as DeriveInput);
239    derive_encode_impl(input).into()
240}
241
242fn derive_encode_impl(input: DeriveInput) -> proc_macro2::TokenStream {
243    let name = input.ident;
244    let generics = input.generics;
245    let Data::Struct(data) = input.data else {
246        return syn::Error::new_spanned(name, "MmdbEncode only supports structs")
247            .to_compile_error();
248    };
249    let Fields::Named(fields) = data.fields else {
250        return syn::Error::new_spanned(name, "MmdbEncode requires named fields")
251            .to_compile_error();
252    };
253
254    let mut inserts = Vec::new();
255    for field in &fields.named {
256        match is_network_field(field) {
257            Ok(true) => continue,
258            Ok(false) => {}
259            Err(error) => return error.to_compile_error(),
260        }
261        let ident = field.ident.as_ref().expect("named field");
262        let key = ident.to_string();
263        inserts.push(quote! {
264            if let Some(value) = ::libmaxminddb_rs::EncodeField::encode_optional_field(&self.#ident)? {
265                map.insert(#key.to_owned(), value);
266            }
267        });
268    }
269    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
270    quote! {
271        impl #impl_generics ::libmaxminddb_rs::MmdbEncode for #name #ty_generics #where_clause {
272            fn encode(&self) -> ::libmaxminddb_rs::Result<::libmaxminddb_rs::Value> {
273                let mut map = ::std::collections::BTreeMap::new();
274                #(#inserts)*
275                Ok(::libmaxminddb_rs::Value::Map(map))
276            }
277        }
278
279        impl #impl_generics ::libmaxminddb_rs::EncodeField for #name #ty_generics #where_clause {
280            fn encode_field(&self) -> ::libmaxminddb_rs::Result<::libmaxminddb_rs::Value> {
281                <Self as ::libmaxminddb_rs::MmdbEncode>::encode(self)
282            }
283        }
284    }
285}
286
287#[proc_macro_derive(MmdbRecord, attributes(mmdb))]
288/// Derives the network key for one-object writer insertion.
289///
290/// Mark exactly one `IpNetwork` field with `#[mmdb(network)]`; that field is
291/// excluded from the encoded payload when paired with `MmdbEncode`.
292///
293/// # Examples
294///
295/// ```ignore
296/// use libmaxminddb_rs::{IpNetwork, MmdbEncode, MmdbRecord};
297/// #[derive(MmdbEncode, MmdbRecord)]
298/// struct Record { #[mmdb(network)] network: IpNetwork, asn: u32 }
299/// ```
300pub fn derive_record(input: TokenStream) -> TokenStream {
301    let input = parse_macro_input!(input as DeriveInput);
302    derive_record_impl(input).into()
303}
304
305fn derive_record_impl(input: DeriveInput) -> proc_macro2::TokenStream {
306    let name = input.ident;
307    let generics = input.generics;
308    let Data::Struct(data) = input.data else {
309        return syn::Error::new_spanned(name, "MmdbRecord only supports structs")
310            .to_compile_error();
311    };
312    let Fields::Named(fields) = data.fields else {
313        return syn::Error::new_spanned(name, "MmdbRecord requires named fields")
314            .to_compile_error();
315    };
316
317    let mut network_field = None;
318    for field in &fields.named {
319        match is_network_field(field) {
320            Ok(true) if network_field.is_none() => network_field = field.ident.clone(),
321            Ok(true) => {
322                return syn::Error::new_spanned(
323                    field,
324                    "MmdbRecord requires exactly one #[mmdb(network)] field",
325                )
326                .to_compile_error();
327            }
328            Ok(false) => {}
329            Err(error) => return error.to_compile_error(),
330        }
331    }
332    let Some(network_field) = network_field else {
333        return syn::Error::new_spanned(name, "MmdbRecord requires one #[mmdb(network)] field")
334            .to_compile_error();
335    };
336    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
337    quote! {
338        impl #impl_generics ::libmaxminddb_rs::MmdbRecord for #name #ty_generics #where_clause {
339            fn network(&self) -> ::libmaxminddb_rs::IpNetwork {
340                ::core::clone::Clone::clone(&self.#network_field)
341            }
342        }
343    }
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349    use syn::parse_quote;
350
351    /// The `syn::parse_quote!` inputs below are plain `DeriveInput`s, so the
352    /// expansion helpers run directly in the test harness without needing a
353    /// `proc_macro::TokenStream` (which cannot be constructed off the
354    /// compiler's proc-macro bridge).
355    fn tokens(input: DeriveInput) -> String {
356        derive_decode_impl(input).to_string()
357    }
358
359    fn tokens_encode(input: DeriveInput) -> String {
360        derive_encode_impl(input).to_string()
361    }
362
363    fn tokens_record(input: DeriveInput) -> String {
364        derive_record_impl(input).to_string()
365    }
366
367    #[test]
368    fn decode_plain_struct_without_generics() {
369        let out = tokens(parse_quote! {
370            struct Plain {
371                a: u32,
372                b: String,
373            }
374        });
375        assert!(
376            out.contains("impl < '__mmdb > :: libmaxminddb_rs :: MmdbDecode < '__mmdb > for Plain")
377        );
378        assert!(
379            out.contains(
380                "impl < '__mmdb > :: libmaxminddb_rs :: DecodeField < '__mmdb > for Plain"
381            )
382        );
383        assert!(out.contains(":: libmaxminddb_rs :: ValueRef < '__mmdb >"));
384        assert!(!out.contains("compile_error"));
385    }
386
387    #[test]
388    fn decode_struct_with_multiple_lifetimes() {
389        let out = tokens(parse_quote! {
390            struct Multi<'a, 'b> {
391                first: &'a str,
392                second: &'b str,
393            }
394        });
395        assert!(out.contains(
396            "impl < 'a , 'b > :: libmaxminddb_rs :: MmdbDecode < 'a > for Multi < 'a , 'b >"
397        ));
398        assert!(!out.contains("compile_error"));
399    }
400
401    #[test]
402    fn decode_enum_is_rejected() {
403        let out = tokens(parse_quote! {
404            enum Wrong {}
405        });
406        assert!(out.contains("MmdbDecode only supports structs"));
407        assert!(out.contains("compile_error"));
408    }
409
410    #[test]
411    fn decode_unamed_fields_are_rejected() {
412        let out = tokens(parse_quote! {
413            struct Tuple(u32);
414        });
415        assert!(out.contains("MmdbDecode requires named fields"));
416    }
417
418    #[test]
419    fn decode_type_generics_are_rejected() {
420        let out = tokens(parse_quote! {
421            struct Generic<T> {
422                x: T,
423            }
424        });
425        assert!(out.contains("MmdbDecode currently supports lifetime generics only"));
426    }
427
428    #[test]
429    fn encode_plain_struct() {
430        let out = tokens_encode(parse_quote! {
431            struct Out<'a> {
432                a: &'a str,
433                b: Option<u32>,
434            }
435        });
436        assert!(out.contains("impl < 'a > :: libmaxminddb_rs :: MmdbEncode for Out < 'a >"));
437        assert!(out.contains("impl < 'a > :: libmaxminddb_rs :: EncodeField for Out < 'a >"));
438        assert!(out.contains("BTreeMap"));
439        assert!(!out.contains("compile_error"));
440    }
441
442    #[test]
443    fn encode_skips_network_field() {
444        let out = tokens_encode(parse_quote! {
445            struct Entry {
446                #[mmdb(network)]
447                network: String,
448                payload: u32,
449            }
450        });
451        assert!(!out.contains("network"));
452        assert!(out.contains("payload"));
453        assert!(!out.contains("compile_error"));
454    }
455
456    #[test]
457    fn encode_enum_is_rejected() {
458        let out = tokens_encode(parse_quote! {
459            enum Wrong {}
460        });
461        assert!(out.contains("MmdbEncode only supports structs"));
462    }
463
464    #[test]
465    fn encode_unamed_fields_are_rejected() {
466        let out = tokens_encode(parse_quote! {
467            struct Tuple(u32);
468        });
469        assert!(out.contains("MmdbEncode requires named fields"));
470    }
471
472    #[test]
473    fn encode_unknown_mmdb_attribute_is_rejected() {
474        let out = tokens_encode(parse_quote! {
475            struct Bad {
476                #[mmdb(other)]
477                a: u32,
478            }
479        });
480        assert!(out.contains("unsupported mmdb attribute"));
481    }
482
483    #[test]
484    fn record_plain_struct() {
485        let out = tokens_record(parse_quote! {
486            struct R {
487                #[mmdb(network)]
488                network: String,
489                a: u32,
490            }
491        });
492        assert!(out.contains("impl :: libmaxminddb_rs :: MmdbRecord for R"));
493        assert!(out.contains(". network"));
494        assert!(!out.contains("compile_error"));
495    }
496
497    #[test]
498    fn record_enum_is_rejected() {
499        let out = tokens_record(parse_quote! {
500            enum Wrong {}
501        });
502        assert!(out.contains("MmdbRecord only supports structs"));
503    }
504
505    #[test]
506    fn record_unamed_fields_are_rejected() {
507        let out = tokens_record(parse_quote! {
508            struct Tuple(u32);
509        });
510        assert!(out.contains("MmdbRecord requires named fields"));
511    }
512
513    #[test]
514    fn record_requires_exactly_one_network_field() {
515        let out = tokens_record(parse_quote! {
516            struct Bad {
517                #[mmdb(network)]
518                network: String,
519                #[mmdb(network)]
520                also_network: String,
521            }
522        });
523        assert!(out.contains("MmdbRecord requires exactly one #[mmdb(network)] field"));
524    }
525
526    #[test]
527    fn record_requires_a_network_field() {
528        let out = tokens_record(parse_quote! {
529            struct Bad {
530                a: u32,
531            }
532        });
533        assert!(out.contains("MmdbRecord requires one #[mmdb(network)] field"));
534    }
535
536    #[test]
537    fn record_unknown_mmdb_attribute_is_rejected() {
538        let out = tokens_record(parse_quote! {
539            struct Bad {
540                #[mmdb(other)]
541                a: u32,
542            }
543        });
544        assert!(out.contains("unsupported mmdb attribute"));
545    }
546
547    #[test]
548    fn malformed_network_attribute_is_rejected() {
549        let input: DeriveInput = parse_quote! {
550            struct Bad {
551                #[mmdb(network = true)]
552                network: String,
553            }
554        };
555        let out = tokens_record(input.clone());
556        assert!(out.contains("compile_error"));
557        assert!(!out.contains("impl :: libmaxminddb_rs :: MmdbRecord"));
558        assert!(tokens_encode(input).contains("compile_error"));
559    }
560}