sz_orm_core/
stream_api.rs1use crate::model::Model;
35use crate::pool::{Connection, QueryStreamItem};
36use crate::query::QueryBuilder;
37use crate::DbError;
38
39use std::collections::HashMap;
40use std::pin::Pin;
41
42use futures::{stream, Stream, StreamExt};
43
44pub type RowResult = HashMap<String, crate::value::Value>;
46
47pub trait StreamApiExt<M: Model> {
51 fn stream_buffered<'a, 'b: 'a, C: Connection + Send + 'b>(
56 self,
57 conn: &'b mut C,
58 ) -> Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>;
59}
60
61impl<M: Model> StreamApiExt<M> for QueryBuilder<M> {
62 fn stream_buffered<'a, 'b: 'a, C: Connection + Send + 'b>(
63 self,
64 conn: &'b mut C,
65 ) -> Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>> {
66 let (sql, params) = self.build_select_with_params();
67
68 let st = stream::once(async move {
69 match conn.query_with_params(&sql, ¶ms).await {
70 Ok(rows) => rows.into_iter().map(Ok).collect::<Vec<_>>(),
71 Err(e) => vec![Err(e)],
72 }
73 })
74 .flat_map(stream::iter);
75
76 Box::pin(st)
77 }
78}
79
80pub fn stream_cursor<'a, 'b: 'a, C: Connection + Send + 'b>(
90 conn: &'b mut C,
91 sql: &'b str,
92 params: Vec<crate::value::Value>,
93 batch_size: usize,
94) -> Pin<Box<dyn Stream<Item = QueryStreamItem> + Send + 'a>> {
95 let _ = params; conn.query_stream_cursor(sql, batch_size)
97}
98
99#[cfg(test)]
104mod tests {
105 use super::*;
106 use crate::db_type::DbType;
107 use crate::dialect::get_dialect;
108 use crate::mock::MockConnection;
109 use crate::model::Model;
110 use crate::query::QueryBuilder;
111 use crate::value::Value;
112 use futures::StreamExt;
113
114 #[derive(Debug, Clone, Default)]
115 #[allow(dead_code)]
116 struct TestUser {
117 id: i64,
118 name: String,
119 }
120
121 impl Model for TestUser {
122 type PrimaryKey = i64;
123 fn table_name() -> &'static str {
124 "users"
125 }
126 fn pk(&self) -> Self::PrimaryKey {
127 self.id
128 }
129 fn set_pk(&mut self, pk: Self::PrimaryKey) {
130 self.id = pk;
131 }
132 }
133
134 #[tokio::test]
135 async fn test_stream_buffered_yields_rows() {
136 let mut mock = MockConnection::new();
137 mock.expect_any().with_rows(vec![
138 vec![("id", Value::I64(1))],
139 vec![("id", Value::I64(2))],
140 ]);
141
142 let dialect = get_dialect(DbType::MySQL).unwrap();
143 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
144
145 let mut stream = query.stream_buffered(&mut mock);
146 let mut rows = Vec::new();
147 while let Some(row) = stream.next().await {
148 rows.push(row.unwrap());
149 }
150
151 assert_eq!(rows.len(), 2);
152 }
153
154 #[tokio::test]
155 async fn test_stream_buffered_error_propagation() {
156 let mut mock = MockConnection::new().with_fallback(crate::mock::FallbackBehavior::Error);
157
158 let dialect = get_dialect(DbType::MySQL).unwrap();
159 let query = QueryBuilder::<TestUser>::new(dialect).table("nonexistent");
160
161 let mut stream = query.stream_buffered(&mut mock);
162 let result = stream.next().await;
163
164 assert!(result.is_some());
165 assert!(result.unwrap().is_err());
166 }
167
168 #[tokio::test]
169 async fn test_stream_buffered_empty() {
170 let mut mock = MockConnection::new();
171 mock.expect_any().with_rows(vec![]);
172
173 let dialect = get_dialect(DbType::MySQL).unwrap();
174 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
175
176 let mut stream = query.stream_buffered(&mut mock);
177 let mut count = 0;
178 while stream.next().await.is_some() {
179 count += 1;
180 }
181
182 assert_eq!(count, 0);
183 }
184}