use crate::model::Model;
use crate::pool::{Connection, QueryStreamItem};
use crate::query::QueryBuilder;
use crate::DbError;
use std::collections::HashMap;
use std::pin::Pin;
use futures::{sink::SinkExt, stream, Stream, StreamExt};
pub type RowResult = HashMap<String, crate::value::Value>;
pub trait StreamApiExt<M: Model> {
fn stream_buffered<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
) -> Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>;
fn stream_with_backpressure<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
buffer_size: usize,
) -> Result<Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>, DbError>;
}
impl<M: Model> StreamApiExt<M> for QueryBuilder<M> {
fn stream_buffered<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
) -> Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>> {
let (sql, params) = self.build_select_with_params();
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 stream_with_backpressure<'a, 'b: 'a, C: Connection + Send + 'b>(
self,
conn: &'b mut C,
buffer_size: usize,
) -> Result<Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>, DbError> {
if buffer_size == 0 {
return Err(DbError::InvalidInput("buffer_size 必须大于 0".to_string()));
}
let (sql, params) = self.build_select_with_params();
let (mut tx, rx) =
futures::channel::mpsc::channel::<Result<RowResult, DbError>>(buffer_size);
let producer = async move {
match conn.query_with_params(&sql, ¶ms).await {
Ok(rows) => {
for row in rows {
if tx.send(Ok(row)).await.is_err() {
break;
}
}
}
Err(e) => {
let _ = tx.send(Err(e)).await;
}
}
};
let receiver_stream = BackpressureStream {
rx,
producer: Some(Box::pin(producer)),
};
Ok(Box::pin(receiver_stream))
}
}
struct BackpressureStream<'a> {
rx: futures::channel::mpsc::Receiver<Result<RowResult, DbError>>,
producer: Option<Pin<Box<dyn std::future::Future<Output = ()> + Send + 'a>>>,
}
impl<'a> Stream for BackpressureStream<'a> {
type Item = Result<RowResult, DbError>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = self.get_mut();
if let Some(mut producer) = this.producer.take() {
match std::future::Future::poll(std::pin::Pin::new(&mut producer), cx) {
std::task::Poll::Ready(()) => {}
std::task::Poll::Pending => {
this.producer = Some(producer);
}
}
}
Stream::poll_next(std::pin::Pin::new(&mut this.rx), cx)
}
}
pub fn stream_cursor<'a, 'b: 'a, C: Connection + Send + 'b>(
conn: &'b mut C,
sql: &'b str,
params: Vec<crate::value::Value>,
batch_size: usize,
) -> Pin<Box<dyn Stream<Item = QueryStreamItem> + Send + 'a>> {
let _ = params; conn.query_stream_cursor(sql, batch_size)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db_type::DbType;
use crate::dialect::get_dialect;
use crate::mock::MockConnection;
use crate::model::Model;
use crate::query::QueryBuilder;
use crate::value::Value;
use futures::StreamExt;
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
struct TestUser {
id: i64,
name: String,
}
impl Model for TestUser {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"users"
}
fn pk(&self) -> Self::PrimaryKey {
self.id
}
fn set_pk(&mut self, pk: Self::PrimaryKey) {
self.id = pk;
}
}
#[tokio::test]
async fn test_stream_buffered_yields_rows() {
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![
vec![("id", Value::I64(1))],
vec![("id", Value::I64(2))],
]);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("users");
let mut stream = query.stream_buffered(&mut mock);
let mut rows = Vec::new();
while let Some(row) = stream.next().await {
rows.push(row.unwrap());
}
assert_eq!(rows.len(), 2);
}
#[tokio::test]
async fn test_stream_buffered_error_propagation() {
let mut mock = MockConnection::new().with_fallback(crate::mock::FallbackBehavior::Error);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("nonexistent");
let mut stream = query.stream_buffered(&mut mock);
let result = stream.next().await;
assert!(result.is_some());
assert!(result.unwrap().is_err());
}
#[tokio::test]
async fn test_stream_buffered_empty() {
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![]);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("users");
let mut stream = query.stream_buffered(&mut mock);
let mut count = 0;
while stream.next().await.is_some() {
count += 1;
}
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_stream_backpressure_zero_buffer() {
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![]);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("users");
let result = query.stream_with_backpressure(&mut mock, 0);
assert!(result.is_err());
let err = result.err().unwrap();
assert!(matches!(err, crate::DbError::InvalidInput(_)));
}
#[tokio::test]
async fn test_stream_backpressure_basic() {
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![
vec![("id", Value::I64(1))],
vec![("id", Value::I64(2))],
vec![("id", Value::I64(3))],
]);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("users");
let stream = query.stream_with_backpressure(&mut mock, 10).unwrap();
let mut stream = Box::pin(stream);
let mut rows = Vec::new();
while let Some(row) = stream.next().await {
rows.push(row.unwrap());
}
assert_eq!(rows.len(), 3);
}
#[tokio::test]
async fn test_stream_backpressure_empty() {
let mut mock = MockConnection::new();
mock.expect_any().with_rows(vec![]);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("users");
let stream = query.stream_with_backpressure(&mut mock, 100).unwrap();
let mut stream = Box::pin(stream);
let mut count = 0;
while stream.next().await.is_some() {
count += 1;
}
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_stream_backpressure_error_propagation() {
let mut mock = MockConnection::new().with_fallback(crate::mock::FallbackBehavior::Error);
let dialect = get_dialect(DbType::MySQL).unwrap();
let query = QueryBuilder::<TestUser>::new(dialect).table("nonexistent");
let stream = query.stream_with_backpressure(&mut mock, 10).unwrap();
let mut stream = Box::pin(stream);
let result = stream.next().await;
assert!(result.is_some());
assert!(result.unwrap().is_err());
}
}