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::{stream, Stream, StreamExt};
43
44/// 流式查询结果行类型
45pub type RowResult = HashMap<String, crate::value::Value>;
46
47/// Stream API 扩展 trait
48///
49/// 为 `QueryBuilder<M>` 提供 `stream_buffered` 兼容版方法。
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
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, &params).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
80/// 真游标流式查询
81///
82/// 委托 `conn.query_stream_cursor` 实现真游标逐行 fetch。
83/// drop 时关闭 DB 游标,连接归还连接池。
84///
85/// # 错误传播
86///
87/// - 游标 fetch 失败 → yield `Some(Err(DbError::ConnectionError))`
88/// - 游标打开失败 → yield `Some(Err(DbError))`
89pub 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; // 真游标在 query_stream_cursor 内部处理参数
96    conn.query_stream_cursor(sql, batch_size)
97}
98
99// ============================================================================
100// 单元测试
101// ============================================================================
102
103#[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}