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