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};
292
293 let st = stream::once(async move {
294 match conn.query(&sql).await {
295 Ok(rows) => rows.into_iter().map(Ok).collect::<Vec<_>>(),
296 Err(e) => vec![Err(e)],
297 }
298 })
299 .flat_map(stream::iter);
300
301 Box::pin(st)
302 }
303}
304
305fn extract_count(rows: &[std::collections::HashMap<String, Value>]) -> Option<u64> {
310 rows.first().and_then(|row| {
311 if let Some(v) = row.get("total") {
312 return value_to_u64(v);
313 }
314 if let Some(v) = row.get("COUNT(*)") {
315 return value_to_u64(v);
316 }
317 row.values().next().and_then(value_to_u64)
318 })
319}
320
321fn value_to_u64(v: &Value) -> Option<u64> {
322 match v {
323 Value::I64(n) => Some(*n as u64),
324 Value::I32(n) => Some(*n as u64),
325 Value::U64(n) => Some(*n),
326 Value::U32(n) => Some(*n as u64),
327 Value::F64(n) => Some(*n as u64),
328 _ => None,
329 }
330}
331
332#[cfg(test)]
337mod tests {
338 use super::*;
339 use crate::db_type::DbType;
340 use crate::dialect::get_dialect;
341 use crate::mock::MockConnection;
342
343 #[derive(Debug, Clone, Default)]
344 struct PagTestModel;
345 impl Model for PagTestModel {
346 type PrimaryKey = i64;
347 fn table_name() -> &'static str {
348 "pag_test"
349 }
350 fn pk_name() -> &'static str {
351 "id"
352 }
353 fn pk(&self) -> Self::PrimaryKey {
354 0
355 }
356 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
357 fn timestamp_fields() -> Option<crate::model::TimestampFields> {
358 None
359 }
360 fn soft_delete_field() -> Option<&'static str> {
361 None
362 }
363 }
364
365 #[derive(Debug, Clone, Default)]
366 #[allow(dead_code)]
367 struct PagRow {
368 id: i64,
369 }
370 impl FromQueryResult for PagRow {
371 fn from_query_result(
372 row: &std::collections::HashMap<String, Value>,
373 ) -> Result<Self, String> {
374 let id = row.get("id").and_then(|v| v.as_i64()).unwrap_or(0);
375 Ok(PagRow { id })
376 }
377 }
378
379 #[tokio::test]
380 async fn test_paginator_trait_exists() {
381 fn _assert<M: Model, Q: PaginatorTrait<M>>() {}
382 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
383 }
384
385 #[tokio::test]
386 async fn test_paginator_with_mock() {
387 let dialect = get_dialect(DbType::MySQL).unwrap();
388 let mut mock = MockConnection::new();
389
390 let _ = mock
391 .expect_any()
392 .with_rows(vec![vec![("total", Value::I64(100))]]);
393 let _ = mock.expect_any().with_rows(vec![
394 vec![("id", Value::I64(1))],
395 vec![("id", Value::I64(2))],
396 ]);
397
398 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
399 .where_eq("status", Value::from("active"))
400 .paginate::<PagRow, _>(1, 20, &mut mock)
401 .await
402 .unwrap();
403
404 assert_eq!(result.total, 100);
405 assert_eq!(result.page, 1);
406 assert_eq!(result.page_size, 20);
407 assert_eq!(result.total_pages(), 5);
408 assert!(result.has_next());
409 assert!(!result.has_prev());
410 }
411
412 #[tokio::test]
413 async fn test_paginator_empty_result() {
414 let dialect = get_dialect(DbType::MySQL).unwrap();
415 let mut mock = MockConnection::new();
416
417 mock.expect_any()
418 .with_rows(vec![vec![("total", Value::I64(0))]]);
419 mock.expect_any().with_rows(vec![]);
420
421 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
422 .paginate::<PagRow, _>(1, 20, &mut mock)
423 .await
424 .unwrap();
425
426 assert_eq!(result.total, 0);
427 assert!(result.is_empty());
428 assert_eq!(result.total_pages(), 0);
429 }
430
431 #[tokio::test]
432 async fn test_paginator_page_beyond_range() {
433 let dialect = get_dialect(DbType::MySQL).unwrap();
434 let mut mock = MockConnection::new();
435
436 mock.expect_any()
437 .with_rows(vec![vec![("total", Value::I64(10))]]);
438 mock.expect_any().with_rows(vec![]);
439
440 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
441 .paginate::<PagRow, _>(99, 20, &mut mock)
442 .await
443 .unwrap();
444
445 assert_eq!(result.total, 10);
446 assert_eq!(result.page, 99);
447 assert!(result.is_empty());
448 }
449
450 #[tokio::test]
460 async fn test_find_page_alias_of_paginate() {
461 let dialect = get_dialect(DbType::MySQL).unwrap();
462 let mut mock = MockConnection::new();
463
464 let _ = mock
465 .expect_any()
466 .with_rows(vec![vec![("total", Value::I64(50))]]);
467 let _ = mock.expect_any().with_rows(vec![
468 vec![("id", Value::I64(21))],
469 vec![("id", Value::I64(22))],
470 ]);
471
472 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
473 .where_eq("status", Value::from("active"))
474 .find_page::<PagRow, _>(3, 10, &mut mock)
475 .await
476 .unwrap();
477
478 assert_eq!(result.total, 50);
479 assert_eq!(result.page, 3);
480 assert_eq!(result.page_size, 10);
481 assert_eq!(result.total_pages(), 5);
482 assert!(result.has_next());
483 assert!(result.has_prev());
484 }
485
486 #[tokio::test]
488 async fn test_find_page_empty_result() {
489 let dialect = get_dialect(DbType::MySQL).unwrap();
490 let mut mock = MockConnection::new();
491
492 mock.expect_any()
493 .with_rows(vec![vec![("total", Value::I64(0))]]);
494
495 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
496 .find_page::<PagRow, _>(1, 20, &mut mock)
497 .await
498 .unwrap();
499
500 assert_eq!(result.total, 0);
501 assert!(result.is_empty());
502 }
503
504 #[tokio::test]
507 async fn test_paginator_total_zero_skips_data_query() {
508 let dialect = get_dialect(DbType::MySQL).unwrap();
509 let mut mock = MockConnection::new();
510
511 mock.expect_any()
513 .with_rows(vec![vec![("total", Value::I64(0))]]);
514 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
517 .paginate::<PagRow, _>(1, 20, &mut mock)
518 .await
519 .unwrap();
520
521 assert_eq!(result.total, 0);
522 assert!(result.is_empty());
523 assert_eq!(mock.executed_sql().len(), 1);
526 }
527
528 #[tokio::test]
531 async fn test_paginator_offset_equals_total_returns_empty() {
532 let dialect = get_dialect(DbType::MySQL).unwrap();
533 let mut mock = MockConnection::new();
534
535 mock.expect_any()
536 .with_rows(vec![vec![("total", Value::I64(10))]]);
537 mock.expect_any().with_rows(vec![]);
539
540 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
541 .paginate::<PagRow, _>(2, 10, &mut mock)
542 .await
543 .unwrap();
544
545 assert_eq!(result.total, 10);
546 assert_eq!(result.page, 2);
547 assert!(result.is_empty());
548 assert_eq!(mock.executed_sql().len(), 1);
551 }
552
553 #[tokio::test]
556 async fn test_paginator_page_size_zero_offset_zero() {
557 let dialect = get_dialect(DbType::MySQL).unwrap();
558 let mut mock = MockConnection::new();
559
560 mock.expect_any()
561 .with_rows(vec![vec![("total", Value::I64(100))]]);
562 mock.expect_any()
563 .with_rows(vec![vec![("id", Value::I64(1))]]);
564
565 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
566 .paginate::<PagRow, _>(3, 0, &mut mock)
567 .await
568 .unwrap();
569
570 assert_eq!(result.total, 100);
571 let data_sql = mock.executed_sql().get(1).unwrap();
573 assert!(
574 data_sql.contains("OFFSET 0"),
575 "expected OFFSET 0, got: {data_sql}"
576 );
577 }
578
579 #[tokio::test]
583 async fn test_paginator_offset_calculation_exact() {
584 let dialect = get_dialect(DbType::MySQL).unwrap();
585 let mut mock = MockConnection::new();
586
587 mock.expect_any()
588 .with_rows(vec![vec![("total", Value::I64(100))]]);
589 mock.expect_any().with_rows(vec![
590 vec![("id", Value::I64(11))],
591 vec![("id", Value::I64(12))],
592 ]);
593
594 let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
595 .paginate::<PagRow, _>(3, 5, &mut mock)
596 .await
597 .unwrap();
598
599 assert_eq!(result.total, 100);
600 assert_eq!(result.page, 3);
601 assert_eq!(result.page_size, 5);
602 assert_eq!(result.items.len(), 2);
603
604 let data_sql = mock.executed_sql().get(1).unwrap();
606 assert!(
607 data_sql.contains("OFFSET 10"),
608 "expected OFFSET 10 in SQL, got: {data_sql}"
609 );
610 assert!(data_sql.contains("LIMIT 5"));
611 }
612
613 #[tokio::test]
618 async fn test_paginator_builder_trait_exists() {
619 use super::PaginatorBuilderTrait;
620 fn _assert<M: Model, Q: PaginatorBuilderTrait<M>>() {}
621 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
622 }
623
624 #[tokio::test]
625 async fn test_paginator_builder_fetch_page() {
626 use super::PaginatorBuilderTrait;
627
628 let dialect = get_dialect(DbType::MySQL).unwrap();
629 let mut mock = MockConnection::new();
630
631 let _ = mock.expect_any().with_rows(vec![
633 vec![("id", Value::I64(1))],
634 vec![("id", Value::I64(2))],
635 vec![("id", Value::I64(3))],
636 ]);
637
638 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
639 .where_eq("status", Value::from("active"))
640 .paginate_with(&mut mock, 20);
641 p.set_total(100);
642
643 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(1).await.unwrap();
644
645 assert_eq!(result.total, 100);
646 assert_eq!(result.page, 1);
647 assert_eq!(result.page_size, 20);
648 assert_eq!(result.items.len(), 3);
649 }
650
651 #[tokio::test]
655 async fn test_fetch_page_offset_multiplication() {
656 use super::PaginatorBuilderTrait;
657
658 let dialect = get_dialect(DbType::MySQL).unwrap();
659 let mut mock = MockConnection::new();
660
661 mock.expect_any().with_rows(vec![
662 vec![("id", Value::I64(11))],
663 vec![("id", Value::I64(12))],
664 ]);
665
666 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
667 .where_eq("status", Value::from("active"))
668 .paginate_with(&mut mock, 5);
669 p.set_total(100);
670
671 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
672
673 assert_eq!(result.items.len(), 2);
674 let data_sql = mock.executed_sql().first().unwrap();
676 assert!(
677 data_sql.contains("OFFSET 10"),
678 "expected OFFSET 10 in fetch_page SQL, got: {data_sql}"
679 );
680 }
681
682 #[tokio::test]
685 async fn test_fetch_page_page_size_zero_offset_zero() {
686 use super::PaginatorBuilderTrait;
687
688 let dialect = get_dialect(DbType::MySQL).unwrap();
689 let mut mock = MockConnection::new();
690
691 mock.expect_any()
692 .with_rows(vec![vec![("id", Value::I64(1))]]);
693
694 let mut p = QueryBuilder::<PagTestModel>::new(dialect)
695 .where_eq("status", Value::from("active"))
696 .paginate_with(&mut mock, 0);
697 p.set_total(100);
698
699 let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
700
701 assert_eq!(result.items.len(), 1);
702 let data_sql = mock.executed_sql().first().unwrap();
705 assert!(
706 data_sql.contains("OFFSET 0"),
707 "expected OFFSET 0, got: {data_sql}"
708 );
709 }
710
711 #[tokio::test]
716 async fn test_stream_trait_exists() {
717 use super::StreamQueryTrait;
718 fn _assert<M: Model, Q: StreamQueryTrait<M>>() {}
719 _assert::<PagTestModel, QueryBuilder<PagTestModel>>();
720 }
721
722 #[tokio::test]
723 async fn test_stream_returns_rows() {
724 use super::StreamQueryTrait;
725 use futures::StreamExt;
726
727 let dialect = get_dialect(DbType::MySQL).unwrap();
728 let mut mock = MockConnection::new();
729
730 let _ = mock.expect_any().with_rows(vec![
731 vec![("id", Value::I64(1))],
732 vec![("id", Value::I64(2))],
733 ]);
734
735 let mut stream = QueryBuilder::<PagTestModel>::new(dialect)
736 .where_eq("status", Value::from("active"))
737 .stream(&mut mock);
738
739 let mut count = 0;
740 while let Some(result) = stream.next().await {
741 let row: RowResult = result.unwrap();
742 assert!(row.contains_key("id"));
743 count += 1;
744 }
745 assert_eq!(count, 2);
746 }
747
748 #[test]
754 fn test_value_to_u64_i32() {
755 assert_eq!(value_to_u64(&Value::I32(42)), Some(42u64));
756 assert_eq!(value_to_u64(&Value::I32(-1)), Some(u64::MAX)); }
758
759 #[test]
761 fn test_value_to_u64_u64() {
762 assert_eq!(value_to_u64(&Value::U64(999)), Some(999u64));
763 }
764
765 #[test]
767 fn test_value_to_u64_u32() {
768 assert_eq!(value_to_u64(&Value::U32(77)), Some(77u64));
769 }
770
771 #[test]
773 fn test_value_to_u64_f64() {
774 assert_eq!(value_to_u64(&Value::F64(2.71)), Some(2u64));
775 assert_eq!(value_to_u64(&Value::F64(0.0)), Some(0u64));
776 }
777
778 #[test]
780 fn test_value_to_u64_i64() {
781 assert_eq!(value_to_u64(&Value::I64(123456789)), Some(123456789u64));
782 }
783
784 #[test]
786 fn test_value_to_u64_unknown() {
787 assert_eq!(value_to_u64(&Value::Bool(true)), None);
788 assert_eq!(value_to_u64(&Value::String("x".into())), None);
789 }
790}