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