use crate::model::Model;
use crate::pool::Connection;
use crate::query::QueryBuilder;
use crate::value::{FromQueryResult, Value};
use crate::DbError;
use std::future::Future;
use std::pin::Pin;
pub use crate::repository::PageResult;
pub trait PaginatorTrait<M: Model>: Sized {
fn paginate<T, C>(
self,
page: u64,
page_size: u64,
conn: &mut C,
) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
where
T: FromQueryResult + Send + 'static,
C: Connection + Send;
fn find_page<T, C>(
self,
page: u64,
page_size: u64,
conn: &mut C,
) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
where
T: FromQueryResult + Send + 'static,
C: Connection + Send,
{
self.paginate::<T, C>(page, page_size, conn)
}
}
impl<M: Model> PaginatorTrait<M> for QueryBuilder<M> {
fn paginate<T, C>(
self,
page: u64,
page_size: u64,
conn: &mut C,
) -> Pin<Box<dyn Future<Output = Result<PageResult<T>, DbError>> + Send + '_>>
where
T: FromQueryResult + Send + 'static,
C: Connection + Send,
{
let count_sql = self.build_count();
let offset = if page_size == 0 {
0
} else {
((page.saturating_sub(1)) * page_size) as usize
};
let limit = page_size as usize;
let (select_sql, params) = self.limit(limit).offset(offset).build_select_with_params();
Box::pin(async move {
let count_rows = conn.query(&count_sql).await?;
let total = extract_count(&count_rows).unwrap_or(0);
let data_rows = if total == 0 || offset as u64 >= total {
Vec::new()
} else {
let rows = conn.query_with_params(&select_sql, ¶ms).await?;
let mut items = Vec::with_capacity(rows.len());
for row in &rows {
match T::from_query_result(row) {
Ok(item) => items.push(item),
Err(e) => return Err(DbError::Internal(e)),
}
}
items
};
Ok(PageResult::new(data_rows, total, page, page_size))
})
}
}
pub struct Paginator<'a, C>
where
C: Connection,
{
conn: &'a mut C,
page_size: u64,
sql: String,
params: Vec<Value>,
total: u64,
}
impl<'a, C> Paginator<'a, C>
where
C: Connection,
{
pub async fn fetch_page<T>(&mut self, page: u64) -> Result<PageResult<T>, DbError>
where
T: FromQueryResult + Send + 'static,
{
let offset = if self.page_size == 0 {
0
} else {
((page.saturating_sub(1)) * self.page_size) as usize
};
let limit = self.page_size as usize;
let select_sql = format!("{} LIMIT {} OFFSET {}", self.sql, limit, offset);
let rows = self
.conn
.query_with_params(&select_sql, &self.params)
.await?;
let mut items = Vec::with_capacity(rows.len());
for row in &rows {
match T::from_query_result(row) {
Ok(item) => items.push(item),
Err(e) => return Err(DbError::Internal(e)),
}
}
Ok(PageResult::new(items, self.total, page, self.page_size))
}
pub fn set_total(&mut self, total: u64) {
self.total = total;
}
}
pub trait PaginatorBuilderTrait<M: Model> {
fn paginate_with<'a, C>(self, conn: &'a mut C, page_size: u64) -> Paginator<'a, C>
where
C: Connection;
}
impl<M: Model> PaginatorBuilderTrait<M> for QueryBuilder<M> {
fn paginate_with<'a, C>(self, conn: &'a mut C, page_size: u64) -> Paginator<'a, C>
where
C: Connection,
{
let (sql, params) = self.build_select_with_params();
Paginator {
conn,
page_size,
sql,
params,
total: 0,
}
}
}
pub type RowResult = std::collections::HashMap<String, Value>;
pub trait StreamQueryTrait<M: Model> {
fn stream<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
) -> Pin<Box<dyn futures::Stream<Item = Result<RowResult, DbError>> + Send + 'a>>;
}
impl<M: Model> StreamQueryTrait<M> for QueryBuilder<M> {
fn stream<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
) -> Pin<Box<dyn futures::Stream<Item = Result<RowResult, DbError>> + Send + 'a>> {
let (sql, params) = self.build_select_with_params();
use futures::{stream, StreamExt};
let st = stream::once(async move {
match conn.query_with_params(&sql, ¶ms).await {
Ok(rows) => rows.into_iter().map(Ok).collect::<Vec<_>>(),
Err(e) => vec![Err(e)],
}
})
.flat_map(stream::iter);
Box::pin(st)
}
}
fn extract_count(rows: &[std::collections::HashMap<String, Value>]) -> Option<u64> {
rows.first().and_then(|row| {
if let Some(v) = row.get("total") {
return value_to_u64(v);
}
if let Some(v) = row.get("COUNT(*)") {
return value_to_u64(v);
}
row.values().next().and_then(value_to_u64)
})
}
fn value_to_u64(v: &Value) -> Option<u64> {
match v {
Value::I64(n) => Some(*n as u64),
Value::I32(n) => Some(*n as u64),
Value::U64(n) => Some(*n),
Value::U32(n) => Some(*n as u64),
Value::F64(n) => Some(*n as u64),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db_type::DbType;
use crate::dialect::get_dialect;
use crate::mock::MockConnection;
#[derive(Debug, Clone, Default)]
struct PagTestModel;
impl Model for PagTestModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"pag_test"
}
fn pk_name() -> &'static str {
"id"
}
fn pk(&self) -> Self::PrimaryKey {
0
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
fn timestamp_fields() -> Option<crate::model::TimestampFields> {
None
}
fn soft_delete_field() -> Option<&'static str> {
None
}
}
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
struct PagRow {
id: i64,
}
impl FromQueryResult for PagRow {
fn from_query_result(
row: &std::collections::HashMap<String, Value>,
) -> Result<Self, String> {
let id = row.get("id").and_then(|v| v.as_i64()).unwrap_or(0);
Ok(PagRow { id })
}
}
#[tokio::test]
async fn test_paginator_trait_exists() {
fn _assert<M: Model, Q: PaginatorTrait<M>>() {}
_assert::<PagTestModel, QueryBuilder<PagTestModel>>();
}
#[tokio::test]
async fn test_paginator_with_mock() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
let _ = mock
.expect_any()
.with_rows(vec![vec![("total", Value::I64(100))]]);
let _ = mock.expect_any().with_rows(vec![
vec![("id", Value::I64(1))],
vec![("id", Value::I64(2))],
]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.paginate::<PagRow, _>(1, 20, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 100);
assert_eq!(result.page, 1);
assert_eq!(result.page_size, 20);
assert_eq!(result.total_pages(), 5);
assert!(result.has_next());
assert!(!result.has_prev());
}
#[tokio::test]
async fn test_paginator_empty_result() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(0))]]);
mock.expect_any().with_rows(vec![]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(1, 20, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 0);
assert!(result.is_empty());
assert_eq!(result.total_pages(), 0);
}
#[tokio::test]
async fn test_paginator_page_beyond_range() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(10))]]);
mock.expect_any().with_rows(vec![]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(99, 20, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 10);
assert_eq!(result.page, 99);
assert!(result.is_empty());
}
#[tokio::test]
async fn test_find_page_alias_of_paginate() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
let _ = mock
.expect_any()
.with_rows(vec![vec![("total", Value::I64(50))]]);
let _ = mock.expect_any().with_rows(vec![
vec![("id", Value::I64(21))],
vec![("id", Value::I64(22))],
]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.find_page::<PagRow, _>(3, 10, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 50);
assert_eq!(result.page, 3);
assert_eq!(result.page_size, 10);
assert_eq!(result.total_pages(), 5);
assert!(result.has_next());
assert!(result.has_prev());
}
#[tokio::test]
async fn test_find_page_empty_result() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(0))]]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.find_page::<PagRow, _>(1, 20, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 0);
assert!(result.is_empty());
}
#[tokio::test]
async fn test_paginator_total_zero_skips_data_query() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(0))]]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(1, 20, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 0);
assert!(result.is_empty());
assert_eq!(mock.executed_sql().len(), 1);
}
#[tokio::test]
async fn test_paginator_offset_equals_total_returns_empty() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(10))]]);
mock.expect_any().with_rows(vec![]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(2, 10, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 10);
assert_eq!(result.page, 2);
assert!(result.is_empty());
assert_eq!(mock.executed_sql().len(), 1);
}
#[tokio::test]
async fn test_paginator_page_size_zero_offset_zero() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(100))]]);
mock.expect_any()
.with_rows(vec![vec![("id", Value::I64(1))]]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(3, 0, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 100);
let data_sql = mock.executed_sql().get(1).unwrap();
assert!(
data_sql.contains("OFFSET 0"),
"expected OFFSET 0, got: {data_sql}"
);
}
#[tokio::test]
async fn test_paginator_offset_calculation_exact() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("total", Value::I64(100))]]);
mock.expect_any().with_rows(vec![
vec![("id", Value::I64(11))],
vec![("id", Value::I64(12))],
]);
let result: PageResult<PagRow> = QueryBuilder::<PagTestModel>::new(dialect)
.paginate::<PagRow, _>(3, 5, &mut mock)
.await
.unwrap();
assert_eq!(result.total, 100);
assert_eq!(result.page, 3);
assert_eq!(result.page_size, 5);
assert_eq!(result.items.len(), 2);
let data_sql = mock.executed_sql().get(1).unwrap();
assert!(
data_sql.contains("OFFSET 10"),
"expected OFFSET 10 in SQL, got: {data_sql}"
);
assert!(data_sql.contains("LIMIT 5"));
}
#[tokio::test]
async fn test_paginator_builder_trait_exists() {
use super::PaginatorBuilderTrait;
fn _assert<M: Model, Q: PaginatorBuilderTrait<M>>() {}
_assert::<PagTestModel, QueryBuilder<PagTestModel>>();
}
#[tokio::test]
async fn test_paginator_builder_fetch_page() {
use super::PaginatorBuilderTrait;
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
let _ = mock.expect_any().with_rows(vec![
vec![("id", Value::I64(1))],
vec![("id", Value::I64(2))],
vec![("id", Value::I64(3))],
]);
let mut p = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.paginate_with(&mut mock, 20);
p.set_total(100);
let result: PageResult<PagRow> = p.fetch_page::<PagRow>(1).await.unwrap();
assert_eq!(result.total, 100);
assert_eq!(result.page, 1);
assert_eq!(result.page_size, 20);
assert_eq!(result.items.len(), 3);
}
#[tokio::test]
async fn test_fetch_page_offset_multiplication() {
use super::PaginatorBuilderTrait;
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![
vec![("id", Value::I64(11))],
vec![("id", Value::I64(12))],
]);
let mut p = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.paginate_with(&mut mock, 5);
p.set_total(100);
let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
assert_eq!(result.items.len(), 2);
let data_sql = mock.executed_sql().first().unwrap();
assert!(
data_sql.contains("OFFSET 10"),
"expected OFFSET 10 in fetch_page SQL, got: {data_sql}"
);
}
#[tokio::test]
async fn test_fetch_page_page_size_zero_offset_zero() {
use super::PaginatorBuilderTrait;
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
mock.expect_any()
.with_rows(vec![vec![("id", Value::I64(1))]]);
let mut p = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.paginate_with(&mut mock, 0);
p.set_total(100);
let result: PageResult<PagRow> = p.fetch_page::<PagRow>(3).await.unwrap();
assert_eq!(result.items.len(), 1);
let data_sql = mock.executed_sql().first().unwrap();
assert!(
data_sql.contains("OFFSET 0"),
"expected OFFSET 0, got: {data_sql}"
);
}
#[tokio::test]
async fn test_stream_trait_exists() {
use super::StreamQueryTrait;
fn _assert<M: Model, Q: StreamQueryTrait<M>>() {}
_assert::<PagTestModel, QueryBuilder<PagTestModel>>();
}
#[tokio::test]
async fn test_stream_returns_rows() {
use super::StreamQueryTrait;
use futures::StreamExt;
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut mock = MockConnection::new();
let _ = mock.expect_any().with_rows(vec![
vec![("id", Value::I64(1))],
vec![("id", Value::I64(2))],
]);
let mut stream = QueryBuilder::<PagTestModel>::new(dialect)
.where_eq("status", Value::from("active"))
.stream(&mut mock);
let mut count = 0;
while let Some(result) = stream.next().await {
let row: RowResult = result.unwrap();
assert!(row.contains_key("id"));
count += 1;
}
assert_eq!(count, 2);
}
#[test]
fn test_value_to_u64_i32() {
assert_eq!(value_to_u64(&Value::I32(42)), Some(42u64));
assert_eq!(value_to_u64(&Value::I32(-1)), Some(u64::MAX)); }
#[test]
fn test_value_to_u64_u64() {
assert_eq!(value_to_u64(&Value::U64(999)), Some(999u64));
}
#[test]
fn test_value_to_u64_u32() {
assert_eq!(value_to_u64(&Value::U32(77)), Some(77u64));
}
#[test]
fn test_value_to_u64_f64() {
assert_eq!(value_to_u64(&Value::F64(2.71)), Some(2u64));
assert_eq!(value_to_u64(&Value::F64(0.0)), Some(0u64));
}
#[test]
fn test_value_to_u64_i64() {
assert_eq!(value_to_u64(&Value::I64(123456789)), Some(123456789u64));
}
#[test]
fn test_value_to_u64_unknown() {
assert_eq!(value_to_u64(&Value::Bool(true)), None);
assert_eq!(value_to_u64(&Value::String("x".into())), None);
}
}