Skip to main content

sea_orm_codegen/entity/writer/
dense.rs

1use super::*;
2use crate::{Entity, Relation, RelationType};
3use heck::ToSnakeCase;
4use sea_query::ForeignKeyAction;
5
6fn relation_column_name(a: &syn::Ident) -> String {
7    let a = a.to_string();
8    let b = a.to_snake_case();
9    if a != b.to_upper_camel_case() {
10        // if roundtrip fails, use original
11        a
12    } else {
13        b
14    }
15}
16
17fn relation_column_list(punctuated: Vec<String>) -> String {
18    let len = punctuated.len();
19    let punctuated = punctuated.join(", ");
20    match len {
21        0..=1 => punctuated,
22        _ => format!("({punctuated})"),
23    }
24}
25
26fn foreign_key_action_attr(
27    attr: TokenStream,
28    action: Option<&ForeignKeyAction>,
29) -> Option<TokenStream> {
30    action.map(|action| {
31        let action = Relation::get_foreign_key_action(action);
32        quote!(, #attr = #action)
33    })
34}
35
36struct DenseRelationField<'a> {
37    entity: &'a Entity,
38    rel: &'a Relation,
39    via_entities: &'a [syn::Ident],
40}
41
42impl DenseRelationField<'_> {
43    fn field_tokens(&self) -> Option<TokenStream> {
44        let (field, target_entity) = if self.rel.self_referencing {
45            let table_name = self.entity.get_table_name_snake_case_ident();
46            let suffix = if self.rel.num_suffix > 0 {
47                format!("_{}", self.rel.num_suffix)
48            } else {
49                String::new()
50            };
51            let field = format_ident!("{table_name}{suffix}");
52
53            (field, quote!(Entity))
54        } else {
55            if !self.rel.impl_related {
56                return None;
57            }
58
59            let to_entity = self.rel.get_module_name()?;
60            if self.via_entities.contains(&to_entity) {
61                return None;
62            }
63
64            let field = match self.rel.rel_type {
65                RelationType::HasMany => {
66                    let to_entity = to_entity.to_string();
67                    let pluralized = pluralizer::pluralize(&to_entity, 2, false);
68                    format_ident!("{pluralized}")
69                }
70                RelationType::HasOne | RelationType::BelongsTo => to_entity.clone(),
71            };
72            let field = if self.rel.num_suffix == 0 {
73                field
74            } else {
75                format_ident!("{field}_{}", self.rel.num_suffix)
76            };
77            (field, quote!(super::#to_entity::Entity))
78        };
79
80        let rel_field_type = match self.rel.rel_type {
81            RelationType::BelongsTo => {
82                let is_optional = !self.rel.columns.is_empty()
83                    && self.rel.columns.iter().all(|name| {
84                        self.entity
85                            .columns
86                            .iter()
87                            .find(|column| column.name == *name)
88                            .is_some_and(|column| column.not_null)
89                    });
90                if is_optional {
91                    quote!(BelongsTo<#target_entity>)
92                } else {
93                    quote!(BelongsTo<Option<#target_entity>>)
94                }
95            }
96            RelationType::HasOne => quote!(HasOne<#target_entity>),
97            RelationType::HasMany => quote!(HasMany<#target_entity>),
98        };
99        let sea_orm_attr = if self.rel.self_referencing {
100            let (from, to) = self.rel.get_src_ref_columns(
101                relation_column_name,
102                relation_column_name,
103                relation_column_list,
104            );
105            let on_update = foreign_key_action_attr(quote!(on_update), self.rel.on_update.as_ref());
106            let on_delete = foreign_key_action_attr(quote!(on_delete), self.rel.on_delete.as_ref());
107            let relation_enum = self.rel.get_enum_name().to_string();
108
109            quote!(#[sea_orm(self_ref, relation_enum = #relation_enum, from = #from, to = #to #on_update #on_delete)])
110        } else {
111            match self.rel.rel_type {
112                RelationType::HasOne => quote!(#[sea_orm(has_one)]),
113                RelationType::HasMany => quote!(#[sea_orm(has_many)]),
114                RelationType::BelongsTo => {
115                    let (from, to) = self.rel.get_src_ref_columns(
116                        relation_column_name,
117                        relation_column_name,
118                        relation_column_list,
119                    );
120                    let on_update =
121                        foreign_key_action_attr(quote!(on_update), self.rel.on_update.as_ref());
122                    let on_delete =
123                        foreign_key_action_attr(quote!(on_delete), self.rel.on_delete.as_ref());
124                    let relation_enum = if self.rel.num_suffix > 0 {
125                        let relation_enum = self.rel.get_enum_name().to_string();
126                        Some(quote!(relation_enum = #relation_enum,))
127                    } else {
128                        None
129                    };
130
131                    quote!(#[sea_orm(belongs_to, #relation_enum from = #from, to = #to #on_update #on_delete)])
132                }
133            }
134        };
135
136        Some(quote! {
137            #sea_orm_attr
138            pub #field: #rel_field_type
139        })
140    }
141}
142
143impl EntityWriter {
144    #[allow(clippy::too_many_arguments)]
145    pub fn gen_dense_code_blocks(
146        entity: &Entity,
147        with_serde: &WithSerde,
148        column_option: &ColumnOption,
149        schema_name: &Option<String>,
150        serde_skip_deserializing_primary_key: bool,
151        serde_skip_hidden_column: bool,
152        model_extra_derives: &TokenStream,
153        model_extra_attributes: &TokenStream,
154        _column_extra_derives: &TokenStream,
155        _seaography: bool,
156        impl_active_model_behavior: bool,
157    ) -> Vec<TokenStream> {
158        let mut imports = Self::gen_import(with_serde);
159        let active_enums = Self::gen_import_active_enum(entity);
160        imports.extend(active_enums.imports);
161        let mut code_blocks = vec![
162            imports,
163            Self::gen_dense_model_struct(
164                entity,
165                with_serde,
166                column_option,
167                schema_name,
168                serde_skip_deserializing_primary_key,
169                serde_skip_hidden_column,
170                model_extra_derives,
171                model_extra_attributes,
172                &active_enums.type_idents,
173            ),
174        ];
175        if impl_active_model_behavior {
176            code_blocks.push(Self::impl_active_model_behavior());
177        }
178        code_blocks
179    }
180
181    #[allow(clippy::too_many_arguments)]
182    pub fn gen_dense_model_struct(
183        entity: &Entity,
184        with_serde: &WithSerde,
185        column_option: &ColumnOption,
186        schema_name: &Option<String>,
187        serde_skip_deserializing_primary_key: bool,
188        serde_skip_hidden_column: bool,
189        model_extra_derives: &TokenStream,
190        model_extra_attributes: &TokenStream,
191        active_enum_type_idents: &ActiveEnumTypeIdents,
192    ) -> TokenStream {
193        let table_name = entity.table_name.as_str();
194        let column_names_snake_case = entity.get_column_names_snake_case();
195        let column_rs_types = Self::get_column_rs_types_with_enum_idents(
196            entity,
197            column_option,
198            active_enum_type_idents,
199        );
200        let if_eq_needed = entity.get_eq_needed();
201        let primary_keys: Vec<String> = entity
202            .primary_keys
203            .iter()
204            .map(|pk| pk.name.clone())
205            .collect();
206        let attrs: Vec<TokenStream> = entity
207            .columns
208            .iter()
209            .map(|col| {
210                let mut attrs: Punctuated<_, Comma> = Punctuated::new();
211                let is_primary_key = primary_keys.contains(&col.name);
212                if !col.is_snake_case_name() {
213                    let column_name = &col.name;
214                    attrs.push(quote! { column_name = #column_name });
215                }
216                if is_primary_key {
217                    attrs.push(quote! { primary_key });
218                    if !col.auto_increment {
219                        attrs.push(quote! { auto_increment = false });
220                    }
221                }
222                if let Some(ts) = col.get_col_type_attrs() {
223                    attrs.extend([ts]);
224                    if !col.not_null {
225                        attrs.push(quote! { nullable });
226                    }
227                };
228                if col.unique {
229                    attrs.push(quote! { unique });
230                } else if let Some(unique_key) = &col.unique_key {
231                    attrs.push(quote! { unique_key = #unique_key });
232                }
233                let mut ts = quote! {};
234                if !attrs.is_empty() {
235                    for (i, attr) in attrs.into_iter().enumerate() {
236                        if i > 0 {
237                            ts = quote! { #ts, };
238                        }
239                        ts = quote! { #ts #attr };
240                    }
241                    ts = quote! { #[sea_orm(#ts)] };
242                }
243                let serde_attribute = col.get_serde_attribute(
244                    is_primary_key,
245                    serde_skip_deserializing_primary_key,
246                    serde_skip_hidden_column,
247                );
248                ts = quote! {
249                    #ts
250                    #serde_attribute
251                };
252                ts
253            })
254            .collect();
255        let schema_name = match Self::gen_schema_name(schema_name) {
256            Some(schema_name) => quote! {
257                schema_name = #schema_name,
258            },
259            None => quote! {},
260        };
261        let extra_derive = with_serde.extra_derive();
262
263        let mut compound_objects: Punctuated<_, Comma> = Punctuated::new();
264
265        let via_entities = entity.get_conjunct_relations_via_snake_case();
266        for rel in entity.relations.iter() {
267            let relation_field = DenseRelationField {
268                entity,
269                rel,
270                via_entities: &via_entities,
271            };
272            if let Some(field) = relation_field.field_tokens() {
273                compound_objects.push(field);
274            }
275        }
276        for (to_entity, via_entity) in entity
277            .get_conjunct_relations_to_snake_case()
278            .into_iter()
279            .zip(via_entities)
280        {
281            let field = format_ident!(
282                "{}",
283                pluralizer::pluralize(&to_entity.to_string(), 2, false)
284            );
285            let via_entity = via_entity.to_string();
286            compound_objects.push(quote! {
287                #[sea_orm(has_many, via = #via_entity)]
288                pub #field: HasMany<super::#to_entity::Entity>
289            });
290        }
291
292        if !compound_objects.is_empty() {
293            compound_objects.push_punct(<syn::Token![,]>::default());
294        }
295
296        quote! {
297            #[sea_orm::model]
298            #[derive(Clone, Debug, PartialEq #if_eq_needed, DeriveEntityModel #extra_derive #model_extra_derives)]
299            #[sea_orm(
300                #schema_name
301                table_name = #table_name
302            )]
303            #model_extra_attributes
304            pub struct Model {
305                #(
306                    #attrs
307                    pub #column_names_snake_case: #column_rs_types,
308                )*
309                #compound_objects
310            }
311        }
312    }
313
314    #[allow(dead_code)]
315    fn gen_dense_related_entity(entity: &Entity) -> TokenStream {
316        let via_entities = entity.get_conjunct_relations_via_snake_case();
317
318        let related_modules = entity.get_related_entity_modules();
319        let related_attrs = entity.get_related_entity_attrs();
320        let related_enum_names = entity.get_related_entity_enum_name();
321
322        let items: Vec<_> = related_modules
323            .into_iter()
324            .zip(related_attrs)
325            .zip(related_enum_names)
326            .filter_map(|((related_module, related_attr), related_enum_name)| {
327                if !via_entities.contains(&related_module) {
328                    // skip junctions
329                    Some(quote!(#related_attr #related_enum_name))
330                } else {
331                    None
332                }
333            })
334            .collect();
335
336        quote! {
337            #[derive(Copy, Clone, Debug, EnumIter, DeriveRelatedEntity)]
338            pub enum RelatedEntity {
339                #(#items),*
340            }
341        }
342    }
343}
344
345#[cfg(test)]
346mod test {
347    #[test]
348    #[ignore]
349    fn test_name() {
350        panic!("{}", pluralizer::pluralize("filling", 2, false));
351    }
352}