#![feature(let_chains)]
#![forbid(unsafe_code)]
use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::{parse_macro_input, Data, DeriveInput, Fields};
mod parser;
#[proc_macro_derive(Schema, attributes(schema))]
pub fn schema_macro(item: TokenStream) -> TokenStream {
const INTEGER_TYPES: [&str; 6] = ["u64", "i64", "u32", "i32", "u16", "i16"];
let input = parse_macro_input!(item as DeriveInput);
let name = input.ident;
let mut type_name = name.to_string();
let mut primary_key_name = String::from("id");
let mut reader_name = String::from("main");
let mut writer_name = String::from("main");
let mut distribution_column = None;
for attr in input.attrs.iter() {
for (key, value) in parser::parse_attr(attr).into_iter() {
if let Some(value) = value {
if key == "type_name" {
type_name = value;
} else if key == "primary_key" {
primary_key_name = value;
} else if key == "reader_name" {
reader_name = value;
} else if key == "writer_name" {
writer_name = value;
} else if key == "distribution_column" {
distribution_column = Some(value);
}
}
}
}
let mut columns = Vec::new();
if let Data::Struct(data) = input.data && let Fields::Named(fields) = data.fields {
for field in fields.named.into_iter() {
let mut type_name = parser::get_type_name(&field.ty);
if let Some(ident) = field.ident && !type_name.is_empty() {
let name = ident.to_string();
let mut default_value = None;
let mut not_null = false;
let mut index_type = None;
for attr in field.attrs.iter() {
for (key, value) in parser::parse_attr(attr).into_iter() {
if key == "type_name" {
if let Some(value) = value {
type_name = value;
}
} else if key == "not_null" {
not_null = true;
} else if key == "default" {
default_value = value;
} else if key == "index" {
index_type = value;
}
}
}
if type_name.starts_with("Option") {
not_null = false;
} else if type_name == "Uuid" {
not_null = true;
} else if INTEGER_TYPES.contains(&type_name.as_str()) {
default_value = default_value.or_else(|| Some("0".to_owned()));
}
let quote_value = match default_value {
Some(value) => {
if value.contains("::") {
if let Some((type_name, type_fn)) = value.split_once("::") {
let type_name_ident = format_ident!("{}", type_name);
let type_fn_ident = format_ident!("{}", type_fn);
quote! { Some(<#type_name_ident>::#type_fn_ident()) }
} else {
quote! { Some(#value) }
}
} else {
quote! { Some(#value) }
}
}
None => quote! { None },
};
let quote_index = match index_type {
Some(index) => quote! { Some(#index) },
None => quote! { None },
};
let column = quote! {
zino_core::database::Column::new(#name, #type_name, #quote_value, #not_null, #quote_index)
};
columns.push(column);
}
}
}
let type_name_lowercase = type_name.to_ascii_lowercase();
let type_name_uppercase = type_name.to_ascii_uppercase();
let quote_distribution_column = match distribution_column {
Some(column_name) => quote! { Some(#column_name) },
None => quote! { None },
};
let schema_primary_key = format_ident!("{}", primary_key_name);
let schema_columns = format_ident!("{}_COLUMNS", type_name_uppercase);
let schema_reader = format_ident!("{}_READER", type_name_uppercase);
let schema_writer = format_ident!("{}_WRITER", type_name_uppercase);
let columns_len = columns.len();
let output = quote! {
use std::sync::{LazyLock, OnceLock};
use zino_core::database::{Column, ConnectionPool, Schema};
static #schema_columns: LazyLock<[Column; #columns_len]> = LazyLock::new(|| {
[#(#columns),*]
});
static #schema_reader: OnceLock<&ConnectionPool> = OnceLock::new();
static #schema_writer: OnceLock<&ConnectionPool> = OnceLock::new();
impl Schema for #name {
const TYPE_NAME: &'static str = #type_name_lowercase;
const PRIMARY_KEY_NAME: &'static str = #primary_key_name;
const READER_NAME: &'static str = #reader_name;
const WRITER_NAME: &'static str = #writer_name;
const DISTRIBUTION_COLUMN: Option<&'static str> = #quote_distribution_column;
#[inline]
fn columns() -> &'static [Column<'static>] {
LazyLock::force(&#schema_columns).as_slice()
}
#[inline]
fn primary_key(&self) -> String {
self.#schema_primary_key.to_string()
}
async fn get_reader() -> Option<&'static ConnectionPool> {
match #schema_reader.get() {
Some(connection_pool) => Some(*connection_pool),
None => {
let connection_pool = Self::init_reader().ok()?;
let _ = Self::create_table().await.ok()?;
let _ = Self::create_indexes().await.ok()?;
let _ = #schema_reader.set(connection_pool).ok()?;
Some(connection_pool)
},
}
}
async fn get_writer() -> Option<&'static ConnectionPool> {
match #schema_writer.get() {
Some(connection_pool) => Some(*connection_pool),
None => {
let connection_pool = Self::init_writer().ok()?;
let _ = Self::create_table().await.ok()?;
let _ = Self::create_indexes().await.ok()?;
let _ = #schema_writer.set(connection_pool).ok()?;
Some(connection_pool)
},
}
}
}
impl PartialEq for #name {
#[inline]
fn eq(&self, other: &Self) -> bool {
self.#schema_primary_key == other.#schema_primary_key
}
}
impl Eq for #name {}
};
TokenStream::from(output)
}