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