extern crate proc_macro;
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::quote;
use syn::{
parse_macro_input, Attribute, Data, DeriveInput, Expr, Fields, Lit, Meta, Type,
punctuated::Punctuated, spanned::Spanned, token::Comma,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FieldKind {
PrimaryKey,
Timestamps,
TimestampField,
Persist,
Skip,
}
#[derive(Debug, Clone)]
struct FieldIndex {
name: String,
unique: bool,
partners: Vec<String>,
}
struct FieldInfo {
ident: syn::Ident,
column: Option<String>,
kind: FieldKind,
ty: Type,
optional: bool,
indexes: Vec<FieldIndex>,
}
impl FieldInfo {
fn column_name(&self) -> String {
self.column
.clone()
.unwrap_or_else(|| self.ident.to_string())
}
}
struct ModelConfig {
table_name: String,
primary_key: String,
}
impl ModelConfig {
fn parse(attrs: &[Attribute]) -> syn::Result<Self> {
let mut table_name: Option<String> = None;
let mut primary_key: Option<String> = None;
for attr in attrs {
if !attr.path().is_ident("model") {
continue;
}
let Meta::List(list) = &attr.meta else { continue };
let nested = list.parse_args_with(Punctuated::<Meta, Comma>::parse_terminated)?;
for item in nested {
match item {
Meta::NameValue(nv) if nv.path.is_ident("table_name") => {
table_name = Some(expr_to_string(&nv.value)?);
}
Meta::NameValue(nv) if nv.path.is_ident("primary_key") => {
primary_key = Some(expr_to_string(&nv.value)?);
}
_ => {}
}
}
}
let table_name = table_name.ok_or_else(|| {
syn::Error::new(
Span::call_site(),
"#[model(table_name = \"...\")] is required on #[derive(Model)] structs",
)
})?;
Ok(Self {
table_name,
primary_key: primary_key.unwrap_or_else(|| "id".to_string()),
})
}
}
fn expr_to_string(expr: &Expr) -> syn::Result<String> {
match expr {
Expr::Lit(lit) => match &lit.lit {
Lit::Str(s) => Ok(s.value()),
_ => Err(syn::Error::new_spanned(lit, "expected a string literal")),
},
_ => Err(syn::Error::new_spanned(expr, "expected a string literal")),
}
}
fn collect_fields(input: &DeriveInput, pk_name: &str) -> syn::Result<Vec<FieldInfo>> {
let data = match &input.data {
Data::Struct(s) => s,
Data::Enum(_) | Data::Union(_) => {
return Err(syn::Error::new(
input.ident.span(),
"#[derive(Model)] is only supported on structs",
))
}
};
let named = match &data.fields {
Fields::Named(named) => &named.named,
Fields::Unnamed(_) | Fields::Unit => {
return Err(syn::Error::new(
input.ident.span(),
"#[derive(Model)] requires named struct fields",
))
}
};
let mut out = Vec::new();
for field in named {
let ident = field
.ident
.clone()
.ok_or_else(|| syn::Error::new(field.span(), "named field expected"))?;
let ty = field.ty.clone();
let optional = field_type_is_option(&ty);
let mut skip = false;
let mut column: Option<String> = None;
let mut indexes: Vec<FieldIndex> = Vec::new();
for attr in &field.attrs {
if !attr.path().is_ident("model") {
continue;
}
let Meta::List(list) = &attr.meta else { continue };
let nested = list.parse_args_with(Punctuated::<Meta, Comma>::parse_terminated)?;
for item in nested {
match item {
Meta::Path(path) if path.is_ident("skip") => skip = true,
Meta::Path(path) if path.is_ident("index") => {
indexes.push(FieldIndex {
name: String::new(),
unique: false,
partners: Vec::new(),
});
}
Meta::Path(path) if path.is_ident("uniqueIndex") => {
indexes.push(FieldIndex {
name: String::new(),
unique: true,
partners: Vec::new(),
});
}
Meta::NameValue(nv) if nv.path.is_ident("column") => {
column = Some(expr_to_string(&nv.value)?);
}
Meta::NameValue(nv) if nv.path.is_ident("index") => {
indexes.push(FieldIndex {
name: expr_to_string(&nv.value)?,
unique: false,
partners: Vec::new(),
});
}
Meta::NameValue(nv) if nv.path.is_ident("uniqueIndex") => {
indexes.push(FieldIndex {
name: expr_to_string(&nv.value)?,
unique: true,
partners: Vec::new(),
});
}
_ => {}
}
}
}
let kind = if skip {
FieldKind::Skip
} else if field_type_is_timestamps(&ty) {
FieldKind::Timestamps
} else if ident.to_string() == pk_name {
FieldKind::PrimaryKey
} else if is_standalone_timestamp_field(&ident, &ty) {
FieldKind::TimestampField
} else if type_tag(&ty).is_some() {
FieldKind::Persist
} else {
FieldKind::Skip
};
out.push(FieldInfo {
ident,
column,
kind,
ty,
optional,
indexes,
});
}
Ok(out)
}
fn field_type_is_timestamps(ty: &Type) -> bool {
let Type::Path(tp) = ty else { return false };
matches!(tp.path.segments.last(), Some(s) if s.ident == "Timestamps")
}
fn field_type_is_option(ty: &Type) -> bool {
let Type::Path(tp) = ty else { return false };
matches!(tp.path.segments.last(), Some(s) if s.ident == "Option")
}
fn is_standalone_timestamp_field(ident: &syn::Ident, ty: &Type) -> bool {
let name = ident.to_string();
let valid_name = matches!(name.as_str(), "created_at" | "updated_at" | "deleted_at");
valid_name && type_tag(ty) == Some("datetime")
}
fn inner_type(ty: &Type) -> &Type {
if let Type::Path(tp) = ty {
if let Some(s) = tp.path.segments.last() {
if s.ident == "Option" {
if let syn::PathArguments::AngleBracketed(args) = &s.arguments {
if let Some(syn::GenericArgument::Type(inner)) = args.args.first() {
return inner;
}
}
}
}
}
ty
}
fn type_tag(ty: &Type) -> Option<&'static str> {
let ty = inner_type(ty);
let Type::Path(tp) = ty else { return None };
let last = tp.path.segments.last()?;
let name = last.ident.to_string();
let tag = match name.as_str() {
"String" => "string",
"bool" => "bool",
"i8" => "i8",
"i16" => "i16",
"i32" => "i32",
"i64" => "i64",
"f32" => "f32",
"f64" => "f64",
"DateTime" => "datetime",
"Uuid" => "uuid",
"Vec" => {
let is_u8 = matches!(
&last.arguments,
syn::PathArguments::AngleBracketed(args)
if args.args.len() == 1
&& matches!(
args.args.first(),
Some(syn::GenericArgument::Type(Type::Path(tp)))
if tp.path.is_ident("u8")
)
);
if is_u8 {
"bytes"
} else {
return None;
}
}
_ => return None,
};
Some(tag)
}
fn gen_columns_entries(fields: &[&FieldInfo]) -> Vec<proc_macro2::TokenStream> {
fields
.iter()
.filter_map(|f| {
let col = f.column_name();
let fname = &f.ident;
let tag = type_tag(&f.ty)?;
let expr = match tag {
"string" => {
if f.optional {
quote! { match &self.#fname { Some(v) => SqlValue::String(v.clone()), None => SqlValue::Null } }
} else {
quote! { SqlValue::String(self.#fname.clone()) }
}
}
"bool" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::Bool(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::Bool(self.#fname) }
}
}
"i8" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::I8(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::I8(self.#fname) }
}
}
"i16" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::I16(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::I16(self.#fname) }
}
}
"i32" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::I32(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::I32(self.#fname) }
}
}
"i64" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::I64(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::I64(self.#fname) }
}
}
"f32" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::F32(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::F32(self.#fname) }
}
}
"f64" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::F64(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::F64(self.#fname) }
}
}
"datetime" => {
if f.optional {
quote! { match self.#fname { Some(v) => SqlValue::DateTime(v), None => SqlValue::Null } }
} else {
quote! { SqlValue::DateTime(self.#fname) }
}
}
"uuid" => {
if f.optional {
quote! { match &self.#fname { Some(v) => SqlValue::String(v.to_string()), None => SqlValue::Null } }
} else {
quote! { SqlValue::String(self.#fname.to_string()) }
}
}
"bytes" => {
if f.optional {
quote! { match &self.#fname { Some(v) => SqlValue::Bytes(v.clone()), None => SqlValue::Null } }
} else {
quote! { SqlValue::Bytes(self.#fname.clone()) }
}
}
_ => return None,
};
Some(quote! { (#col, #expr) })
})
.collect()
}
fn gen_from_row_expr(f: &FieldInfo) -> Option<proc_macro2::TokenStream> {
let col = f.column_name();
let tag = type_tag(&f.ty)?;
let inner = match tag {
"string" => Some(quote! {
match row.get(#col)? {
SqlValue::String(s) => s.clone(),
SqlValue::Json(s) => s.clone(),
SqlValue::I64(i) => i.to_string(),
SqlValue::I32(i) => i.to_string(),
SqlValue::I16(i) => i.to_string(),
SqlValue::I8(i) => i.to_string(),
SqlValue::Bool(b) => b.to_string(),
_ => return None,
}
}),
"bool" => Some(quote! {
match row.get(#col)? {
SqlValue::Bool(v) => *v,
SqlValue::I32(1) => true,
SqlValue::I64(1) => true,
SqlValue::I32(0) => false,
SqlValue::I64(0) => false,
_ => return None,
}
}),
"i8" => Some(quote! {
match row.get(#col)? {
SqlValue::I8(v) => *v,
SqlValue::I16(v) => *v as i8,
SqlValue::I32(v) => *v as i8,
SqlValue::I64(v) => *v as i8,
_ => return None,
}
}),
"i16" => Some(quote! {
match row.get(#col)? {
SqlValue::I16(v) => *v,
SqlValue::I8(v) => *v as i16,
SqlValue::I32(v) => *v as i16,
SqlValue::I64(v) => *v as i16,
_ => return None,
}
}),
"i32" => Some(quote! {
match row.get(#col)? {
SqlValue::I32(v) => *v,
SqlValue::I8(v) => *v as i32,
SqlValue::I16(v) => *v as i32,
SqlValue::I64(v) => *v as i32,
_ => return None,
}
}),
"i64" => Some(quote! {
match row.get(#col)? {
SqlValue::I64(v) => *v,
SqlValue::I32(v) => *v as i64,
SqlValue::I16(v) => *v as i64,
SqlValue::I8(v) => *v as i64,
_ => return None,
}
}),
"f32" => Some(quote! {
match row.get(#col)? {
SqlValue::F32(v) => *v,
SqlValue::F64(v) => *v as f32,
_ => return None,
}
}),
"f64" => Some(quote! {
match row.get(#col)? {
SqlValue::F64(v) => *v,
SqlValue::F32(v) => *v as f64,
_ => return None,
}
}),
"datetime" => Some(quote! {
match row.get(#col)? {
SqlValue::DateTime(v) => *v,
SqlValue::String(s) => {
chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S")
.map(|dt| dt.and_utc())
.ok()?
}
_ => return None,
}
}),
"uuid" => Some(quote! {
Uuid::parse_str(row.get(#col)?.as_str()?).ok()?
}),
"bytes" => Some(quote! {
match row.get(#col)? {
SqlValue::Bytes(v) => v.clone(),
_ => return None,
}
}),
_ => None,
}?;
if f.optional {
Some(quote! {
match row.get(#col) {
Some(SqlValue::Null) | None => None,
_ => Some(#inner),
}
})
} else {
Some(inner)
}
}
fn tag_to_column_type(tag: &str) -> Option<&'static str> {
match tag {
"i8" | "i16" | "i32" => Some("Integer"),
"i64" => Some("BigInteger"),
"string" => Some("String"),
"bool" => Some("Boolean"),
"f32" => Some("Float"),
"f64" => Some("Double"),
"datetime" => Some("DateTime"),
"uuid" => Some("Uuid"),
"bytes" => Some("Binary"),
_ => None,
}
}
fn gen_schema_impl(
table_name: &str,
fields: &[&FieldInfo],
) -> proc_macro2::TokenStream {
let mut column_toks: Vec<proc_macro2::TokenStream> = Vec::new();
for f in fields {
if f.kind == FieldKind::Skip {
continue;
}
let col = f.column_name();
let optional = f.optional;
let (col_type, is_pk) = match f.kind {
FieldKind::PrimaryKey => (
tag_to_column_type(type_tag(&f.ty).unwrap_or("string")).or(Some("Integer")),
true,
),
FieldKind::Persist => (type_tag(&f.ty).and_then(tag_to_column_type), false),
FieldKind::TimestampField => (Some("DateTime"), false),
FieldKind::Timestamps => {
for ts_col in ["created_at", "updated_at", "deleted_at"] {
column_toks.push(quote! {
ColumnDefinition::new(#ts_col, ColumnType::DateTime).nullable(true)
});
}
continue;
}
FieldKind::Skip => continue,
};
let Some(col_type) = col_type else {
continue;
};
let col_type_tok = syn::Ident::new(col_type, Span::call_site());
let is_unique = f.indexes.iter().any(|ix| ix.unique && ix.partners.is_empty());
let mut col_def = quote! {
ColumnDefinition::new(#col, ColumnType::#col_type_tok).nullable(#optional)
};
if is_pk {
col_def = quote! { #col_def .primary_key() };
if col_type == "Integer" || col_type == "BigInteger" {
col_def = quote! { #col_def .auto_increment() };
}
}
if is_unique {
col_def = quote! { #col_def .unique() };
}
column_toks.push(col_def);
}
let mut index_defs: Vec<proc_macro2::TokenStream> = Vec::new();
let mut groups: std::collections::HashMap<String, (bool, Vec<String>)> =
std::collections::HashMap::new();
for f in fields {
if f.kind == FieldKind::Skip {
continue;
}
let col = f.column_name();
for ix in &f.indexes {
let key = if ix.name.is_empty() {
col.clone()
} else {
ix.name.clone()
};
groups
.entry(key.clone())
.or_insert((ix.unique, Vec::new()))
.1
.push(col.clone());
if ix.unique {
groups.get_mut(&key).unwrap().0 = true;
}
}
}
let mut sorted_keys: Vec<&String> = groups.keys().collect();
sorted_keys.sort();
for key in sorted_keys {
let (unique, cols) = &groups[key];
let name = if key.starts_with("idx_") {
key.clone()
} else {
format!("idx_{}_{}", table_name, key)
};
let cols_lit: Vec<syn::LitStr> = cols
.iter()
.map(|c| syn::LitStr::new(c, Span::call_site()))
.collect();
if *unique {
index_defs.push(quote! {
IndexDefinition::new(#name, &[#(#cols_lit),*]).unique()
});
} else {
index_defs.push(quote! {
IndexDefinition::new(#name, &[#(#cols_lit),*])
});
}
}
quote! {
fn schema() -> Option<TableDefinition> {
Some(TableDefinition::new(#table_name)
#( .add_column(#column_toks) )*
#( .add_index(#index_defs) )*
)
}
}
}
fn build_model_impl(
input: &DeriveInput,
config: &ModelConfig,
fields: &[FieldInfo],
) -> proc_macro2::TokenStream {
let ident = &input.ident;
let table_name = &config.table_name;
let pk = fields.iter().find(|f| f.kind == FieldKind::PrimaryKey);
let ts = fields.iter().find(|f| f.kind == FieldKind::Timestamps);
let ts_field_created = fields
.iter()
.find(|f| f.kind == FieldKind::TimestampField && f.ident == "created_at");
let ts_field_updated = fields
.iter()
.find(|f| f.kind == FieldKind::TimestampField && f.ident == "updated_at");
let ts_field_deleted = fields
.iter()
.find(|f| f.kind == FieldKind::TimestampField && f.ident == "deleted_at");
let persist: Vec<&FieldInfo> = fields
.iter()
.filter(|f| f.kind == FieldKind::Persist)
.collect();
let (id_fn, set_id_fn) = match pk {
Some(pk) => {
let pk_ident = &pk.ident;
let pk_ty = inner_type(&pk.ty);
let pk_tag = type_tag(pk_ty).unwrap_or("");
let (id_expr, set_expr) = match pk_tag {
"string" => (
quote! {
if self.#pk_ident.is_empty() { None } else { Some(self.#pk_ident.clone()) }
},
quote! { self.#pk_ident = id; },
),
"uuid" => (
quote! {
if self.#pk_ident.is_nil() { None } else { Some(self.#pk_ident.to_string()) }
},
quote! { self.#pk_ident = Uuid::parse_str(&id).unwrap_or_else(|_| Uuid::nil()); },
),
"i64" => (
quote! {
if self.#pk_ident > 0 { Some(self.#pk_ident.to_string()) } else { None }
},
quote! { self.#pk_ident = id.parse().unwrap_or(0); },
),
"i32" => (
quote! {
if self.#pk_ident > 0 { Some(self.#pk_ident.to_string()) } else { None }
},
quote! { self.#pk_ident = id.parse().unwrap_or(0); },
),
"i16" => (
quote! {
if self.#pk_ident > 0 { Some(self.#pk_ident.to_string()) } else { None }
},
quote! { self.#pk_ident = id.parse().unwrap_or(0); },
),
"i8" => (
quote! {
if self.#pk_ident > 0 { Some(self.#pk_ident.to_string()) } else { None }
},
quote! { self.#pk_ident = id.parse().unwrap_or(0); },
),
_ => (
quote! {
let s = self.#pk_ident.to_string();
if s.is_empty() { None } else { Some(s) }
},
quote! { self.#pk_ident = id.into(); },
),
};
(
quote! { fn id(&self) -> Option<String> { #id_expr } },
quote! { fn set_id(&mut self, id: String) { #set_expr } },
)
}
None => (
quote! {
fn id(&self) -> Option<String> { None }
},
quote! {
fn set_id(&mut self, _id: String) {}
},
),
};
let ts_impl = if let Some(ts) = ts {
let t = &ts.ident;
quote! {
fn created_at(&self) -> Option<DateTime<Utc>> { self.#t.created_at }
fn updated_at(&self) -> Option<DateTime<Utc>> { self.#t.updated_at }
fn deleted_at(&self) -> Option<DateTime<Utc>> { self.#t.deleted_at }
fn set_created_at(&mut self, timestamp: DateTime<Utc>) { self.#t.created_at = Some(timestamp); }
fn set_updated_at(&mut self, timestamp: DateTime<Utc>) { self.#t.updated_at = Some(timestamp); }
fn set_deleted_at(&mut self, timestamp: Option<DateTime<Utc>>) { self.#t.deleted_at = timestamp; }
}
} else {
let created = ts_field_created.map(|f| &f.ident);
let updated = ts_field_updated.map(|f| &f.ident);
let deleted = ts_field_deleted.map(|f| &f.ident);
let (created_get, created_set) = match created {
Some(ci) => (
quote! { self.#ci },
quote! { self.#ci = Some(timestamp); },
),
None => (quote! { None }, quote! {}),
};
let (updated_get, updated_set) = match updated {
Some(ui) => (
quote! { self.#ui },
quote! { self.#ui = Some(timestamp); },
),
None => (quote! { None }, quote! {}),
};
let (deleted_get, deleted_set) = match deleted {
Some(di) => (
quote! { self.#di },
quote! { self.#di = timestamp; },
),
None => (quote! { None }, quote! {}),
};
quote! {
fn created_at(&self) -> Option<DateTime<Utc>> { #created_get }
fn updated_at(&self) -> Option<DateTime<Utc>> { #updated_get }
fn deleted_at(&self) -> Option<DateTime<Utc>> { #deleted_get }
fn set_created_at(&mut self, timestamp: DateTime<Utc>) { #created_set }
fn set_updated_at(&mut self, timestamp: DateTime<Utc>) { #updated_set }
fn set_deleted_at(&mut self, timestamp: Option<DateTime<Utc>>) { #deleted_set }
}
};
let column_entries = gen_columns_entries(&persist);
let mut literals: Vec<proc_macro2::TokenStream> = Vec::new();
for f in fields {
let fident = &f.ident;
match f.kind {
FieldKind::Skip => {
literals.push(quote! { #fident: Default::default() });
}
FieldKind::Timestamps => {
literals.push(quote! {
#fident: {
let mut ts = Timestamps::new();
ts.created_at = row.get("created_at").and_then(|v| match v {
SqlValue::DateTime(dt) => Some(*dt),
_ => None,
});
ts.updated_at = row.get("updated_at").and_then(|v| match v {
SqlValue::DateTime(dt) => Some(*dt),
_ => None,
});
ts.deleted_at = row.get("deleted_at").and_then(|v| match v {
SqlValue::DateTime(dt) => Some(*dt),
_ => None,
});
ts
}
});
}
FieldKind::PrimaryKey | FieldKind::Persist | FieldKind::TimestampField => {
if let Some(expr) = gen_from_row_expr(f) {
literals.push(quote! { #fident: #expr });
} else {
literals.push(quote! { #fident: Default::default() });
}
}
}
}
let field_refs: Vec<&FieldInfo> = fields.iter().collect();
let schema_impl = gen_schema_impl(&config.table_name, &field_refs);
quote! {
#[automatically_derived]
impl Model for #ident {
fn table_name() -> &'static str {
#table_name
}
#id_fn
#set_id_fn
#ts_impl
#schema_impl
fn columns(&self) -> Vec<(&'static str, SqlValue)> {
vec![
#(#column_entries,)*
]
}
fn from_row(row: &Row) -> Option<Self> {
Some(Self {
#(#literals,)*
})
}
}
}
}
#[proc_macro_derive(Model, attributes(model))]
pub fn derive_model(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
match expand(&input) {
Ok(ts) => ts.into(),
Err(e) => e.to_compile_error().into(),
}
}
fn expand(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
let config = ModelConfig::parse(&input.attrs)?;
let fields = collect_fields(input, &config.primary_key)?;
let model_impl = build_model_impl(input, &config, &fields);
Ok(quote! {
const _: () = {
use torm::orm::model::Timestamps;
use torm::db::db_types::{Row, SqlValue};
use torm::chrono::{DateTime, Utc};
use torm::Uuid;
use torm::orm::migration::{
TableDefinition, ColumnDefinition, ColumnType, IndexDefinition,
};
#model_impl
};
})
}