use proc_macro2::TokenStream;
use quote::quote;
use syn::DeriveInput;
#[derive(Clone)]
pub struct RelationInfo {
pub field_name: syn::Ident,
pub related_entity: syn::Path,
pub foreign_key_field: syn::Ident,
}
pub fn parse_relations(input: &DeriveInput) -> Vec<RelationInfo> {
let mut relations = Vec::new();
for attr in &input.attrs {
if !attr.path().is_ident("ormada") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("relations") {
let content;
syn::parenthesized!(content in meta.input);
while !content.is_empty() {
let field_name: syn::Ident = content.parse()?;
let _: syn::Token![=] = content.parse()?;
let entity_path: syn::LitStr = content.parse()?;
let related_entity: syn::Path = syn::parse_str(&entity_path.value())
.map_err(|e| syn::Error::new(entity_path.span(), e))?;
let fk_field_name = format!("{}_id", field_name);
let foreign_key_field = syn::Ident::new(&fk_field_name, field_name.span());
relations.push(RelationInfo {
field_name,
related_entity,
foreign_key_field,
});
if content.peek(syn::Token![,]) {
let _: syn::Token![,] = content.parse()?;
}
}
}
Ok(())
});
}
relations
}
pub fn generate_model_with_relations(
fields: &syn::punctuated::Punctuated<syn::Field, syn::token::Comma>,
relations: &[RelationInfo],
) -> TokenStream {
let original_fields: Vec<_> = fields
.iter()
.map(|field| {
let name = &field.ident;
let ty = &field.ty;
let vis = &field.vis;
quote! {
#vis #name: #ty
}
})
.collect();
let relation_fields: Vec<_> = relations
.iter()
.map(|rel| {
let field_name = &rel.field_name;
let related_entity = &rel.related_entity;
quote! {
pub #field_name: ::core::option::Option<
<#related_entity as ::ormada::traits::WithRelationsTrait>::ModelWithRelations
>
}
})
.collect();
let accessor_methods: Vec<_> = relations
.iter()
.map(|rel| {
let field_name = &rel.field_name;
let related_entity = &rel.related_entity;
let doc = format!("Get the {} relation if it was prefetched", field_name);
quote! {
#[doc = #doc]
pub fn #field_name(&self) -> ::core::option::Option<
&<#related_entity as ::ormada::traits::WithRelationsTrait>::ModelWithRelations
> {
self.#field_name.as_ref()
}
}
})
.collect();
quote! {
#[derive(Clone, Debug)]
pub struct ModelWithRelations {
#(#original_fields,)*
#(#relation_fields,)*
}
impl ModelWithRelations {
#(#accessor_methods)*
}
}
}
pub fn generate_from_impl(
fields: &syn::punctuated::Punctuated<syn::Field, syn::token::Comma>,
relations: &[RelationInfo],
) -> TokenStream {
let field_copies: Vec<_> = fields
.iter()
.map(|field| {
let name = field.ident.as_ref().unwrap();
quote! { #name: model.#name }
})
.collect();
let relation_nones: Vec<_> = relations
.iter()
.map(|rel| {
let field_name = &rel.field_name;
quote! { #field_name: ::core::option::Option::None }
})
.collect();
quote! {
impl ::core::convert::From<Model> for ModelWithRelations {
fn from(model: Model) -> Self {
Self {
#(#field_copies,)*
#(#relation_nones,)*
}
}
}
}
}
pub fn generate_trait_impl(relations: &[RelationInfo], fields: &[&syn::Field]) -> TokenStream {
let relations_type = if relations.is_empty() {
quote! { () }
} else if relations.len() == 1 {
let related_entity = &relations[0].related_entity;
quote! { ::rustc_hash::FxHashMap<i32, <#related_entity as ::sea_orm::EntityTrait>::Model> }
} else {
let hashmap_types: Vec<_> = relations
.iter()
.map(|rel| {
let related_entity = &rel.related_entity;
quote! {
::rustc_hash::FxHashMap<i32, <#related_entity as ::sea_orm::EntityTrait>::Model>
}
})
.collect();
quote! { ( #(#hashmap_types),* ) }
};
let relation_lookups: Vec<_> = relations
.iter()
.enumerate()
.map(|(idx, rel)| {
let field_name = &rel.field_name;
let fk_field = &rel.foreign_key_field;
let related_entity = &rel.related_entity;
let access = if relations.len() == 1 {
quote! { relations }
} else {
let index = syn::Index::from(idx);
quote! { &relations.#index }
};
quote! {
#field_name: #access
.get(&model.#fk_field)
.cloned()
.map(|m| <#related_entity as ::ormada::traits::WithRelationsTrait>::from_model_and_relations(m, &()))
}
})
.collect();
let field_copies: Vec<_> = fields
.iter()
.map(|field| {
let name = field.ident.as_ref().unwrap();
quote! { #name: model.#name }
})
.collect();
quote! {
impl ::ormada::traits::WithRelationsTrait for Entity {
type Model = Model;
type ModelWithRelations = ModelWithRelations;
type Relations = #relations_type;
fn from_model_and_relations(
model: Self::Model,
relations: &Self::Relations,
) -> Self::ModelWithRelations {
ModelWithRelations {
#(#field_copies,)*
#(#relation_lookups,)*
}
}
}
}
}
pub fn generate_has_relation_impls(relations: &[RelationInfo]) -> TokenStream {
if relations.is_empty() {
return quote! {};
}
let impls: Vec<_> = relations
.iter()
.map(|rel| {
let related_entity = &rel.related_entity;
let fk_field = &rel.foreign_key_field;
let relation_name = &rel.field_name;
quote! {
impl ::ormada::relations::HasRelation<#related_entity> for Entity {
type RelatedPK = <<#related_entity as ::sea_orm::EntityTrait>::PrimaryKey as ::sea_orm::PrimaryKeyTrait>::ValueType;
fn get_foreign_key(model: &Self::Model) -> Self::RelatedPK {
model.#fk_field
}
async fn load_related<C: ::sea_orm::ConnectionTrait>(
models: &[Self::Model],
db: &C,
) -> ::core::result::Result<
::ormada::prelude::FxHashMap<Self::RelatedPK, <#related_entity as ::sea_orm::EntityTrait>::Model>,
::ormada::error::OrmadaOrmError
> {
use ::sea_orm::{EntityTrait, QueryFilter, ColumnTrait, PrimaryKeyToColumn, ModelTrait, Iterable};
let fk_values: ::std::vec::Vec<Self::RelatedPK> = models
.iter()
.map(|m| m.#fk_field)
.collect();
if fk_values.is_empty() {
return ::core::result::Result::Ok(::ormada::prelude::FxHashMap::default());
}
let pk_cols: ::std::vec::Vec<_> = <#related_entity as ::sea_orm::EntityTrait>::PrimaryKey::iter()
.map(|pk| pk.into_column())
.collect();
let id_column = pk_cols[0];
let related_models = <#related_entity as ::sea_orm::EntityTrait>::find()
.filter(id_column.is_in(fk_values))
.all(db)
.await?;
::core::result::Result::Ok(
::ormada::prelude::FxHashMap::from_iter(
related_models
.into_iter()
.map(|m| (m.id, m))
)
)
}
fn set_related(
model: &mut <Self as ::ormada::traits::WithRelationsTrait>::ModelWithRelations,
related: ::core::option::Option<<#related_entity as ::sea_orm::EntityTrait>::Model>,
) {
model.#relation_name = related.expect("select_related requires non-null FK");
}
fn relation_def() -> ::sea_orm::RelationDef {
<Relation as ::sea_orm::RelationTrait>::def(&Relation::#relation_name)
}
}
}
})
.collect();
quote! {
#(#impls)*
}
}