Skip to main content

toolu_orm_macros/
lib.rs

1mod column_enum;
2mod expand;
3mod from_row;
4mod from_row_expand;
5mod parse;
6mod relational;
7mod view;
8
9use proc_macro::TokenStream;
10use syn::parse::Parser;
11use syn::{parse_macro_input, Expr, ItemStruct, Lit, Meta};
12
13#[proc_macro_attribute]
14pub fn table(attr: TokenStream, item: TokenStream) -> TokenStream {
15  let mut item_struct = parse_macro_input!(item as ItemStruct);
16
17  let (table_name, strict) = match parse_table_attrs(attr) {
18    Ok(attrs) => attrs,
19    Err(e) => return e.to_compile_error().into(),
20  };
21
22  // Parse and strip #[index] / #[unique_index] attrs from the struct
23  let indexes = match parse::parse_index_attrs(&mut item_struct) {
24    Ok(idxs) => idxs,
25    Err(e) => return e.to_compile_error().into(),
26  };
27
28  // Parse and strip #[view] attrs from the struct
29  let views = match parse_view_attrs(&mut item_struct) {
30    Ok(v) => v,
31    Err(e) => return e.to_compile_error().into(),
32  };
33
34  let columns = match parse::parse_struct(&item_struct) {
35    Ok(cols) => cols,
36    Err(e) => return e.to_compile_error().into(),
37  };
38
39  let input = parse::TableInput {
40    table_name,
41    strict,
42    struct_name: item_struct.ident.clone(),
43    columns,
44    indexes,
45  };
46
47  let expanded = expand::expand(&input);
48  let columns_mod = expand::expand_columns_module(&input);
49  let builder_methods = expand::expand_builder_methods(&input);
50  let vis = &item_struct.vis;
51  let view_structs: Vec<proc_macro2::TokenStream> = views
52    .iter()
53    .map(|v| view::generate_view_struct(v, &input.columns, vis))
54    .collect();
55  let clean_struct = parse::strip_column_attrs(item_struct);
56
57  let output = quote::quote! {
58      #clean_struct
59      #expanded
60      #builder_methods
61      #columns_mod
62      #(#view_structs)*
63  };
64
65  output.into()
66}
67
68fn parse_view_attrs(item: &mut ItemStruct) -> syn::Result<Vec<view::ViewInput>> {
69  let mut views = Vec::new();
70  let mut remaining_attrs = Vec::new();
71  for attr in &item.attrs {
72    if attr.path().is_ident("view") {
73      views.push(view::parse_view_attr(attr)?);
74    } else {
75      remaining_attrs.push(attr.clone());
76    }
77  }
78  item.attrs = remaining_attrs;
79  Ok(views)
80}
81
82/// Parses `#[table(name = "table_name")]` or `#[table(name = "table_name", strict = true)]`.
83/// Returns (table_name, strict). Strict defaults to false (standard SQLite compatibility).
84fn parse_table_attrs(attr: TokenStream) -> syn::Result<(String, bool)> {
85  let parser = syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated;
86  let metas = parser.parse(attr)?;
87
88  let mut table_name = None;
89  let mut strict = false; // default: non-strict for SQLite/libsql compatibility
90
91  for meta in &metas {
92    let Meta::NameValue(nv) = meta else {
93      return Err(syn::Error::new_spanned(
94        meta,
95        "expected name = value pairs in #[table(...)]",
96      ));
97    };
98    if nv.path.is_ident("name") {
99      let Expr::Lit(expr_lit) = &nv.value else {
100        return Err(syn::Error::new_spanned(
101          &nv.value,
102          "expected a string literal",
103        ));
104      };
105      let Lit::Str(s) = &expr_lit.lit else {
106        return Err(syn::Error::new_spanned(
107          &expr_lit.lit,
108          "expected a string literal",
109        ));
110      };
111      table_name = Some(s.value());
112    } else if nv.path.is_ident("strict") {
113      let Expr::Lit(expr_lit) = &nv.value else {
114        return Err(syn::Error::new_spanned(
115          &nv.value,
116          "expected a bool literal",
117        ));
118      };
119      let Lit::Bool(b) = &expr_lit.lit else {
120        return Err(syn::Error::new_spanned(
121          &expr_lit.lit,
122          "expected true or false",
123        ));
124      };
125      strict = b.value();
126    } else {
127      return Err(syn::Error::new_spanned(
128        &nv.path,
129        "unknown attribute, expected `name` or `strict`",
130      ));
131    }
132  }
133
134  let Some(name) = table_name else {
135    return Err(syn::Error::new(
136      proc_macro2::Span::call_site(),
137      "missing `name` in #[table(name = \"...\")]",
138    ));
139  };
140
141  Ok((name, strict))
142}
143
144#[proc_macro_derive(ColumnEnum)]
145pub fn derive_column_enum(input: TokenStream) -> TokenStream {
146  let input = parse_macro_input!(input as syn::DeriveInput);
147  match column_enum::expand_column_enum(&input) {
148    Ok(tokens) => tokens.into(),
149    Err(e) => e.to_compile_error().into(),
150  }
151}
152
153#[proc_macro_derive(FromRow, attributes(from_row))]
154pub fn derive_from_row(input: TokenStream) -> TokenStream {
155  let input = parse_macro_input!(input as syn::DeriveInput);
156  match from_row::expand_from_row(&input) {
157    Ok(tokens) => tokens.into(),
158    Err(e) => e.to_compile_error().into(),
159  }
160}
161
162#[proc_macro_derive(Relational, attributes(relational, has_many, belongs_to, many_to_many))]
163pub fn derive_relational(input: TokenStream) -> TokenStream {
164  let input = parse_macro_input!(input as syn::DeriveInput);
165  match relational::parse_relational(&input) {
166    Ok(parsed) => expand::expand_relational(&parsed).into(),
167    Err(e) => e.to_compile_error().into(),
168  }
169}