Skip to main content

sz_orm_core/
queryable.rs

1//! derive(Queryable) — 从 SELECT 结果自动派生结构体(Diesel 风格)
2//!
3//! Diesel 通过 `#[derive(Queryable)]` 让结构体自动从 SQL 行反序列化。
4//! SZ-ORM 在 [`Value`] 之上提供类似的 trait + 派生辅助。
5//!
6//! 由于 proc-macro derive 需要在 `sz-orm-macros` 包中实现,
7//! 此模块提供 trait 定义和运行时反序列化逻辑;
8//! 派生宏 `#[derive(Queryable)]` 在 `sz-orm-macros` 中实现。
9//!
10//! # 设计
11//!
12//! - [`Queryable`] trait:从 `Vec<Value>` 按列顺序填充结构体字段
13//! - [`FromRow`] trait:从 `HashMap<String, Value>` 按列名填充(更鲁棒)
14//! - [`RowDesc`]:行描述(列名 + 列数),用于反序列化校验
15//!
16//! # 用法
17//!
18//! ```ignore
19//! use sz_orm_core::value::Value;
20//! use sz_orm_core::queryable::{Queryable, FromRow, RowDesc};
21//!
22//! #[derive(Debug, Default, PartialEq)]
23//! struct UserRow {
24//!     id: i64,
25//!     name: String,
26//! }
27//!
28//! impl Queryable for UserRow {
29//!     fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
30//!         if values.len() != 2 {
31//!             return Err(QueryError::ColumnCountMismatch {
32//!                 expected: 2,
33//!                 actual: values.len(),
34//!             });
35//!         }
36//!         let id = values[0].as_i64().ok_or(QueryError::TypeMismatch {
37//!             column: 0,
38//!             expected: "i64",
39//!         })?;
40//!         let name = values[1].as_str().ok_or(QueryError::TypeMismatch {
41//!             column: 1,
42//!             expected: "String",
43//!         })?.to_string();
44//!         Ok(UserRow { id, name })
45//!     }
46//! }
47//!
48//! let row = UserRow::from_values(vec![Value::I64(42), Value::String("Alice".into())]).unwrap();
49//! assert_eq!(row.id, 42);
50//! assert_eq!(row.name, "Alice");
51//! ```
52
53use crate::Value;
54use std::collections::HashMap;
55
56/// 反序列化错误
57#[derive(Debug, Clone, PartialEq)]
58pub enum QueryError {
59    /// 列数不匹配
60    ColumnCountMismatch {
61        /// 期望的列数
62        expected: usize,
63        /// 实际的列数
64        actual: usize,
65    },
66    /// 类型不匹配
67    TypeMismatch {
68        /// 列索引(按位置反序列化时)或列名(按名反序列化时)
69        column: std::borrow::Cow<'static, str>,
70        /// 期望的 Rust 类型名
71        expected: &'static str,
72    },
73    /// 缺少列
74    MissingColumn {
75        /// 缺失的列名
76        column: &'static str,
77    },
78    /// 自定义错误
79    Custom(String),
80}
81
82impl std::fmt::Display for QueryError {
83    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84        match self {
85            QueryError::ColumnCountMismatch { expected, actual } => {
86                write!(f, "列数不匹配: 期望 {}, 实际 {}", expected, actual)
87            }
88            QueryError::TypeMismatch { column, expected } => {
89                write!(f, "列 {:?} 类型不匹配, 期望 {}", column, expected)
90            }
91            QueryError::MissingColumn { column } => {
92                write!(f, "缺少列: {}", column)
93            }
94            QueryError::Custom(msg) => write!(f, "{}", msg),
95        }
96    }
97}
98
99impl std::error::Error for QueryError {}
100
101/// 行描述:列名 + 列数
102///
103/// 用于在按位置反序列化([`Queryable`])时提供列名信息,
104/// 或在按名反序列化([`FromRow`])时校验列存在性。
105#[derive(Debug, Clone)]
106pub struct RowDesc {
107    /// 列名列表(按 SELECT 顺序)
108    pub columns: Vec<String>,
109}
110
111impl RowDesc {
112    /// 创建行描述
113    pub fn new(columns: Vec<String>) -> Self {
114        Self { columns }
115    }
116
117    /// 列数
118    pub fn len(&self) -> usize {
119        self.columns.len()
120    }
121
122    /// 是否为空
123    pub fn is_empty(&self) -> bool {
124        self.columns.is_empty()
125    }
126
127    /// 查找列索引(按名)
128    pub fn index_of(&self, name: &str) -> Option<usize> {
129        self.columns.iter().position(|c| c == name)
130    }
131}
132
133/// 从 `Vec<Value>` 按列顺序反序列化(Diesel 风格)
134///
135/// 适用于 `SELECT id, name FROM users` 这种列顺序已知的查询。
136/// 列顺序由 SQL 决定,结构体字段顺序需与之对应。
137pub trait Queryable: Sized {
138    /// 从按 SELECT 顺序排列的值列表构造实例
139    fn from_values(values: Vec<Value>) -> Result<Self, QueryError>;
140
141    /// 从带行描述的值列表构造(默认实现忽略描述)
142    fn from_values_with_desc(values: Vec<Value>, desc: &RowDesc) -> Result<Self, QueryError> {
143        if values.len() != desc.len() {
144            return Err(QueryError::ColumnCountMismatch {
145                expected: desc.len(),
146                actual: values.len(),
147            });
148        }
149        Self::from_values(values)
150    }
151}
152
153/// 从 `HashMap<String, Value>` 按列名反序列化(更鲁棒)
154///
155/// 适用于列顺序不固定或查询使用 `*` 的场景。
156/// 按列名查找,不受 SQL 列顺序影响。
157pub trait FromRow: Sized {
158    /// 从列名到值的映射构造实例
159    fn from_row(row: HashMap<String, Value>) -> Result<Self, QueryError>;
160}
161
162// ---- 基础类型的 Queryable 实现 ----
163
164/// 单列查询结果(如 `SELECT COUNT(*)`)
165impl Queryable for Value {
166    fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
167        if values.len() != 1 {
168            return Err(QueryError::ColumnCountMismatch {
169                expected: 1,
170                actual: values.len(),
171            });
172        }
173        Ok(values.into_iter().next().expect("len==1 verified above"))
174    }
175}
176
177/// 双列查询结果
178impl Queryable for (Value, Value) {
179    fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
180        if values.len() != 2 {
181            return Err(QueryError::ColumnCountMismatch {
182                expected: 2,
183                actual: values.len(),
184            });
185        }
186        let mut iter = values.into_iter();
187        Ok((iter.next().expect("len==2 verified above"), iter.next().expect("len==2 verified above")))
188    }
189}
190
191/// 三列查询结果
192impl Queryable for (Value, Value, Value) {
193    fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
194        if values.len() != 3 {
195            return Err(QueryError::ColumnCountMismatch {
196                expected: 3,
197                actual: values.len(),
198            });
199        }
200        let mut iter = values.into_iter();
201        Ok((
202            iter.next().expect("len==3 verified above"),
203            iter.next().expect("len==3 verified above"),
204            iter.next().expect("len==3 verified above"),
205        ))
206    }
207}
208
209// ---- 辅助函数:从 Value 提取 Rust 类型 ----
210
211/// 提取 i64(支持 I8/I16/I32/I64/U8/U16/U32/U64/F32/F64/Bool/String 转换)
212pub fn value_as_i64(v: &Value) -> Option<i64> {
213    v.as_i64()
214}
215
216/// 提取 f64
217pub fn value_as_f64(v: &Value) -> Option<f64> {
218    v.as_f64()
219}
220
221/// 提取 String
222pub fn value_as_string(v: &Value) -> Option<String> {
223    v.as_str().map(|s| s.to_string())
224}
225
226/// 提取 bool
227pub fn value_as_bool(v: &Value) -> Option<bool> {
228    v.as_bool()
229}
230
231/// 提取 `Option<i64>`(Null 返回 None)
232pub fn value_as_nullable_i64(v: &Value) -> Option<i64> {
233    if v.is_null() {
234        None
235    } else {
236        v.as_i64()
237    }
238}
239
240/// 提取 `Option<String>`(Null 返回 None)
241pub fn value_as_nullable_string(v: &Value) -> Option<String> {
242    if v.is_null() {
243        None
244    } else {
245        v.as_str().map(|s| s.to_string())
246    }
247}
248
249#[cfg(test)]
250mod tests {
251    use super::*;
252
253    // ---- 测试用结构体 ----
254
255    #[derive(Debug, Default, PartialEq)]
256    struct UserRow {
257        id: i64,
258        name: String,
259    }
260
261    impl Queryable for UserRow {
262        fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
263            if values.len() != 2 {
264                return Err(QueryError::ColumnCountMismatch {
265                    expected: 2,
266                    actual: values.len(),
267                });
268            }
269            let id = values[0].as_i64().ok_or(QueryError::TypeMismatch {
270                column: "0".into(),
271                expected: "i64",
272            })?;
273            let name = values[1]
274                .as_str()
275                .ok_or(QueryError::TypeMismatch {
276                    column: "1".into(),
277                    expected: "String",
278                })?
279                .to_string();
280            Ok(UserRow { id, name })
281        }
282    }
283
284    impl FromRow for UserRow {
285        fn from_row(row: HashMap<String, Value>) -> Result<Self, QueryError> {
286            let id = row
287                .get("id")
288                .ok_or(QueryError::MissingColumn { column: "id" })?
289                .as_i64()
290                .ok_or(QueryError::TypeMismatch {
291                    column: "id".into(),
292                    expected: "i64",
293                })?;
294            let name = row
295                .get("name")
296                .ok_or(QueryError::MissingColumn { column: "name" })?
297                .as_str()
298                .ok_or(QueryError::TypeMismatch {
299                    column: "name".into(),
300                    expected: "String",
301                })?
302                .to_string();
303            Ok(UserRow { id, name })
304        }
305    }
306
307    // ---- QueryError 测试 ----
308
309    #[test]
310    fn test_query_error_display() {
311        let e = QueryError::ColumnCountMismatch {
312            expected: 3,
313            actual: 2,
314        };
315        assert!(format!("{}", e).contains("3"));
316        assert!(format!("{}", e).contains("2"));
317
318        let e = QueryError::TypeMismatch {
319            column: "age".into(),
320            expected: "i64",
321        };
322        assert!(format!("{}", e).contains("age"));
323
324        let e = QueryError::MissingColumn { column: "id" };
325        assert!(format!("{}", e).contains("id"));
326
327        let e = QueryError::Custom("custom".into());
328        assert_eq!(format!("{}", e), "custom");
329    }
330
331    // ---- RowDesc 测试 ----
332
333    #[test]
334    fn test_row_desc_basic() {
335        let desc = RowDesc::new(vec!["id".into(), "name".into(), "age".into()]);
336        assert_eq!(desc.len(), 3);
337        assert!(!desc.is_empty());
338        assert_eq!(desc.index_of("name"), Some(1));
339        assert_eq!(desc.index_of("missing"), None);
340    }
341
342    #[test]
343    fn test_row_desc_empty() {
344        let desc = RowDesc::new(vec![]);
345        assert!(desc.is_empty());
346        assert_eq!(desc.len(), 0);
347    }
348
349    // ---- Queryable for UserRow 测试 ----
350
351    #[test]
352    fn test_user_row_from_values_success() {
353        let row =
354            UserRow::from_values(vec![Value::I64(42), Value::String("Alice".into())]).unwrap();
355        assert_eq!(row.id, 42);
356        assert_eq!(row.name, "Alice");
357    }
358
359    #[test]
360    fn test_user_row_from_values_count_mismatch() {
361        let result = UserRow::from_values(vec![Value::I64(42)]);
362        assert!(matches!(
363            result,
364            Err(QueryError::ColumnCountMismatch {
365                expected: 2,
366                actual: 1
367            })
368        ));
369    }
370
371    #[test]
372    fn test_user_row_from_values_type_mismatch() {
373        let result = UserRow::from_values(vec![
374            Value::String("not_an_int".into()),
375            Value::String("Alice".into()),
376        ]);
377        assert!(matches!(result, Err(QueryError::TypeMismatch { .. })));
378    }
379
380    #[test]
381    fn test_user_row_from_values_with_desc() {
382        let desc = RowDesc::new(vec!["id".into(), "name".into()]);
383        let row =
384            UserRow::from_values_with_desc(vec![Value::I64(1), Value::String("Bob".into())], &desc)
385                .unwrap();
386        assert_eq!(row.id, 1);
387        assert_eq!(row.name, "Bob");
388    }
389
390    #[test]
391    fn test_user_row_from_values_with_desc_mismatch() {
392        let desc = RowDesc::new(vec!["id".into(), "name".into(), "age".into()]);
393        let result =
394            UserRow::from_values_with_desc(vec![Value::I64(1), Value::String("Bob".into())], &desc);
395        assert!(matches!(
396            result,
397            Err(QueryError::ColumnCountMismatch { .. })
398        ));
399    }
400
401    // ---- FromRow for UserRow 测试 ----
402
403    #[test]
404    fn test_user_row_from_row_success() {
405        let mut map = HashMap::new();
406        map.insert("id".into(), Value::I64(99));
407        map.insert("name".into(), Value::String("Charlie".into()));
408        let row = UserRow::from_row(map).unwrap();
409        assert_eq!(row.id, 99);
410        assert_eq!(row.name, "Charlie");
411    }
412
413    #[test]
414    fn test_user_row_from_row_missing_column() {
415        let mut map = HashMap::new();
416        map.insert("id".into(), Value::I64(99));
417        // 缺少 name
418        let result = UserRow::from_row(map);
419        assert!(matches!(
420            result,
421            Err(QueryError::MissingColumn { column: "name" })
422        ));
423    }
424
425    #[test]
426    fn test_user_row_from_row_extra_columns_ignored() {
427        let mut map = HashMap::new();
428        map.insert("id".into(), Value::I64(1));
429        map.insert("name".into(), Value::String("X".into()));
430        map.insert("extra".into(), Value::String("ignored".into()));
431        let row = UserRow::from_row(map).unwrap();
432        assert_eq!(row.id, 1);
433    }
434
435    // ---- 基础类型 Queryable 实现 ----
436
437    #[test]
438    fn test_value_queryable_single() {
439        let v = Value::from_values(vec![Value::I64(42)]).unwrap();
440        assert_eq!(v.as_i64(), Some(42));
441    }
442
443    #[test]
444    fn test_value_queryable_count_mismatch() {
445        let result = Value::from_values(vec![Value::I64(1), Value::I64(2)]);
446        assert!(matches!(
447            result,
448            Err(QueryError::ColumnCountMismatch { .. })
449        ));
450    }
451
452    #[test]
453    fn test_tuple_2_queryable() {
454        let (a, b) =
455            <(Value, Value)>::from_values(vec![Value::I64(1), Value::String("hello".into())])
456                .unwrap();
457        assert_eq!(a.as_i64(), Some(1));
458        assert_eq!(b.as_str(), Some("hello"));
459    }
460
461    #[test]
462    fn test_tuple_3_queryable() {
463        let (a, b, c) = <(Value, Value, Value)>::from_values(vec![
464            Value::I64(1),
465            Value::String("two".into()),
466            Value::F64(3.5),
467        ])
468        .unwrap();
469        assert_eq!(a.as_i64(), Some(1));
470        assert_eq!(b.as_str(), Some("two"));
471        assert_eq!(c.as_f64(), Some(3.5));
472    }
473
474    // ---- 辅助函数测试 ----
475
476    #[test]
477    fn test_value_helpers() {
478        assert_eq!(value_as_i64(&Value::I64(42)), Some(42));
479        assert_eq!(value_as_i64(&Value::String("42".into())), Some(42));
480        assert_eq!(value_as_f64(&Value::F64(3.5)), Some(3.5));
481        assert_eq!(
482            value_as_string(&Value::String("hi".into())),
483            Some("hi".into())
484        );
485        assert_eq!(value_as_bool(&Value::Bool(true)), Some(true));
486    }
487
488    #[test]
489    fn test_nullable_helpers() {
490        assert_eq!(value_as_nullable_i64(&Value::Null), None);
491        assert_eq!(value_as_nullable_i64(&Value::I64(42)), Some(42));
492        assert_eq!(value_as_nullable_string(&Value::Null), None);
493        assert_eq!(
494            value_as_nullable_string(&Value::String("hi".into())),
495            Some("hi".into())
496        );
497    }
498
499    // ---- 完整流程测试 ----
500
501    #[test]
502    fn test_full_flow_queryable() {
503        // 模拟 SELECT id, name FROM users
504        let values = vec![Value::I64(1), Value::String("Alice".into())];
505        let row = UserRow::from_values(values).unwrap();
506        assert_eq!(
507            row,
508            UserRow {
509                id: 1,
510                name: "Alice".into()
511            }
512        );
513    }
514
515    #[test]
516    fn test_full_flow_from_row_with_extra_data() {
517        // 模拟 SELECT * FROM users(带额外列)
518        let mut map = HashMap::new();
519        map.insert("id".into(), Value::I64(7));
520        map.insert("name".into(), Value::String("Bob".into()));
521        map.insert("email".into(), Value::String("bob@example.com".into()));
522        map.insert("created_at".into(), Value::String("2026-01-01".into()));
523
524        let row = UserRow::from_row(map).unwrap();
525        assert_eq!(row.id, 7);
526        assert_eq!(row.name, "Bob");
527    }
528}