Skip to main content

forte_macros/
lib.rs

1//! # Integer Key Formatting for PK/SK
2//!
3//! When integer types are used as PK (Partition Key) or SK (Sort Key) fields,
4//! they are zero-padded to their maximum decimal digit width so that
5//! lexicographic (string) sort order matches numeric sort order.
6//!
7//! ## Unsigned types
8//!
9//! | Type          | Width | Example              |
10//! |---------------|-------|----------------------|
11//! | `u8`          | 3     | `042`                |
12//! | `u16`         | 5     | `00042`              |
13//! | `u32`         | 10    | `0000000042`         |
14//! | `u64`/`usize` | 20   | `00000000000000000042` |
15//!
16//! ## Signed types (offset encoding)
17//!
18//! Signed integers are converted to unsigned by adding an offset equal to
19//! `|T::MIN|` (i.e. `2^(bits-1)`), then zero-padded to the same width as the
20//! corresponding unsigned type. This maps the full signed range onto `0..=U::MAX`
21//! while preserving numeric order.
22//!
23//! | Type          | Offset           | Width | MIN → | 0 →  | MAX →  |
24//! |---------------|------------------|-------|-------|------|--------|
25//! | `i8`          | 128              | 3     | `000` | `128` | `255` |
26//! | `i16`         | 32768            | 5     | `00000` | `32768` | `65535` |
27//! | `i32`         | 2147483648       | 10    | `0000000000` | `2147483648` | `4294967295` |
28//! | `i64`/`isize` | 9223372036854775808 | 20 | `00000000000000000000` | `09223372036854775808` | `18446744073709551615` |
29//!
30//! `usize` is always treated as `u64`, and `isize` as `i64`.
31
32use proc_macro::TokenStream;
33use quote::{format_ident, quote, quote_spanned};
34use syn::{Fields, ItemFn, ItemStruct, parse_macro_input, spanned::Spanned};
35
36#[proc_macro_attribute]
37pub fn test(_attr: TokenStream, item: TokenStream) -> TokenStream {
38    let input = parse_macro_input!(item as ItemFn);
39
40    if input.sig.asyncness.is_none() {
41        return quote_spanned! { input.sig.fn_token.span()=>
42            compile_error!("fn must be `async fn`");
43        }
44        .into();
45    }
46
47    if !input.sig.inputs.is_empty() {
48        return quote_spanned! { input.sig.inputs.span()=>
49            compile_error!("arguments to test functions are not supported");
50        }
51        .into();
52    }
53
54    let name = input.sig.ident;
55    let attrs = input.attrs;
56    let output = input.sig.output;
57    let block = input.block;
58    quote! {
59        #(#attrs)*
60        async fn #name() #output {
61            #block
62        }
63
64        ::forte_sdk::inventory::submit! {
65            ::forte_sdk::test_harness::RegisteredTest {
66                module_path: ::core::module_path!(),
67                name: ::core::stringify!(#name),
68                run: || ::std::boxed::Box::pin(async {
69                    ::forte_sdk::test_harness::TestOutcome::into_outcome(#name().await)
70                }),
71            }
72        }
73    }
74    .into()
75}
76
77#[proc_macro_attribute]
78pub fn cache_static(_attr: TokenStream, item: TokenStream) -> TokenStream {
79    let input = parse_macro_input!(item as syn::ItemFn);
80    quote!(#input).into()
81}
82
83fn format_placeholder(ty: &syn::Type) -> String {
84    if let syn::Type::Path(type_path) = ty
85        && let Some(segment) = type_path.path.segments.last()
86    {
87        match segment.ident.to_string().as_str() {
88            "u8" | "i8" => return "{:03}".to_string(),
89            "u16" | "i16" => return "{:05}".to_string(),
90            "u32" | "i32" => return "{:010}".to_string(),
91            "u64" | "i64" | "usize" | "isize" => return "{:020}".to_string(),
92            _ => {}
93        }
94    }
95    "{}".to_string()
96}
97
98fn wrap_expr(ty: &syn::Type, expr: proc_macro2::TokenStream) -> proc_macro2::TokenStream {
99    if let syn::Type::Path(type_path) = ty
100        && let Some(segment) = type_path.path.segments.last()
101    {
102        match segment.ident.to_string().as_str() {
103            "i8" => return quote! { (#expr as u8).wrapping_add(128u8) },
104            "i16" => return quote! { (#expr as u16).wrapping_add(32768u16) },
105            "i32" => return quote! { (#expr as u32).wrapping_add(2147483648u32) },
106            "i64" | "isize" => {
107                return quote! { (#expr as u64).wrapping_add(9223372036854775808u64) };
108            }
109            _ => {}
110        }
111    }
112    expr
113}
114
115fn is_string_type(ty: &syn::Type) -> bool {
116    if let syn::Type::Path(type_path) = ty
117        && let Some(segment) = type_path.path.segments.last()
118    {
119        return segment.ident == "String";
120    }
121    false
122}
123
124fn make_generics(
125    pk_is_string: &[bool],
126    sk_is_string: &[bool],
127) -> (
128    Vec<Option<proc_macro2::Ident>>,
129    Vec<Option<proc_macro2::Ident>>,
130) {
131    let mut counter = 0usize;
132    let pk = pk_is_string
133        .iter()
134        .map(|&s| {
135            if s {
136                let ident = format_ident!("__T{}", counter);
137                counter += 1;
138                Some(ident)
139            } else {
140                None
141            }
142        })
143        .collect();
144    let sk = sk_is_string
145        .iter()
146        .map(|&s| {
147            if s {
148                let ident = format_ident!("__T{}", counter);
149                counter += 1;
150                Some(ident)
151            } else {
152                None
153            }
154        })
155        .collect();
156    (pk, sk)
157}
158
159#[proc_macro_attribute]
160pub fn forte_doc(_attr: TokenStream, item: TokenStream) -> TokenStream {
161    let input = parse_macro_input!(item as ItemStruct);
162
163    let name = &input.ident;
164    let vis = &input.vis;
165    let get_name = format_ident!("{}Get", name);
166    let put_name = format_ident!("{}Put", name);
167    let query_name = format_ident!("{}Query", name);
168    let delete_name = format_ident!("{}Delete", name);
169
170    let fields = match &input.fields {
171        Fields::Named(fields) => &fields.named,
172        _ => panic!("forte_doc only supports named fields"),
173    };
174
175    let pk_fields: Vec<_> = fields
176        .iter()
177        .filter(|f| f.attrs.iter().any(|a| a.path().is_ident("pk")))
178        .collect();
179
180    let sk_fields: Vec<_> = fields
181        .iter()
182        .filter(|f| f.attrs.iter().any(|a| a.path().is_ident("sk")))
183        .collect();
184
185    let pk_field_names: Vec<_> = pk_fields.iter().map(|f| &f.ident).collect();
186    let pk_field_types: Vec<_> = pk_fields.iter().map(|f| &f.ty).collect();
187    let sk_field_names: Vec<_> = sk_fields.iter().map(|f| &f.ident).collect();
188    let sk_field_types: Vec<_> = sk_fields.iter().map(|f| &f.ty).collect();
189
190    let pk_is_string: Vec<bool> = pk_field_types.iter().map(|ty| is_string_type(ty)).collect();
191    let sk_is_string: Vec<bool> = sk_field_types.iter().map(|ty| is_string_type(ty)).collect();
192
193    let (gpk, gsk) = make_generics(&pk_is_string, &sk_is_string);
194    let all_generics: Vec<_> = gpk
195        .iter()
196        .chain(gsk.iter())
197        .filter_map(|g| g.as_ref())
198        .collect();
199
200    let generic_def = if all_generics.is_empty() {
201        quote! {}
202    } else {
203        quote! { <#(#all_generics: AsRef<str>),*> }
204    };
205    let generic_use = if all_generics.is_empty() {
206        quote! {}
207    } else {
208        quote! { <#(#all_generics),*> }
209    };
210
211    let query_generics: Vec<_> = gpk.iter().filter_map(|g| g.as_ref()).collect();
212    let query_generic_def = if query_generics.is_empty() {
213        quote! {}
214    } else {
215        quote! { <#(#query_generics: AsRef<str>),*> }
216    };
217    let query_generic_use = if query_generics.is_empty() {
218        quote! {}
219    } else {
220        quote! { <#(#query_generics),*> }
221    };
222
223    let get_pk_fields: Vec<_> = pk_field_names
224        .iter()
225        .zip(pk_field_types.iter())
226        .zip(gpk.iter())
227        .map(|((name, ty), gp)| {
228            let field_name = name.as_ref().unwrap();
229            if let Some(g) = gp {
230                quote! { pub #field_name: #g }
231            } else {
232                quote! { pub #field_name: #ty }
233            }
234        })
235        .collect();
236
237    let get_sk_fields: Vec<_> = sk_field_names
238        .iter()
239        .zip(sk_field_types.iter())
240        .zip(gsk.iter())
241        .map(|((name, ty), gp)| {
242            let field_name = name.as_ref().unwrap();
243            if let Some(g) = gp {
244                quote! { pub #field_name: #g }
245            } else {
246                quote! { pub #field_name: #ty }
247            }
248        })
249        .collect();
250
251    let query_pk_fields: Vec<_> = pk_field_names
252        .iter()
253        .zip(pk_field_types.iter())
254        .zip(gpk.iter())
255        .map(|((name, ty), gp)| {
256            let field_name = name.as_ref().unwrap();
257            if let Some(g) = gp {
258                quote! { pub #field_name: #g }
259            } else {
260                quote! { pub #field_name: #ty }
261            }
262        })
263        .collect();
264
265    let query_sk_fields: Vec<_> = sk_field_names
266        .iter()
267        .zip(sk_field_types.iter())
268        .map(|(name, ty)| {
269            let field_name = name.as_ref().unwrap();
270            quote! { pub #field_name: Option<#ty> }
271        })
272        .collect();
273
274    let query_pk_str = if pk_fields.is_empty() {
275        let name_str = name.to_string();
276        quote! { #name_str.to_string() }
277    } else {
278        let name_str = name.to_string();
279        let pk_format_parts: Vec<_> = pk_field_names
280            .iter()
281            .zip(pk_field_types.iter())
282            .map(|(n, ty)| {
283                let name_str = n.as_ref().unwrap().to_string();
284                format!("{}={}", name_str, format_placeholder(ty))
285            })
286            .collect();
287        let pk_format_string = format!("{}/{}", name_str, pk_format_parts.join("&"));
288        let pk_format_args: Vec<_> = pk_field_names
289            .iter()
290            .zip(pk_field_types.iter())
291            .zip(pk_is_string.iter())
292            .map(|((n, ty), &is_str)| {
293                let field_name = n.as_ref().unwrap();
294                if is_str {
295                    quote! { self.#field_name.as_ref() }
296                } else {
297                    wrap_expr(ty, quote! { self.#field_name })
298                }
299            })
300            .collect();
301        quote! { format!(#pk_format_string, #(#pk_format_args),*) }
302    };
303
304    let query_sk_build: Vec<_> = sk_field_names
305        .iter()
306        .zip(sk_field_types.iter())
307        .map(|(n, ty)| {
308            let name_str = n.as_ref().unwrap().to_string();
309            let field_name = n.as_ref().unwrap();
310            let fmt = format!("{}={}", name_str, format_placeholder(ty));
311            let val_expr = wrap_expr(ty, quote! { *v });
312            quote! {
313                if let Some(v) = &self.#field_name {
314                    parts.push(format!(#fmt, #val_expr));
315                } else {
316                    break 'build;
317                }
318            }
319        })
320        .collect();
321
322    let pk_str = if pk_fields.is_empty() {
323        let name_str = name.to_string();
324        quote! { #name_str.to_string() }
325    } else {
326        let name_str = name.to_string();
327        let pk_format_parts: Vec<_> = pk_field_names
328            .iter()
329            .zip(pk_field_types.iter())
330            .map(|(n, ty)| {
331                let name_str = n.as_ref().unwrap().to_string();
332                format!("{}={}", name_str, format_placeholder(ty))
333            })
334            .collect();
335        let pk_format_string = format!("{}/{}", name_str, pk_format_parts.join("&"));
336        let pk_format_args: Vec<_> = pk_field_names
337            .iter()
338            .zip(pk_field_types.iter())
339            .zip(pk_is_string.iter())
340            .map(|((n, ty), &is_str)| {
341                let field_name = n.as_ref().unwrap();
342                if is_str {
343                    quote! { self.#field_name.as_ref() }
344                } else {
345                    wrap_expr(ty, quote! { self.#field_name })
346                }
347            })
348            .collect();
349        quote! { format!(#pk_format_string, #(#pk_format_args),*) }
350    };
351
352    let sk_format_parts: Vec<_> = sk_field_names
353        .iter()
354        .zip(sk_field_types.iter())
355        .map(|(n, ty)| {
356            let name_str = n.as_ref().unwrap().to_string();
357            format!("{}={}", name_str, format_placeholder(ty))
358        })
359        .collect();
360    let sk_format_string = sk_format_parts.join("&");
361    let sk_format_args: Vec<_> = sk_field_names
362        .iter()
363        .zip(sk_field_types.iter())
364        .zip(sk_is_string.iter())
365        .map(|((n, ty), &is_str)| {
366            let field_name = n.as_ref().unwrap();
367            if is_str {
368                quote! { self.#field_name.as_ref() }
369            } else {
370                wrap_expr(ty, quote! { self.#field_name })
371            }
372        })
373        .collect();
374
375    let put_pk_str = if pk_fields.is_empty() {
376        let name_str = name.to_string();
377        quote! { #name_str.to_string() }
378    } else {
379        let name_str = name.to_string();
380        let pk_format_parts: Vec<_> = pk_field_names
381            .iter()
382            .zip(pk_field_types.iter())
383            .map(|(n, ty)| {
384                let name_str = n.as_ref().unwrap().to_string();
385                format!("{}={}", name_str, format_placeholder(ty))
386            })
387            .collect();
388        let pk_format_string = format!("{}/{}", name_str, pk_format_parts.join("&"));
389        let pk_format_args: Vec<_> = pk_field_names
390            .iter()
391            .zip(pk_field_types.iter())
392            .map(|(n, ty)| {
393                let field_name = n.as_ref().unwrap();
394                wrap_expr(ty, quote! { self.0.#field_name })
395            })
396            .collect();
397        quote! { format!(#pk_format_string, #(#pk_format_args),*) }
398    };
399
400    let doc_pk_str = if pk_fields.is_empty() {
401        let name_str = name.to_string();
402        quote! { #name_str.to_string() }
403    } else {
404        let name_str = name.to_string();
405        let pk_format_parts: Vec<_> = pk_field_names
406            .iter()
407            .zip(pk_field_types.iter())
408            .map(|(n, ty)| {
409                let name_str = n.as_ref().unwrap().to_string();
410                format!("{}={}", name_str, format_placeholder(ty))
411            })
412            .collect();
413        let pk_format_string = format!("{}/{}", name_str, pk_format_parts.join("&"));
414        let pk_format_args: Vec<_> = pk_field_names
415            .iter()
416            .zip(pk_field_types.iter())
417            .zip(pk_is_string.iter())
418            .map(|((n, ty), &is_str)| {
419                let field_name = n.as_ref().unwrap();
420                if is_str {
421                    quote! { self.#field_name.as_str() }
422                } else {
423                    wrap_expr(ty, quote! { self.#field_name })
424                }
425            })
426            .collect();
427        quote! { format!(#pk_format_string, #(#pk_format_args),*) }
428    };
429
430    let put_sk_format_args: Vec<_> = sk_field_names
431        .iter()
432        .zip(sk_field_types.iter())
433        .map(|(n, ty)| {
434            let field_name = n.as_ref().unwrap();
435            wrap_expr(ty, quote! { self.0.#field_name })
436        })
437        .collect();
438
439    let doc_sk_format_args: Vec<_> = sk_field_names
440        .iter()
441        .zip(sk_field_types.iter())
442        .zip(sk_is_string.iter())
443        .map(|((n, ty), &is_str)| {
444            let field_name = n.as_ref().unwrap();
445            if is_str {
446                quote! { self.#field_name.as_str() }
447            } else {
448                wrap_expr(ty, quote! { self.#field_name })
449            }
450        })
451        .collect();
452
453    let clean_fields: Vec<_> = fields
454        .iter()
455        .map(|f| {
456            let mut f = f.clone();
457            f.attrs
458                .retain(|a| !a.path().is_ident("pk") && !a.path().is_ident("sk"));
459            f
460        })
461        .collect();
462
463    let expanded = quote! {
464        #[derive(serde::Serialize, serde::Deserialize, Clone)]
465        #vis struct #name {
466            #(#clean_fields,)*
467        }
468
469        impl doc_db::Document for #name {
470            fn key(&self) -> doc_db::DocKey {
471                let pk = #doc_pk_str;
472                let sk = format!(#sk_format_string, #(#doc_sk_format_args),*);
473                doc_db::DocKey::new(pk, sk)
474            }
475        }
476
477        #vis struct #put_name(pub #name);
478
479        impl doc_db::DbRequest for #put_name {
480            type Output = ();
481            fn prepare(self) -> doc_db::Prepared<Self::Output> {
482                let pk = #put_pk_str;
483                let sk = format!(#sk_format_string, #(#put_sk_format_args),*);
484                let data = serde_json::to_vec(&self.0).expect("failed to serialize");
485                doc_db::Prepared {
486                    ops: vec![doc_db::DbOp::Put { pk, sk, data }],
487                    parse: Box::new(|iter| {
488                        match iter.next().ok_or_else(|| anyhow::anyhow!("missing result"))? {
489                            doc_db::DbResult::Done => Ok(()),
490                            _ => anyhow::bail!("unexpected result type"),
491                        }
492                    }),
493                }
494            }
495        }
496
497        #vis struct #get_name #generic_def {
498            #(#get_pk_fields,)*
499            #(#get_sk_fields,)*
500        }
501
502        impl #generic_def doc_db::DocGet for #get_name #generic_use {
503            type Doc = #name;
504
505            fn key(&self) -> doc_db::DocKey {
506                let pk = #pk_str;
507                let sk = format!(#sk_format_string, #(#sk_format_args),*);
508                doc_db::DocKey::new(pk, sk)
509            }
510        }
511
512        impl #generic_def doc_db::DbRequest for #get_name #generic_use {
513            type Output = Option<#name>;
514            fn prepare(self) -> doc_db::Prepared<Self::Output> {
515                let pk = #pk_str;
516                let sk = format!(#sk_format_string, #(#sk_format_args),*);
517                doc_db::Prepared {
518                    ops: vec![doc_db::DbOp::Get { pk, sk }],
519                    parse: Box::new(|iter| {
520                        match iter.next().ok_or_else(|| anyhow::anyhow!("missing result"))? {
521                            doc_db::DbResult::Single(opt) => {
522                                opt.map(|data| serde_json::from_slice(&data))
523                                    .transpose()
524                                    .map_err(Into::into)
525                            }
526                            _ => anyhow::bail!("unexpected result type"),
527                        }
528                    }),
529                }
530            }
531        }
532
533        #vis struct #query_name #query_generic_def {
534            #(#query_pk_fields,)*
535            #(#query_sk_fields,)*
536            pub limit: Option<usize>,
537        }
538
539        impl #query_generic_def doc_db::DbRequest for #query_name #query_generic_use {
540            type Output = Vec<#name>;
541            fn prepare(self) -> doc_db::Prepared<Self::Output> {
542                let pk = #query_pk_str;
543                let after_sk: Option<String> = {
544                    let mut parts: Vec<String> = Vec::new();
545                    'build: {
546                        #(#query_sk_build)*
547                    }
548                    if parts.is_empty() { None } else { Some(parts.join("&")) }
549                };
550                let limit = self.limit;
551                doc_db::Prepared {
552                    ops: vec![doc_db::DbOp::Query { pk, after_sk, limit }],
553                    parse: Box::new(|iter| {
554                        match iter.next().ok_or_else(|| anyhow::anyhow!("missing result"))? {
555                            doc_db::DbResult::Multiple(items) => {
556                                items.into_iter()
557                                    .map(|(_sk, data)| serde_json::from_slice(&data))
558                                    .collect::<Result<Vec<_>, _>>()
559                                    .map_err(Into::into)
560                            }
561                            _ => anyhow::bail!("unexpected result type"),
562                        }
563                    }),
564                }
565            }
566        }
567
568        #vis struct #delete_name #generic_def {
569            #(#get_pk_fields,)*
570            #(#get_sk_fields,)*
571        }
572
573        impl #generic_def doc_db::DbRequest for #delete_name #generic_use {
574            type Output = ();
575            fn prepare(self) -> doc_db::Prepared<Self::Output> {
576                let pk = #pk_str;
577                let sk = format!(#sk_format_string, #(#sk_format_args),*);
578                doc_db::Prepared {
579                    ops: vec![doc_db::DbOp::Delete { pk, sk }],
580                    parse: Box::new(|iter| {
581                        match iter.next().ok_or_else(|| anyhow::anyhow!("missing result"))? {
582                            doc_db::DbResult::Done => Ok(()),
583                            _ => anyhow::bail!("unexpected result type"),
584                        }
585                    }),
586                }
587            }
588        }
589    };
590
591    TokenStream::from(expanded)
592}