Skip to main content

sea_orm/entity/
compound.rs

1#![allow(missing_docs)]
2use super::{ColumnTrait, EntityTrait, PrimaryKeyToColumn, PrimaryKeyTrait};
3use crate::{
4    ConnectionTrait, DbErr, IntoSimpleExpr, ItemsAndPagesNumber, Iterable, ModelTrait, QueryFilter,
5    QueryOrder,
6};
7use sea_query::{IntoValueTuple, Order, TableRef};
8use std::marker::PhantomData;
9
10mod belongs_to;
11mod has_many;
12mod has_one;
13
14pub use belongs_to::{BelongsTo, BelongsToCardinality};
15pub use has_many::{HasMany, Iter as HasManyIter};
16pub use has_one::HasOne;
17
18#[async_trait::async_trait]
19pub trait EntityLoaderTrait<E: EntityTrait>: QueryFilter + QueryOrder + Clone {
20    /// The return type of this loader
21    type ModelEx: ModelTrait<Entity = E>;
22
23    /// Find a model by primary key
24    fn filter_by_id<T>(mut self, values: T) -> Self
25    where
26        T: Into<<E::PrimaryKey as PrimaryKeyTrait>::ValueType>,
27    {
28        let mut keys = E::PrimaryKey::iter();
29        for v in values.into().into_value_tuple() {
30            if let Some(key) = keys.next() {
31                let col = key.into_column();
32                self.filter_mut(col.eq(v));
33            } else {
34                unreachable!("primary key arity mismatch");
35            }
36        }
37        self
38    }
39
40    /// Apply order by primary key to the query statement
41    fn order_by_id_asc(self) -> Self {
42        self.order_by_id(Order::Asc)
43    }
44
45    /// Apply order by primary key to the query statement
46    fn order_by_id_desc(self) -> Self {
47        self.order_by_id(Order::Desc)
48    }
49
50    /// Apply order by primary key to the query statement
51    fn order_by_id(mut self, order: Order) -> Self {
52        for key in E::PrimaryKey::iter() {
53            let col = key.into_column();
54            <Self as QueryOrder>::query(&mut self)
55                .order_by_expr(col.into_simple_expr(), order.clone());
56        }
57        self
58    }
59
60    /// Paginate query.
61    fn paginate<'db, C: ConnectionTrait>(
62        self,
63        db: &'db C,
64        page_size: u64,
65    ) -> EntityLoaderPaginator<'db, C, E, Self> {
66        EntityLoaderPaginator {
67            loader: self,
68            page: 0,
69            page_size,
70            db,
71            phantom: PhantomData,
72        }
73    }
74
75    #[doc(hidden)]
76    async fn fetch<C: ConnectionTrait>(
77        self,
78        db: &C,
79        page: u64,
80        page_size: u64,
81    ) -> Result<Vec<Self::ModelEx>, DbErr>;
82
83    #[doc(hidden)]
84    async fn num_items<C: ConnectionTrait>(self, db: &C, page_size: u64) -> Result<u64, DbErr>;
85}
86
87#[derive(Debug)]
88pub struct EntityLoaderPaginator<'db, C, E, L>
89where
90    C: ConnectionTrait,
91    E: EntityTrait,
92    L: EntityLoaderTrait<E>,
93{
94    pub(crate) loader: L,
95    pub(crate) page: u64,
96    pub(crate) page_size: u64,
97    pub(crate) db: &'db C,
98    pub(crate) phantom: PhantomData<E>,
99}
100
101/// Just a marker trait on EntityReverse
102pub trait EntityReverse {
103    type Entity: EntityTrait;
104}
105
106/// Subject to change, not yet stable
107#[derive(Debug, Copy, Clone, PartialEq)]
108pub struct EntityLoaderWithSelf<R: EntityTrait, S: EntityTrait>(pub R, pub S);
109
110/// Subject to change, not yet stable
111#[derive(Debug, Copy, Clone, PartialEq)]
112pub struct EntityLoaderWithSelfRev<R: EntityTrait, S: EntityReverse>(pub R, pub S);
113
114#[derive(Debug, Clone, PartialEq)]
115pub enum LoadTarget {
116    TableRef(TableRef),
117    TableRefRev(TableRef),
118    Relation(String),
119}
120
121impl<'db, C, E, L> EntityLoaderPaginator<'db, C, E, L>
122where
123    C: ConnectionTrait,
124    E: EntityTrait,
125    L: EntityLoaderTrait<E>,
126{
127    /// Fetch a specific page; page index starts from zero
128    pub async fn fetch_page(&self, page: u64) -> Result<Vec<L::ModelEx>, DbErr> {
129        self.loader
130            .clone()
131            .fetch(self.db, page, self.page_size)
132            .await
133    }
134
135    /// Fetch the current page
136    pub async fn fetch(&self) -> Result<Vec<L::ModelEx>, DbErr> {
137        self.fetch_page(self.page).await
138    }
139
140    /// Get the total number of items
141    pub async fn num_items(&self) -> Result<u64, DbErr> {
142        self.loader.clone().num_items(self.db, self.page_size).await
143    }
144
145    /// Get the total number of pages
146    pub async fn num_pages(&self) -> Result<u64, DbErr> {
147        let num_items = self.num_items().await?;
148        let num_pages = self.compute_pages_number(num_items);
149        Ok(num_pages)
150    }
151
152    /// Get the total number of items and pages
153    pub async fn num_items_and_pages(&self) -> Result<ItemsAndPagesNumber, DbErr> {
154        let number_of_items = self.num_items().await?;
155        let number_of_pages = self.compute_pages_number(number_of_items);
156
157        Ok(ItemsAndPagesNumber {
158            number_of_items,
159            number_of_pages,
160        })
161    }
162
163    /// Compute the number of pages for the current page
164    #[allow(clippy::manual_is_multiple_of)]
165    fn compute_pages_number(&self, num_items: u64) -> u64 {
166        (num_items / self.page_size) + (num_items % self.page_size > 0) as u64
167    }
168
169    /// Increment the page counter
170    pub fn next(&mut self) {
171        self.page += 1;
172    }
173
174    /// Get current page number
175    pub fn cur_page(&self) -> u64 {
176        self.page
177    }
178
179    /// Fetch one page and increment the page counter
180    pub async fn fetch_and_next(&mut self) -> Result<Option<Vec<L::ModelEx>>, DbErr> {
181        let vec = self.fetch().await?;
182        self.next();
183        let opt = if !vec.is_empty() { Some(vec) } else { None };
184        Ok(opt)
185    }
186}
187
188#[cfg(test)]
189mod test {
190    use crate::ModelTrait;
191    use crate::tests_cfg::cake;
192
193    #[test]
194    fn test_model_ex_convert() {
195        let cake = cake::Model {
196            id: 12,
197            name: "hello".into(),
198        };
199        let cake_ex: cake::ModelEx = cake.clone().into();
200
201        assert_eq!(cake, cake_ex);
202        assert_eq!(cake_ex, cake);
203        assert_eq!(cake.id, cake_ex.id);
204        assert_eq!(cake.name, cake_ex.name);
205
206        assert_eq!(cake_ex.get(cake::Column::Id), 12i32.into());
207        assert_eq!(cake_ex.get(cake::Column::Name), "hello".into());
208
209        assert_eq!(cake::Model::from(cake_ex), cake);
210    }
211}