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