sz-orm-core 1.0.0

Core ORM engine: Model trait, ActiveRecord, QueryBuilder, Pool, Transaction, migration, and SQL dialect abstraction
Documentation
//! Diesel 风格 schema.rs 自动生成
//!
//! 从数据库表元数据生成 `typed_query!` 表声明,配合宏做编译期列名校验。
//!
//! # 设计
//!
//! Diesel 的 `diesel print-schema` 命令从数据库反向生成 `schema.rs` 文件,
//! 包含 `table!` 宏声明,让 SQL 列名错误在编译期被捕获。
//!
//! 本模块提供等价功能,生成 SZ-ORM 的 `typed_query!` 声明:
//!
//! ```ignore
//! // 生成的 schema.rs 内容
//! use sz_orm_core::typed_query;
//!
//! typed_query! {
//!     table users {
//!         id: i64,
//!         name: String,
//!         email: String,
//!         created_at: String,
//!     }
//! }
//!
//! typed_query! {
//!     table orders {
//!         order_id: i64,
//!         user_id: i64,
//!         total: f64,
//!     }
//! }
//! ```

use std::fmt::Write as _;

/// 单列的元数据(足够生成 typed_query! 声明)
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ColumnSchema {
    /// 列名
    pub name: String,
    /// Rust 类型名(如 "i64"、"String"、"f64"、"`Option<String>`")
    pub rust_type: String,
}

/// 单表的元数据
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TableSchema {
    /// 表名
    pub name: String,
    /// 所有列
    pub columns: Vec<ColumnSchema>,
}

/// 生成器:把表元数据列表转换成 schema.rs 文件内容
pub struct SchemaGenerator {
    /// 文件头注释(默认包含 SZ-ORM 自动生成提示)
    header: String,
    /// 是否生成 `use` 语句
    emit_use: bool,
}

impl Default for SchemaGenerator {
    fn default() -> Self {
        Self::new()
    }
}

impl SchemaGenerator {
    /// 创建新的生成器
    pub fn new() -> Self {
        Self {
            header: format!(
                "// Auto-generated by sz-orm-cli generate schema at {}\n\
                 // DO NOT EDIT MANUALLY — re-run the command to refresh.\n\
                 //\n\
                 // This file contains typed_query! table declarations\n\
                 // enabling compile-time column-name verification.\n",
                chrono::Utc::now().format("%Y-%m-%d %H:%M:%S UTC")
            ),
            emit_use: true,
        }
    }

    /// 自定义文件头注释
    pub fn with_header(mut self, header: impl Into<String>) -> Self {
        self.header = header.into();
        self
    }

    /// 是否生成 `use sz_orm_core::typed_query;` 语句
    pub fn emit_use(mut self, emit: bool) -> Self {
        self.emit_use = emit;
        self
    }

    /// 生成完整的 schema.rs 文件内容
    pub fn generate(&self, tables: &[TableSchema]) -> String {
        let mut out = String::new();

        // 文件头
        writeln!(out, "{}", self.header).unwrap();
        writeln!(out).unwrap();

        // use 语句
        if self.emit_use {
            writeln!(out, "use sz_orm_core::typed_query;").unwrap();
            writeln!(out).unwrap();
        }

        // 每张表生成一个 typed_query! 声明
        for (idx, table) in tables.iter().enumerate() {
            if idx > 0 {
                writeln!(out).unwrap();
            }
            write!(out, "{}", self.render_table(table)).unwrap();
        }

        out
    }

    /// 渲染单张表的 typed_query! 声明
    fn render_table(&self, table: &TableSchema) -> String {
        let mut out = String::new();
        writeln!(out, "typed_query! {{").unwrap();
        writeln!(out, "    table {} {{", table.name).unwrap();
        for col in &table.columns {
            writeln!(out, "        {}: {},", col.name, col.rust_type).unwrap();
        }
        writeln!(out, "    }}").unwrap();
        writeln!(out, "}}").unwrap();
        out
    }
}

/// 把 SQL 类型字符串映射到 Rust 类型字符串
///
/// 用于 `generate schema` 命令从 DB 元数据生成 typed_query! 声明
pub fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
    let base = match sql_type.to_lowercase() {
        // 整数(按长度优先匹配,避免 "int" 误匹配 "bigint")
        s if s.contains("tinyint") => "i8",
        s if s.contains("smallint") || s.contains("int2") => "i16",
        s if s.contains("bigint") || s.contains("int8") => "i64",
        s if s.contains("int") || s.contains("serial") => "i32",
        // 浮点(更具体的先匹配:float8/float4 在 float 之前)
        s if s.contains("float8") || s.contains("double") => "f64",
        s if s.contains("float4") || s.contains("real") => "f32",
        s if s.contains("float") => "f32",
        s if s.contains("decimal") || s.contains("numeric") => "f64",
        // 布尔
        s if s.contains("bool") => "bool",
        // 字节
        s if s.contains("blob") || s.contains("bytea") || s.contains("binary") => "Vec<u8>",
        // 时间
        s if s.contains("date") && !s.contains("datetime") && !s.contains("timestamp") => "String",
        s if s.contains("datetime") || s.contains("timestamp") => "String",
        s if s.contains("time") => "String",
        // JSON
        s if s.contains("json") => "String",
        // UUID
        s if s.contains("uuid") => "String",
        // 字符串
        s if s.contains("char") || s.contains("text") || s.contains("varchar") => "String",
        // 默认
        _ => "String",
    };

    if nullable {
        format!("Option<{}>", base)
    } else {
        base.to_string()
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_sql_type_to_rust_int() {
        assert_eq!(sql_type_to_rust("INT", false), "i32");
        assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
        assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
        assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
    }

    #[test]
    fn test_sql_type_to_rust_float() {
        assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
        assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
        assert_eq!(sql_type_to_rust("DECIMAL(10,2)", false), "f64");
    }

    #[test]
    fn test_sql_type_to_rust_bool() {
        assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
        assert_eq!(sql_type_to_rust("TINYINT(1)", false), "i8");
    }

    #[test]
    fn test_sql_type_to_rust_string() {
        assert_eq!(sql_type_to_rust("VARCHAR(255)", false), "String");
        assert_eq!(sql_type_to_rust("TEXT", false), "String");
        assert_eq!(sql_type_to_rust("CHAR(36)", false), "String");
    }

    #[test]
    fn test_sql_type_to_rust_binary() {
        assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
        assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
    }

    #[test]
    fn test_sql_type_to_rust_nullable() {
        assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
        assert_eq!(sql_type_to_rust("VARCHAR(255)", true), "Option<String>");
    }

    #[test]
    fn test_sql_type_to_rust_pg_types() {
        assert_eq!(sql_type_to_rust("int8", false), "i64");
        assert_eq!(sql_type_to_rust("int2", false), "i16");
        assert_eq!(sql_type_to_rust("float8", false), "f64");
        assert_eq!(sql_type_to_rust("float4", false), "f32");
        assert_eq!(sql_type_to_rust("bytea", false), "Vec<u8>");
    }

    #[test]
    fn test_sql_type_to_rust_json_uuid() {
        assert_eq!(sql_type_to_rust("JSON", false), "String");
        assert_eq!(sql_type_to_rust("JSONB", false), "String");
        assert_eq!(sql_type_to_rust("UUID", false), "String");
    }

    #[test]
    fn test_schema_generator_single_table() {
        let gen = SchemaGenerator::new().emit_use(false);
        let tables = vec![TableSchema {
            name: "users".to_string(),
            columns: vec![
                ColumnSchema {
                    name: "id".to_string(),
                    rust_type: "i64".to_string(),
                },
                ColumnSchema {
                    name: "name".to_string(),
                    rust_type: "String".to_string(),
                },
            ],
        }];
        let output = gen.generate(&tables);

        assert!(output.contains("table users {"));
        assert!(output.contains("id: i64,"));
        assert!(output.contains("name: String,"));
        assert!(output.contains("typed_query! {"));
    }

    #[test]
    fn test_schema_generator_multiple_tables() {
        let gen = SchemaGenerator::new().emit_use(false);
        let tables = vec![
            TableSchema {
                name: "users".to_string(),
                columns: vec![ColumnSchema {
                    name: "id".to_string(),
                    rust_type: "i64".to_string(),
                }],
            },
            TableSchema {
                name: "orders".to_string(),
                columns: vec![ColumnSchema {
                    name: "order_id".to_string(),
                    rust_type: "i64".to_string(),
                }],
            },
        ];
        let output = gen.generate(&tables);

        assert!(output.contains("table users {"));
        assert!(output.contains("table orders {"));
        // 两个表声明之间应有空行
        let users_end = output.find("}").unwrap();
        let orders_start = output.find("table orders").unwrap();
        let between = &output[users_end..orders_start];
        assert!(between.contains("\n\n"));
    }

    #[test]
    fn test_schema_generator_with_use_statement() {
        let gen = SchemaGenerator::new().emit_use(true);
        let tables = vec![TableSchema {
            name: "t".to_string(),
            columns: vec![],
        }];
        let output = gen.generate(&tables);

        assert!(output.contains("use sz_orm_core::typed_query;"));
    }

    #[test]
    fn test_schema_generator_header() {
        let gen = SchemaGenerator::new();
        let output = gen.generate(&[]);
        assert!(output.contains("Auto-generated"));
        assert!(output.contains("DO NOT EDIT MANUALLY"));
    }

    #[test]
    fn test_schema_generator_custom_header() {
        let gen = SchemaGenerator::new().with_header("// Custom header\n");
        let output = gen.generate(&[]);
        assert!(output.starts_with("// Custom header"));
        assert!(!output.contains("Auto-generated"));
    }

    #[test]
    fn test_schema_generator_empty_tables() {
        let gen = SchemaGenerator::new().emit_use(false);
        let output = gen.generate(&[]);
        // 空表列表也应生成(仅有 header,不应有 typed_query! 块)
        assert!(output.contains("Auto-generated"));
        // 应该没有 typed_query! 块(header 里只是描述性文字,不算块)
        // 通过检查 "typed_query! {" 来判断是否有实际声明
        assert!(!output.contains("typed_query! {"));
    }

    #[test]
    fn test_schema_generator_option_type() {
        let gen = SchemaGenerator::new().emit_use(false);
        let tables = vec![TableSchema {
            name: "products".to_string(),
            columns: vec![ColumnSchema {
                name: "price".to_string(),
                rust_type: "Option<f64>".to_string(),
            }],
        }];
        let output = gen.generate(&tables);

        assert!(output.contains("price: Option<f64>,"));
    }

    #[test]
    fn test_schema_generator_compound_type() {
        let gen = SchemaGenerator::new().emit_use(false);
        let tables = vec![TableSchema {
            name: "files".to_string(),
            columns: vec![ColumnSchema {
                name: "content".to_string(),
                rust_type: "Vec<u8>".to_string(),
            }],
        }];
        let output = gen.generate(&tables);

        assert!(output.contains("content: Vec<u8>,"));
    }

    #[test]
    fn test_generated_code_is_valid_syntax() {
        // 验证生成的代码包含正确的 typed_query! 调用语法
        let gen = SchemaGenerator::new().emit_use(false);
        let tables = vec![TableSchema {
            name: "typed_validate_test".to_string(),
            columns: vec![
                ColumnSchema {
                    name: "id".to_string(),
                    rust_type: "i64".to_string(),
                },
                ColumnSchema {
                    name: "name".to_string(),
                    rust_type: "String".to_string(),
                },
            ],
        }];
        let output = gen.generate(&tables);

        // 验证生成的代码包含 typed_query! 块的完整结构
        assert!(output.contains("typed_query! {"));
        assert!(output.contains("table typed_validate_test {"));
        assert!(output.contains("id: i64,"));
        assert!(output.contains("name: String,"));
        // 块应正确闭合
        let count_open = output.matches("typed_query! {").count();
        let count_close = output.matches("}\n}").count();
        assert_eq!(count_open, 1);
        assert_eq!(count_close, 1);
    }
}