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 #[::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}