1#![warn(missing_docs)]
2use 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)]
30pub 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 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 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 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))]
226pub 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))]
288pub 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 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}