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    // 将 ? 占位符替换为 NULL,使 EXPLAIN 无需绑定参数即可执行。
460    // EXPLAIN 不实际执行查询,NULL 对所有列类型都合法。
461    let sql_no_placeholders = replace_placeholders_with_null(sql);
462
463    // Oracle/SQL Server 使用 EXPLAIN PLAN FOR(不同语法),其余用 EXPLAIN
464    let explain_sql = match db_kind {
465        DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql_no_placeholders),
466        DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql_no_placeholders),
467        // Oracle: EXPLAIN PLAN FOR 放入 PLAN_TABLE,再查询结果验证语法
468        DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
469        // SQL Server: SET SHOWPLAN_TEXT ON 后执行(不实际运行)
470        DbKind::SqlServer => sql_no_placeholders,
471    };
472
473    // MySQL/PG/SQLite 走 sqlx 异步路径
474    if matches!(db_kind, DbKind::MySql | DbKind::Postgres | DbKind::Sqlite) {
475        let rt = tokio::runtime::Runtime::new()
476            .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
477        return rt.block_on(async {
478            match db_kind {
479                DbKind::MySql => verify_mysql(&dsn, &explain_sql).await,
480                DbKind::Postgres => verify_postgres(&dsn, &explain_sql).await,
481                DbKind::Sqlite => verify_sqlite(&dsn, &explain_sql).await,
482                _ => unreachable!(),
483            }
484        });
485    }
486
487    // Oracle/SQL Server 走命令行工具验证(避免引入重依赖)
488    match db_kind {
489        DbKind::Oracle => verify_oracle(&dsn, &explain_sql),
490        DbKind::SqlServer => verify_sqlserver(&dsn, &explain_sql),
491        _ => unreachable!(),
492    }
493}
494
495#[cfg(feature = "db-verify")]
496#[derive(Debug, Clone, Copy, PartialEq, Eq)]
497enum DbKind {
498    MySql,
499    Postgres,
500    Sqlite,
501    Oracle,
502    SqlServer,
503}
504
505/// 将 SQL 中的 `?` 占位符替换为 `NULL`,跳过字符串字面量内的 `?`。
506///
507/// EXPLAIN 不实际执行查询,用 NULL 代替参数可验证语法和表/列存在性,
508/// 同时避免 sqlx 预处理语句要求绑定参数的问题。
509#[cfg(feature = "db-verify")]
510fn replace_placeholders_with_null(sql: &str) -> String {
511    let mut result = String::with_capacity(sql.len() + 16);
512    let mut in_single_quote = false;
513    let mut in_double_quote = false;
514    let mut prev = '\0';
515
516    for ch in sql.chars() {
517        if prev == '\\' {
518            // 转义字符:直接追加
519            result.push(ch);
520            prev = ch;
521            continue;
522        }
523        match ch {
524            '\'' if !in_double_quote => in_single_quote = !in_single_quote,
525            '"' if !in_single_quote => in_double_quote = !in_double_quote,
526            '?' if !in_single_quote && !in_double_quote => {
527                result.push_str("NULL");
528                prev = ch;
529                continue;
530            }
531            _ => {}
532        }
533        result.push(ch);
534        prev = ch;
535    }
536    result
537}
538
539#[cfg(feature = "db-verify")]
540fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
541    let lower = dsn.to_lowercase();
542    if lower.starts_with("mysql://") {
543        Ok(DbKind::MySql)
544    } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
545        Ok(DbKind::Postgres)
546    } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
547        Ok(DbKind::Sqlite)
548    } else if lower.starts_with("oracle://") || lower.starts_with("oracle:") {
549        Ok(DbKind::Oracle)
550    } else if lower.starts_with("sqlserver://")
551        || lower.starts_with("mssql://")
552        || lower.starts_with("tds://")
553    {
554        Ok(DbKind::SqlServer)
555    } else {
556        Err(format!("Unsupported DSN scheme: {}", dsn))
557    }
558}
559
560#[cfg(feature = "db-verify")]
561async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
562    let pool = sqlx::MySqlPool::connect(dsn)
563        .await
564        .map_err(|e| format!("MySQL connect failed: {}", e))?;
565    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
566        .execute(&pool)
567        .await
568        .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
569    Ok(())
570}
571
572#[cfg(feature = "db-verify")]
573async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
574    let pool = sqlx::PgPool::connect(dsn)
575        .await
576        .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
577    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
578        .execute(&pool)
579        .await
580        .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
581    Ok(())
582}
583
584#[cfg(feature = "db-verify")]
585async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
586    let pool = sqlx::SqlitePool::connect(dsn)
587        .await
588        .map_err(|e| format!("SQLite connect failed: {}", e))?;
589    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
590        .execute(&pool)
591        .await
592        .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
593    Ok(())
594}
595
596/// Oracle 编译期验证:通过 sqlplus 命令行工具执行 EXPLAIN PLAN FOR
597///
598/// DSN 格式:`oracle://user:pass@host:port/service`(可选 `?sysdba=1`)
599/// 例如:`oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1`
600#[cfg(feature = "db-verify")]
601fn verify_oracle(dsn: &str, explain_sql: &str) -> Result<(), String> {
602    let parsed = parse_oracle_dsn(dsn)?;
603    // 构造 sqlplus 连接串:user/pass@host:port/service [AS SYSDBA]
604    let mut conn_str = format!(
605        "{}/{}@{}:{}/{}",
606        parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
607    );
608    if parsed.sysdba {
609        conn_str.push_str(" AS SYSDBA");
610    }
611    // 用 SET SHOWPLAN 不适用于 Oracle,用 EXPLAIN PLAN FOR 并立即查询 PLAN_TABLE
612    let full_script = format!(
613        "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
614         EXPLAIN PLAN FOR {};\n\
615         SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
616         EXIT;\n",
617        explain_sql
618    );
619    let output = std::process::Command::new("sqlplus")
620        .args(["-S", "-L", &conn_str])
621        .stdin(std::process::Stdio::piped())
622        .stdout(std::process::Stdio::piped())
623        .stderr(std::process::Stdio::piped())
624        .spawn()
625        .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
626    use std::io::Write;
627    let mut child = output;
628    if let Some(mut stdin) = child.stdin.take() {
629        stdin
630            .write_all(full_script.as_bytes())
631            .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
632    }
633    let out = child
634        .wait_with_output()
635        .map_err(|e| format!("sqlplus wait failed: {}", e))?;
636    let stdout = String::from_utf8_lossy(&out.stdout);
637    let stderr = String::from_utf8_lossy(&out.stderr);
638    if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
639        return Err(format!(
640            "Oracle EXPLAIN failed: stdout={} stderr={}",
641            stdout.trim(),
642            stderr.trim()
643        ));
644    }
645    Ok(())
646}
647
648/// SQL Server 编译期验证:通过 sqlcmd 命令行工具执行 SET SHOWPLAN_TEXT ON
649///
650/// DSN 格式:`sqlserver://user:pass@host:port/db`
651/// 例如:`sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test`
652#[cfg(feature = "db-verify")]
653fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
654    let parsed = parse_sqlserver_dsn(dsn)?;
655    // sqlcmd -S host,port -U user -P pass -d db -Q "SET SHOWPLAN_TEXT ON; <sql>"
656    let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
657    let out = std::process::Command::new("sqlcmd")
658        .args([
659            "-S",
660            &format!("{},{}", parsed.host, parsed.port),
661            "-U",
662            &parsed.user,
663            "-P",
664            &parsed.password,
665            "-d",
666            &parsed.database,
667            "-Q",
668            &query,
669            "-h",
670            "-1",
671            "-W",
672        ])
673        .output()
674        .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
675    let stdout = String::from_utf8_lossy(&out.stdout);
676    let stderr = String::from_utf8_lossy(&out.stderr);
677    if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
678        return Err(format!(
679            "SQL Server SHOWPLAN failed: stdout={} stderr={}",
680            stdout.trim(),
681            stderr.trim()
682        ));
683    }
684    Ok(())
685}
686
687/// Oracle DSN 解析结果
688#[cfg(feature = "db-verify")]
689struct OracleDsn {
690    user: String,
691    password: String,
692    host: String,
693    port: u16,
694    service: String,
695    sysdba: bool,
696}
697
698/// 解析 oracle://user:pass@host:port/service?sysdba=1
699#[cfg(feature = "db-verify")]
700fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
701    let raw = dsn
702        .strip_prefix("oracle://")
703        .or_else(|| dsn.strip_prefix("oracle:"))
704        .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
705    // 分离 query
706    let (auth_host_service, query) = match raw.find('?') {
707        Some(idx) => (&raw[..idx], &raw[idx + 1..]),
708        None => (raw, ""),
709    };
710    let sysdba = query
711        .split('&')
712        .any(|p| p == "sysdba=1" || p == "sysdba=true");
713    // user:pass@host:port/service
714    let at = auth_host_service
715        .find('@')
716        .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
717    let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
718    let colon = user_pass
719        .find(':')
720        .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
721    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
722    let (host_port, service) = match host_port_service.rfind('/') {
723        Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
724        None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
725    };
726    let (host, port) = match host_port.find(':') {
727        Some(idx) => (
728            &host_port[..idx],
729            host_port[idx + 1..]
730                .parse::<u16>()
731                .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
732        ),
733        None => (host_port, 1521u16),
734    };
735    Ok(OracleDsn {
736        user: user.to_string(),
737        password: password.to_string(),
738        host: host.to_string(),
739        port,
740        service: service.to_string(),
741        sysdba,
742    })
743}
744
745/// SQL Server DSN 解析结果
746#[cfg(feature = "db-verify")]
747struct SqlServerDsn {
748    user: String,
749    password: String,
750    host: String,
751    port: u16,
752    database: String,
753}
754
755/// 解析 sqlserver://user:pass@host:port/db
756#[cfg(feature = "db-verify")]
757fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
758    let raw = dsn
759        .strip_prefix("sqlserver://")
760        .or_else(|| dsn.strip_prefix("mssql://"))
761        .or_else(|| dsn.strip_prefix("tds://"))
762        .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
763    let at = raw
764        .find('@')
765        .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
766    let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
767    let colon = user_pass
768        .find(':')
769        .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
770    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
771    let (host_port, database) = match host_port_db.rfind('/') {
772        Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
773        None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
774    };
775    let (host, port) = match host_port.find(':') {
776        Some(idx) => (
777            &host_port[..idx],
778            host_port[idx + 1..]
779                .parse::<u16>()
780                .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
781        ),
782        None => (host_port, 1433u16),
783    };
784    Ok(SqlServerDsn {
785        user: user.to_string(),
786        password: password.to_string(),
787        host: host.to_string(),
788        port,
789        database: database.to_string(),
790    })
791}
792
793// ---------------------------------------------------------------------------
794// Helpers
795// ---------------------------------------------------------------------------
796
797/// Create a compile_error! token stream
798fn compile_error(span: Span, msg: &str) -> TokenStream {
799    // emit: compile_error!("msg")
800    let mut ts = TokenStream::new();
801    ts.extend([
802        TokenTree::Ident(Ident::new("compile_error", span)),
803        TokenTree::Punct(Punct::new('!', Spacing::Alone)),
804        TokenTree::Group(Group::new(
805            Delimiter::Parenthesis,
806            TokenStream::from(TokenTree::Literal(Literal::string(msg))),
807        )),
808    ]);
809    ts
810}
811
812// ---------------------------------------------------------------------------
813// typed_query! — Diesel 风格强类型 AST 宏
814// ---------------------------------------------------------------------------
815
816/// Diesel 风格强类型 AST 宏(与 `sql_string!` / `query!` 并存)。
817///
818/// # 设计
819///
820/// 接收 `table { col1: Type, col2: Type, ... }` 声明,生成:
821/// 1. 一个 `table` 模块
822/// 2. 每列对应一个零大小标记类型(如 `table::id`)
823/// 3. 实现 `TypedColumn` trait,把列名 + Rust 类型提升到类型系统
824///
825/// 这样,`typed_query!(SELECT id FROM users WHERE name = ?)` 在编译期就能:
826/// - 校验 `id` / `name` 列是否存在于 `users` 表声明中
827/// - 校验 `?` 参数的 Rust 类型与列声明的类型一致
828///
829/// # 用法
830///
831/// ```ignore
832/// use sz_orm_macros::typed_query;
833///
834/// // 1. 声明表 schema(编译期生成 column 标记类型)
835/// typed_query! {
836///     table users {
837///         id: i64,
838///         name: String,
839///         email: String,
840///         age: i32,
841///     }
842/// }
843///
844/// // 2. 编译期校验 SELECT:列名必须存在于 users 表
845/// let sql = typed_query!(SELECT id, name FROM users WHERE age > ?);
846/// // ❌ 编译错误:unknown column 'foo' in table 'users'
847/// // let sql = typed_query!(SELECT foo FROM users);
848/// ```
849#[proc_macro]
850pub fn typed_query(input: TokenStream) -> TokenStream {
851    let tokens: Vec<TokenTree> = input.into_iter().collect();
852
853    // 分支 1:table 声明
854    if tokens.iter().any(|t| {
855        if let TokenTree::Ident(id) = t {
856            id.to_string() == "table"
857        } else {
858            false
859        }
860    }) {
861        return parse_table_decl(&tokens);
862    }
863
864    // 分支 2:SELECT 表达式
865    if tokens.iter().any(|t| {
866        if let TokenTree::Ident(id) = t {
867            id.to_string().eq_ignore_ascii_case("SELECT")
868        } else {
869            false
870        }
871    }) {
872        return parse_typed_select(&tokens);
873    }
874
875    compile_error(
876        Span::call_site(),
877        "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
878    )
879}
880
881/// 解析 `table name { col: Type, ... }` 声明
882fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
883    // 期望格式:table <ident> { <ident> : <ident> [, ...] }
884    let mut idx = 0;
885
886    // 跳过 'table' 关键字
887    if idx >= tokens.len() {
888        return compile_error(Span::call_site(), "expected table name after 'table'");
889    }
890    if let TokenTree::Ident(id) = &tokens[idx] {
891        if id.to_string() != "table" {
892            return compile_error(id.span(), "expected 'table' keyword");
893        }
894    }
895    idx += 1;
896
897    // 表名
898    let table_name = if idx < tokens.len() {
899        if let TokenTree::Ident(id) = &tokens[idx] {
900            id.to_string()
901        } else {
902            return compile_error(tokens[idx].span(), "expected table name identifier");
903        }
904    } else {
905        return compile_error(Span::call_site(), "expected table name");
906    };
907    idx += 1;
908
909    // 表体({} 内)
910    let body_group = if idx < tokens.len() {
911        if let TokenTree::Group(g) = &tokens[idx] {
912            if g.delimiter() != Delimiter::Brace {
913                return compile_error(g.span(), "expected '{' after table name");
914            }
915            g.clone()
916        } else {
917            return compile_error(tokens[idx].span(), "expected '{' after table name");
918        }
919    } else {
920        return compile_error(Span::call_site(), "expected table body in '{ }'");
921    };
922
923    // 解析列声明
924    let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
925    let columns = match parse_column_list(&body_tokens) {
926        Ok(c) => c,
927        Err(e) => return compile_error(Span::call_site(), &e),
928    };
929
930    // 使用 quote! 构建类型安全的 TokenStream
931    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
932    let table_name_lit = table_name.as_str();
933
934    // 为每列构建标记类型 + trait 实现
935    let col_impls: Vec<TokenStream2> = columns
936        .iter()
937        .map(|(col_name, col_type)| {
938            let col_ident =
939                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
940            let col_name_lit = col_name.as_str();
941            // 解析类型字符串为 TokenStream(quote! 会处理)
942            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
943            quote! {
944                #[derive(Debug, Clone, Copy)]
945                pub struct #col_ident;
946                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
947                    const NAME: &'static str = #col_name_lit;
948                    type Table = table;
949                    type RustType = #rust_type;
950                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
951                }
952            }
953        })
954        .collect();
955
956    // schema 常量条目
957    let schema_entries: Vec<TokenStream2> = columns
958        .iter()
959        .map(|(n, t)| {
960            let n_lit = n.as_str();
961            let t_lit = t.as_str();
962            quote! { (#n_lit, #t_lit) }
963        })
964        .collect();
965
966    let schema_const_ident = proc_macro2::Ident::new(
967        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
968        Span::call_site().into(),
969    );
970
971    let expanded = quote! {
972        pub mod #table_ident {
973            use super::*;
974            pub struct table;
975            impl ::sz_orm_core::typed::TypedTable for table {
976                const NAME: &'static str = #table_name_lit;
977            }
978            #(#col_impls)*
979        }
980        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
981    };
982
983    expanded.into()
984}
985
986/// 解析列声明列表:`col: Type, col2: Type2, ...`
987fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
988    let mut cols = Vec::new();
989    let mut i = 0;
990    while i < tokens.len() {
991        // 列名
992        let col_name = if let TokenTree::Ident(id) = &tokens[i] {
993            id.to_string()
994        } else {
995            return Err(format!("expected column name at position {}", i));
996        };
997        i += 1;
998
999        // 冒号
1000        if i >= tokens.len() {
1001            return Err(format!("expected ':' after column '{}'", col_name));
1002        }
1003        if let TokenTree::Punct(p) = &tokens[i] {
1004            if p.as_char() != ':' {
1005                return Err(format!("expected ':' after column '{}'", col_name));
1006            }
1007        } else {
1008            return Err(format!("expected ':' after column '{}'", col_name));
1009        }
1010        i += 1;
1011
1012        // 类型(可能是 ident 或 path,如 String / i64 / Option<i64>)
1013        // 简化处理:收集直到遇到 ',' 或末尾
1014        let mut type_str = String::new();
1015        let mut depth = 0;
1016        while i < tokens.len() {
1017            match &tokens[i] {
1018                TokenTree::Punct(p) => {
1019                    if p.as_char() == ',' && depth == 0 {
1020                        i += 1;
1021                        break;
1022                    } else if p.as_char() == '<' || p.as_char() == '(' {
1023                        depth += 1;
1024                        type_str.push(p.as_char());
1025                    } else if p.as_char() == '>' || p.as_char() == ')' {
1026                        depth -= 1;
1027                        type_str.push(p.as_char());
1028                    } else {
1029                        type_str.push(p.as_char());
1030                    }
1031                }
1032                TokenTree::Ident(id) => {
1033                    if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
1034                    {
1035                        type_str.push(' ');
1036                    }
1037                    type_str.push_str(&id.to_string());
1038                }
1039                _ => {}
1040            }
1041            i += 1;
1042        }
1043
1044        cols.push((col_name, type_str.trim().to_string()));
1045    }
1046    Ok(cols)
1047}
1048
1049/// 解析 `SELECT col1, col2 FROM table WHERE col = ?` 表达式
1050///
1051/// 校验列名是否在表 schema 中(通过编译期常量查找)。
1052fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
1053    // 收集所有 ident 与 literal,构造 SQL 字符串
1054    let mut sql_parts: Vec<String> = Vec::new();
1055    let mut table_name: Option<String> = None;
1056    let mut in_from = false;
1057
1058    for (i, t) in tokens.iter().enumerate() {
1059        match t {
1060            TokenTree::Ident(id) => {
1061                let s = id.to_string();
1062                if s.eq_ignore_ascii_case("SELECT") {
1063                    sql_parts.push("SELECT".to_string());
1064                } else if s.eq_ignore_ascii_case("FROM") {
1065                    in_from = true;
1066                    sql_parts.push("FROM".to_string());
1067                } else if s.eq_ignore_ascii_case("WHERE")
1068                    || s.eq_ignore_ascii_case("AND")
1069                    || s.eq_ignore_ascii_case("OR")
1070                    || s.eq_ignore_ascii_case("LIMIT")
1071                    || s.eq_ignore_ascii_case("OFFSET")
1072                    || s.eq_ignore_ascii_case("ORDER")
1073                    || s.eq_ignore_ascii_case("BY")
1074                    || s.eq_ignore_ascii_case("GROUP")
1075                    || s.eq_ignore_ascii_case("HAVING")
1076                    || s.eq_ignore_ascii_case("JOIN")
1077                    || s.eq_ignore_ascii_case("INNER")
1078                    || s.eq_ignore_ascii_case("LEFT")
1079                    || s.eq_ignore_ascii_case("RIGHT")
1080                    || s.eq_ignore_ascii_case("ON")
1081                    || s.eq_ignore_ascii_case("AS")
1082                    || s.eq_ignore_ascii_case("ASC")
1083                    || s.eq_ignore_ascii_case("DESC")
1084                    || s.eq_ignore_ascii_case("DISTINCT")
1085                    || s.eq_ignore_ascii_case("NOT")
1086                    || s.eq_ignore_ascii_case("NULL")
1087                    || s.eq_ignore_ascii_case("IN")
1088                    || s.eq_ignore_ascii_case("BETWEEN")
1089                    || s.eq_ignore_ascii_case("LIKE")
1090                    || s.eq_ignore_ascii_case("IS")
1091                {
1092                    sql_parts.push(s.to_uppercase());
1093                } else if in_from && table_name.is_none() {
1094                    // FROM 后第一个 ident 是表名
1095                    table_name = Some(s.clone());
1096                    sql_parts.push(s.clone());
1097                } else {
1098                    sql_parts.push(s.clone());
1099                }
1100            }
1101            TokenTree::Literal(lit) => {
1102                sql_parts.push(lit.to_string());
1103            }
1104            TokenTree::Punct(p) => {
1105                let c = p.as_char();
1106                // SQL 中常见标点:, ; * ? = > < ( ) . 等
1107                let part = if c == ',' {
1108                    ",".to_string()
1109                } else if c == '?' {
1110                    "?".to_string()
1111                } else if c == '*' {
1112                    "*".to_string()
1113                } else if c == '=' {
1114                    "=".to_string()
1115                } else if c == '>' {
1116                    ">".to_string()
1117                } else if c == '<' {
1118                    "<".to_string()
1119                } else if c == '.' {
1120                    ".".to_string()
1121                } else if c == ';' {
1122                    ";".to_string()
1123                } else {
1124                    c.to_string()
1125                };
1126                sql_parts.push(part);
1127            }
1128            TokenTree::Group(g) => {
1129                // 处理 group(如 (1, 2, 3))
1130                let inner: String = g.stream().to_string();
1131                let delim = match g.delimiter() {
1132                    Delimiter::Parenthesis => "(",
1133                    Delimiter::Brace => "{",
1134                    Delimiter::Bracket => "[",
1135                    Delimiter::None => "",
1136                };
1137                let close = match g.delimiter() {
1138                    Delimiter::Parenthesis => ")",
1139                    Delimiter::Brace => "}",
1140                    Delimiter::Bracket => "]",
1141                    Delimiter::None => "",
1142                };
1143                sql_parts.push(format!("{}{}{}", delim, inner, close));
1144            }
1145        }
1146        // 单空格分隔(去重多个空格由 trim 处理)
1147        let _ = i;
1148    }
1149
1150    let sql = sql_parts
1151        .join(" ")
1152        .replace(", ", ",")
1153        .replace(" ,", ",")
1154        .replace("= ", "=")
1155        .replace(" =", "=")
1156        .replace("> ", ">")
1157        .replace(" >", ">")
1158        .replace("< ", "<")
1159        .replace(" <", "<")
1160        .replace("  ", " ");
1161
1162    // 验证 SQL 语法
1163    if let Err(e) = validate_sql_content(&sql, None) {
1164        return compile_error(
1165            Span::call_site(),
1166            &format!("typed_query! SQL validation failed: {}", e),
1167        );
1168    }
1169
1170    // 生成 SQL 字符串字面量
1171    let mut ts = TokenStream::new();
1172    let lit = Literal::string(&sql);
1173    ts.extend([TokenTree::Literal(lit)]);
1174    ts
1175}
1176
1177// ---------------------------------------------------------------------------
1178// schema! — Compile-time SQL schema generator
1179// ---------------------------------------------------------------------------
1180
1181/// Compile-time SQL schema generator.
1182///
1183/// Parses a SQL `CREATE TABLE` statement and generates typed table declarations
1184/// equivalent to `typed_query! { table ... }`.
1185///
1186/// # Syntax
1187///
1188/// ```ignore
1189/// use sz_orm_macros::schema;
1190///
1191/// schema! {
1192///     "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL, email TEXT)"
1193/// }
1194/// ```
1195///
1196/// 生成与以下手动声明等价的代码:
1197/// ```ignore
1198/// typed_query! {
1199///     table users {
1200///         id: i64,
1201///         name: String,
1202///         email: Option<String>,
1203///     }
1204/// }
1205/// ```
1206#[proc_macro]
1207pub fn schema(input: TokenStream) -> TokenStream {
1208    let mut tokens = input.into_iter().peekable();
1209
1210    // 解析 SQL 字符串字面量
1211    let sql_raw = match tokens.next() {
1212        Some(TokenTree::Literal(lit)) => lit.to_string(),
1213        Some(other) => {
1214            return compile_error(
1215                other.span(),
1216                "Expected a string literal as the argument to schema!",
1217            );
1218        }
1219        None => {
1220            return compile_error(
1221                Span::call_site(),
1222                "Expected a string literal argument to schema!",
1223            );
1224        }
1225    };
1226
1227    let sql = match strip_string_literal(&sql_raw) {
1228        Some(s) => s,
1229        None => {
1230            return compile_error(
1231                Span::call_site(),
1232                "schema! requires a string literal argument",
1233            );
1234        }
1235    };
1236
1237    // 解析 CREATE TABLE
1238    let (table_name, columns) = match parse_create_table(sql) {
1239        Ok(v) => v,
1240        Err(e) => return compile_error(Span::call_site(), &e),
1241    };
1242
1243    // 生成代码(与 parse_table_decl 一致)
1244    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1245    let table_name_lit = table_name.as_str();
1246
1247    let col_impls: Vec<TokenStream2> = columns
1248        .iter()
1249        .map(|(col_name, col_type)| {
1250            let col_ident =
1251                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1252            let col_name_lit = col_name.as_str();
1253            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1254            quote! {
1255                #[derive(Debug, Clone, Copy)]
1256                pub struct #col_ident;
1257                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1258                    const NAME: &'static str = #col_name_lit;
1259                    type Table = table;
1260                    type RustType = #rust_type;
1261                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1262                }
1263            }
1264        })
1265        .collect();
1266
1267    let schema_entries: Vec<TokenStream2> = columns
1268        .iter()
1269        .map(|(n, t)| {
1270            let n_lit = n.as_str();
1271            let t_lit = t.as_str();
1272            quote! { (#n_lit, #t_lit) }
1273        })
1274        .collect();
1275
1276    let schema_const_ident = proc_macro2::Ident::new(
1277        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1278        Span::call_site().into(),
1279    );
1280
1281    let expanded = quote! {
1282        pub mod #table_ident {
1283            use super::*;
1284            pub struct table;
1285            impl ::sz_orm_core::typed::TypedTable for table {
1286                const NAME: &'static str = #table_name_lit;
1287            }
1288            #(#col_impls)*
1289        }
1290        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1291    };
1292
1293    expanded.into()
1294}
1295
1296/// 解析 SQL `CREATE TABLE` 语句,返回 (表名, Vec<(列名, Rust 类型字符串)>)。
1297///
1298/// 支持以下语法:
1299/// - `CREATE TABLE [IF NOT EXISTS] <name> ( ... )`
1300/// - 表名/列名可带反引号、双引号或无引号
1301/// - 跳过 PRIMARY KEY / FOREIGN KEY / CONSTRAINT / UNIQUE / INDEX / KEY 约束行
1302/// - 列定义按顶层逗号分隔(嵌套括号如 DECIMAL(10,2) 不拆分)
1303fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
1304    let trimmed = sql.trim();
1305    let upper = trimmed.to_uppercase();
1306
1307    // 必须以 CREATE TABLE 开头
1308    if !upper.starts_with("CREATE TABLE") {
1309        return Err("schema! expects a CREATE TABLE statement".to_string());
1310    }
1311
1312    // 跳过 "CREATE TABLE"
1313    let mut rest = &trimmed["CREATE TABLE".len()..];
1314
1315    // 跳过可选的 "IF NOT EXISTS"
1316    let rest_upper = rest.trim_start().to_uppercase();
1317    if rest_upper.starts_with("IF NOT EXISTS") {
1318        rest = &rest.trim_start()["IF NOT EXISTS".len()..];
1319    }
1320
1321    rest = rest.trim_start();
1322
1323    // 解析表名(可能带反引号、双引号或无引号)
1324    let (table_name, after_name) = parse_identifier(rest)?;
1325    let rest = after_name.trim_start();
1326
1327    // 找到列定义起始的 '(' 与匹配的最后一个 ')'
1328    let paren_start = rest
1329        .find('(')
1330        .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
1331    let paren_end = rest
1332        .rfind(')')
1333        .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
1334    if paren_end <= paren_start {
1335        return Err("CREATE TABLE has malformed parentheses".to_string());
1336    }
1337
1338    let cols_str = &rest[paren_start + 1..paren_end];
1339
1340    // 按顶层逗号分隔列定义(注意嵌套括号,如 DECIMAL(10,2))
1341    let col_defs = split_top_level_commas(cols_str);
1342
1343    let mut columns = Vec::new();
1344    for def in col_defs {
1345        let def = def.trim();
1346        if def.is_empty() {
1347            continue;
1348        }
1349
1350        // 跳过约束定义行
1351        let def_upper = def.to_uppercase();
1352        if def_upper.starts_with("PRIMARY KEY")
1353            || def_upper.starts_with("FOREIGN KEY")
1354            || def_upper.starts_with("CONSTRAINT")
1355            || def_upper.starts_with("UNIQUE")
1356            || def_upper.starts_with("INDEX")
1357            || def_upper.starts_with("KEY")
1358        {
1359            continue;
1360        }
1361
1362        // 解析列名
1363        let (col_name, after_col) = parse_identifier(def)?;
1364        let rest = after_col.trim_start();
1365
1366        // 解析类型(取第一个 token,去掉括号参数)
1367        let (sql_type, after_type) = parse_type_token(rest)?;
1368        let rest = after_type.trim();
1369
1370        // 判断 nullability:NOT NULL 或 PRIMARY KEY 隐含 NOT NULL
1371        let rest_upper = rest.to_uppercase();
1372        let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
1373        let rust_type = sql_type_to_rust(&sql_type, !not_null);
1374
1375        columns.push((col_name, rust_type));
1376    }
1377
1378    Ok((table_name, columns))
1379}
1380
1381/// 解析标识符:支持反引号、双引号或无引号。
1382/// 返回 (标识符, 剩余字符串)。
1383fn parse_identifier(s: &str) -> Result<(String, &str), String> {
1384    let s = s.trim_start();
1385    if s.is_empty() {
1386        return Err("expected identifier".to_string());
1387    }
1388
1389    let bytes = s.as_bytes();
1390    match bytes[0] {
1391        b'`' => {
1392            let end = s[1..]
1393                .find('`')
1394                .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
1395            let ident = s[1..1 + end].to_string();
1396            Ok((ident, &s[1 + end + 1..]))
1397        }
1398        b'"' => {
1399            let end = s[1..]
1400                .find('"')
1401                .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
1402            let ident = s[1..1 + end].to_string();
1403            Ok((ident, &s[1 + end + 1..]))
1404        }
1405        _ => {
1406            let end = s
1407                .find(|c: char| !c.is_alphanumeric() && c != '_')
1408                .unwrap_or(s.len());
1409            if end == 0 {
1410                return Err(format!("invalid identifier: '{}'", s));
1411            }
1412            let ident = s[..end].to_string();
1413            Ok((ident, &s[end..]))
1414        }
1415    }
1416}
1417
1418/// 解析类型 token:取第一个标识符,可选跟随括号参数(如 VARCHAR(255) → VARCHAR)。
1419/// 返回 (类型名, 剩余字符串)。
1420fn parse_type_token(s: &str) -> Result<(String, &str), String> {
1421    let s = s.trim_start();
1422    if s.is_empty() {
1423        return Err("expected column type".to_string());
1424    }
1425
1426    let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
1427    if end == 0 {
1428        return Err(format!("invalid type: '{}'", s));
1429    }
1430    let type_name = s[..end].to_string();
1431    let mut rest = &s[end..];
1432
1433    // 跳过可选的括号参数,如 (255) 或 (10,2)
1434    rest = rest.trim_start();
1435    if rest.starts_with('(') {
1436        let close = rest
1437            .find(')')
1438            .ok_or_else(|| "unterminated type parameter list".to_string())?;
1439        rest = &rest[close + 1..];
1440    }
1441
1442    Ok((type_name, rest))
1443}
1444
1445/// 按顶层逗号分隔字符串(不进入嵌套括号)。
1446fn split_top_level_commas(s: &str) -> Vec<String> {
1447    let mut parts = Vec::new();
1448    let mut depth: i32 = 0;
1449    let mut current = String::new();
1450
1451    for ch in s.chars() {
1452        match ch {
1453            '(' => {
1454                depth += 1;
1455                current.push(ch);
1456            }
1457            ')' => {
1458                depth -= 1;
1459                current.push(ch);
1460            }
1461            ',' if depth == 0 => {
1462                parts.push(std::mem::take(&mut current));
1463            }
1464            _ => {
1465                current.push(ch);
1466            }
1467        }
1468    }
1469
1470    if !current.trim().is_empty() {
1471        parts.push(current);
1472    }
1473
1474    parts
1475}
1476
1477/// 将 SQL 类型映射为 Rust 类型字符串。
1478///
1479/// 匹配规则:取类型名第一个 token(去掉括号参数),不区分大小写匹配。
1480/// 未识别的类型默认映射为 `String`。若 `nullable == true`,用 `Option<T>` 包裹。
1481fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
1482    let upper = sql_type.to_uppercase();
1483    let rust = match upper.as_str() {
1484        // 8 字节整数
1485        "BIGINT" | "INT8" => "i64",
1486        // 4 字节整数(INT/INTEGER/INT4/SERIAL)
1487        "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
1488        // 2 字节整数
1489        "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
1490        // 1 字节整数
1491        "TINYINT" => "i8",
1492        // 浮点(4 字节)
1493        "FLOAT" | "REAL" | "FLOAT4" => "f32",
1494        // 浮点(8 字节)/ 定点数
1495        "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
1496        // 布尔
1497        "BOOLEAN" | "BOOL" => "bool",
1498        // 二进制(与 schema_gen::sql_type_to_rust 保持一致)
1499        "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
1500        // 字符串/日期/JSON/UUID(统一映射到 String,运行时再解析)
1501        "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
1502        | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
1503        _ => "String",
1504    };
1505
1506    if nullable {
1507        format!("Option<{}>", rust)
1508    } else {
1509        rust.to_string()
1510    }
1511}
1512
1513// ---------------------------------------------------------------------------
1514// `#[derive(Schema)]` — auto-generate table structure from a struct
1515// ---------------------------------------------------------------------------
1516
1517/// 派生宏:自动从 Rust 结构体生成表结构信息。
1518///
1519/// 解析 `#[table(name = "...")]` 和 `#[column(...)]` 属性,
1520/// 生成 `Schema` trait 实现,便于在运行时反射表名与列信息。
1521///
1522/// # 支持的属性
1523///
1524/// - `#[table(name = "users")]` — 指定表名(默认使用结构体名的蛇形形式)
1525/// - `#[column(name = "user_id")]` — 指定列名(默认使用字段名)
1526/// - `#[column(type = "VARCHAR(255)")]` — 指定 SQL 类型
1527/// - `#[column(primary_key)]` — 标记主键
1528/// - `#[column(nullable)]` — 显式标记允许 NULL
1529/// - `#[column(skip)]` — 跳过此字段,不生成 schema 条目
1530/// - `#[column(default = "0")]` — 标记字段有默认值
1531///
1532/// # 类型推断
1533///
1534/// 字段的 Rust 类型会自动映射为 SQL 类型:
1535/// - `i64`/`u64` → `BIGINT`
1536/// - `i32`/`u32` → `INTEGER`
1537/// - `String` → `TEXT`
1538/// - `f64` → `DOUBLE`
1539/// - `bool` → `BOOLEAN`
1540/// - `Vec<u8>` → `BLOB`
1541/// - `Option<T>` → 与 `T` 相同,但标记为 nullable
1542#[proc_macro_derive(Schema, attributes(table, column))]
1543pub fn derive_schema(input: TokenStream) -> TokenStream {
1544    let input = parse_macro_input!(input as syn::DeriveInput);
1545    derive::derive_schema_impl(input).into()
1546}
1547
1548// ---------------------------------------------------------------------------
1549// `#[derive(Builder)]` — auto-generate builder pattern code
1550// ---------------------------------------------------------------------------
1551
1552/// 派生宏:自动生成构造器模式代码。
1553///
1554/// 为目标结构体生成一个 `XxxBuilder` 类型,包含:
1555/// - `new()` 构造空 builder
1556/// - 每个字段的 setter 方法
1557/// - `build()` 方法返回 `Result<T, String>`
1558///
1559/// # 支持的属性
1560///
1561/// - `#[builder(skip)]` — 跳过此字段(不生成 setter,使用 Default)
1562/// - `#[builder(default = expr)]` — 指定默认值表达式
1563///
1564/// # 示例
1565///
1566/// ```ignore
1567/// use sz_orm_macros::Builder;
1568///
1569/// #[derive(Builder)]
1570/// struct User {
1571///     id: i64,
1572///     name: String,
1573/// }
1574///
1575/// let user = User::builder()
1576///     .id(1)
1577///     .name("Alice".to_string())
1578///     .build()
1579///     .unwrap();
1580/// ```
1581#[proc_macro_derive(Builder, attributes(builder))]
1582pub fn derive_builder(input: TokenStream) -> TokenStream {
1583    let input = parse_macro_input!(input as syn::DeriveInput);
1584    derive::derive_builder_impl(input).into()
1585}
1586
1587// ---------------------------------------------------------------------------
1588// Unit tests — cover helper functions used by both macros
1589// ---------------------------------------------------------------------------
1590
1591#[cfg(test)]
1592mod tests {
1593    use super::*;
1594
1595    // ---- strip_string_literal ----
1596
1597    #[test]
1598    fn test_strip_plain_double_quoted() {
1599        assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
1600    }
1601
1602    #[test]
1603    fn test_strip_raw_double_hash() {
1604        assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
1605    }
1606
1607    #[test]
1608    fn test_strip_raw_double_no_hash() {
1609        assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
1610    }
1611
1612    #[test]
1613    fn test_strip_byte_string() {
1614        assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
1615        assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
1616    }
1617
1618    #[test]
1619    fn test_strip_non_string_returns_none() {
1620        assert_eq!(strip_string_literal("123"), None);
1621        assert_eq!(strip_string_literal("foo"), None);
1622    }
1623
1624    // ---- validate_sql_content ----
1625
1626    #[test]
1627    fn test_validate_select_with_from_ok() {
1628        assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
1629    }
1630
1631    #[test]
1632    fn test_validate_select_missing_from_fails() {
1633        assert!(validate_sql_content("SELECT * users", None).is_err());
1634    }
1635
1636    #[test]
1637    fn test_validate_insert_missing_into_fails() {
1638        assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
1639        assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
1640    }
1641
1642    #[test]
1643    fn test_validate_update_missing_set_fails() {
1644        assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
1645        assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
1646    }
1647
1648    #[test]
1649    fn test_validate_delete_missing_from_fails() {
1650        assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
1651        assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
1652    }
1653
1654    #[test]
1655    fn test_validate_empty_sql_fails() {
1656        assert!(validate_sql_content("", None).is_err());
1657        assert!(validate_sql_content("   ", None).is_err());
1658    }
1659
1660    // ---- balanced parens ----
1661
1662    #[test]
1663    fn test_validate_balanced_parens_ok() {
1664        assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
1665    }
1666
1667    #[test]
1668    fn test_validate_balanced_parens_unbalanced() {
1669        assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
1670        assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
1671    }
1672
1673    // ---- injection patterns ----
1674
1675    #[test]
1676    fn test_validate_no_injection_clean() {
1677        assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
1678    }
1679
1680    #[test]
1681    fn test_validate_no_injection_drop_table() {
1682        assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
1683    }
1684
1685    #[test]
1686    fn test_validate_no_injection_or_1_1() {
1687        // 编译期 SQL 已剥离外层引号,检测模式不再依赖引号字符。
1688        // "' OR '1'='1" 因引号分隔不再匹配 "or 1=1",故不再检测;
1689        // 但不含引号分隔的 "OR 1=1" 仍可被检测。
1690        assert!(validate_no_injection("' OR 1=1").is_err());
1691        assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
1692    }
1693
1694    #[test]
1695    fn test_validate_no_injection_drop_database() {
1696        assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
1697    }
1698
1699    #[test]
1700    fn test_validate_no_injection_information_schema() {
1701        assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
1702    }
1703
1704    #[test]
1705    fn test_validate_no_injection_xp_cmdshell() {
1706        assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
1707    }
1708
1709    #[test]
1710    fn test_validate_no_injection_union_select() {
1711        assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
1712    }
1713
1714    #[test]
1715    fn test_validate_no_injection_comment_dashes() {
1716        assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
1717    }
1718
1719    #[test]
1720    fn test_validate_no_injection_block_comment() {
1721        assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
1722    }
1723
1724    // ---- string literal closure ----
1725
1726    #[test]
1727    fn test_validate_string_literals_closed_ok() {
1728        assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
1729        assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
1730    }
1731
1732    #[test]
1733    fn test_validate_string_literals_closed_unclosed_single() {
1734        assert!(validate_string_literals_closed("'hello").is_err());
1735    }
1736
1737    #[test]
1738    fn test_validate_string_literals_closed_unclosed_double() {
1739        assert!(validate_string_literals_closed(r#""hello"#).is_err());
1740    }
1741
1742    // ---- param count check ----
1743
1744    #[test]
1745    fn test_validate_param_count_match() {
1746        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
1747        assert!(
1748            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
1749        );
1750    }
1751
1752    #[test]
1753    fn test_validate_param_count_mismatch() {
1754        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
1755        assert!(
1756            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
1757        );
1758    }
1759
1760    // ---- db-verify feature: detect_db_kind ----
1761
1762    #[cfg(feature = "db-verify")]
1763    #[test]
1764    fn test_detect_db_kind_mysql() {
1765        assert_eq!(
1766            detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
1767            DbKind::MySql
1768        );
1769    }
1770
1771    #[cfg(feature = "db-verify")]
1772    #[test]
1773    fn test_detect_db_kind_postgres() {
1774        assert_eq!(
1775            detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
1776            DbKind::Postgres
1777        );
1778        assert_eq!(
1779            detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
1780            DbKind::Postgres
1781        );
1782    }
1783
1784    #[cfg(feature = "db-verify")]
1785    #[test]
1786    fn test_detect_db_kind_sqlite() {
1787        assert_eq!(
1788            detect_db_kind("sqlite://path/to/db.db").unwrap(),
1789            DbKind::Sqlite
1790        );
1791        assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
1792    }
1793
1794    #[cfg(feature = "db-verify")]
1795    #[test]
1796    fn test_detect_db_kind_oracle() {
1797        assert_eq!(
1798            detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
1799            DbKind::Oracle
1800        );
1801        assert_eq!(
1802            detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
1803            DbKind::Oracle
1804        );
1805    }
1806
1807    #[cfg(feature = "db-verify")]
1808    #[test]
1809    fn test_detect_db_kind_sqlserver() {
1810        assert_eq!(
1811            detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
1812            DbKind::SqlServer
1813        );
1814        assert_eq!(
1815            detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
1816            DbKind::SqlServer
1817        );
1818        assert_eq!(
1819            detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
1820            DbKind::SqlServer
1821        );
1822    }
1823
1824    #[cfg(feature = "db-verify")]
1825    #[test]
1826    fn test_detect_db_kind_unsupported() {
1827        assert!(detect_db_kind("redis://user:pass@host/db").is_err());
1828        assert!(detect_db_kind("not-a-url").is_err());
1829    }
1830
1831    #[cfg(feature = "db-verify")]
1832    #[test]
1833    fn test_parse_oracle_dsn_basic() {
1834        let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
1835        let p = parse_oracle_dsn(dsn).unwrap();
1836        assert_eq!(p.user, "sys");
1837        assert_eq!(p.password, "test123");
1838        assert_eq!(p.host, "127.0.0.1");
1839        assert_eq!(p.port, 1521);
1840        assert_eq!(p.service, "freepdb1.FALSE");
1841        assert!(p.sysdba);
1842    }
1843
1844    #[cfg(feature = "db-verify")]
1845    #[test]
1846    fn test_parse_oracle_dsn_default_port() {
1847        // 无端口号时默认 1521
1848        let dsn = "oracle://sys:test123@127.0.0.1/FREE";
1849        let p = parse_oracle_dsn(dsn).unwrap();
1850        assert_eq!(p.port, 1521);
1851        assert_eq!(p.service, "FREE");
1852        assert!(!p.sysdba);
1853    }
1854
1855    #[cfg(feature = "db-verify")]
1856    #[test]
1857    fn test_parse_sqlserver_dsn_basic() {
1858        let dsn =
1859            "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
1860        let p = parse_sqlserver_dsn(dsn).unwrap();
1861        assert_eq!(p.user, "test");
1862        assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
1863        assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
1864        assert_eq!(p.port, 22527);
1865        assert_eq!(p.database, "test");
1866    }
1867
1868    #[cfg(feature = "db-verify")]
1869    #[test]
1870    fn test_parse_sqlserver_dsn_default_port() {
1871        let dsn = "mssql://user:pass@host/db";
1872        let p = parse_sqlserver_dsn(dsn).unwrap();
1873        assert_eq!(p.port, 1433);
1874        assert_eq!(p.database, "db");
1875    }
1876
1877    // ---- schema! 宏 parse_create_table 测试 ----
1878
1879    #[test]
1880    fn test_parse_create_table_basic() {
1881        let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
1882        let (table, cols) = parse_create_table(sql).unwrap();
1883        assert_eq!(table, "users");
1884        assert_eq!(
1885            cols,
1886            vec![
1887                ("id".to_string(), "i32".to_string()),
1888                ("name".to_string(), "String".to_string())
1889            ]
1890        );
1891    }
1892
1893    #[test]
1894    fn test_parse_create_table_with_if_not_exists() {
1895        let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
1896        let (table, cols) = parse_create_table(sql).unwrap();
1897        assert_eq!(table, "orders");
1898        assert_eq!(
1899            cols,
1900            vec![
1901                ("id".to_string(), "i64".to_string()),
1902                ("total".to_string(), "f64".to_string())
1903            ]
1904        );
1905    }
1906
1907    #[test]
1908    fn test_parse_create_table_nullable() {
1909        let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
1910        let (_, cols) = parse_create_table(sql).unwrap();
1911        assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
1912        assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
1913    }
1914
1915    #[test]
1916    fn test_parse_create_table_skip_constraints() {
1917        let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
1918        let (_, cols) = parse_create_table(sql).unwrap();
1919        assert_eq!(cols.len(), 2);
1920        assert_eq!(cols[0].0, "id");
1921        assert_eq!(cols[1].0, "name");
1922    }
1923
1924    #[test]
1925    fn test_parse_create_table_varchar_with_len() {
1926        let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
1927        let (_, cols) = parse_create_table(sql).unwrap();
1928        assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
1929        assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
1930    }
1931
1932    #[test]
1933    fn test_sql_type_to_rust_mappings() {
1934        // 整数(按字节宽度严格映射,与 SQL 标准一致)
1935        assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
1936        assert_eq!(sql_type_to_rust("INT8", false), "i64");
1937        assert_eq!(sql_type_to_rust("INT", false), "i32");
1938        assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
1939        assert_eq!(sql_type_to_rust("INT4", false), "i32");
1940        assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
1941        assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
1942        assert_eq!(sql_type_to_rust("INT2", false), "i16");
1943        assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
1944        assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
1945        // 浮点
1946        assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
1947        assert_eq!(sql_type_to_rust("REAL", false), "f32");
1948        assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
1949        assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
1950        assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
1951        assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
1952        assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
1953        assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
1954        // 布尔
1955        assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
1956        assert_eq!(sql_type_to_rust("BOOL", false), "bool");
1957        // 字符串
1958        assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
1959        assert_eq!(sql_type_to_rust("TEXT", false), "String");
1960        assert_eq!(sql_type_to_rust("CHAR", false), "String");
1961        assert_eq!(sql_type_to_rust("UUID", false), "String");
1962        assert_eq!(sql_type_to_rust("DATE", false), "String");
1963        assert_eq!(sql_type_to_rust("DATETIME", false), "String");
1964        assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
1965        assert_eq!(sql_type_to_rust("JSON", false), "String");
1966        assert_eq!(sql_type_to_rust("JSONB", false), "String");
1967        // 二进制
1968        assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
1969        assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
1970        assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
1971        assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
1972        // nullable
1973        assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
1974        assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
1975        assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
1976        assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
1977        // unknown
1978        assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
1979    }
1980
1981    #[test]
1982    fn test_parse_create_table_error_no_create() {
1983        assert!(parse_create_table("SELECT * FROM users").is_err());
1984    }
1985
1986    #[test]
1987    fn test_parse_create_table_error_no_parens() {
1988        assert!(parse_create_table("CREATE TABLE foo").is_err());
1989    }
1990}