use crate::db_type::DbType;
use crate::pool::{Connection, QueryStreamItem};
use crate::DbError;
use std::pin::Pin;
pub fn build_paged_query(
db_type: DbType,
sql: &str,
offset: u64,
batch: u64,
) -> Result<String, DbError> {
let sql = sql.trim().trim_end_matches(';').trim();
if sql.is_empty() {
return Err(DbError::InvalidInput("build_paged_query: SQL 为空".into()));
}
if batch == 0 {
return Err(DbError::InvalidInput(
"build_paged_query: batch 必须大于 0".into(),
));
}
match db_type {
DbType::Oracle => {
let end = offset + batch;
Ok(format!(
"SELECT * FROM (SELECT t.*, ROWNUM AS __sz_rn FROM ({}) t \
WHERE ROWNUM <= {}) WHERE __sz_rn > {}",
sql, end, offset
))
}
DbType::SqlServer => Ok(format!(
"SELECT * FROM ({}) AS __sz_t ORDER BY (SELECT NULL) \
OFFSET {} ROWS FETCH NEXT {} ROWS ONLY",
sql, offset, batch
)),
DbType::MySQL | DbType::PostgreSQL | DbType::Sqlite | DbType::OceanBase => {
Ok(format!("{} LIMIT {} OFFSET {}", sql, batch, offset))
}
_ => Err(DbError::Unsupported(format!(
"build_paged_query: {:?} 方言不支持分页游标",
db_type
))),
}
}
pub fn stream_cursor_paged<'a>(
conn: &'a mut dyn Connection,
sql: &'a str,
db_type: DbType,
batch: u64,
) -> Pin<Box<dyn futures::Stream<Item = QueryStreamItem> + Send + 'a>> {
let stream = futures::stream::unfold(
(conn, 0u64, Vec::<QueryStreamItem>::new(), false),
move |(conn, mut offset, mut page, done)| async move {
if done {
return None;
}
loop {
if let Some(item) = page.pop() {
return Some((item, (conn, offset, page, false)));
}
let paged = match build_paged_query(db_type, sql, offset, batch) {
Ok(s) => s,
Err(e) => return Some((Err(e), (conn, offset, Vec::new(), true))),
};
match conn.query(&paged).await {
Ok(rows) if rows.is_empty() => return None,
Ok(rows) => {
offset += rows.len() as u64;
page = rows.into_iter().map(Ok).rev().collect();
}
Err(e) => return Some((Err(e), (conn, offset, Vec::new(), true))),
}
}
},
);
Box::pin(stream)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pool::QueryRows;
use crate::value::Value;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
#[test]
fn test_build_paged_query_oracle() {
let sql = build_paged_query(
DbType::Oracle,
"SELECT id, name FROM users ORDER BY id",
10,
100,
)
.unwrap();
assert!(sql.contains("SELECT * FROM (SELECT t.*, ROWNUM AS __sz_rn FROM (SELECT id, name FROM users ORDER BY id) t WHERE ROWNUM <= 110) WHERE __sz_rn > 10"), "Oracle 包装 SQL 不正确: {}", sql);
}
#[test]
fn test_build_paged_query_oracle_strips_semicolon() {
let sql = build_paged_query(DbType::Oracle, "SELECT * FROM t;", 0, 50).unwrap();
assert!(sql.contains("FROM (SELECT * FROM t) t"));
}
#[test]
fn test_build_paged_query_sqlserver() {
let sql = build_paged_query(DbType::SqlServer, "SELECT id FROM orders", 20, 50).unwrap();
assert!(
sql.contains("OFFSET 20 ROWS FETCH NEXT 50 ROWS ONLY"),
"MSSQL 包装 SQL 不正确: {}",
sql
);
assert!(sql.contains("ORDER BY (SELECT NULL)"));
}
#[test]
fn test_build_paged_query_mysql_postgres_sqlite() {
for db in [DbType::MySQL, DbType::PostgreSQL, DbType::Sqlite] {
let sql = build_paged_query(db, "SELECT id FROM t WHERE id > 0", 5, 10).unwrap();
assert_eq!(sql, "SELECT id FROM t WHERE id > 0 LIMIT 10 OFFSET 5");
}
}
#[test]
fn test_build_paged_query_unsupported_dialect() {
let r = build_paged_query(DbType::Redis, "SELECT 1", 0, 10);
assert!(r.is_err(), "Redis 方言不应支持分页游标");
let r2 = build_paged_query(DbType::ClickHouse, "SELECT 1", 0, 10);
assert!(r2.is_err());
}
#[test]
fn test_build_paged_query_invalid_args() {
assert!(
build_paged_query(DbType::MySQL, "", 0, 10).is_err(),
"空 SQL 应报错"
);
assert!(
build_paged_query(DbType::MySQL, "SELECT 1", 0, 0).is_err(),
"batch=0 应报错"
);
assert!(
build_paged_query(DbType::MySQL, ";", 0, 10).is_err(),
"仅分号应报错"
);
}
struct PagedStubConn {
rows: Vec<HashMap<String, Value>>,
query_calls: Arc<AtomicUsize>,
}
impl PagedStubConn {
fn new(rows: Vec<HashMap<String, Value>>) -> (Self, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
(
Self {
rows,
query_calls: Arc::clone(&calls),
},
calls,
)
}
fn make_row(id: i64) -> HashMap<String, Value> {
let mut m = HashMap::new();
m.insert("id".to_string(), Value::I64(id));
m
}
}
impl Connection for PagedStubConn {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(0) })
}
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>> {
let limit = sql
.split("LIMIT ")
.nth(1)
.and_then(|s| s.split(' ').next())
.and_then(|s| s.parse::<usize>().ok())
.unwrap_or(0);
let offset = sql
.split("OFFSET ")
.nth(1)
.and_then(|s| s.trim().parse::<usize>().ok())
.unwrap_or(0);
let snapshot: Vec<_> = self.rows.iter().skip(offset).take(limit).cloned().collect();
self.query_calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async move { Ok(snapshot) })
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
}
#[tokio::test]
async fn test_stream_cursor_paged_fetches_in_batches() {
let rows: Vec<_> = (1..=5).map(PagedStubConn::make_row).collect();
let (mut conn, calls) = PagedStubConn::new(rows);
let mut stream = stream_cursor_paged(&mut conn, "SELECT id FROM t", DbType::MySQL, 2);
use futures::StreamExt;
let ids: Vec<i64> = {
let mut out = Vec::new();
while let Some(item) = stream.next().await {
let row = item.unwrap();
if let Value::I64(id) = row["id"] {
out.push(id);
}
}
out
};
assert_eq!(ids, vec![1, 2, 3, 4, 5], "行序应保持");
assert_eq!(calls.load(Ordering::SeqCst), 4);
}
#[tokio::test]
async fn test_stream_cursor_paged_empty() {
let (mut conn, calls) = PagedStubConn::new(Vec::new());
let mut stream = stream_cursor_paged(&mut conn, "SELECT id FROM t", DbType::MySQL, 2);
use futures::StreamExt;
let mut count = 0;
while let Some(item) = stream.next().await {
let _ = item.unwrap();
count += 1;
}
assert_eq!(count, 0, "空结果不应 yield 任何行");
assert_eq!(calls.load(Ordering::SeqCst), 1, "空结果应只查一次");
}
#[tokio::test]
async fn test_stream_cursor_paged_batch_larger_than_rows() {
let rows: Vec<_> = (1..=3).map(PagedStubConn::make_row).collect();
let (mut conn, calls) = PagedStubConn::new(rows);
let mut stream = stream_cursor_paged(&mut conn, "SELECT id FROM t", DbType::MySQL, 10);
use futures::StreamExt;
let mut count = 0;
while let Some(item) = stream.next().await {
let _ = item.unwrap();
count += 1;
}
assert_eq!(count, 3);
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"3 行 batch=10 → 首页 3 行 + 尾页空"
);
}
}