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::{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>>;
}
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)
}
}
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);
}
}