1use crate::model::Model;
35use crate::pool::Connection;
36use crate::query::QueryBuilder;
37use crate::value::{FromQueryResult, Value};
38use crate::DbError;
39use std::future::Future;
40use std::pin::Pin;
41
42pub use crate::repository::PageResult;
44
45pub trait PaginatorTrait<M: Model>: Sized {
56 fn paginate<T, C>(
64 self,
65 page: u64,
66 page_size: u64,
67 conn: &mut C,
68 ) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
69 where
70 T: FromQueryResult + Send + 'static,
71 C: Connection + Send;
72
73 fn find_page<T, C>(
87 self,
88 page: u64,
89 page_size: u64,
90 conn: &mut C,
91 ) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
92 where
93 T: FromQueryResult + Send + 'static,
94 C: Connection + Send,
95 {
96 self.paginate::<T, C>(page, page_size, conn)
97 }
98}
99
100impl<M: Model> PaginatorTrait<M> for QueryBuilder<M> {
101 fn paginate<T, C>(
102 self,
103 page: u64,
104 page_size: u64,
105 conn: &mut C,
106 ) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
107 where
108 T: FromQueryResult + Send + 'static,
109 C: Connection + Send,
110 {
111 let count_sql = self.build_count();
112 let offset = if page_size == 0 {
113 0
114 } else {
115 ((page.saturating_sub(1)) * page_size) as usize
116 };
117 let limit = page_size as usize;
118 let (select_sql, params) = self.limit(limit).offset(offset).build_select_with_params();
119
120 Box::pin(async move {
121 let count_rows = conn.query(&count_sql).await?;
122 let total = extract_count(&count_rows).unwrap_or(0);
123
124 let data_rows = if total == 0 || offset as u64 >= total {
125 Vec::new()
126 } else {
127 let rows = conn.query_with_params(&select_sql, ¶ms).await?;
128 let mut items = Vec::with_capacity(rows.len());
129 for row in &rows {
130 match T::from_query_result(row) {
131 Ok(item) => items.push(item),
132 Err(e) => return Err(DbError::Internal(e)),
133 }
134 }
135 items
136 };
137
138 Ok(PageResult::new(data_rows, total, page, page_size))
139 })
140 }
141}
142
143pub struct Paginator<'a, C>
159where
160 C: Connection,
161{
162 conn: &'a mut C,
163 page_size: u64,
164 sql: String,
165 params: Vec<Value>,
166 total: u64,
167}
168
169impl<'a, C> Paginator<'a, C>
170where
171 C: Connection,
172{
173 pub async fn fetch_page<T>(&mut self, page: u64) -> Result<PageResult<T>, DbError>
182 where
183 T: FromQueryResult + Send + 'static,
184 {
185 let offset = if self.page_size == 0 {
186 0
187 } else {
188 ((page.saturating_sub(1)) * self.page_size) as usize
189 };
190 let limit = self.page_size as usize;
191
192 let select_sql = format!("{} LIMIT {} OFFSET {}", self.sql, limit, offset);
193
194 let rows = self
195 .conn
196 .query_with_params(&select_sql, &self.params)
197 .await?;
198 let mut items = Vec::with_capacity(rows.len());
199 for row in &rows {
200 match T::from_query_result(row) {
201 Ok(item) => items.push(item),
202 Err(e) => return Err(DbError::Internal(e)),
203 }
204 }
205
206 Ok(PageResult::new(items, self.total, page, self.page_size))
207 }
208
209 pub fn set_total(&mut self, total: u64) {
211 self.total = total;
212 }
213}
214
215pub trait PaginatorBuilderTrait<M: Model> {
228 fn paginate_with<'a, C>(self, conn: &'a mut C, page_size: u64) -> Paginator<'a, C>
230 where
231 C: Connection;
232}
233
234impl<M: Model> PaginatorBuilderTrait<M> for QueryBuilder<M> {
235 fn paginate_with<'a, C>(self, conn: &'a mut C, page_size: u64) -> Paginator<'a, C>
236 where
237 C: Connection,
238 {
239 let (sql, params) = self.build_select_with_params();
240 Paginator {
241 conn,
242 page_size,
243 sql,
244 params,
245 total: 0,
246 }
247 }
248}
249
250pub type RowResult = std::collections::HashMap<String, Value>;
256
257pub trait StreamQueryTrait<M: Model> {
274 fn stream<'a, 'b: 'a, C: Connection + Send + 'b>(
276 self,
277 conn: &'b mut C,
278 ) -> Pin<Box<dyn futures::Stream<Item = Result<RowResult, DbError>> + Send + 'a>>;
279}
280
281impl<M: Model> StreamQueryTrait<M> for QueryBuilder<M> {
282 fn stream<'a, 'b: 'a, C: Connection + Send + 'b>(
283 self,
284 conn: &'b mut C,
285 ) -> Pin<Box<dyn futures::Stream<Item = Result<RowResult, DbError>> + Send + 'a>> {
286 let (sql, params) = self.build_select_with_params();
287
288 use futures::{stream, StreamExt};
289
290 let st = stream::once(async move {
291 match conn.query_with_params(&sql, ¶ms).await {
292 Ok(rows) => rows.into_iter().map(Ok).collect::<Vec<_>>(),
293 Err(e) => vec![Err(e)],
294 }
295 })
296 .flat_map(stream::iter);
297
298 Box::pin(st)
299 }
300}
301
302fn extract_count(rows: &[std::collections::HashMap<String, Value>]) -> Option<u64> {
307 rows.first().and_then(|row| {
308 if let Some(v) = row.get("total") {
309 return value_to_u64(v);
310 }
311 if let Some(v) = row.get("COUNT(*)") {
312 return value_to_u64(v);
313 }
314 row.values().next().and_then(value_to_u64)
315 })
316}
317
318fn value_to_u64(v: &Value) -> Option<u64> {
319 match v {
320 Value::I64(n) => Some(*n as u64),
321 Value::I32(n) => Some(*n as u64),
322 Value::U64(n) => Some(*n),
323 Value::U32(n) => Some(*n as u64),
324 Value::F64(n) => Some(*n as u64),
325 _ => None,
326 }
327}
328
329#[cfg(test)]
334mod tests {
335 use super::*;
336 use crate::db_type::DbType;
337 use crate::dialect::get_dialect;
338 use crate::mock::MockConnection;
339
340 #[derive(Debug, Clone, Default)]
341 struct PagTestModel;
342 impl Model for PagTestModel {
343 type PrimaryKey = i64;
344 fn table_name() -> &'static str {
345 "pag_test"
346 }
347 fn pk_name() -> &'static str {
348 "id"
349 }
350 fn pk(&self) -> Self::PrimaryKey {
351 0
352 }
353 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
354 fn timestamp_fields() -> Option<crate::model::TimestampFields> {
355 None
356 }
357 fn soft_delete_field() -> Option<&'static str> {
358 None
359 }
360 }
361
362 #[derive(Debug, Clone, Default)]
363 #[allow(dead_code)]
364 struct PagRow {
365 id: i64,
366 }
367 impl FromQueryResult for PagRow {
368 fn from_query_result(
369 row: &std::collections::HashMap<String, Value>,
370 ) -> Result<Self, String> {
371 let id = row.get("id").and_then(|v| v.as_i64()).unwrap_or(0);
372 Ok(PagRow { id })
373 }
374 }
375
376 #[tokio::test]
377 async fn test_paginator_trait_exists() {
378 fn _assert<M: Model, Q: PaginatorTrait<M>>() {}
379 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
380 }
381
382 #[tokio::test]
383 async fn test_paginator_with_mock() {
384 let dialect = get_dialect(DbType::MySQL).unwrap();
385 let mut mock = MockConnection::new();
386
387 let _ = mock
388 .expect_any()
389 .with_rows(vec![vec![("total", Value::I64(100))]]);
390 let _ = mock.expect_any().with_rows(vec![
391 vec![("id", Value::I64(1))],
392 vec![("id", Value::I64(2))],
393 ]);
394
395 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
396 .where_eq("status", Value::from("active"))
397 .paginate::<PagRow, _>(1, 20, &mut mock)
398 .await
399 .unwrap();
400
401 assert_eq!(result.total, 100);
402 assert_eq!(result.page, 1);
403 assert_eq!(result.page_size, 20);
404 assert_eq!(result.total_pages(), 5);
405 assert!(result.has_next());
406 assert!(!result.has_prev());
407 }
408
409 #[tokio::test]
410 async fn test_paginator_empty_result() {
411 let dialect = get_dialect(DbType::MySQL).unwrap();
412 let mut mock = MockConnection::new();
413
414 mock.expect_any()
415 .with_rows(vec![vec![("total", Value::I64(0))]]);
416 mock.expect_any().with_rows(vec![]);
417
418 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
419 .paginate::<PagRow, _>(1, 20, &mut mock)
420 .await
421 .unwrap();
422
423 assert_eq!(result.total, 0);
424 assert!(result.is_empty());
425 assert_eq!(result.total_pages(), 0);
426 }
427
428 #[tokio::test]
429 async fn test_paginator_page_beyond_range() {
430 let dialect = get_dialect(DbType::MySQL).unwrap();
431 let mut mock = MockConnection::new();
432
433 mock.expect_any()
434 .with_rows(vec![vec![("total", Value::I64(10))]]);
435 mock.expect_any().with_rows(vec![]);
436
437 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
438 .paginate::<PagRow, _>(99, 20, &mut mock)
439 .await
440 .unwrap();
441
442 assert_eq!(result.total, 10);
443 assert_eq!(result.page, 99);
444 assert!(result.is_empty());
445 }
446
447 #[tokio::test]
457 async fn test_find_page_alias_of_paginate() {
458 let dialect = get_dialect(DbType::MySQL).unwrap();
459 let mut mock = MockConnection::new();
460
461 let _ = mock
462 .expect_any()
463 .with_rows(vec![vec![("total", Value::I64(50))]]);
464 let _ = mock.expect_any().with_rows(vec![
465 vec![("id", Value::I64(21))],
466 vec![("id", Value::I64(22))],
467 ]);
468
469 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
470 .where_eq("status", Value::from("active"))
471 .find_page::<PagRow, _>(3, 10, &mut mock)
472 .await
473 .unwrap();
474
475 assert_eq!(result.total, 50);
476 assert_eq!(result.page, 3);
477 assert_eq!(result.page_size, 10);
478 assert_eq!(result.total_pages(), 5);
479 assert!(result.has_next());
480 assert!(result.has_prev());
481 }
482
483 #[tokio::test]
485 async fn test_find_page_empty_result() {
486 let dialect = get_dialect(DbType::MySQL).unwrap();
487 let mut mock = MockConnection::new();
488
489 mock.expect_any()
490 .with_rows(vec![vec![("total", Value::I64(0))]]);
491
492 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
493 .find_page::<PagRow, _>(1, 20, &mut mock)
494 .await
495 .unwrap();
496
497 assert_eq!(result.total, 0);
498 assert!(result.is_empty());
499 }
500
501 #[tokio::test]
504 async fn test_paginator_total_zero_skips_data_query() {
505 let dialect = get_dialect(DbType::MySQL).unwrap();
506 let mut mock = MockConnection::new();
507
508 mock.expect_any()
510 .with_rows(vec![vec![("total", Value::I64(0))]]);
511 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
514 .paginate::<PagRow, _>(1, 20, &mut mock)
515 .await
516 .unwrap();
517
518 assert_eq!(result.total, 0);
519 assert!(result.is_empty());
520 assert_eq!(mock.executed_sql().len(), 1);
523 }
524
525 #[tokio::test]
528 async fn test_paginator_offset_equals_total_returns_empty() {
529 let dialect = get_dialect(DbType::MySQL).unwrap();
530 let mut mock = MockConnection::new();
531
532 mock.expect_any()
533 .with_rows(vec![vec![("total", Value::I64(10))]]);
534 mock.expect_any().with_rows(vec![]);
536
537 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
538 .paginate::<PagRow, _>(2, 10, &mut mock)
539 .await
540 .unwrap();
541
542 assert_eq!(result.total, 10);
543 assert_eq!(result.page, 2);
544 assert!(result.is_empty());
545 assert_eq!(mock.executed_sql().len(), 1);
548 }
549
550 #[tokio::test]
553 async fn test_paginator_page_size_zero_offset_zero() {
554 let dialect = get_dialect(DbType::MySQL).unwrap();
555 let mut mock = MockConnection::new();
556
557 mock.expect_any()
558 .with_rows(vec![vec![("total", Value::I64(100))]]);
559 mock.expect_any()
560 .with_rows(vec![vec![("id", Value::I64(1))]]);
561
562 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
563 .paginate::<PagRow, _>(3, 0, &mut mock)
564 .await
565 .unwrap();
566
567 assert_eq!(result.total, 100);
568 let data_sql = mock.executed_sql().get(1).unwrap();
570 assert!(
571 data_sql.contains("OFFSET 0"),
572 "expected OFFSET 0, got: {data_sql}"
573 );
574 }
575
576 #[tokio::test]
580 async fn test_paginator_offset_calculation_exact() {
581 let dialect = get_dialect(DbType::MySQL).unwrap();
582 let mut mock = MockConnection::new();
583
584 mock.expect_any()
585 .with_rows(vec![vec![("total", Value::I64(100))]]);
586 mock.expect_any().with_rows(vec![
587 vec![("id", Value::I64(11))],
588 vec![("id", Value::I64(12))],
589 ]);
590
591 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
592 .paginate::<PagRow, _>(3, 5, &mut mock)
593 .await
594 .unwrap();
595
596 assert_eq!(result.total, 100);
597 assert_eq!(result.page, 3);
598 assert_eq!(result.page_size, 5);
599 assert_eq!(result.items.len(), 2);
600
601 let data_sql = mock.executed_sql().get(1).unwrap();
603 assert!(
604 data_sql.contains("OFFSET 10"),
605 "expected OFFSET 10 in SQL, got: {data_sql}"
606 );
607 assert!(data_sql.contains("LIMIT 5"));
608 }
609
610 #[tokio::test]
615 async fn test_paginator_builder_trait_exists() {
616 use super::PaginatorBuilderTrait;
617 fn _assert<M: Model, Q: PaginatorBuilderTrait<M>>() {}
618 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
619 }
620
621 #[tokio::test]
622 async fn test_paginator_builder_fetch_page() {
623 use super::PaginatorBuilderTrait;
624
625 let dialect = get_dialect(DbType::MySQL).unwrap();
626 let mut mock = MockConnection::new();
627
628 let _ = mock.expect_any().with_rows(vec![
630 vec![("id", Value::I64(1))],
631 vec![("id", Value::I64(2))],
632 vec![("id", Value::I64(3))],
633 ]);
634
635 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
636 .where_eq("status", Value::from("active"))
637 .paginate_with(&mut mock, 20);
638 p.set_total(100);
639
640 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(1).await.unwrap();
641
642 assert_eq!(result.total, 100);
643 assert_eq!(result.page, 1);
644 assert_eq!(result.page_size, 20);
645 assert_eq!(result.items.len(), 3);
646 }
647
648 #[tokio::test]
652 async fn test_fetch_page_offset_multiplication() {
653 use super::PaginatorBuilderTrait;
654
655 let dialect = get_dialect(DbType::MySQL).unwrap();
656 let mut mock = MockConnection::new();
657
658 mock.expect_any().with_rows(vec![
659 vec![("id", Value::I64(11))],
660 vec![("id", Value::I64(12))],
661 ]);
662
663 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
664 .where_eq("status", Value::from("active"))
665 .paginate_with(&mut mock, 5);
666 p.set_total(100);
667
668 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
669
670 assert_eq!(result.items.len(), 2);
671 let data_sql = mock.executed_sql().first().unwrap();
673 assert!(
674 data_sql.contains("OFFSET 10"),
675 "expected OFFSET 10 in fetch_page SQL, got: {data_sql}"
676 );
677 }
678
679 #[tokio::test]
682 async fn test_fetch_page_page_size_zero_offset_zero() {
683 use super::PaginatorBuilderTrait;
684
685 let dialect = get_dialect(DbType::MySQL).unwrap();
686 let mut mock = MockConnection::new();
687
688 mock.expect_any()
689 .with_rows(vec![vec![("id", Value::I64(1))]]);
690
691 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
692 .where_eq("status", Value::from("active"))
693 .paginate_with(&mut mock, 0);
694 p.set_total(100);
695
696 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
697
698 assert_eq!(result.items.len(), 1);
699 let data_sql = mock.executed_sql().first().unwrap();
702 assert!(
703 data_sql.contains("OFFSET 0"),
704 "expected OFFSET 0, got: {data_sql}"
705 );
706 }
707
708 #[tokio::test]
713 async fn test_stream_trait_exists() {
714 use super::StreamQueryTrait;
715 fn _assert<M: Model, Q: StreamQueryTrait<M>>() {}
716 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
717 }
718
719 #[tokio::test]
720 async fn test_stream_returns_rows() {
721 use super::StreamQueryTrait;
722 use futures::StreamExt;
723
724 let dialect = get_dialect(DbType::MySQL).unwrap();
725 let mut mock = MockConnection::new();
726
727 let _ = mock.expect_any().with_rows(vec![
728 vec![("id", Value::I64(1))],
729 vec![("id", Value::I64(2))],
730 ]);
731
732 let mut stream = QueryBuilder::<PagTestModel>::new(dialect)
733 .where_eq("status", Value::from("active"))
734 .stream(&mut mock);
735
736 let mut count = 0;
737 while let Some(result) = stream.next().await {
738 let row: RowResult = result.unwrap();
739 assert!(row.contains_key("id"));
740 count += 1;
741 }
742 assert_eq!(count, 2);
743 }
744
745 #[test]
751 fn test_value_to_u64_i32() {
752 assert_eq!(value_to_u64(&Value::I32(42)), Some(42u64));
753 assert_eq!(value_to_u64(&Value::I32(-1)), Some(u64::MAX)); }
755
756 #[test]
758 fn test_value_to_u64_u64() {
759 assert_eq!(value_to_u64(&Value::U64(999)), Some(999u64));
760 }
761
762 #[test]
764 fn test_value_to_u64_u32() {
765 assert_eq!(value_to_u64(&Value::U32(77)), Some(77u64));
766 }
767
768 #[test]
770 fn test_value_to_u64_f64() {
771 assert_eq!(value_to_u64(&Value::F64(2.71)), Some(2u64));
772 assert_eq!(value_to_u64(&Value::F64(0.0)), Some(0u64));
773 }
774
775 #[test]
777 fn test_value_to_u64_i64() {
778 assert_eq!(value_to_u64(&Value::I64(123456789)), Some(123456789u64));
779 }
780
781 #[test]
783 fn test_value_to_u64_unknown() {
784 assert_eq!(value_to_u64(&Value::Bool(true)), None);
785 assert_eq!(value_to_u64(&Value::String("x".into())), None);
786 }
787}