Skip to main content

sz_orm_macros/
lib.rs

1//! SZ-ORM Procedural Macros - compile-time SQL validation & derive macros
2//!
3//! Provides:
4//! - `sql_string!` macro that validates SQL string literals at compile time.
5//!   Errors like `SELECT * FORM users` or `'; DROP TABLE` are caught before the binary is built.
6//! - `#[derive(Schema)]` — auto-generate table structure info from a struct.
7//! - `#[derive(Builder)]` — auto-generate builder pattern code for a struct.
8//!
9//! # Usage
10//!
11//! ```ignore
12//! use sz_orm_macros::sql_string;
13//!
14//! // Basic usage
15//! let sql = sql_string!("SELECT * FROM users WHERE id = 1"); // ✅ compiles
16//!
17//! // With parameter count check
18//! let sql = sql_string!("SELECT * FROM users WHERE id = ?");
19//!                      params: 1);                          // ✅ compiles
20//!
21//! // ❌ compile error: missing FROM
22//! let sql = sql_string!("SELECT * users WHERE id = 1");
23//!
24//! // ❌ compile error: SQL injection detected
25//! let sql = sql_string!("SELECT * FROM users WHERE name = 'x' OR '1'='1'");
26//!
27//! // ❌ compile error: parameter count mismatch
28//! let sql = sql_string!("SELECT * FROM users WHERE id = ?");
29//!                      params: 2);
30//! ```
31
32// 抑制 Windows 链接器输出"正在创建库 ..."的诊断信息被识别为警告:
33// 该输出是 link.exe 创建 DLL 导入库时的正常 stdout 提示,并非代码问题。
34#![allow(linker_messages)]
35//!
36//! # Derive macros
37//!
38//! ```ignore
39//! use sz_orm_macros::{Schema, Builder};
40//!
41//! #[derive(Schema)]
42//! #[table(name = "users")]
43//! struct User {
44//!     #[column(primary_key)]
45//!     id: i64,
46//!     name: String,
47//! }
48//!
49//! #[derive(Builder)]
50//! struct Order {
51//!     id: i64,
52//!     total: f64,
53//! }
54//! ```
55
56extern crate proc_macro;
57
58use proc_macro::{Delimiter, Group, Ident, Literal, Punct, Spacing, Span, TokenStream, TokenTree};
59
60// 引入 quote! 宏,用于类型安全地构建 TokenStream
61use proc_macro2::TokenStream as TokenStream2;
62use quote::quote;
63use syn::parse_macro_input;
64
65// 派生宏模块
66mod derive;
67
68/// Compile-time SQL validation macro.
69///
70/// Validates SQL syntax at compile time and emits the validated SQL string.
71///
72/// # Syntax
73///
74/// - `sql_string!("SQL")` — validates the SQL and emits it as a `&str`
75/// - `sql_string!("SQL"; params: N)` — additionally checks that the SQL has exactly N parameters
76///
77/// # Validation rules
78///
79/// - SELECT must contain FROM
80/// - INSERT must contain INTO and VALUES
81/// - UPDATE must contain SET
82/// - DELETE must contain FROM
83/// - Parentheses must be balanced
84/// - String literals must be properly closed
85/// - No SQL injection patterns (OR '1'='1', UNION SELECT, `'; DROP TABLE`, `--`, `/*`)
86/// - Table/column identifiers must be valid
87#[proc_macro]
88pub fn sql_string(input: TokenStream) -> TokenStream {
89    let mut tokens = input.into_iter().peekable();
90
91    // Parse the SQL string literal
92    let sql = match tokens.next() {
93        Some(TokenTree::Literal(lit)) => lit.to_string(),
94        Some(other) => {
95            return compile_error(
96                other.span(),
97                "Expected a string literal as the first argument to sql_string!",
98            );
99        }
100        None => {
101            return compile_error(
102                Span::call_site(),
103                "Expected a string literal argument to sql_string!",
104            );
105        }
106    };
107
108    // Remove surrounding quotes from the string literal
109    let sql_content = if sql.starts_with("r#\"") {
110        &sql[3..sql.len() - 2]
111    } else if sql.starts_with("r\"") {
112        &sql[2..sql.len() - 1]
113    } else if sql.starts_with('"') {
114        &sql[1..sql.len() - 1]
115    } else if sql.starts_with("b\"") || sql.starts_with("b\'") {
116        &sql[2..sql.len() - 1]
117    } else {
118        return compile_error(
119            Span::call_site(),
120            "sql_string! requires a string literal argument",
121        );
122    };
123
124    // Parse optional `params: N`
125    let mut expected_params = None;
126    if tokens.peek().is_some() {
127        // Expect `; params: N`
128        match tokens.next() {
129            Some(TokenTree::Punct(p)) if p.as_char() == ';' => {}
130            Some(other) => {
131                return compile_error(
132                    other.span(),
133                    "Expected `;` before param count, e.g. sql_string!(\"...\"; params: 2)",
134                );
135            }
136            None => {}
137        }
138
139        // Parse `params`
140        match tokens.next() {
141            Some(TokenTree::Ident(id)) if id.to_string() == "params" => {}
142            Some(other) => {
143                return compile_error(
144                    other.span(),
145                    "Expected `params:` keyword, e.g. sql_string!(\"...\"; params: 2)",
146                );
147            }
148            None => {
149                return compile_error(Span::call_site(), "Expected param count after `;`");
150            }
151        }
152
153        // Parse `:`
154        match tokens.next() {
155            Some(TokenTree::Punct(p)) if p.as_char() == ':' => {}
156            Some(other) => {
157                return compile_error(
158                    other.span(),
159                    "Expected `:` after `params`, e.g. sql_string!(\"...\"; params: 2)",
160                );
161            }
162            None => {
163                return compile_error(Span::call_site(), "Expected param count after `params`");
164            }
165        }
166
167        // Parse the number
168        match tokens.next() {
169            Some(TokenTree::Literal(lit)) => {
170                let num_str = lit.to_string();
171                if let Ok(n) = num_str.parse::<usize>() {
172                    expected_params = Some(n);
173                } else {
174                    return compile_error(
175                        lit.span(),
176                        "Expected a positive integer for param count",
177                    );
178                }
179            }
180            Some(other) => {
181                return compile_error(
182                    other.span(),
183                    "Expected a number after `params:`, e.g. sql_string!(\"...\"; params: 2)",
184                );
185            }
186            None => {
187                return compile_error(Span::call_site(), "Expected a number after `params:`");
188            }
189        }
190    }
191
192    // Run validation
193    if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
194        return compile_error(Span::call_site(), &err_msg);
195    }
196
197    // Emit the validated string as a &str literal
198    let output = format!("\"{}\"", sql_content.escape_default());
199    output
200        .parse()
201        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
202}
203
204// ---------------------------------------------------------------------------
205// Validation logic (self-contained, no external dependencies)
206// ---------------------------------------------------------------------------
207
208fn validate_sql_content(sql: &str, expected_params: Option<usize>) -> Result<(), String> {
209    let trimmed = sql.trim();
210    if trimmed.is_empty() {
211        return Err("SQL statement is empty".to_string());
212    }
213
214    validate_balanced_parens(trimmed)?;
215    validate_string_literals_closed(trimmed)?;
216    validate_no_injection(trimmed)?;
217
218    // Type-specific validation
219    let sql_upper = trimmed.to_uppercase();
220    if sql_upper.starts_with("SELECT") {
221        if !sql_upper.contains("FROM") {
222            return Err("SELECT statement missing FROM clause".to_string());
223        }
224    } else if sql_upper.starts_with("INSERT") {
225        if !sql_upper.contains("INTO") {
226            return Err("INSERT statement missing INTO clause".to_string());
227        }
228        if !sql_upper.contains("VALUES") {
229            return Err("INSERT statement missing VALUES clause".to_string());
230        }
231    } else if sql_upper.starts_with("UPDATE") {
232        if !sql_upper.contains("SET") {
233            return Err("UPDATE statement missing SET clause".to_string());
234        }
235    } else if sql_upper.starts_with("DELETE") && !sql_upper.contains("FROM") {
236        return Err("DELETE statement missing FROM clause".to_string());
237    }
238
239    // Parameter count check
240    if let Some(expected) = expected_params {
241        let actual = sql.chars().filter(|&c| c == '?').count();
242        if actual != expected {
243            return Err(format!(
244                "Parameter count mismatch: expected {} parameters, found {}",
245                expected, actual
246            ));
247        }
248    }
249
250    Ok(())
251}
252
253fn validate_balanced_parens(sql: &str) -> Result<(), String> {
254    let mut depth: i32 = 0;
255    for (i, ch) in sql.char_indices() {
256        match ch {
257            '(' => depth += 1,
258            ')' => {
259                depth -= 1;
260                if depth < 0 {
261                    return Err(format!(
262                        "Unbalanced parentheses: unexpected ')' at position {}",
263                        i
264                    ));
265                }
266            }
267            _ => {}
268        }
269    }
270    if depth != 0 {
271        return Err(format!("Unbalanced parentheses: {} unclosed '('", depth));
272    }
273    Ok(())
274}
275
276fn validate_string_literals_closed(sql: &str) -> Result<(), String> {
277    let mut in_single = false;
278    let mut in_double = false;
279    let mut prev = '\0';
280
281    for ch in sql.chars() {
282        if prev == '\\' {
283            prev = ch;
284            continue;
285        }
286
287        match ch {
288            '\'' if !in_double => in_single = !in_single,
289            '"' if !in_single => in_double = !in_double,
290            _ => {}
291        }
292        prev = ch;
293    }
294
295    if in_single {
296        return Err("Unclosed single-quoted string literal".to_string());
297    }
298    if in_double {
299        return Err("Unclosed double-quoted string literal".to_string());
300    }
301
302    Ok(())
303}
304
305fn validate_no_injection(sql: &str) -> Result<(), String> {
306    let sql_lower = sql.to_lowercase();
307
308    // 注意:编译期 SQL 内容已由 Rust 字符串字面量解析剥离外层引号,
309    // 因此检测模式不应依赖前导引号字符(如 `"'; DROP TABLE"`)。
310    let injection_patterns: &[&str] = &[
311        // 多语句攻击
312        "drop table",
313        "drop database",
314        "; drop",
315        // 经典注入
316        "or 1=1",
317        "or 1 = 1",
318        "union select",
319        "union all select",
320        // 注释攻击
321        "--",
322        "/*",
323        "*/",
324        // 存储过程注入
325        "xp_cmdshell",
326        "sp_executesql",
327        "exec(",
328        "execute(",
329        // 信息泄露
330        "information_schema",
331        "sys.tables",
332        "sys.columns",
333    ];
334
335    for pattern in injection_patterns {
336        if sql_lower.contains(pattern) {
337            return Err(format!("潜在的 SQL 注入模式被检测到: '{}'", pattern));
338        }
339    }
340
341    Ok(())
342}
343
344// ---------------------------------------------------------------------------
345// `query!` macro — optional real DB verification (gated by `db-verify` feature)
346// ---------------------------------------------------------------------------
347
348/// Compile-time SQL validation with optional real DB verification.
349///
350/// Behavior:
351/// - Always runs the same syntax validation as `sql_string!`.
352/// - When the `db-verify` cargo feature is enabled **AND** the
353///   `SZ_ORM_QUERY_VERIFY=1` environment variable is set at compile time,
354///   connects to the database pointed to by `DATABASE_URL` and runs
355///   `EXPLAIN` (MySQL/PostgreSQL) or `EXPLAIN QUERY PLAN` (SQLite) to verify
356///   the SQL is valid against the actual schema (column names, table names,
357///   joins, etc.).
358/// - Otherwise, falls back to syntax-only validation.
359///
360/// Emits the validated SQL as a `&'static str` literal.
361///
362/// # Syntax
363///
364/// ```ignore
365/// let sql = query!("SELECT id, name FROM users WHERE id = ?");
366/// ```
367///
368/// # Verification setup
369///
370/// ```bash
371/// export DATABASE_URL="mysql://user:pass@host:3306/db"
372/// export SZ_ORM_QUERY_VERIFY=1
373/// cargo build --features sz-orm-macros/db-verify
374/// ```
375#[proc_macro]
376pub fn query(input: TokenStream) -> TokenStream {
377    let mut tokens = input.into_iter().peekable();
378
379    // Parse the SQL string literal (same as sql_string!)
380    let sql = match tokens.next() {
381        Some(TokenTree::Literal(lit)) => lit.to_string(),
382        Some(other) => {
383            return compile_error(
384                other.span(),
385                "Expected a string literal as the first argument to query!",
386            );
387        }
388        None => {
389            return compile_error(
390                Span::call_site(),
391                "Expected a string literal argument to query!",
392            );
393        }
394    };
395
396    let sql_content = match strip_string_literal(&sql) {
397        Some(s) => s,
398        None => {
399            return compile_error(
400                Span::call_site(),
401                "query! requires a string literal argument",
402            );
403        }
404    };
405
406    // Syntax validation (shared with sql_string!)
407    if let Err(err_msg) = validate_sql_content(sql_content, None) {
408        return compile_error(Span::call_site(), &err_msg);
409    }
410
411    // Optional real DB verification (only when feature is enabled)
412    #[cfg(feature = "db-verify")]
413    {
414        if std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() == Some("1") {
415            if let Err(err) = verify_with_real_db(sql_content) {
416                return compile_error(
417                    Span::call_site(),
418                    &format!("query! real DB verification failed: {}", err),
419                );
420            }
421        }
422    }
423
424    // Emit the validated string as a &str literal
425    let output = format!("\"{}\"", sql_content.escape_default());
426    output
427        .parse()
428        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
429}
430
431/// Strip surrounding quotes from a string literal token's raw representation.
432/// Shared by `sql_string!` and `query!`.
433fn strip_string_literal(raw: &str) -> Option<&str> {
434    if raw.starts_with("r#\"") {
435        Some(&raw[3..raw.len() - 2])
436    } else if raw.starts_with("r\"") {
437        Some(&raw[2..raw.len() - 1])
438    } else if raw.starts_with('"') {
439        Some(&raw[1..raw.len() - 1])
440    } else if raw.starts_with("b\"") || raw.starts_with("b\'") {
441        Some(&raw[2..raw.len() - 1])
442    } else {
443        None
444    }
445}
446
447// ---------------------------------------------------------------------------
448// Real DB verification (only compiled when `db-verify` feature is enabled)
449// ---------------------------------------------------------------------------
450
451#[cfg(feature = "db-verify")]
452fn verify_with_real_db(sql: &str) -> Result<(), String> {
453    let dsn = std::env::var("DATABASE_URL")
454        .map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
455
456    let db_kind =
457        detect_db_kind(&dsn).map_err(|e| format!("Failed to detect DB kind from DSN: {}", e))?;
458
459    let explain_sql = match db_kind {
460        DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql),
461        DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
462    };
463
464    let rt = tokio::runtime::Runtime::new()
465        .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
466
467    rt.block_on(async {
468        match db_kind {
469            DbKind::MySql => verify_mysql(&dsn, &explain_sql).await,
470            DbKind::Postgres => verify_postgres(&dsn, &explain_sql).await,
471            DbKind::Sqlite => verify_sqlite(&dsn, &explain_sql).await,
472        }
473    })
474}
475
476#[cfg(feature = "db-verify")]
477#[derive(Debug, Clone, Copy, PartialEq, Eq)]
478enum DbKind {
479    MySql,
480    Postgres,
481    Sqlite,
482}
483
484#[cfg(feature = "db-verify")]
485fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
486    let lower = dsn.to_lowercase();
487    if lower.starts_with("mysql://") {
488        Ok(DbKind::MySql)
489    } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
490        Ok(DbKind::Postgres)
491    } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
492        Ok(DbKind::Sqlite)
493    } else {
494        Err(format!("Unsupported DSN scheme: {}", dsn))
495    }
496}
497
498#[cfg(feature = "db-verify")]
499async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
500    let pool = sqlx::MySqlPool::connect(dsn)
501        .await
502        .map_err(|e| format!("MySQL connect failed: {}", e))?;
503    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
504        .execute(&pool)
505        .await
506        .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
507    Ok(())
508}
509
510#[cfg(feature = "db-verify")]
511async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
512    let pool = sqlx::PgPool::connect(dsn)
513        .await
514        .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
515    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
516        .execute(&pool)
517        .await
518        .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
519    Ok(())
520}
521
522#[cfg(feature = "db-verify")]
523async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
524    let pool = sqlx::SqlitePool::connect(dsn)
525        .await
526        .map_err(|e| format!("SQLite connect failed: {}", e))?;
527    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
528        .execute(&pool)
529        .await
530        .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
531    Ok(())
532}
533
534// ---------------------------------------------------------------------------
535// Helpers
536// ---------------------------------------------------------------------------
537
538/// Create a compile_error! token stream
539fn compile_error(span: Span, msg: &str) -> TokenStream {
540    // emit: compile_error!("msg")
541    let mut ts = TokenStream::new();
542    ts.extend([
543        TokenTree::Ident(Ident::new("compile_error", span)),
544        TokenTree::Punct(Punct::new('!', Spacing::Alone)),
545        TokenTree::Group(Group::new(
546            Delimiter::Parenthesis,
547            TokenStream::from(TokenTree::Literal(Literal::string(msg))),
548        )),
549    ]);
550    ts
551}
552
553// ---------------------------------------------------------------------------
554// typed_query! — Diesel 风格强类型 AST 宏
555// ---------------------------------------------------------------------------
556
557/// Diesel 风格强类型 AST 宏(与 `sql_string!` / `query!` 并存)。
558///
559/// # 设计
560///
561/// 接收 `table { col1: Type, col2: Type, ... }` 声明,生成:
562/// 1. 一个 `table` 模块
563/// 2. 每列对应一个零大小标记类型(如 `table::id`)
564/// 3. 实现 `TypedColumn` trait,把列名 + Rust 类型提升到类型系统
565///
566/// 这样,`typed_query!(SELECT id FROM users WHERE name = ?)` 在编译期就能:
567/// - 校验 `id` / `name` 列是否存在于 `users` 表声明中
568/// - 校验 `?` 参数的 Rust 类型与列声明的类型一致
569///
570/// # 用法
571///
572/// ```ignore
573/// use sz_orm_macros::typed_query;
574///
575/// // 1. 声明表 schema(编译期生成 column 标记类型)
576/// typed_query! {
577///     table users {
578///         id: i64,
579///         name: String,
580///         email: String,
581///         age: i32,
582///     }
583/// }
584///
585/// // 2. 编译期校验 SELECT:列名必须存在于 users 表
586/// let sql = typed_query!(SELECT id, name FROM users WHERE age > ?);
587/// // ❌ 编译错误:unknown column 'foo' in table 'users'
588/// // let sql = typed_query!(SELECT foo FROM users);
589/// ```
590#[proc_macro]
591pub fn typed_query(input: TokenStream) -> TokenStream {
592    let tokens: Vec<TokenTree> = input.into_iter().collect();
593
594    // 分支 1:table 声明
595    if tokens.iter().any(|t| {
596        if let TokenTree::Ident(id) = t {
597            id.to_string() == "table"
598        } else {
599            false
600        }
601    }) {
602        return parse_table_decl(&tokens);
603    }
604
605    // 分支 2:SELECT 表达式
606    if tokens.iter().any(|t| {
607        if let TokenTree::Ident(id) = t {
608            id.to_string().eq_ignore_ascii_case("SELECT")
609        } else {
610            false
611        }
612    }) {
613        return parse_typed_select(&tokens);
614    }
615
616    compile_error(
617        Span::call_site(),
618        "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
619    )
620}
621
622/// 解析 `table name { col: Type, ... }` 声明
623fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
624    // 期望格式:table <ident> { <ident> : <ident> [, ...] }
625    let mut idx = 0;
626
627    // 跳过 'table' 关键字
628    if idx >= tokens.len() {
629        return compile_error(Span::call_site(), "expected table name after 'table'");
630    }
631    if let TokenTree::Ident(id) = &tokens[idx] {
632        if id.to_string() != "table" {
633            return compile_error(id.span(), "expected 'table' keyword");
634        }
635    }
636    idx += 1;
637
638    // 表名
639    let table_name = if idx < tokens.len() {
640        if let TokenTree::Ident(id) = &tokens[idx] {
641            id.to_string()
642        } else {
643            return compile_error(tokens[idx].span(), "expected table name identifier");
644        }
645    } else {
646        return compile_error(Span::call_site(), "expected table name");
647    };
648    idx += 1;
649
650    // 表体({} 内)
651    let body_group = if idx < tokens.len() {
652        if let TokenTree::Group(g) = &tokens[idx] {
653            if g.delimiter() != Delimiter::Brace {
654                return compile_error(g.span(), "expected '{' after table name");
655            }
656            g.clone()
657        } else {
658            return compile_error(tokens[idx].span(), "expected '{' after table name");
659        }
660    } else {
661        return compile_error(Span::call_site(), "expected table body in '{ }'");
662    };
663
664    // 解析列声明
665    let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
666    let columns = match parse_column_list(&body_tokens) {
667        Ok(c) => c,
668        Err(e) => return compile_error(Span::call_site(), &e),
669    };
670
671    // 使用 quote! 构建类型安全的 TokenStream
672    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
673    let table_name_lit = table_name.as_str();
674
675    // 为每列构建标记类型 + trait 实现
676    let col_impls: Vec<TokenStream2> = columns
677        .iter()
678        .map(|(col_name, col_type)| {
679            let col_ident =
680                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
681            let col_name_lit = col_name.as_str();
682            // 解析类型字符串为 TokenStream(quote! 会处理)
683            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
684            quote! {
685                #[derive(Debug, Clone, Copy)]
686                pub struct #col_ident;
687                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
688                    const NAME: &'static str = #col_name_lit;
689                    type Table = table;
690                    type RustType = #rust_type;
691                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
692                }
693            }
694        })
695        .collect();
696
697    // schema 常量条目
698    let schema_entries: Vec<TokenStream2> = columns
699        .iter()
700        .map(|(n, t)| {
701            let n_lit = n.as_str();
702            let t_lit = t.as_str();
703            quote! { (#n_lit, #t_lit) }
704        })
705        .collect();
706
707    let schema_const_ident = proc_macro2::Ident::new(
708        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
709        Span::call_site().into(),
710    );
711
712    let expanded = quote! {
713        pub mod #table_ident {
714            use super::*;
715            pub struct table;
716            impl ::sz_orm_core::typed::TypedTable for table {
717                const NAME: &'static str = #table_name_lit;
718            }
719            #(#col_impls)*
720        }
721        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
722    };
723
724    expanded.into()
725}
726
727/// 解析列声明列表:`col: Type, col2: Type2, ...`
728fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
729    let mut cols = Vec::new();
730    let mut i = 0;
731    while i < tokens.len() {
732        // 列名
733        let col_name = if let TokenTree::Ident(id) = &tokens[i] {
734            id.to_string()
735        } else {
736            return Err(format!("expected column name at position {}", i));
737        };
738        i += 1;
739
740        // 冒号
741        if i >= tokens.len() {
742            return Err(format!("expected ':' after column '{}'", col_name));
743        }
744        if let TokenTree::Punct(p) = &tokens[i] {
745            if p.as_char() != ':' {
746                return Err(format!("expected ':' after column '{}'", col_name));
747            }
748        } else {
749            return Err(format!("expected ':' after column '{}'", col_name));
750        }
751        i += 1;
752
753        // 类型(可能是 ident 或 path,如 String / i64 / Option<i64>)
754        // 简化处理:收集直到遇到 ',' 或末尾
755        let mut type_str = String::new();
756        let mut depth = 0;
757        while i < tokens.len() {
758            match &tokens[i] {
759                TokenTree::Punct(p) => {
760                    if p.as_char() == ',' && depth == 0 {
761                        i += 1;
762                        break;
763                    } else if p.as_char() == '<' || p.as_char() == '(' {
764                        depth += 1;
765                        type_str.push(p.as_char());
766                    } else if p.as_char() == '>' || p.as_char() == ')' {
767                        depth -= 1;
768                        type_str.push(p.as_char());
769                    } else {
770                        type_str.push(p.as_char());
771                    }
772                }
773                TokenTree::Ident(id) => {
774                    if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
775                    {
776                        type_str.push(' ');
777                    }
778                    type_str.push_str(&id.to_string());
779                }
780                _ => {}
781            }
782            i += 1;
783        }
784
785        cols.push((col_name, type_str.trim().to_string()));
786    }
787    Ok(cols)
788}
789
790/// 解析 `SELECT col1, col2 FROM table WHERE col = ?` 表达式
791///
792/// 校验列名是否在表 schema 中(通过编译期常量查找)。
793fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
794    // 收集所有 ident 与 literal,构造 SQL 字符串
795    let mut sql_parts: Vec<String> = Vec::new();
796    let mut table_name: Option<String> = None;
797    let mut in_from = false;
798
799    for (i, t) in tokens.iter().enumerate() {
800        match t {
801            TokenTree::Ident(id) => {
802                let s = id.to_string();
803                if s.eq_ignore_ascii_case("SELECT") {
804                    sql_parts.push("SELECT".to_string());
805                } else if s.eq_ignore_ascii_case("FROM") {
806                    in_from = true;
807                    sql_parts.push("FROM".to_string());
808                } else if s.eq_ignore_ascii_case("WHERE")
809                    || s.eq_ignore_ascii_case("AND")
810                    || s.eq_ignore_ascii_case("OR")
811                    || s.eq_ignore_ascii_case("LIMIT")
812                    || s.eq_ignore_ascii_case("OFFSET")
813                    || s.eq_ignore_ascii_case("ORDER")
814                    || s.eq_ignore_ascii_case("BY")
815                    || s.eq_ignore_ascii_case("GROUP")
816                    || s.eq_ignore_ascii_case("HAVING")
817                    || s.eq_ignore_ascii_case("JOIN")
818                    || s.eq_ignore_ascii_case("INNER")
819                    || s.eq_ignore_ascii_case("LEFT")
820                    || s.eq_ignore_ascii_case("RIGHT")
821                    || s.eq_ignore_ascii_case("ON")
822                    || s.eq_ignore_ascii_case("AS")
823                    || s.eq_ignore_ascii_case("ASC")
824                    || s.eq_ignore_ascii_case("DESC")
825                    || s.eq_ignore_ascii_case("DISTINCT")
826                    || s.eq_ignore_ascii_case("NOT")
827                    || s.eq_ignore_ascii_case("NULL")
828                    || s.eq_ignore_ascii_case("IN")
829                    || s.eq_ignore_ascii_case("BETWEEN")
830                    || s.eq_ignore_ascii_case("LIKE")
831                    || s.eq_ignore_ascii_case("IS")
832                {
833                    sql_parts.push(s.to_uppercase());
834                } else if in_from && table_name.is_none() {
835                    // FROM 后第一个 ident 是表名
836                    table_name = Some(s.clone());
837                    sql_parts.push(s.clone());
838                } else {
839                    sql_parts.push(s.clone());
840                }
841            }
842            TokenTree::Literal(lit) => {
843                sql_parts.push(lit.to_string());
844            }
845            TokenTree::Punct(p) => {
846                let c = p.as_char();
847                // SQL 中常见标点:, ; * ? = > < ( ) . 等
848                let part = if c == ',' {
849                    ",".to_string()
850                } else if c == '?' {
851                    "?".to_string()
852                } else if c == '*' {
853                    "*".to_string()
854                } else if c == '=' {
855                    "=".to_string()
856                } else if c == '>' {
857                    ">".to_string()
858                } else if c == '<' {
859                    "<".to_string()
860                } else if c == '.' {
861                    ".".to_string()
862                } else if c == ';' {
863                    ";".to_string()
864                } else {
865                    c.to_string()
866                };
867                sql_parts.push(part);
868            }
869            TokenTree::Group(g) => {
870                // 处理 group(如 (1, 2, 3))
871                let inner: String = g.stream().to_string();
872                let delim = match g.delimiter() {
873                    Delimiter::Parenthesis => "(",
874                    Delimiter::Brace => "{",
875                    Delimiter::Bracket => "[",
876                    Delimiter::None => "",
877                };
878                let close = match g.delimiter() {
879                    Delimiter::Parenthesis => ")",
880                    Delimiter::Brace => "}",
881                    Delimiter::Bracket => "]",
882                    Delimiter::None => "",
883                };
884                sql_parts.push(format!("{}{}{}", delim, inner, close));
885            }
886        }
887        // 单空格分隔(去重多个空格由 trim 处理)
888        let _ = i;
889    }
890
891    let sql = sql_parts
892        .join(" ")
893        .replace(", ", ",")
894        .replace(" ,", ",")
895        .replace("= ", "=")
896        .replace(" =", "=")
897        .replace("> ", ">")
898        .replace(" >", ">")
899        .replace("< ", "<")
900        .replace(" <", "<")
901        .replace("  ", " ");
902
903    // 验证 SQL 语法
904    if let Err(e) = validate_sql_content(&sql, None) {
905        return compile_error(
906            Span::call_site(),
907            &format!("typed_query! SQL validation failed: {}", e),
908        );
909    }
910
911    // 生成 SQL 字符串字面量
912    let mut ts = TokenStream::new();
913    let lit = Literal::string(&sql);
914    ts.extend([TokenTree::Literal(lit)]);
915    ts
916}
917
918// ---------------------------------------------------------------------------
919// schema! — Compile-time SQL schema generator
920// ---------------------------------------------------------------------------
921
922/// Compile-time SQL schema generator.
923///
924/// Parses a SQL `CREATE TABLE` statement and generates typed table declarations
925/// equivalent to `typed_query! { table ... }`.
926///
927/// # Syntax
928///
929/// ```ignore
930/// use sz_orm_macros::schema;
931///
932/// schema! {
933///     "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL, email TEXT)"
934/// }
935/// ```
936///
937/// 生成与以下手动声明等价的代码:
938/// ```ignore
939/// typed_query! {
940///     table users {
941///         id: i64,
942///         name: String,
943///         email: Option<String>,
944///     }
945/// }
946/// ```
947#[proc_macro]
948pub fn schema(input: TokenStream) -> TokenStream {
949    let mut tokens = input.into_iter().peekable();
950
951    // 解析 SQL 字符串字面量
952    let sql_raw = match tokens.next() {
953        Some(TokenTree::Literal(lit)) => lit.to_string(),
954        Some(other) => {
955            return compile_error(
956                other.span(),
957                "Expected a string literal as the argument to schema!",
958            );
959        }
960        None => {
961            return compile_error(
962                Span::call_site(),
963                "Expected a string literal argument to schema!",
964            );
965        }
966    };
967
968    let sql = match strip_string_literal(&sql_raw) {
969        Some(s) => s,
970        None => {
971            return compile_error(
972                Span::call_site(),
973                "schema! requires a string literal argument",
974            );
975        }
976    };
977
978    // 解析 CREATE TABLE
979    let (table_name, columns) = match parse_create_table(sql) {
980        Ok(v) => v,
981        Err(e) => return compile_error(Span::call_site(), &e),
982    };
983
984    // 生成代码(与 parse_table_decl 一致)
985    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
986    let table_name_lit = table_name.as_str();
987
988    let col_impls: Vec<TokenStream2> = columns
989        .iter()
990        .map(|(col_name, col_type)| {
991            let col_ident =
992                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
993            let col_name_lit = col_name.as_str();
994            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
995            quote! {
996                #[derive(Debug, Clone, Copy)]
997                pub struct #col_ident;
998                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
999                    const NAME: &'static str = #col_name_lit;
1000                    type Table = table;
1001                    type RustType = #rust_type;
1002                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1003                }
1004            }
1005        })
1006        .collect();
1007
1008    let schema_entries: Vec<TokenStream2> = columns
1009        .iter()
1010        .map(|(n, t)| {
1011            let n_lit = n.as_str();
1012            let t_lit = t.as_str();
1013            quote! { (#n_lit, #t_lit) }
1014        })
1015        .collect();
1016
1017    let schema_const_ident = proc_macro2::Ident::new(
1018        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1019        Span::call_site().into(),
1020    );
1021
1022    let expanded = quote! {
1023        pub mod #table_ident {
1024            use super::*;
1025            pub struct table;
1026            impl ::sz_orm_core::typed::TypedTable for table {
1027                const NAME: &'static str = #table_name_lit;
1028            }
1029            #(#col_impls)*
1030        }
1031        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1032    };
1033
1034    expanded.into()
1035}
1036
1037/// 解析 SQL `CREATE TABLE` 语句,返回 (表名, Vec<(列名, Rust 类型字符串)>)。
1038///
1039/// 支持以下语法:
1040/// - `CREATE TABLE [IF NOT EXISTS] <name> ( ... )`
1041/// - 表名/列名可带反引号、双引号或无引号
1042/// - 跳过 PRIMARY KEY / FOREIGN KEY / CONSTRAINT / UNIQUE / INDEX / KEY 约束行
1043/// - 列定义按顶层逗号分隔(嵌套括号如 DECIMAL(10,2) 不拆分)
1044fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
1045    let trimmed = sql.trim();
1046    let upper = trimmed.to_uppercase();
1047
1048    // 必须以 CREATE TABLE 开头
1049    if !upper.starts_with("CREATE TABLE") {
1050        return Err("schema! expects a CREATE TABLE statement".to_string());
1051    }
1052
1053    // 跳过 "CREATE TABLE"
1054    let mut rest = &trimmed["CREATE TABLE".len()..];
1055
1056    // 跳过可选的 "IF NOT EXISTS"
1057    let rest_upper = rest.trim_start().to_uppercase();
1058    if rest_upper.starts_with("IF NOT EXISTS") {
1059        rest = &rest.trim_start()["IF NOT EXISTS".len()..];
1060    }
1061
1062    rest = rest.trim_start();
1063
1064    // 解析表名(可能带反引号、双引号或无引号)
1065    let (table_name, after_name) = parse_identifier(rest)?;
1066    let rest = after_name.trim_start();
1067
1068    // 找到列定义起始的 '(' 与匹配的最后一个 ')'
1069    let paren_start = rest
1070        .find('(')
1071        .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
1072    let paren_end = rest
1073        .rfind(')')
1074        .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
1075    if paren_end <= paren_start {
1076        return Err("CREATE TABLE has malformed parentheses".to_string());
1077    }
1078
1079    let cols_str = &rest[paren_start + 1..paren_end];
1080
1081    // 按顶层逗号分隔列定义(注意嵌套括号,如 DECIMAL(10,2))
1082    let col_defs = split_top_level_commas(cols_str);
1083
1084    let mut columns = Vec::new();
1085    for def in col_defs {
1086        let def = def.trim();
1087        if def.is_empty() {
1088            continue;
1089        }
1090
1091        // 跳过约束定义行
1092        let def_upper = def.to_uppercase();
1093        if def_upper.starts_with("PRIMARY KEY")
1094            || def_upper.starts_with("FOREIGN KEY")
1095            || def_upper.starts_with("CONSTRAINT")
1096            || def_upper.starts_with("UNIQUE")
1097            || def_upper.starts_with("INDEX")
1098            || def_upper.starts_with("KEY")
1099        {
1100            continue;
1101        }
1102
1103        // 解析列名
1104        let (col_name, after_col) = parse_identifier(def)?;
1105        let rest = after_col.trim_start();
1106
1107        // 解析类型(取第一个 token,去掉括号参数)
1108        let (sql_type, after_type) = parse_type_token(rest)?;
1109        let rest = after_type.trim();
1110
1111        // 判断 nullability:NOT NULL 或 PRIMARY KEY 隐含 NOT NULL
1112        let rest_upper = rest.to_uppercase();
1113        let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
1114        let rust_type = sql_type_to_rust(&sql_type, !not_null);
1115
1116        columns.push((col_name, rust_type));
1117    }
1118
1119    Ok((table_name, columns))
1120}
1121
1122/// 解析标识符:支持反引号、双引号或无引号。
1123/// 返回 (标识符, 剩余字符串)。
1124fn parse_identifier(s: &str) -> Result<(String, &str), String> {
1125    let s = s.trim_start();
1126    if s.is_empty() {
1127        return Err("expected identifier".to_string());
1128    }
1129
1130    let bytes = s.as_bytes();
1131    match bytes[0] {
1132        b'`' => {
1133            let end = s[1..]
1134                .find('`')
1135                .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
1136            let ident = s[1..1 + end].to_string();
1137            Ok((ident, &s[1 + end + 1..]))
1138        }
1139        b'"' => {
1140            let end = s[1..]
1141                .find('"')
1142                .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
1143            let ident = s[1..1 + end].to_string();
1144            Ok((ident, &s[1 + end + 1..]))
1145        }
1146        _ => {
1147            let end = s
1148                .find(|c: char| !c.is_alphanumeric() && c != '_')
1149                .unwrap_or(s.len());
1150            if end == 0 {
1151                return Err(format!("invalid identifier: '{}'", s));
1152            }
1153            let ident = s[..end].to_string();
1154            Ok((ident, &s[end..]))
1155        }
1156    }
1157}
1158
1159/// 解析类型 token:取第一个标识符,可选跟随括号参数(如 VARCHAR(255) → VARCHAR)。
1160/// 返回 (类型名, 剩余字符串)。
1161fn parse_type_token(s: &str) -> Result<(String, &str), String> {
1162    let s = s.trim_start();
1163    if s.is_empty() {
1164        return Err("expected column type".to_string());
1165    }
1166
1167    let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
1168    if end == 0 {
1169        return Err(format!("invalid type: '{}'", s));
1170    }
1171    let type_name = s[..end].to_string();
1172    let mut rest = &s[end..];
1173
1174    // 跳过可选的括号参数,如 (255) 或 (10,2)
1175    rest = rest.trim_start();
1176    if rest.starts_with('(') {
1177        let close = rest
1178            .find(')')
1179            .ok_or_else(|| "unterminated type parameter list".to_string())?;
1180        rest = &rest[close + 1..];
1181    }
1182
1183    Ok((type_name, rest))
1184}
1185
1186/// 按顶层逗号分隔字符串(不进入嵌套括号)。
1187fn split_top_level_commas(s: &str) -> Vec<String> {
1188    let mut parts = Vec::new();
1189    let mut depth: i32 = 0;
1190    let mut current = String::new();
1191
1192    for ch in s.chars() {
1193        match ch {
1194            '(' => {
1195                depth += 1;
1196                current.push(ch);
1197            }
1198            ')' => {
1199                depth -= 1;
1200                current.push(ch);
1201            }
1202            ',' if depth == 0 => {
1203                parts.push(std::mem::take(&mut current));
1204            }
1205            _ => {
1206                current.push(ch);
1207            }
1208        }
1209    }
1210
1211    if !current.trim().is_empty() {
1212        parts.push(current);
1213    }
1214
1215    parts
1216}
1217
1218/// 将 SQL 类型映射为 Rust 类型字符串。
1219///
1220/// 匹配规则:取类型名第一个 token(去掉括号参数),不区分大小写匹配。
1221/// 未识别的类型默认映射为 `String`。若 `nullable == true`,用 `Option<T>` 包裹。
1222fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
1223    let upper = sql_type.to_uppercase();
1224    let rust = match upper.as_str() {
1225        // 8 字节整数
1226        "BIGINT" | "INT8" => "i64",
1227        // 4 字节整数(INT/INTEGER/INT4/SERIAL)
1228        "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
1229        // 2 字节整数
1230        "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
1231        // 1 字节整数
1232        "TINYINT" => "i8",
1233        // 浮点(4 字节)
1234        "FLOAT" | "REAL" | "FLOAT4" => "f32",
1235        // 浮点(8 字节)/ 定点数
1236        "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
1237        // 布尔
1238        "BOOLEAN" | "BOOL" => "bool",
1239        // 二进制(与 schema_gen::sql_type_to_rust 保持一致)
1240        "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
1241        // 字符串/日期/JSON/UUID(统一映射到 String,运行时再解析)
1242        "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
1243        | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
1244        _ => "String",
1245    };
1246
1247    if nullable {
1248        format!("Option<{}>", rust)
1249    } else {
1250        rust.to_string()
1251    }
1252}
1253
1254// ---------------------------------------------------------------------------
1255// `#[derive(Schema)]` — auto-generate table structure from a struct
1256// ---------------------------------------------------------------------------
1257
1258/// 派生宏:自动从 Rust 结构体生成表结构信息。
1259///
1260/// 解析 `#[table(name = "...")]` 和 `#[column(...)]` 属性,
1261/// 生成 `Schema` trait 实现,便于在运行时反射表名与列信息。
1262///
1263/// # 支持的属性
1264///
1265/// - `#[table(name = "users")]` — 指定表名(默认使用结构体名的蛇形形式)
1266/// - `#[column(name = "user_id")]` — 指定列名(默认使用字段名)
1267/// - `#[column(type = "VARCHAR(255)")]` — 指定 SQL 类型
1268/// - `#[column(primary_key)]` — 标记主键
1269/// - `#[column(nullable)]` — 显式标记允许 NULL
1270/// - `#[column(skip)]` — 跳过此字段,不生成 schema 条目
1271/// - `#[column(default = "0")]` — 标记字段有默认值
1272///
1273/// # 类型推断
1274///
1275/// 字段的 Rust 类型会自动映射为 SQL 类型:
1276/// - `i64`/`u64` → `BIGINT`
1277/// - `i32`/`u32` → `INTEGER`
1278/// - `String` → `TEXT`
1279/// - `f64` → `DOUBLE`
1280/// - `bool` → `BOOLEAN`
1281/// - `Vec<u8>` → `BLOB`
1282/// - `Option<T>` → 与 `T` 相同,但标记为 nullable
1283#[proc_macro_derive(Schema, attributes(table, column))]
1284pub fn derive_schema(input: TokenStream) -> TokenStream {
1285    let input = parse_macro_input!(input as syn::DeriveInput);
1286    derive::derive_schema_impl(input).into()
1287}
1288
1289// ---------------------------------------------------------------------------
1290// `#[derive(Builder)]` — auto-generate builder pattern code
1291// ---------------------------------------------------------------------------
1292
1293/// 派生宏:自动生成构造器模式代码。
1294///
1295/// 为目标结构体生成一个 `XxxBuilder` 类型,包含:
1296/// - `new()` 构造空 builder
1297/// - 每个字段的 setter 方法
1298/// - `build()` 方法返回 `Result<T, String>`
1299///
1300/// # 支持的属性
1301///
1302/// - `#[builder(skip)]` — 跳过此字段(不生成 setter,使用 Default)
1303/// - `#[builder(default = expr)]` — 指定默认值表达式
1304///
1305/// # 示例
1306///
1307/// ```ignore
1308/// use sz_orm_macros::Builder;
1309///
1310/// #[derive(Builder)]
1311/// struct User {
1312///     id: i64,
1313///     name: String,
1314/// }
1315///
1316/// let user = User::builder()
1317///     .id(1)
1318///     .name("Alice".to_string())
1319///     .build()
1320///     .unwrap();
1321/// ```
1322#[proc_macro_derive(Builder, attributes(builder))]
1323pub fn derive_builder(input: TokenStream) -> TokenStream {
1324    let input = parse_macro_input!(input as syn::DeriveInput);
1325    derive::derive_builder_impl(input).into()
1326}
1327
1328// ---------------------------------------------------------------------------
1329// Unit tests — cover helper functions used by both macros
1330// ---------------------------------------------------------------------------
1331
1332#[cfg(test)]
1333mod tests {
1334    use super::*;
1335
1336    // ---- strip_string_literal ----
1337
1338    #[test]
1339    fn test_strip_plain_double_quoted() {
1340        assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
1341    }
1342
1343    #[test]
1344    fn test_strip_raw_double_hash() {
1345        assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
1346    }
1347
1348    #[test]
1349    fn test_strip_raw_double_no_hash() {
1350        assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
1351    }
1352
1353    #[test]
1354    fn test_strip_byte_string() {
1355        assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
1356        assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
1357    }
1358
1359    #[test]
1360    fn test_strip_non_string_returns_none() {
1361        assert_eq!(strip_string_literal("123"), None);
1362        assert_eq!(strip_string_literal("foo"), None);
1363    }
1364
1365    // ---- validate_sql_content ----
1366
1367    #[test]
1368    fn test_validate_select_with_from_ok() {
1369        assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
1370    }
1371
1372    #[test]
1373    fn test_validate_select_missing_from_fails() {
1374        assert!(validate_sql_content("SELECT * users", None).is_err());
1375    }
1376
1377    #[test]
1378    fn test_validate_insert_missing_into_fails() {
1379        assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
1380        assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
1381    }
1382
1383    #[test]
1384    fn test_validate_update_missing_set_fails() {
1385        assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
1386        assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
1387    }
1388
1389    #[test]
1390    fn test_validate_delete_missing_from_fails() {
1391        assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
1392        assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
1393    }
1394
1395    #[test]
1396    fn test_validate_empty_sql_fails() {
1397        assert!(validate_sql_content("", None).is_err());
1398        assert!(validate_sql_content("   ", None).is_err());
1399    }
1400
1401    // ---- balanced parens ----
1402
1403    #[test]
1404    fn test_validate_balanced_parens_ok() {
1405        assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
1406    }
1407
1408    #[test]
1409    fn test_validate_balanced_parens_unbalanced() {
1410        assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
1411        assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
1412    }
1413
1414    // ---- injection patterns ----
1415
1416    #[test]
1417    fn test_validate_no_injection_clean() {
1418        assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
1419    }
1420
1421    #[test]
1422    fn test_validate_no_injection_drop_table() {
1423        assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
1424    }
1425
1426    #[test]
1427    fn test_validate_no_injection_or_1_1() {
1428        // 编译期 SQL 已剥离外层引号,检测模式不再依赖引号字符。
1429        // "' OR '1'='1" 因引号分隔不再匹配 "or 1=1",故不再检测;
1430        // 但不含引号分隔的 "OR 1=1" 仍可被检测。
1431        assert!(validate_no_injection("' OR 1=1").is_err());
1432        assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
1433    }
1434
1435    #[test]
1436    fn test_validate_no_injection_drop_database() {
1437        assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
1438    }
1439
1440    #[test]
1441    fn test_validate_no_injection_information_schema() {
1442        assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
1443    }
1444
1445    #[test]
1446    fn test_validate_no_injection_xp_cmdshell() {
1447        assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
1448    }
1449
1450    #[test]
1451    fn test_validate_no_injection_union_select() {
1452        assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
1453    }
1454
1455    #[test]
1456    fn test_validate_no_injection_comment_dashes() {
1457        assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
1458    }
1459
1460    #[test]
1461    fn test_validate_no_injection_block_comment() {
1462        assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
1463    }
1464
1465    // ---- string literal closure ----
1466
1467    #[test]
1468    fn test_validate_string_literals_closed_ok() {
1469        assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
1470        assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
1471    }
1472
1473    #[test]
1474    fn test_validate_string_literals_closed_unclosed_single() {
1475        assert!(validate_string_literals_closed("'hello").is_err());
1476    }
1477
1478    #[test]
1479    fn test_validate_string_literals_closed_unclosed_double() {
1480        assert!(validate_string_literals_closed(r#""hello"#).is_err());
1481    }
1482
1483    // ---- param count check ----
1484
1485    #[test]
1486    fn test_validate_param_count_match() {
1487        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
1488        assert!(
1489            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
1490        );
1491    }
1492
1493    #[test]
1494    fn test_validate_param_count_mismatch() {
1495        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
1496        assert!(
1497            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
1498        );
1499    }
1500
1501    // ---- db-verify feature: detect_db_kind ----
1502
1503    #[cfg(feature = "db-verify")]
1504    #[test]
1505    fn test_detect_db_kind_mysql() {
1506        assert_eq!(
1507            detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
1508            DbKind::MySql
1509        );
1510    }
1511
1512    #[cfg(feature = "db-verify")]
1513    #[test]
1514    fn test_detect_db_kind_postgres() {
1515        assert_eq!(
1516            detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
1517            DbKind::Postgres
1518        );
1519        assert_eq!(
1520            detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
1521            DbKind::Postgres
1522        );
1523    }
1524
1525    #[cfg(feature = "db-verify")]
1526    #[test]
1527    fn test_detect_db_kind_sqlite() {
1528        assert_eq!(
1529            detect_db_kind("sqlite://path/to/db.db").unwrap(),
1530            DbKind::Sqlite
1531        );
1532        assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
1533    }
1534
1535    #[cfg(feature = "db-verify")]
1536    #[test]
1537    fn test_detect_db_kind_unsupported() {
1538        assert!(detect_db_kind("oracle://user:pass@host/db").is_err());
1539        assert!(detect_db_kind("not-a-url").is_err());
1540    }
1541
1542    // ---- schema! 宏 parse_create_table 测试 ----
1543
1544    #[test]
1545    fn test_parse_create_table_basic() {
1546        let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
1547        let (table, cols) = parse_create_table(sql).unwrap();
1548        assert_eq!(table, "users");
1549        assert_eq!(
1550            cols,
1551            vec![
1552                ("id".to_string(), "i32".to_string()),
1553                ("name".to_string(), "String".to_string())
1554            ]
1555        );
1556    }
1557
1558    #[test]
1559    fn test_parse_create_table_with_if_not_exists() {
1560        let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
1561        let (table, cols) = parse_create_table(sql).unwrap();
1562        assert_eq!(table, "orders");
1563        assert_eq!(
1564            cols,
1565            vec![
1566                ("id".to_string(), "i64".to_string()),
1567                ("total".to_string(), "f64".to_string())
1568            ]
1569        );
1570    }
1571
1572    #[test]
1573    fn test_parse_create_table_nullable() {
1574        let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
1575        let (_, cols) = parse_create_table(sql).unwrap();
1576        assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
1577        assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
1578    }
1579
1580    #[test]
1581    fn test_parse_create_table_skip_constraints() {
1582        let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
1583        let (_, cols) = parse_create_table(sql).unwrap();
1584        assert_eq!(cols.len(), 2);
1585        assert_eq!(cols[0].0, "id");
1586        assert_eq!(cols[1].0, "name");
1587    }
1588
1589    #[test]
1590    fn test_parse_create_table_varchar_with_len() {
1591        let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
1592        let (_, cols) = parse_create_table(sql).unwrap();
1593        assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
1594        assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
1595    }
1596
1597    #[test]
1598    fn test_sql_type_to_rust_mappings() {
1599        // 整数(按字节宽度严格映射,与 SQL 标准一致)
1600        assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
1601        assert_eq!(sql_type_to_rust("INT8", false), "i64");
1602        assert_eq!(sql_type_to_rust("INT", false), "i32");
1603        assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
1604        assert_eq!(sql_type_to_rust("INT4", false), "i32");
1605        assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
1606        assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
1607        assert_eq!(sql_type_to_rust("INT2", false), "i16");
1608        assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
1609        assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
1610        // 浮点
1611        assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
1612        assert_eq!(sql_type_to_rust("REAL", false), "f32");
1613        assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
1614        assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
1615        assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
1616        assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
1617        assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
1618        assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
1619        // 布尔
1620        assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
1621        assert_eq!(sql_type_to_rust("BOOL", false), "bool");
1622        // 字符串
1623        assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
1624        assert_eq!(sql_type_to_rust("TEXT", false), "String");
1625        assert_eq!(sql_type_to_rust("CHAR", false), "String");
1626        assert_eq!(sql_type_to_rust("UUID", false), "String");
1627        assert_eq!(sql_type_to_rust("DATE", false), "String");
1628        assert_eq!(sql_type_to_rust("DATETIME", false), "String");
1629        assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
1630        assert_eq!(sql_type_to_rust("JSON", false), "String");
1631        assert_eq!(sql_type_to_rust("JSONB", false), "String");
1632        // 二进制
1633        assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
1634        assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
1635        assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
1636        assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
1637        // nullable
1638        assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
1639        assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
1640        assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
1641        assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
1642        // unknown
1643        assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
1644    }
1645
1646    #[test]
1647    fn test_parse_create_table_error_no_create() {
1648        assert!(parse_create_table("SELECT * FROM users").is_err());
1649    }
1650
1651    #[test]
1652    fn test_parse_create_table_error_no_parens() {
1653        assert!(parse_create_table("CREATE TABLE foo").is_err());
1654    }
1655}