1#![warn(missing_docs)]
2use 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))]
89pub 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 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 let key = decode_field_key(f)?;
169 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 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))]
292pub 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))]
354pub 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 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}