use std::fmt::Write as _;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ColumnSchema {
pub name: String,
pub rust_type: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TableSchema {
pub name: String,
pub columns: Vec<ColumnSchema>,
}
pub struct SchemaGenerator {
header: String,
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
}
pub fn emit_use(mut self, emit: bool) -> Self {
self.emit_use = emit;
self
}
pub fn generate(&self, tables: &[TableSchema]) -> String {
let mut out = String::new();
writeln!(out, "{}", self.header).unwrap();
writeln!(out).unwrap();
if self.emit_use {
writeln!(out, "use sz_orm_core::typed_query;").unwrap();
writeln!(out).unwrap();
}
for (idx, table) in tables.iter().enumerate() {
if idx > 0 {
writeln!(out).unwrap();
}
write!(out, "{}", self.render_table(table)).unwrap();
}
out
}
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
}
}
pub fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
let base = match sql_type.to_lowercase() {
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",
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",
s if s.contains("json") => "String",
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(&[]);
assert!(output.contains("Auto-generated"));
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() {
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);
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);
}
}