use proc_macro2::{Span, TokenStream};
use syn::parse::{Parse, ParseStream};
use syn::punctuated::Punctuated;
use syn::{Ident, Result, Token, braced};
use super::types::ScalarType;
#[derive(Debug)]
pub struct Schema {
pub entities: Vec<EntityDef>,
}
#[derive(Debug, Clone)]
pub struct EntityAttrs {
pub table_name: Option<String>,
pub has_created_at: bool,
pub has_updated_at: bool,
pub primary_key: Option<Vec<String>>,
}
impl Default for EntityAttrs {
fn default() -> Self {
Self {
table_name: None,
has_created_at: true,
has_updated_at: true,
primary_key: None,
}
}
}
#[derive(Debug, Default, Clone)]
pub struct FieldAttrs {
pub unique: bool,
pub column_name: Option<String>,
pub indexed: bool,
}
#[derive(Debug)]
pub struct EntityDef {
pub attrs: EntityAttrs,
pub name: Ident,
pub fields: Vec<FieldDef>,
pub span: Span,
}
#[derive(Debug)]
pub struct FieldDef {
pub attrs: FieldAttrs,
pub name: Ident,
pub ty: RawFieldType,
pub span: Span,
}
#[derive(Debug)]
pub enum RawFieldType {
Scalar { scalar: ScalarType, optional: bool },
Vec { inner: Ident },
Unknown { name: Ident, optional: bool },
}
impl Parse for Schema {
fn parse(input: ParseStream) -> Result<Self> {
let mut entities = Vec::new();
while !input.is_empty() {
entities.push(input.parse()?);
}
if entities.is_empty() {
return Err(syn::Error::new(
Span::call_site(),
"schema! macro requires at least one entity definition",
));
}
Ok(Schema { entities })
}
}
impl Parse for EntityDef {
fn parse(input: ParseStream) -> Result<Self> {
let attrs = parse_entity_attrs(input)?;
let name: Ident = input.parse()?;
let span = name.span();
let content;
braced!(content in input);
let fields_punctuated: Punctuated<FieldDef, Token![,]> =
content.parse_terminated(FieldDef::parse, Token![,])?;
let fields: Vec<FieldDef> = fields_punctuated.into_iter().collect();
if attrs.primary_key.is_none() {
for field in &fields {
if field.name == "id" {
return Err(syn::Error::new(
field.name.span(),
"field 'id' is reserved and automatically generated. Use #[primary_key(...)] to define a custom primary key",
));
}
}
}
let mut seen_fields = std::collections::HashSet::new();
for field in &fields {
let field_name = field.name.to_string();
if !seen_fields.insert(field_name.clone()) {
return Err(syn::Error::new(
field.name.span(),
format!("duplicate field name '{}'", field_name),
));
}
}
Ok(EntityDef {
attrs,
name,
fields,
span,
})
}
}
fn parse_entity_attrs(input: ParseStream) -> Result<EntityAttrs> {
let mut attrs = EntityAttrs::default();
while input.peek(Token![#]) {
input.parse::<Token![#]>()?;
let content;
syn::bracketed!(content in input);
let attr_name: Ident = content.parse()?;
let attr_name_str = attr_name.to_string();
match attr_name_str.as_str() {
"table_name" => {
content.parse::<Token![=]>()?;
let value: syn::LitStr = content.parse()?;
attrs.table_name = Some(value.value());
}
"timestamps" => {
let inner;
syn::parenthesized!(inner in content);
let ts_type: Ident = inner.parse()?;
let ts_str = ts_type.to_string();
match ts_str.as_str() {
"created_at" => {
attrs.has_created_at = true;
attrs.has_updated_at = false;
}
"updated_at" => {
attrs.has_created_at = false;
attrs.has_updated_at = true;
}
"none" => {
attrs.has_created_at = false;
attrs.has_updated_at = false;
}
_ => {
return Err(syn::Error::new(
ts_type.span(),
format!(
"unknown timestamps option '{}'. Supported: created_at, updated_at, none",
ts_str
),
));
}
}
}
"primary_key" => {
let inner;
syn::parenthesized!(inner in content);
let columns: Punctuated<Ident, Token![,]> =
inner.parse_terminated(Ident::parse, Token![,])?;
let pk_cols: Vec<String> = columns.iter().map(|c| c.to_string()).collect();
if pk_cols.is_empty() {
return Err(syn::Error::new(
attr_name.span(),
"primary_key requires at least one column",
));
}
attrs.primary_key = Some(pk_cols);
}
_ => {
return Err(syn::Error::new(
attr_name.span(),
format!(
"unknown entity attribute '{}'. Supported: table_name, timestamps, primary_key",
attr_name_str
),
));
}
}
}
Ok(attrs)
}
impl Parse for FieldDef {
fn parse(input: ParseStream) -> Result<Self> {
let attrs = parse_field_attrs(input)?;
let name: Ident = input.parse()?;
let span = name.span();
input.parse::<Token![:]>()?;
let ty = parse_field_type(input)?;
Ok(FieldDef {
attrs,
name,
ty,
span,
})
}
}
fn parse_field_attrs(input: ParseStream) -> Result<FieldAttrs> {
let mut attrs = FieldAttrs::default();
while input.peek(Token![#]) {
input.parse::<Token![#]>()?;
let content;
syn::bracketed!(content in input);
let attr_name: Ident = content.parse()?;
let attr_name_str = attr_name.to_string();
match attr_name_str.as_str() {
"unique" => {
attrs.unique = true;
}
"index" => {
attrs.indexed = true;
}
"column" => {
content.parse::<Token![=]>()?;
let value: syn::LitStr = content.parse()?;
attrs.column_name = Some(value.value());
}
_ => {
return Err(syn::Error::new(
attr_name.span(),
format!(
"unknown field attribute '{}'. Supported: unique, index, column",
attr_name_str
),
));
}
}
}
Ok(attrs)
}
fn parse_field_type(input: ParseStream) -> Result<RawFieldType> {
if input.peek(Ident) {
let ident: Ident = input.parse()?;
let ident_str = ident.to_string();
if ident_str == "Option" {
input.parse::<Token![<]>()?;
let inner_type = parse_inner_type(input)?;
input.parse::<Token![>]>()?;
return match inner_type {
InnerType::Scalar(scalar) => Ok(RawFieldType::Scalar {
scalar,
optional: true,
}),
InnerType::Ident(name) => Ok(RawFieldType::Unknown {
name,
optional: true,
}),
};
}
if ident_str == "Vec" {
input.parse::<Token![<]>()?;
let inner: Ident = input.parse()?;
input.parse::<Token![>]>()?;
if inner == "u8" {
return Ok(RawFieldType::Scalar {
scalar: ScalarType::Bytes,
optional: false,
});
}
return Ok(RawFieldType::Vec { inner });
}
if ScalarType::is_unsupported_unsigned(&ident_str) {
return Err(syn::Error::new(
ident.span(),
format!(
"unsigned integer type '{}' is not supported in schema fields. \
Most SQL databases lack a native unsigned integer type and \
silently widening to a signed column truncates values above \
the signed maximum. Use a signed type ('i32' or 'i64') and \
validate the upper bound at the application layer, or model \
the column as 'Decimal' if the unsigned range is required.",
ident_str
),
));
}
if let Some(scalar) = ScalarType::from_ident(&ident_str) {
return Ok(RawFieldType::Scalar {
scalar,
optional: false,
});
}
Ok(RawFieldType::Unknown {
name: ident,
optional: false,
})
} else {
Err(syn::Error::new(input.span(), "expected type"))
}
}
enum InnerType {
Scalar(ScalarType),
Ident(Ident),
}
fn parse_inner_type(input: ParseStream) -> Result<InnerType> {
let ident: Ident = input.parse()?;
let ident_str = ident.to_string();
if ident_str == "Vec" {
input.parse::<Token![<]>()?;
let inner: Ident = input.parse()?;
input.parse::<Token![>]>()?;
if inner == "u8" {
return Ok(InnerType::Scalar(ScalarType::Bytes));
}
return Ok(InnerType::Ident(ident));
}
if ScalarType::is_unsupported_unsigned(&ident_str) {
return Err(syn::Error::new(
ident.span(),
format!(
"unsigned integer type '{}' is not supported in schema fields. \
Most SQL databases lack a native unsigned integer type and \
silently widening to a signed column truncates values above \
the signed maximum. Use a signed type ('i32' or 'i64') and \
validate the upper bound at the application layer, or model \
the column as 'Decimal' if the unsigned range is required.",
ident_str
),
));
}
if let Some(scalar) = ScalarType::from_ident(&ident_str) {
Ok(InnerType::Scalar(scalar))
} else {
Ok(InnerType::Ident(ident))
}
}
pub fn parse_schema(input: TokenStream) -> Result<Schema> {
syn::parse2(input)
}
#[cfg(test)]
mod tests {
use super::*;
use quote::quote;
#[test]
fn test_parse_simple_entity() {
let input = quote! {
User {
email: String,
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(schema.entities.len(), 1);
assert_eq!(schema.entities[0].name.to_string(), "User");
assert_eq!(schema.entities[0].fields.len(), 2);
}
#[test]
fn test_parse_multiple_entities() {
let input = quote! {
User {
email: String,
}
Post {
title: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(schema.entities.len(), 2);
}
#[test]
fn test_parse_vec_field() {
let input = quote! {
User {
posts: Vec<Post>,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
assert!(matches!(field.ty, RawFieldType::Vec { .. }));
}
#[test]
fn test_parse_option_field() {
let input = quote! {
Post {
author: Option<User>,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
assert!(matches!(
field.ty,
RawFieldType::Unknown { optional: true, .. }
));
}
#[test]
fn test_reserved_field_error() {
let input = quote! {
User {
id: i32,
}
};
let result = parse_schema(input);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("reserved"));
}
#[test]
fn test_unsigned_integer_field_rejected() {
for ty in ["u8", "u16", "u32", "u64"] {
let input: TokenStream = format!("Entity {{ count: {ty}, }}").parse().unwrap();
let result = parse_schema(input);
assert!(
result.is_err(),
"expected '{ty}' to be rejected by parse_schema"
);
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("unsigned integer type"),
"expected error message for '{ty}' to mention 'unsigned integer type', got: {msg}"
);
assert!(
msg.contains(ty),
"expected error message for '{ty}' to name the offending type, got: {msg}"
);
}
}
#[test]
fn test_unsigned_integer_field_inside_option_rejected() {
let input = quote! {
Entity {
count: Option<u64>,
}
};
let result = parse_schema(input);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(msg.contains("unsigned integer type"));
assert!(msg.contains("u64"));
}
#[test]
fn test_duplicate_field_error() {
let input = quote! {
User {
email: String,
email: String,
}
};
let result = parse_schema(input);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("duplicate"));
}
#[test]
fn test_parse_table_name_attr() {
let input = quote! {
#[table_name = "people"]
Person {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].attrs.table_name,
Some("people".to_string())
);
}
#[test]
fn test_parse_unique_attr() {
let input = quote! {
User {
#[unique]
email: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].fields[0].attrs.unique);
}
#[test]
fn test_parse_column_attr() {
let input = quote! {
User {
#[column = "email_address"]
email: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].fields[0].attrs.column_name,
Some("email_address".to_string())
);
}
#[test]
fn test_parse_multiple_field_attrs() {
let input = quote! {
User {
#[unique]
#[column = "user_email"]
email: String,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
assert!(field.attrs.unique);
assert_eq!(field.attrs.column_name, Some("user_email".to_string()));
}
#[test]
fn test_unknown_entity_attr_error() {
let input = quote! {
#[unknown_attr = "value"]
User {
email: String,
}
};
let result = parse_schema(input);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("unknown entity attribute")
);
}
#[test]
fn test_unknown_field_attr_error() {
let input = quote! {
User {
#[unknown]
email: String,
}
};
let result = parse_schema(input);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("unknown field attribute")
);
}
#[test]
fn test_parse_timestamps_created_at_only() {
let input = quote! {
#[timestamps(created_at)]
User {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].attrs.has_created_at);
assert!(!schema.entities[0].attrs.has_updated_at);
}
#[test]
fn test_parse_timestamps_updated_at_only() {
let input = quote! {
#[timestamps(updated_at)]
User {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(!schema.entities[0].attrs.has_created_at);
assert!(schema.entities[0].attrs.has_updated_at);
}
#[test]
fn test_parse_timestamps_none() {
let input = quote! {
#[timestamps(none)]
User {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(!schema.entities[0].attrs.has_created_at);
assert!(!schema.entities[0].attrs.has_updated_at);
}
#[test]
fn test_parse_index_attr() {
let input = quote! {
User {
#[index]
email: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].fields[0].attrs.indexed);
}
#[test]
fn test_parse_combined_field_attrs() {
let input = quote! {
User {
#[unique]
#[index]
#[column = "user_email"]
email: String,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
assert!(field.attrs.unique);
assert!(field.attrs.indexed);
assert_eq!(field.attrs.column_name, Some("user_email".to_string()));
}
#[test]
fn test_default_timestamps_enabled() {
let input = quote! {
User {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].attrs.has_created_at);
assert!(schema.entities[0].attrs.has_updated_at);
}
#[test]
fn test_parse_primary_key_single() {
let input = quote! {
#[primary_key(user_id)]
UserPreference {
user_id: i32,
theme: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].attrs.primary_key,
Some(vec!["user_id".to_string()])
);
}
#[test]
fn test_parse_primary_key_composite() {
let input = quote! {
#[primary_key(user_id, role_id)]
UsersRole {
user_id: i32,
role_id: i32,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].attrs.primary_key,
Some(vec!["user_id".to_string(), "role_id".to_string()])
);
}
#[test]
fn test_parse_primary_key_three_columns() {
let input = quote! {
#[primary_key(a_id, b_id, c_id)]
ThreeWayJoin {
a_id: i32,
b_id: i32,
c_id: i32,
}
};
let schema = parse_schema(input).unwrap();
let pk = schema.entities[0].attrs.primary_key.as_ref().unwrap();
assert_eq!(pk.len(), 3);
assert_eq!(pk, &["a_id", "b_id", "c_id"]);
}
#[test]
fn test_parse_primary_key_with_table_name() {
let input = quote! {
#[table_name = "users_roles"]
#[primary_key(user_id, role_id)]
UsersRole {
user_id: i32,
role_id: i32,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].attrs.table_name,
Some("users_roles".to_string())
);
assert_eq!(
schema.entities[0].attrs.primary_key,
Some(vec!["user_id".to_string(), "role_id".to_string()])
);
}
#[test]
fn test_parse_primary_key_with_timestamps_none() {
let input = quote! {
#[primary_key(user_id, role_id)]
#[timestamps(none)]
UsersRole {
user_id: i32,
role_id: i32,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].attrs.primary_key.is_some());
assert!(!schema.entities[0].attrs.has_created_at);
assert!(!schema.entities[0].attrs.has_updated_at);
}
#[test]
fn test_id_allowed_with_custom_primary_key() {
let input = quote! {
#[primary_key(id)]
LegacyTable {
id: i64,
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert_eq!(
schema.entities[0].attrs.primary_key,
Some(vec!["id".to_string()])
);
assert_eq!(schema.entities[0].fields.len(), 2);
}
#[test]
fn test_id_rejected_without_custom_primary_key() {
let input = quote! {
User {
id: i32,
name: String,
}
};
let result = parse_schema(input);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("reserved"));
}
#[test]
fn test_parse_bytes_field() {
let input = quote! {
User {
avatar: Vec<u8>,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
if let RawFieldType::Scalar { scalar, .. } = &field.ty {
assert_eq!(*scalar, ScalarType::Bytes);
} else {
panic!("Expected scalar type");
}
}
#[test]
fn test_parse_option_bytes_field() {
let input = quote! {
User {
avatar: Option<Vec<u8>>,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
if let RawFieldType::Scalar { scalar, optional } = &field.ty {
assert_eq!(*scalar, ScalarType::Bytes);
assert!(optional);
} else {
panic!("Expected scalar type");
}
}
#[test]
fn test_parse_time_field() {
let input = quote! {
Event {
start_time: Time,
}
};
let schema = parse_schema(input).unwrap();
let field = &schema.entities[0].fields[0];
if let RawFieldType::Scalar { scalar, .. } = &field.ty {
assert_eq!(*scalar, ScalarType::Time);
} else {
panic!("Expected scalar type");
}
}
#[test]
fn test_default_no_primary_key() {
let input = quote! {
User {
name: String,
}
};
let schema = parse_schema(input).unwrap();
assert!(schema.entities[0].attrs.primary_key.is_none());
}
}