1use 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}