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::{sink::SinkExt, 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 fn stream_with_backpressure<'a, 'b: 'a, C: Connection + Send + 'b>(
78 self,
79 conn: &'b mut C,
80 buffer_size: usize,
81 ) -> Result<Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>, DbError>;
82}
83
84impl<M: Model> StreamApiExt<M> for QueryBuilder<M> {
85 fn stream_buffered<'a, 'b: 'a, C: Connection + Send + 'b>(
86 self,
87 conn: &'b mut C,
88 ) -> Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>> {
89 let (sql, params) = self.build_select_with_params();
90
91 let st = stream::once(async move {
92 match conn.query_with_params(&sql, ¶ms).await {
93 Ok(rows) => rows.into_iter().map(Ok).collect::<Vec<_>>(),
94 Err(e) => vec![Err(e)],
95 }
96 })
97 .flat_map(stream::iter);
98
99 Box::pin(st)
100 }
101
102 fn stream_with_backpressure<'a, 'b: 'a, C: Connection + Send + 'b>(
103 self,
104 conn: &'b mut C,
105 buffer_size: usize,
106 ) -> Result<Pin<Box<dyn Stream<Item = Result<RowResult, DbError>> + Send + 'a>>, DbError> {
107 if buffer_size == 0 {
108 return Err(DbError::InvalidInput("buffer_size 必须大于 0".to_string()));
109 }
110
111 let (sql, params) = self.build_select_with_params();
112 let (mut tx, rx) =
113 futures::channel::mpsc::channel::<Result<RowResult, DbError>>(buffer_size);
114
115 let producer = async move {
116 match conn.query_with_params(&sql, ¶ms).await {
117 Ok(rows) => {
118 for row in rows {
119 if tx.send(Ok(row)).await.is_err() {
120 break;
121 }
122 }
123 }
124 Err(e) => {
125 let _ = tx.send(Err(e)).await;
126 }
127 }
128 };
129
130 let receiver_stream = BackpressureStream {
131 rx,
132 producer: Some(Box::pin(producer)),
133 };
134
135 Ok(Box::pin(receiver_stream))
136 }
137}
138
139struct BackpressureStream<'a> {
141 rx: futures::channel::mpsc::Receiver<Result<RowResult, DbError>>,
142 producer: Option<Pin<Box<dyn std::future::Future<Output = ()> + Send + 'a>>>,
143}
144
145impl<'a> Stream for BackpressureStream<'a> {
146 type Item = Result<RowResult, DbError>;
147
148 fn poll_next(
149 self: std::pin::Pin<&mut Self>,
150 cx: &mut std::task::Context<'_>,
151 ) -> std::task::Poll<Option<Self::Item>> {
152 let this = self.get_mut();
153
154 if let Some(mut producer) = this.producer.take() {
155 match std::future::Future::poll(std::pin::Pin::new(&mut producer), cx) {
156 std::task::Poll::Ready(()) => {}
157 std::task::Poll::Pending => {
158 this.producer = Some(producer);
159 }
160 }
161 }
162
163 Stream::poll_next(std::pin::Pin::new(&mut this.rx), cx)
164 }
165}
166
167pub fn stream_cursor<'a, 'b: 'a, C: Connection + Send + 'b>(
177 conn: &'b mut C,
178 sql: &'b str,
179 params: Vec<crate::value::Value>,
180 batch_size: usize,
181) -> Pin<Box<dyn Stream<Item = QueryStreamItem> + Send + 'a>> {
182 let _ = params; conn.query_stream_cursor(sql, batch_size)
184}
185
186#[cfg(test)]
191mod tests {
192 use super::*;
193 use crate::db_type::DbType;
194 use crate::dialect::get_dialect;
195 use crate::mock::MockConnection;
196 use crate::model::Model;
197 use crate::query::QueryBuilder;
198 use crate::value::Value;
199 use futures::StreamExt;
200
201 #[derive(Debug, Clone, Default)]
202 #[allow(dead_code)]
203 struct TestUser {
204 id: i64,
205 name: String,
206 }
207
208 impl Model for TestUser {
209 type PrimaryKey = i64;
210 fn table_name() -> &'static str {
211 "users"
212 }
213 fn pk(&self) -> Self::PrimaryKey {
214 self.id
215 }
216 fn set_pk(&mut self, pk: Self::PrimaryKey) {
217 self.id = pk;
218 }
219 }
220
221 #[tokio::test]
222 async fn test_stream_buffered_yields_rows() {
223 let mut mock = MockConnection::new();
224 mock.expect_any().with_rows(vec![
225 vec![("id", Value::I64(1))],
226 vec![("id", Value::I64(2))],
227 ]);
228
229 let dialect = get_dialect(DbType::MySQL).unwrap();
230 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
231
232 let mut stream = query.stream_buffered(&mut mock);
233 let mut rows = Vec::new();
234 while let Some(row) = stream.next().await {
235 rows.push(row.unwrap());
236 }
237
238 assert_eq!(rows.len(), 2);
239 }
240
241 #[tokio::test]
242 async fn test_stream_buffered_error_propagation() {
243 let mut mock = MockConnection::new().with_fallback(crate::mock::FallbackBehavior::Error);
244
245 let dialect = get_dialect(DbType::MySQL).unwrap();
246 let query = QueryBuilder::<TestUser>::new(dialect).table("nonexistent");
247
248 let mut stream = query.stream_buffered(&mut mock);
249 let result = stream.next().await;
250
251 assert!(result.is_some());
252 assert!(result.unwrap().is_err());
253 }
254
255 #[tokio::test]
256 async fn test_stream_buffered_empty() {
257 let mut mock = MockConnection::new();
258 mock.expect_any().with_rows(vec![]);
259
260 let dialect = get_dialect(DbType::MySQL).unwrap();
261 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
262
263 let mut stream = query.stream_buffered(&mut mock);
264 let mut count = 0;
265 while stream.next().await.is_some() {
266 count += 1;
267 }
268
269 assert_eq!(count, 0);
270 }
271
272 #[tokio::test]
273 async fn test_stream_backpressure_zero_buffer() {
274 let mut mock = MockConnection::new();
275 mock.expect_any().with_rows(vec![]);
276
277 let dialect = get_dialect(DbType::MySQL).unwrap();
278 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
279
280 let result = query.stream_with_backpressure(&mut mock, 0);
281 assert!(result.is_err());
282 let err = result.err().unwrap();
283 assert!(matches!(err, crate::DbError::InvalidInput(_)));
284 }
285
286 #[tokio::test]
287 async fn test_stream_backpressure_basic() {
288 let mut mock = MockConnection::new();
289 mock.expect_any().with_rows(vec![
290 vec![("id", Value::I64(1))],
291 vec![("id", Value::I64(2))],
292 vec![("id", Value::I64(3))],
293 ]);
294
295 let dialect = get_dialect(DbType::MySQL).unwrap();
296 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
297
298 let stream = query.stream_with_backpressure(&mut mock, 10).unwrap();
299 let mut stream = Box::pin(stream);
300 let mut rows = Vec::new();
301 while let Some(row) = stream.next().await {
302 rows.push(row.unwrap());
303 }
304
305 assert_eq!(rows.len(), 3);
306 }
307
308 #[tokio::test]
309 async fn test_stream_backpressure_empty() {
310 let mut mock = MockConnection::new();
311 mock.expect_any().with_rows(vec![]);
312
313 let dialect = get_dialect(DbType::MySQL).unwrap();
314 let query = QueryBuilder::<TestUser>::new(dialect).table("users");
315
316 let stream = query.stream_with_backpressure(&mut mock, 100).unwrap();
317 let mut stream = Box::pin(stream);
318 let mut count = 0;
319 while stream.next().await.is_some() {
320 count += 1;
321 }
322
323 assert_eq!(count, 0);
324 }
325
326 #[tokio::test]
327 async fn test_stream_backpressure_error_propagation() {
328 let mut mock = MockConnection::new().with_fallback(crate::mock::FallbackBehavior::Error);
329
330 let dialect = get_dialect(DbType::MySQL).unwrap();
331 let query = QueryBuilder::<TestUser>::new(dialect).table("nonexistent");
332
333 let stream = query.stream_with_backpressure(&mut mock, 10).unwrap();
334 let mut stream = Box::pin(stream);
335 let result = stream.next().await;
336
337 assert!(result.is_some());
338 assert!(result.unwrap().is_err());
339 }
340}