Skip to main content

sz_orm_core/
stream_api.rs

1//! Stream API — 异步流式查询
2//!
3//! # 概述
4//!
5//! v2.1.0 改造 `stream` 为真游标逐行产出,峰值内存 ≤ 50 MB(100 万行)。
6//! `stream_buffered` 保留旧实现(全量收集后逐行 yield)作为兼容逃生舱。
7//!
8//! # 向后兼容
9//!
10//! - `stream` trait 签名不变(ADR-v2.1.0-003),仅改 impl 实现
11//! - `stream_buffered` 行为与 v2.0.0 `stream` 完全一致
12//!
13//! # 示例
14//!
15//! ```ignore
16//! use sz_orm_core::paginator::StreamQueryTrait;
17//! use futures::StreamExt;
18//!
19//! // 真游标流(推荐)
20//! let mut stream = query.stream(&mut conn);
21//! while let Some(row) = stream.next().await {
22//!     let row = row?;
23//!     // 处理...
24//! }
25//!
26//! // 兼容版(全量收集后逐行 yield)
27//! let mut stream = query.stream_buffered(&mut conn);
28//! while let Some(row) = stream.next().await {
29//!     let row = row?;
30//!     // 处理...
31//! }
32//! ```
33
34use 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
44/// 流式查询结果行类型
45pub type RowResult = HashMap<String, crate::value::Value>;
46
47/// Stream API 扩展 trait
48///
49/// 为 `QueryBuilder<M>` 提供 `stream_buffered` 兼容版方法和 `stream_with_backpressure` 背压方法。
50pub trait StreamApiExt<M: Model> {
51    /// 兼容版流式查询(全量收集后逐行 yield)
52    ///
53    /// 保留 v2.0.0 `stream` 的行为,作为逃生舱。
54    /// 推荐使用 `stream`(真游标,低内存)。
55    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    /// 背压流式查询(v2.2.0 B-4)
61    ///
62    /// 创建有界缓冲通道,缓冲区满时生产者阻塞(背压)。
63    ///
64    /// # 参数
65    ///
66    /// - `conn`:数据库连接
67    /// - `buffer_size`:缓冲区容量(必须 > 0)
68    ///
69    /// # 错误
70    ///
71    /// - `buffer_size == 0` → 返回 `Err(DbError::InvalidInput)`
72    ///
73    /// ```ignore
74    /// let stream = query.stream_with_backpressure(&mut conn, 1000)?;
75    /// // 缓冲区容量 1000,满时生产者阻塞
76    /// ```
77    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, &params).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, &params).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
139/// 背压流(v2.2.0 B-4):合并生产者 future 和接收器
140struct 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
167/// 真游标流式查询
168///
169/// 委托 `conn.query_stream_cursor` 实现真游标逐行 fetch。
170/// drop 时关闭 DB 游标,连接归还连接池。
171///
172/// # 错误传播
173///
174/// - 游标 fetch 失败 → yield `Some(Err(DbError::ConnectionError))`
175/// - 游标打开失败 → yield `Some(Err(DbError))`
176pub 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; // 真游标在 query_stream_cursor 内部处理参数
183    conn.query_stream_cursor(sql, batch_size)
184}
185
186// ============================================================================
187// 单元测试
188// ============================================================================
189
190#[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}