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#[cfg(feature = "db-verify")]
61use sqlx::Row as _;
62
63// 引入 quote! 宏,用于类型安全地构建 TokenStream
64use proc_macro2::TokenStream as TokenStream2;
65use quote::quote;
66use syn::parse_macro_input;
67
68// 派生宏模块
69mod derive;
70
71/// Compile-time SQL validation macro.
72///
73/// Validates SQL syntax at compile time and emits the validated SQL string.
74///
75/// # Syntax
76///
77/// - `sql_string!("SQL")` — validates the SQL and emits it as a `&str`
78/// - `sql_string!("SQL"; params: N)` — additionally checks that the SQL has exactly N parameters
79///
80/// # Validation rules
81///
82/// - SELECT must contain FROM
83/// - INSERT must contain INTO and VALUES
84/// - UPDATE must contain SET
85/// - DELETE must contain FROM
86/// - Parentheses must be balanced
87/// - String literals must be properly closed
88/// - No SQL injection patterns (OR '1'='1', UNION SELECT, `'; DROP TABLE`, `--`, `/*`)
89/// - Table/column identifiers must be valid
90#[proc_macro]
91pub fn sql_string(input: TokenStream) -> TokenStream {
92    let mut tokens = input.into_iter().peekable();
93
94    // Parse the SQL string literal
95    let sql = match tokens.next() {
96        Some(TokenTree::Literal(lit)) => lit.to_string(),
97        Some(other) => {
98            return compile_error(
99                other.span(),
100                "Expected a string literal as the first argument to sql_string!",
101            );
102        }
103        None => {
104            return compile_error(
105                Span::call_site(),
106                "Expected a string literal argument to sql_string!",
107            );
108        }
109    };
110
111    // Remove surrounding quotes from the string literal
112    let sql_content = if sql.starts_with("r#\"") {
113        &sql[3..sql.len() - 2]
114    } else if sql.starts_with("r\"") {
115        &sql[2..sql.len() - 1]
116    } else if sql.starts_with('"') {
117        &sql[1..sql.len() - 1]
118    } else if sql.starts_with("b\"") || sql.starts_with("b\'") {
119        &sql[2..sql.len() - 1]
120    } else {
121        return compile_error(
122            Span::call_site(),
123            "sql_string! requires a string literal argument",
124        );
125    };
126
127    // Parse optional `params: N`
128    let mut expected_params = None;
129    if tokens.peek().is_some() {
130        // Expect `; params: N`
131        match tokens.next() {
132            Some(TokenTree::Punct(p)) if p.as_char() == ';' => {}
133            Some(other) => {
134                return compile_error(
135                    other.span(),
136                    "Expected `;` before param count, e.g. sql_string!(\"...\"; params: 2)",
137                );
138            }
139            None => {}
140        }
141
142        // Parse `params`
143        match tokens.next() {
144            Some(TokenTree::Ident(id)) if id.to_string() == "params" => {}
145            Some(other) => {
146                return compile_error(
147                    other.span(),
148                    "Expected `params:` keyword, e.g. sql_string!(\"...\"; params: 2)",
149                );
150            }
151            None => {
152                return compile_error(Span::call_site(), "Expected param count after `;`");
153            }
154        }
155
156        // Parse `:`
157        match tokens.next() {
158            Some(TokenTree::Punct(p)) if p.as_char() == ':' => {}
159            Some(other) => {
160                return compile_error(
161                    other.span(),
162                    "Expected `:` after `params`, e.g. sql_string!(\"...\"; params: 2)",
163                );
164            }
165            None => {
166                return compile_error(Span::call_site(), "Expected param count after `params`");
167            }
168        }
169
170        // Parse the number
171        match tokens.next() {
172            Some(TokenTree::Literal(lit)) => {
173                let num_str = lit.to_string();
174                if let Ok(n) = num_str.parse::<usize>() {
175                    expected_params = Some(n);
176                } else {
177                    return compile_error(
178                        lit.span(),
179                        "Expected a positive integer for param count",
180                    );
181                }
182            }
183            Some(other) => {
184                return compile_error(
185                    other.span(),
186                    "Expected a number after `params:`, e.g. sql_string!(\"...\"; params: 2)",
187                );
188            }
189            None => {
190                return compile_error(Span::call_site(), "Expected a number after `params:`");
191            }
192        }
193    }
194
195    // Run validation
196    if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
197        return compile_error(Span::call_site(), &err_msg);
198    }
199
200    // Emit the validated string as a &str literal
201    let output = format!("\"{}\"", sql_content.escape_default());
202    output
203        .parse()
204        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
205}
206
207// ---------------------------------------------------------------------------
208// Validation logic (self-contained, no external dependencies)
209// ---------------------------------------------------------------------------
210
211fn validate_sql_content(sql: &str, expected_params: Option<usize>) -> Result<(), String> {
212    let trimmed = sql.trim();
213    if trimmed.is_empty() {
214        return Err("SQL statement is empty".to_string());
215    }
216
217    validate_balanced_parens(trimmed)?;
218    validate_string_literals_closed(trimmed)?;
219    validate_no_injection(trimmed)?;
220
221    // Type-specific validation
222    let sql_upper = trimmed.to_uppercase();
223    if sql_upper.starts_with("SELECT") {
224        if !sql_upper.contains("FROM") {
225            return Err("SELECT statement missing FROM clause".to_string());
226        }
227    } else if sql_upper.starts_with("INSERT") {
228        if !sql_upper.contains("INTO") {
229            return Err("INSERT statement missing INTO clause".to_string());
230        }
231        if !sql_upper.contains("VALUES") {
232            return Err("INSERT statement missing VALUES clause".to_string());
233        }
234    } else if sql_upper.starts_with("UPDATE") {
235        if !sql_upper.contains("SET") {
236            return Err("UPDATE statement missing SET clause".to_string());
237        }
238    } else if sql_upper.starts_with("DELETE") && !sql_upper.contains("FROM") {
239        return Err("DELETE statement missing FROM clause".to_string());
240    }
241
242    // Parameter count check
243    if let Some(expected) = expected_params {
244        let actual = sql.chars().filter(|&c| c == '?').count();
245        if actual != expected {
246            return Err(format!(
247                "Parameter count mismatch: expected {} parameters, found {}",
248                expected, actual
249            ));
250        }
251    }
252
253    Ok(())
254}
255
256fn validate_balanced_parens(sql: &str) -> Result<(), String> {
257    let mut depth: i32 = 0;
258    for (i, ch) in sql.char_indices() {
259        match ch {
260            '(' => depth += 1,
261            ')' => {
262                depth -= 1;
263                if depth < 0 {
264                    return Err(format!(
265                        "Unbalanced parentheses: unexpected ')' at position {}",
266                        i
267                    ));
268                }
269            }
270            _ => {}
271        }
272    }
273    if depth != 0 {
274        return Err(format!("Unbalanced parentheses: {} unclosed '('", depth));
275    }
276    Ok(())
277}
278
279fn validate_string_literals_closed(sql: &str) -> Result<(), String> {
280    let mut in_single = false;
281    let mut in_double = false;
282    let mut prev = '\0';
283
284    for ch in sql.chars() {
285        if prev == '\\' {
286            prev = ch;
287            continue;
288        }
289
290        match ch {
291            '\'' if !in_double => in_single = !in_single,
292            '"' if !in_single => in_double = !in_double,
293            _ => {}
294        }
295        prev = ch;
296    }
297
298    if in_single {
299        return Err("Unclosed single-quoted string literal".to_string());
300    }
301    if in_double {
302        return Err("Unclosed double-quoted string literal".to_string());
303    }
304
305    Ok(())
306}
307
308fn validate_no_injection(sql: &str) -> Result<(), String> {
309    let sql_lower = sql.to_lowercase();
310
311    // 注意:编译期 SQL 内容已由 Rust 字符串字面量解析剥离外层引号,
312    // 因此检测模式不应依赖前导引号字符(如 `"'; DROP TABLE"`)。
313    let injection_patterns: &[&str] = &[
314        // 多语句攻击
315        "drop table",
316        "drop database",
317        "; drop",
318        // 经典注入
319        "or 1=1",
320        "or 1 = 1",
321        "union select",
322        "union all select",
323        // 注释攻击
324        "--",
325        "/*",
326        "*/",
327        // 存储过程注入
328        "xp_cmdshell",
329        "sp_executesql",
330        "exec(",
331        "execute(",
332        // 信息泄露
333        "information_schema",
334        "sys.tables",
335        "sys.columns",
336    ];
337
338    for pattern in injection_patterns {
339        if sql_lower.contains(pattern) {
340            return Err(format!("潜在的 SQL 注入模式被检测到: '{}'", pattern));
341        }
342    }
343
344    Ok(())
345}
346
347// ---------------------------------------------------------------------------
348// `query!` macro — optional real DB verification (gated by `db-verify` feature)
349// ---------------------------------------------------------------------------
350
351/// Compile-time SQL validation with optional real DB verification.
352///
353/// Behavior:
354/// - Always runs the same syntax validation as `sql_string!`.
355/// - When the `db-verify` cargo feature is enabled **AND** the
356///   `SZ_ORM_QUERY_VERIFY=1` environment variable is set at compile time,
357///   connects to the database pointed to by `DATABASE_URL` and runs
358///   `EXPLAIN` (MySQL/PostgreSQL) or `EXPLAIN QUERY PLAN` (SQLite) to verify
359///   the SQL is valid against the actual schema (column names, table names,
360///   joins, etc.).
361/// - Otherwise, falls back to syntax-only validation.
362///
363/// Emits a [`sz_orm_core::queryable::Query`] object wrapping the validated SQL.
364///
365/// # Syntax
366///
367/// ```ignore
368/// use sz_orm_core::queryable::Query;
369/// let q = query!("SELECT id, name FROM users WHERE id = ?");
370/// let rows = q.fetch_all(&mut conn).await?;
371/// ```
372///
373/// # Verification setup
374///
375/// ```bash
376/// export DATABASE_URL="mysql://user:pass@host:3306/db"
377/// export SZ_ORM_QUERY_VERIFY=1
378/// cargo build --features sz-orm-macros/db-verify
379/// ```
380#[proc_macro]
381pub fn query(input: TokenStream) -> TokenStream {
382    let mut tokens = input.into_iter().peekable();
383
384    // P0-1:支持可选类型参数 `query!(T, "SQL")` → `QueryAs::<T>::new(sql)`
385    // 若无类型参数则保持 `query!("SQL")` → `Query::new(sql)`
386    let type_param: Option<TokenStream2> = match tokens.peek() {
387        Some(TokenTree::Ident(_)) | Some(TokenTree::Punct(_)) => {
388            // 收集类型路径(如 `User` 或 `crate::User`)
389            let mut ty_tokens = Vec::new();
390            while let Some(tok) = tokens.peek() {
391                match tok {
392                    TokenTree::Punct(p) if p.as_char() == ',' => break,
393                    TokenTree::Punct(p) if p.as_char() == ':' => {
394                        ty_tokens.push(tokens.next().unwrap());
395                        // 消耗 `:`
396                        if let Some(TokenTree::Punct(p2)) = tokens.peek() {
397                            if p2.as_char() == ':' {
398                                ty_tokens.push(tokens.next().unwrap());
399                            }
400                        }
401                    }
402                    _ => ty_tokens.push(tokens.next().unwrap()),
403                }
404            }
405            // 确认下一个 token 是逗号(类型参数分隔符)
406            match tokens.peek() {
407                Some(TokenTree::Punct(p)) if p.as_char() == ',' => {
408                    tokens.next(); // 消耗逗号
409                    let ts: proc_macro::TokenStream = ty_tokens.into_iter().collect();
410                    Some(TokenStream2::from(ts))
411                }
412                _ => None, // 不是类型参数,回退
413            }
414        }
415        _ => None,
416    };
417
418    // Parse the SQL string literal
419    let sql = match tokens.next() {
420        Some(TokenTree::Literal(lit)) => lit.to_string(),
421        Some(other) => {
422            return compile_error(
423                other.span(),
424                if type_param.is_some() {
425                    "query!(T, \"SQL\"): expected a string literal as the second argument"
426                } else {
427                    "Expected a string literal as the first argument to query!"
428                },
429            );
430        }
431        None => {
432            return compile_error(
433                Span::call_site(),
434                if type_param.is_some() {
435                    "query!(T, \"SQL\"): missing SQL string argument"
436                } else {
437                    "Expected a string literal argument to query!"
438                },
439            );
440        }
441    };
442
443    let sql_content = match strip_string_literal(&sql) {
444        Some(s) => s,
445        None => {
446            return compile_error(
447                Span::call_site(),
448                "query! requires a string literal argument",
449            );
450        }
451    };
452
453    // Syntax validation (shared with sql_string!)
454    if let Err(err_msg) = validate_sql_content(sql_content, None) {
455        return compile_error(Span::call_site(), &err_msg);
456    }
457
458    // Optional real DB verification (only when feature is enabled)
459    #[cfg(feature = "db-verify")]
460    let verify_cols: Option<Vec<(String, String)>> = {
461        match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
462            // 模式 1:连真 DB 执行 EXPLAIN 验证(需 DATABASE_URL),并获取 SELECT 列的实际类型
463            Some("1") => match verify_with_real_db(sql_content) {
464                Ok(cols) => Some(cols),
465                Err(err) => {
466                    return compile_error(
467                        Span::call_site(),
468                        &format!("query! real DB verification failed: {}", err),
469                    )
470                }
471            },
472            // 模式 cache:从离线缓存文件查找(无需 DB,适合 CI)
473            Some("cache") => {
474                if let Err(err) = verify_with_cache(sql_content) {
475                    return compile_error(
476                        Span::call_site(),
477                        &format!("query! offline cache verification failed: {}", err),
478                    );
479                }
480                None
481            }
482            _ => None,
483        }
484    };
485    #[cfg(not(feature = "db-verify"))]
486    let _verify_cols: Option<Vec<(String, String)>> = None;
487
488    // Emit the appropriate query object
489    let escaped = sql_content.escape_default().to_string();
490    let base = if let Some(ref ty) = type_param {
491        // query!(T, "SQL") → QueryAs::<T>::new("SQL")
492        format!(
493            "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
494            ty, escaped
495        )
496    } else {
497        // query!("SQL") → Query::new("SQL")
498        format!("::sz_orm_core::queryable::Query::new(\"{}\")", escaped)
499    };
500    // db-verify 通过且有类型参数时,附加编译期类型验证块(P0-2)
501    #[cfg(feature = "db-verify")]
502    let output = match (&verify_cols, &type_param) {
503        (Some(cols), Some(ty)) if !cols.is_empty() => {
504            gen_compile_time_type_check(&ty.to_string(), sql_content, cols, &base)
505        }
506        _ => base,
507    };
508    #[cfg(not(feature = "db-verify"))]
509    let output = base;
510    output
511        .parse()
512        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query! output"))
513}
514
515/// Strip surrounding quotes from a string literal token's raw representation.
516/// Shared by `sql_string!` and `query!`.
517fn strip_string_literal(raw: &str) -> Option<&str> {
518    if raw.starts_with("r#\"") {
519        Some(&raw[3..raw.len() - 2])
520    } else if raw.starts_with("r\"") {
521        Some(&raw[2..raw.len() - 1])
522    } else if raw.starts_with('"') {
523        Some(&raw[1..raw.len() - 1])
524    } else if raw.starts_with("b\"") || raw.starts_with("b\'") {
525        Some(&raw[2..raw.len() - 1])
526    } else {
527        None
528    }
529}
530
531// ---------------------------------------------------------------------------
532// Real DB verification (only compiled when `db-verify` feature is enabled)
533// ---------------------------------------------------------------------------
534
535#[cfg(feature = "db-verify")]
536fn verify_with_real_db(sql: &str) -> Result<Vec<(String, String)>, String> {
537    let dsn = std::env::var("DATABASE_URL")
538        .map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
539
540    let db_kind =
541        detect_db_kind(&dsn).map_err(|e| format!("Failed to detect DB kind from DSN: {}", e))?;
542
543    // 将 ? 占位符替换为 NULL,使 EXPLAIN 无需绑定参数即可执行。
544    // EXPLAIN 不实际执行查询,NULL 对所有列类型都合法。
545    let sql_no_placeholders = replace_placeholders_with_null(sql);
546
547    // Oracle/SQL Server 使用 EXPLAIN PLAN FOR(不同语法),其余用 EXPLAIN
548    let explain_sql = match db_kind {
549        DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql_no_placeholders),
550        DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql_no_placeholders),
551        // Oracle: EXPLAIN PLAN FOR 放入 PLAN_TABLE,再查询结果验证语法
552        DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
553        // SQL Server: SET SHOWPLAN_TEXT ON 后执行(不实际运行)
554        DbKind::SqlServer => sql_no_placeholders,
555    };
556
557    // MySQL/PG/SQLite 走 sqlx 异步路径
558    if matches!(db_kind, DbKind::MySql | DbKind::Postgres | DbKind::Sqlite) {
559        let rt = tokio::runtime::Runtime::new()
560            .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
561        return rt.block_on(async {
562            // 1. EXPLAIN 语法验证
563            if let DbKind::MySql = db_kind {
564                verify_mysql(&dsn, &explain_sql).await?;
565            } else if let DbKind::Postgres = db_kind {
566                verify_postgres(&dsn, &explain_sql).await?;
567            } else {
568                // 由外层 if matches! 保证此处必为 Sqlite
569                verify_sqlite(&dsn, &explain_sql).await?;
570            }
571            // 2. 列名/类型验证(Gap 1 修复)
572            verify_columns(&dsn, db_kind, sql).await?;
573            // 3. 获取 SELECT 列的实际 DB 类型(供编译期类型验证使用)
574            //    SQLite/Oracle/SQL Server 返回空列表(跳过类型级验证)
575            fetch_column_types(&dsn, db_kind, sql).await
576        });
577    }
578
579    // Oracle/SQL Server 走命令行工具验证(避免引入重依赖)
580    if let DbKind::Oracle = db_kind {
581        verify_oracle(&dsn, &explain_sql).map(|_| Vec::new())
582    } else {
583        // 由外层 if matches! 保证此处必为 SqlServer
584        verify_sqlserver(&dsn, &explain_sql).map(|_| Vec::new())
585    }
586}
587
588/// 离线缓存验证:从 `SZ_ORM_SQLX_CACHE` 指定的 JSON 文件中查找已验证的 SQL。
589///
590/// 缓存文件格式为 JSON 字符串数组,每行一条已验证 SQL:
591/// ```json
592/// ["SELECT `id`, `name` FROM `users` WHERE `id` = ?", ...]
593/// ```
594///
595/// 生成方式:在有 DB 的环境中运行 `cargo build --features db-verify`(`SZ_ORM_QUERY_VERIFY=1`),
596/// 或使用 `cargo sz-orm prepare` 工具扫描项目中的 `query!` 宏并生成缓存。
597///
598/// CI 中只需设置 `SZ_ORM_QUERY_VERIFY=cache` + `SZ_ORM_SQLX_CACHE=.sz-orm/query-cache.json`
599/// 即可在不连接 DB 的情况下完成编译期 SQL 验证。
600#[cfg(feature = "db-verify")]
601fn verify_with_cache(sql: &str) -> Result<(), String> {
602    let cache_path = std::env::var("SZ_ORM_SQLX_CACHE").map_err(|_| {
603        "SZ_ORM_SQLX_CACHE not set. \
604             Set it to the path of a JSON file containing verified SQL statements, \
605             e.g. SZ_ORM_SQLX_CACHE=.sz-orm/query-cache.json"
606            .to_string()
607    })?;
608
609    let cache_content = std::fs::read_to_string(&cache_path).map_err(|e| {
610        format!(
611            "Failed to read cache file '{}': {}. \
612             Run `cargo sz-orm prepare` or build with SZ_ORM_QUERY_VERIFY=1 to generate it.",
613            cache_path, e
614        )
615    })?;
616
617    // 支持两种格式:JSON 数组 或 每行一条 SQL 的文本文件
618    let verified: Vec<String> = serde_json::from_str(&cache_content).unwrap_or_else(|_| {
619        cache_content
620            .lines()
621            .map(|l| l.trim().to_string())
622            .filter(|l| !l.is_empty() && !l.starts_with('#'))
623            .collect()
624    });
625
626    if verified.iter().any(|v| v.trim() == sql.trim()) {
627        Ok(())
628    } else {
629        Err(format!(
630            "SQL not found in offline cache ({} entries): \"{}\". \
631             Add it to the cache by running with SZ_ORM_QUERY_VERIFY=1 first.",
632            verified.len(),
633            truncate_sql(sql, 80)
634        ))
635    }
636}
637
638/// 截断 SQL 用于错误消息显示
639#[cfg(feature = "db-verify")]
640fn truncate_sql(sql: &str, max: usize) -> String {
641    if sql.len() <= max {
642        sql.to_string()
643    } else {
644        format!("{}...", &sql[..max])
645    }
646}
647
648#[cfg(feature = "db-verify")]
649#[derive(Debug, Clone, Copy, PartialEq, Eq)]
650enum DbKind {
651    MySql,
652    Postgres,
653    Sqlite,
654    Oracle,
655    SqlServer,
656}
657
658/// 将 SQL 中的 `?` 占位符替换为 `NULL`,跳过字符串字面量内的 `?`。
659///
660/// EXPLAIN 不实际执行查询,用 NULL 代替参数可验证语法和表/列存在性,
661/// 同时避免 sqlx 预处理语句要求绑定参数的问题。
662#[cfg(feature = "db-verify")]
663fn replace_placeholders_with_null(sql: &str) -> String {
664    let mut result = String::with_capacity(sql.len() + 16);
665    let mut in_single_quote = false;
666    let mut in_double_quote = false;
667    let mut prev = '\0';
668
669    for ch in sql.chars() {
670        if prev == '\\' {
671            // 转义字符:直接追加
672            result.push(ch);
673            prev = ch;
674            continue;
675        }
676        match ch {
677            '\'' if !in_double_quote => in_single_quote = !in_single_quote,
678            '"' if !in_single_quote => in_double_quote = !in_double_quote,
679            '?' if !in_single_quote && !in_double_quote => {
680                result.push_str("NULL");
681                prev = ch;
682                continue;
683            }
684            _ => {}
685        }
686        result.push(ch);
687        prev = ch;
688    }
689    result
690}
691
692#[cfg(feature = "db-verify")]
693fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
694    let lower = dsn.to_lowercase();
695    if lower.starts_with("mysql://") {
696        Ok(DbKind::MySql)
697    } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
698        Ok(DbKind::Postgres)
699    } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
700        Ok(DbKind::Sqlite)
701    } else if lower.starts_with("oracle://") || lower.starts_with("oracle:") {
702        Ok(DbKind::Oracle)
703    } else if lower.starts_with("sqlserver://")
704        || lower.starts_with("mssql://")
705        || lower.starts_with("tds://")
706    {
707        Ok(DbKind::SqlServer)
708    } else {
709        Err(format!("Unsupported DSN scheme: {}", dsn))
710    }
711}
712
713#[cfg(feature = "db-verify")]
714async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
715    let pool = sqlx::MySqlPool::connect(dsn)
716        .await
717        .map_err(|e| format!("MySQL connect failed: {}", e))?;
718    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
719        .execute(&pool)
720        .await
721        .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
722    Ok(())
723}
724
725#[cfg(feature = "db-verify")]
726async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
727    let pool = sqlx::PgPool::connect(dsn)
728        .await
729        .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
730    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
731        .execute(&pool)
732        .await
733        .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
734    Ok(())
735}
736
737#[cfg(feature = "db-verify")]
738async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
739    let pool = sqlx::SqlitePool::connect(dsn)
740        .await
741        .map_err(|e| format!("SQLite connect failed: {}", e))?;
742    sqlx::query(sqlx::AssertSqlSafe(explain_sql))
743        .execute(&pool)
744        .await
745        .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
746    Ok(())
747}
748
749// ========================================================================
750// Gap 1 修复:列名/类型验证(在 EXPLAIN 语法验证通过后执行)
751// ========================================================================
752
753/// 从 SQL 中提取表名和列引用,并查询 information_schema 验证列存在性。
754///
755/// EXPLAIN 已验证语法和表存在性;此函数进一步验证:
756/// - SELECT/WHERE/ORDER BY/GROUP BY 中引用的列是否存在于对应表中
757/// - 不验证表别名限定的列(由 EXPLAIN 负责)
758/// - 仅对 `*` 以外的显式列名做验证
759///
760/// # 支持的 DB
761///
762/// MySQL / PostgreSQL / SQLite(Oracle/SQL Server 跳过此步骤)
763#[cfg(feature = "db-verify")]
764async fn verify_columns(dsn: &str, db_kind: DbKind, sql: &str) -> Result<(), String> {
765    // SQLite 的 information_schema 支持有限,跳过
766    if matches!(db_kind, DbKind::Sqlite | DbKind::Oracle | DbKind::SqlServer) {
767        return Ok(());
768    }
769
770    let tables = extract_tables(sql);
771    let columns = extract_columns(sql);
772
773    if tables.is_empty() || columns.is_empty() {
774        return Ok(());
775    }
776
777    match db_kind {
778        DbKind::MySql => verify_columns_mysql(dsn, &tables, &columns, sql).await,
779        DbKind::Postgres => verify_columns_postgres(dsn, &tables, &columns, sql).await,
780        _ => Ok(()),
781    }
782}
783
784/// 从 SQL 的 FROM 子句中提取表名(支持别名和 JOIN)
785#[cfg(feature = "db-verify")]
786fn extract_tables(sql: &str) -> Vec<String> {
787    let mut tables = Vec::new();
788    let upper = sql.to_uppercase();
789
790    // 查找 FROM ... WHERE/ORDER/GROUP/LIMIT/HAVING/JOIN 之间的内容
791    let from_idx = match upper.find("FROM") {
792        Some(i) => i,
793        None => return tables,
794    };
795
796    let end_patterns = ["WHERE", "ORDER", "GROUP", "LIMIT", "HAVING", "UNION"];
797    let end_idx = end_patterns
798        .iter()
799        .filter_map(|p| {
800            // 按单词边界查找,避免 "ORDER" 误匹配 "user_id" 等标识符中的子串
801            let mut search_start = 0;
802            while let Some(i) = upper[search_start..].find(*p) {
803                let abs_i = search_start + i;
804                let before = upper[..abs_i].chars().last().unwrap_or(' ');
805                let after = upper[abs_i + p.len()..].chars().next().unwrap_or(' ');
806                if !before.is_alphanumeric()
807                    && !after.is_alphanumeric()
808                    && before != '_'
809                    && after != '_'
810                {
811                    return Some(abs_i);
812                }
813                search_start = abs_i + p.len();
814            }
815            None
816        })
817        .filter(|&i| i > from_idx)
818        .min()
819        .unwrap_or(sql.len());
820
821    let from_clause = &sql[from_idx + 4..end_idx];
822
823    // 按逗号、换行、JOIN 关键字分割(忽略大小写)
824    let join_split = {
825        let lower = from_clause.to_lowercase();
826        let mut result = String::with_capacity(from_clause.len());
827        let mut i = 0;
828        let bytes = from_clause.as_bytes();
829        let lower_bytes = lower.as_bytes();
830        while i < bytes.len() {
831            let mut matched = false;
832            for join_kw in &[
833                " join ",
834                " inner join ",
835                " left join ",
836                " right join ",
837                " left outer join ",
838                " right outer join ",
839                " cross join ",
840                " full join ",
841                " full outer join ",
842            ] {
843                let kw = join_kw.as_bytes();
844                if i + kw.len() <= bytes.len() && &lower_bytes[i..i + kw.len()] == kw {
845                    result.push(',');
846                    i += kw.len();
847                    matched = true;
848                    break;
849                }
850            }
851            if !matched {
852                result.push(bytes[i] as char);
853                i += 1;
854            }
855        }
856        result
857    };
858    let parts: Vec<&str> = join_split.split([',', '\n']).collect();
859
860    for part in parts {
861        let part = part.trim();
862        if part.is_empty() {
863            continue;
864        }
865        // 取第一个词作为表名(忽略别名)
866        let table_word = part
867            .split_whitespace()
868            .next()
869            .unwrap_or(part)
870            .trim_end_matches([',', ';']);
871        // 去除反引号/双引号
872        let clean = table_word.trim_matches(|c| c == '`' || c == '"');
873        if !clean.is_empty()
874            && !matches!(
875                clean.to_uppercase().as_str(),
876                "INNER"
877                    | "LEFT"
878                    | "RIGHT"
879                    | "OUTER"
880                    | "CROSS"
881                    | "FULL"
882                    | "NATURAL"
883                    | "ON"
884                    | "USING"
885                    | "AS"
886            )
887        {
888            tables.push(clean.to_lowercase());
889        }
890    }
891
892    tables
893}
894
895/// 从 SQL 中提取未限定的列名引用(SELECT/WHERE/ORDER BY/GROUP BY 中)
896#[cfg(feature = "db-verify")]
897fn extract_columns(sql: &str) -> Vec<String> {
898    let mut columns = Vec::new();
899    let upper = sql.to_uppercase();
900
901    // 收集各子句中的标识符
902    let mut collect_from_segment = |segment: &str| {
903        // 简单的标识符提取:匹配 `\w+` 模式的词
904        // 排除 SQL 关键字和已限定的列(table.col)
905        let keywords = [
906            "SELECT",
907            "FROM",
908            "WHERE",
909            "AND",
910            "OR",
911            "NOT",
912            "IN",
913            "IS",
914            "NULL",
915            "LIKE",
916            "BETWEEN",
917            "AS",
918            "ON",
919            "JOIN",
920            "INNER",
921            "LEFT",
922            "RIGHT",
923            "OUTER",
924            "CROSS",
925            "FULL",
926            "NATURAL",
927            "ORDER",
928            "BY",
929            "GROUP",
930            "HAVING",
931            "LIMIT",
932            "OFFSET",
933            "ASC",
934            "DESC",
935            "DISTINCT",
936            "COUNT",
937            "SUM",
938            "AVG",
939            "MIN",
940            "MAX",
941            "CASE",
942            "WHEN",
943            "THEN",
944            "ELSE",
945            "END",
946            "COALESCE",
947            "NULLIF",
948            "CAST",
949            "TRUE",
950            "FALSE",
951            "INSERT",
952            "INTO",
953            "VALUES",
954            "UPDATE",
955            "SET",
956            "DELETE",
957            "CREATE",
958            "TABLE",
959            "INDEX",
960            "IF",
961            "EXISTS",
962            "PRIMARY",
963            "KEY",
964            "REFERENCES",
965            "FOREIGN",
966        ];
967
968        for word in segment.split(|c: char| !c.is_alphanumeric() && c != '_') {
969            if word.is_empty() || word.len() < 2 {
970                continue;
971            }
972            let w = word.to_uppercase();
973            if keywords.contains(&w.as_str()) {
974                continue;
975            }
976            // 跳过纯数字
977            if word.chars().all(|c| c.is_ascii_digit()) {
978                continue;
979            }
980            // 跳过已限定的列(table.col)— 这些由 EXPLAIN 验证
981            // 检查前面是否有 `.`
982            let pos = segment.find(word).unwrap_or(0);
983            if pos > 0 && segment.chars().nth(pos - 1) == Some('.') {
984                continue;
985            }
986            // 跳过 *
987            if word == "*" {
988                continue;
989            }
990            let lower = word.to_lowercase();
991            if !columns.contains(&lower) {
992                columns.push(lower);
993            }
994        }
995    };
996
997    // 收集 SELECT 列(FROM 之前)
998    if let Some(from_idx) = upper.find("FROM") {
999        if let Some(sel_idx) = upper.find("SELECT") {
1000            let sel_segment = &sql[sel_idx + 6..from_idx];
1001            collect_from_segment(sel_segment);
1002        }
1003    }
1004
1005    // 收集 WHERE 列
1006    if let Some(where_idx) = upper.find("WHERE") {
1007        let end_idx = ["ORDER", "GROUP", "LIMIT", "HAVING", "UNION"]
1008            .iter()
1009            .filter_map(|p| upper.find(p))
1010            .filter(|&i| i > where_idx)
1011            .min()
1012            .unwrap_or(sql.len());
1013        collect_from_segment(&sql[where_idx + 5..end_idx]);
1014    }
1015
1016    // 收集 ORDER BY 列
1017    if let Some(order_idx) = upper.find("ORDER BY") {
1018        let end_idx = ["GROUP", "LIMIT", "HAVING", "UNION"]
1019            .iter()
1020            .filter_map(|p| upper.find(p))
1021            .filter(|&i| i > order_idx)
1022            .min()
1023            .unwrap_or(sql.len());
1024        collect_from_segment(&sql[order_idx + 8..end_idx]);
1025    }
1026
1027    columns
1028}
1029
1030#[cfg(feature = "db-verify")]
1031async fn verify_columns_mysql(
1032    dsn: &str,
1033    tables: &[String],
1034    columns: &[String],
1035    sql: &str,
1036) -> Result<(), String> {
1037    let pool = sqlx::MySqlPool::connect(dsn)
1038        .await
1039        .map_err(|e| format!("MySQL connect failed: {}", e))?;
1040
1041    for col in columns {
1042        // 查询 information_schema.COLUMNS
1043        let rows = sqlx::query(
1044            "SELECT TABLE_NAME, COLUMN_NAME FROM INFORMATION_SCHEMA.COLUMNS \
1045             WHERE TABLE_SCHEMA = DATABASE() AND COLUMN_NAME = ?",
1046        )
1047        .bind(col)
1048        .fetch_all(&pool)
1049        .await
1050        .map_err(|e| format!("MySQL column lookup failed for '{}': {}", col, e))?;
1051
1052        if rows.is_empty() {
1053            // 尝试检查是否是数据库函数(如 NOW, COUNT 等)
1054            if is_sql_function(col) {
1055                continue;
1056            }
1057            return Err(format!(
1058                "query! column verification failed: column '{}' not found in any table of the current database. \
1059                 SQL: {}",
1060                col,
1061                truncate_sql(sql, 120)
1062            ));
1063        }
1064
1065        // 验证列至少存在于一个 FROM 表中(如果有表信息)
1066        if !tables.is_empty() {
1067            let found_in_table = rows.iter().any(|row| {
1068                let table_name: String = row.get("TABLE_NAME");
1069                tables.iter().any(|t| t == &table_name.to_lowercase())
1070            });
1071            if !found_in_table {
1072                let available: Vec<String> = rows.iter().map(|r| r.get("TABLE_NAME")).collect();
1073                return Err(format!(
1074                    "query! column verification failed: column '{}' exists but not in FROM table(s) {:?}. \
1075                     Found in: {:?}. SQL: {}",
1076                    col,
1077                    tables,
1078                    available,
1079                    truncate_sql(sql, 120)
1080                ));
1081            }
1082        }
1083    }
1084
1085    Ok(())
1086}
1087
1088#[cfg(feature = "db-verify")]
1089async fn verify_columns_postgres(
1090    dsn: &str,
1091    tables: &[String],
1092    columns: &[String],
1093    sql: &str,
1094) -> Result<(), String> {
1095    let pool = sqlx::PgPool::connect(dsn)
1096        .await
1097        .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
1098
1099    for col in columns {
1100        let rows = sqlx::query(
1101            "SELECT TABLE_NAME, COLUMN_NAME FROM INFORMATION_SCHEMA.COLUMNS \
1102             WHERE TABLE_CATALOG = CURRENT_CATALOG AND COLUMN_NAME = $1",
1103        )
1104        .bind(col)
1105        .fetch_all(&pool)
1106        .await
1107        .map_err(|e| format!("PostgreSQL column lookup failed for '{}': {}", col, e))?;
1108
1109        if rows.is_empty() && !is_sql_function(col) {
1110            return Err(format!(
1111                "query! column verification failed: column '{}' not found in any table of the current database. \
1112                 SQL: {}",
1113                col,
1114                truncate_sql(sql, 120)
1115            ));
1116        }
1117
1118        if !tables.is_empty() && !rows.is_empty() {
1119            let found_in_table = rows.iter().any(|row| {
1120                let table_name: String = row.get("TABLE_NAME");
1121                tables.iter().any(|t| t == &table_name.to_lowercase())
1122            });
1123            if !found_in_table {
1124                let available: Vec<String> = rows.iter().map(|r| r.get("TABLE_NAME")).collect();
1125                return Err(format!(
1126                    "query! column verification failed: column '{}' exists but not in FROM table(s) {:?}. \
1127                     Found in: {:?}. SQL: {}",
1128                    col, tables, available,
1129                    truncate_sql(sql, 120)
1130                ));
1131            }
1132        }
1133    }
1134
1135    Ok(())
1136}
1137
1138// ---------------------------------------------------------------------------
1139// 列类型获取(P0-2)
1140//
1141// 获取 SELECT 列的实际 DB 类型(列名 → 类型名),由 `query_as!`/`query!(T, ...)`
1142// 宏嵌入到生成的编译期验证代码中:用户代码在 const 上下文中将实际类型与
1143// 结构体 `__sz_orm_column_types()` 期望值对比,不匹配即编译失败。
1144// ---------------------------------------------------------------------------
1145
1146/// 获取 SELECT 列的实际 DB 类型列表 `(列名, DATA_TYPE/udt_name)`。
1147///
1148/// 仅 MySQL/PostgreSQL 支持(SQLite/Oracle/SQL Server 返回空列表,跳过类型级验证)。
1149#[cfg(feature = "db-verify")]
1150async fn fetch_column_types(
1151    dsn: &str,
1152    db_kind: DbKind,
1153    sql: &str,
1154) -> Result<Vec<(String, String)>, String> {
1155    if !matches!(db_kind, DbKind::MySql | DbKind::Postgres) {
1156        return Ok(Vec::new());
1157    }
1158
1159    let tables = extract_tables(sql);
1160    let columns = extract_columns(sql);
1161    if tables.is_empty() || columns.is_empty() {
1162        return Ok(Vec::new());
1163    }
1164
1165    match db_kind {
1166        DbKind::MySql => fetch_column_types_mysql(dsn, &tables, &columns).await,
1167        DbKind::Postgres => fetch_column_types_postgres(dsn, &tables, &columns).await,
1168        _ => Ok(Vec::new()),
1169    }
1170}
1171
1172#[cfg(feature = "db-verify")]
1173async fn fetch_column_types_mysql(
1174    dsn: &str,
1175    tables: &[String],
1176    columns: &[String],
1177) -> Result<Vec<(String, String)>, String> {
1178    let pool = sqlx::MySqlPool::connect(dsn)
1179        .await
1180        .map_err(|e| format!("MySQL connect failed for type fetch: {}", e))?;
1181
1182    let mut result = Vec::new();
1183    for col in columns {
1184        let rows = sqlx::query(
1185            "SELECT TABLE_NAME, COLUMN_NAME, DATA_TYPE \
1186             FROM INFORMATION_SCHEMA.COLUMNS \
1187             WHERE TABLE_SCHEMA = DATABASE() AND COLUMN_NAME = ?",
1188        )
1189        .bind(col)
1190        .fetch_all(&pool)
1191        .await
1192        .map_err(|e| format!("MySQL type lookup failed for '{}': {}", col, e))?;
1193
1194        // 取 FROM 表中的类型(多表同名列时取第一个匹配)
1195        let ty = rows
1196            .iter()
1197            .find(|row| {
1198                let tn: String = row.get("TABLE_NAME");
1199                tables.iter().any(|t| t == &tn.to_lowercase())
1200            })
1201            .and_then(|r| r.try_get::<String, _>("DATA_TYPE").ok());
1202        if let Some(ty) = ty {
1203            result.push((col.to_lowercase(), ty));
1204        }
1205    }
1206    Ok(result)
1207}
1208
1209#[cfg(feature = "db-verify")]
1210async fn fetch_column_types_postgres(
1211    dsn: &str,
1212    tables: &[String],
1213    columns: &[String],
1214) -> Result<Vec<(String, String)>, String> {
1215    let pool = sqlx::PgPool::connect(dsn)
1216        .await
1217        .map_err(|e| format!("PostgreSQL connect failed for type fetch: {}", e))?;
1218
1219    let mut result = Vec::new();
1220    for col in columns {
1221        let rows = sqlx::query(
1222            "SELECT TABLE_NAME, COLUMN_NAME, udt_name \
1223             FROM INFORMATION_SCHEMA.COLUMNS \
1224             WHERE TABLE_CATALOG = CURRENT_CATALOG AND COLUMN_NAME = $1",
1225        )
1226        .bind(col)
1227        .fetch_all(&pool)
1228        .await
1229        .map_err(|e| format!("PostgreSQL type lookup failed for '{}': {}", col, e))?;
1230
1231        let ty = rows
1232            .iter()
1233            .find(|row| {
1234                let tn: String = row.get("TABLE_NAME");
1235                tables.iter().any(|t| t == &tn.to_lowercase())
1236            })
1237            .and_then(|r| r.try_get::<String, _>("udt_name").ok());
1238        if let Some(ty) = ty {
1239            result.push((col.to_lowercase(), ty));
1240        }
1241    }
1242    Ok(result)
1243}
1244
1245/// 生成编译期类型验证代码块:`{ const _: () = { ...检查... }; <查询表达式> }`。
1246///
1247/// 验证逻辑在 const 上下文中执行(`panic!` 触发即编译失败,实现真正的编译期拦截):
1248/// 1. 列数必须与结构体字段数一致;
1249/// 2. 每个 SELECT 列名必须存在于结构体字段中(与 `__sz_orm_column_types()` 对比);
1250/// 3. 每个列的实际 DB 类型必须与结构体字段类型兼容(`__sz_orm_const_types_compatible`)。
1251///
1252/// `record_type` 为 `query_as!` 第一个参数(如 `User` / `crate::User`),
1253/// 生成的代码通过 `<记录类型>::__sz_orm_column_types()` 引用 derive 宏生成的
1254/// const fn(因此记录类型必须 `#[derive(FromQueryResult)]`)。
1255#[cfg(feature = "db-verify")]
1256fn gen_compile_time_type_check(
1257    record_type: &str,
1258    sql: &str,
1259    cols: &[(String, String)],
1260    query_expr: &str,
1261) -> String {
1262    let n = cols.len();
1263    let sql_esc = sql.escape_default().to_string();
1264    let mut checks = String::new();
1265    checks.push_str(&format!(
1266        "if exp.len() != {} {{ panic!(\"sz-orm compile-time type check failed for `{}`: SELECT returns {} columns but struct field count differs\"); }}",
1267        n, sql_esc, n
1268    ));
1269    for (i, (name, ty)) in cols.iter().enumerate() {
1270        let name_esc = name.escape_default().to_string();
1271        let ty_esc = ty.escape_default().to_string();
1272        checks.push_str(&format!(
1273            "if !::sz_orm_core::__sz_orm_const_str_eq(exp[{}].0, \"{}\") {{ panic!(\"sz-orm compile-time type check failed for `{}`: SELECT column #{} `{}` not found in struct fields\"); }}",
1274            i, name_esc, sql_esc, i, name_esc
1275        ));
1276        checks.push_str(&format!(
1277            "if !::sz_orm_core::__sz_orm_const_types_compatible(\"{}\", exp[{}].1) {{ panic!(\"sz-orm compile-time type check failed for `{}`: column `{}` type mismatch (db type `{}` not compatible with struct field type)\"); }}",
1278            ty_esc, i, sql_esc, name_esc, ty_esc
1279        ));
1280    }
1281    format!(
1282        "{{ const _: () = {{ let exp = <{}>::__sz_orm_column_types(); {} }}; {} }}",
1283        record_type, checks, query_expr
1284    )
1285}
1286
1287/// 常见 SQL 函数名列表(不应作为列名验证)
1288#[cfg(feature = "db-verify")]
1289fn is_sql_function(name: &str) -> bool {
1290    matches!(
1291        name.to_uppercase().as_str(),
1292        "NOW"
1293            | "CURRENT_TIMESTAMP"
1294            | "CURRENT_DATE"
1295            | "CURRENT_TIME"
1296            | "COUNT"
1297            | "SUM"
1298            | "AVG"
1299            | "MIN"
1300            | "MAX"
1301            | "COALESCE"
1302            | "NULLIF"
1303            | "CAST"
1304            | "CONVERT"
1305            | "IFNULL"
1306            | "NVL"
1307            | "UPPER"
1308            | "LOWER"
1309            | "LENGTH"
1310            | "TRIM"
1311            | "SUBSTRING"
1312            | "CONCAT"
1313            | "REPLACE"
1314            | "ROUND"
1315            | "CEIL"
1316            | "FLOOR"
1317            | "ABS"
1318            | "MOD"
1319            | "POWER"
1320            | "SQRT"
1321            | "LOG"
1322            | "EXP"
1323            | "DATE"
1324            | "YEAR"
1325            | "MONTH"
1326            | "DAY"
1327            | "HOUR"
1328            | "MINUTE"
1329            | "SECOND"
1330            | "NOW()"
1331            | "UUID"
1332            | "RANDOM"
1333            | "MD5"
1334            | "TRUE"
1335            | "FALSE"
1336            | "NULL"
1337    )
1338}
1339
1340/// Oracle 编译期验证:通过 sqlplus 命令行工具执行 EXPLAIN PLAN FOR
1341///
1342/// DSN 格式:`oracle://user:pass@host:port/service`(可选 `?sysdba=1`)
1343/// 例如:`oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1`
1344#[cfg(feature = "db-verify")]
1345fn verify_oracle(dsn: &str, explain_sql: &str) -> Result<(), String> {
1346    let parsed = parse_oracle_dsn(dsn)?;
1347    // 构造 sqlplus 连接串:user/pass@host:port/service [AS SYSDBA]
1348    let mut conn_str = format!(
1349        "{}/{}@{}:{}/{}",
1350        parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
1351    );
1352    if parsed.sysdba {
1353        conn_str.push_str(" AS SYSDBA");
1354    }
1355    // 用 SET SHOWPLAN 不适用于 Oracle,用 EXPLAIN PLAN FOR 并立即查询 PLAN_TABLE
1356    let full_script = format!(
1357        "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
1358         EXPLAIN PLAN FOR {};\n\
1359         SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
1360         EXIT;\n",
1361        explain_sql
1362    );
1363    let output = std::process::Command::new("sqlplus")
1364        .args(["-S", "-L", &conn_str])
1365        .stdin(std::process::Stdio::piped())
1366        .stdout(std::process::Stdio::piped())
1367        .stderr(std::process::Stdio::piped())
1368        .spawn()
1369        .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
1370    use std::io::Write;
1371    let mut child = output;
1372    if let Some(mut stdin) = child.stdin.take() {
1373        stdin
1374            .write_all(full_script.as_bytes())
1375            .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
1376    }
1377    let out = child
1378        .wait_with_output()
1379        .map_err(|e| format!("sqlplus wait failed: {}", e))?;
1380    let stdout = String::from_utf8_lossy(&out.stdout);
1381    let stderr = String::from_utf8_lossy(&out.stderr);
1382    if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
1383        return Err(format!(
1384            "Oracle EXPLAIN failed: stdout={} stderr={}",
1385            stdout.trim(),
1386            stderr.trim()
1387        ));
1388    }
1389    Ok(())
1390}
1391
1392/// SQL Server 编译期验证:通过 sqlcmd 命令行工具执行 SET SHOWPLAN_TEXT ON
1393///
1394/// DSN 格式:`sqlserver://user:pass@host:port/db`
1395/// 例如:`sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test`
1396#[cfg(feature = "db-verify")]
1397fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
1398    let parsed = parse_sqlserver_dsn(dsn)?;
1399    // sqlcmd -S host,port -U user -P pass -d db -Q "SET SHOWPLAN_TEXT ON; <sql>"
1400    let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
1401    let out = std::process::Command::new("sqlcmd")
1402        .args([
1403            "-S",
1404            &format!("{},{}", parsed.host, parsed.port),
1405            "-U",
1406            &parsed.user,
1407            "-P",
1408            &parsed.password,
1409            "-d",
1410            &parsed.database,
1411            "-Q",
1412            &query,
1413            "-h",
1414            "-1",
1415            "-W",
1416        ])
1417        .output()
1418        .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
1419    let stdout = String::from_utf8_lossy(&out.stdout);
1420    let stderr = String::from_utf8_lossy(&out.stderr);
1421    if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
1422        return Err(format!(
1423            "SQL Server SHOWPLAN failed: stdout={} stderr={}",
1424            stdout.trim(),
1425            stderr.trim()
1426        ));
1427    }
1428    Ok(())
1429}
1430
1431/// Oracle DSN 解析结果
1432#[cfg(feature = "db-verify")]
1433struct OracleDsn {
1434    user: String,
1435    password: String,
1436    host: String,
1437    port: u16,
1438    service: String,
1439    sysdba: bool,
1440}
1441
1442/// 解析 oracle://user:pass@host:port/service?sysdba=1
1443#[cfg(feature = "db-verify")]
1444fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
1445    let raw = dsn
1446        .strip_prefix("oracle://")
1447        .or_else(|| dsn.strip_prefix("oracle:"))
1448        .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
1449    // 分离 query
1450    let (auth_host_service, query) = match raw.find('?') {
1451        Some(idx) => (&raw[..idx], &raw[idx + 1..]),
1452        None => (raw, ""),
1453    };
1454    let sysdba = query
1455        .split('&')
1456        .any(|p| p == "sysdba=1" || p == "sysdba=true");
1457    // user:pass@host:port/service
1458    let at = auth_host_service
1459        .find('@')
1460        .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
1461    let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
1462    let colon = user_pass
1463        .find(':')
1464        .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
1465    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1466    let (host_port, service) = match host_port_service.rfind('/') {
1467        Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
1468        None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
1469    };
1470    let (host, port) = match host_port.find(':') {
1471        Some(idx) => (
1472            &host_port[..idx],
1473            host_port[idx + 1..]
1474                .parse::<u16>()
1475                .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
1476        ),
1477        None => (host_port, 1521u16),
1478    };
1479    Ok(OracleDsn {
1480        user: user.to_string(),
1481        password: password.to_string(),
1482        host: host.to_string(),
1483        port,
1484        service: service.to_string(),
1485        sysdba,
1486    })
1487}
1488
1489/// SQL Server DSN 解析结果
1490#[cfg(feature = "db-verify")]
1491struct SqlServerDsn {
1492    user: String,
1493    password: String,
1494    host: String,
1495    port: u16,
1496    database: String,
1497}
1498
1499/// 解析 sqlserver://user:pass@host:port/db
1500#[cfg(feature = "db-verify")]
1501fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
1502    let raw = dsn
1503        .strip_prefix("sqlserver://")
1504        .or_else(|| dsn.strip_prefix("mssql://"))
1505        .or_else(|| dsn.strip_prefix("tds://"))
1506        .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
1507    let at = raw
1508        .find('@')
1509        .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
1510    let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
1511    let colon = user_pass
1512        .find(':')
1513        .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
1514    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1515    let (host_port, database) = match host_port_db.rfind('/') {
1516        Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
1517        None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
1518    };
1519    let (host, port) = match host_port.find(':') {
1520        Some(idx) => (
1521            &host_port[..idx],
1522            host_port[idx + 1..]
1523                .parse::<u16>()
1524                .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
1525        ),
1526        None => (host_port, 1433u16),
1527    };
1528    Ok(SqlServerDsn {
1529        user: user.to_string(),
1530        password: password.to_string(),
1531        host: host.to_string(),
1532        port,
1533        database: database.to_string(),
1534    })
1535}
1536
1537// ---------------------------------------------------------------------------
1538// Helpers
1539// ---------------------------------------------------------------------------
1540
1541/// Create a compile_error! token stream
1542fn compile_error(span: Span, msg: &str) -> TokenStream {
1543    // emit: compile_error!("msg")
1544    let mut ts = TokenStream::new();
1545    ts.extend([
1546        TokenTree::Ident(Ident::new("compile_error", span)),
1547        TokenTree::Punct(Punct::new('!', Spacing::Alone)),
1548        TokenTree::Group(Group::new(
1549            Delimiter::Parenthesis,
1550            TokenStream::from(TokenTree::Literal(Literal::string(msg))),
1551        )),
1552    ]);
1553    ts
1554}
1555
1556// ---------------------------------------------------------------------------
1557// typed_query! — Diesel 风格强类型 AST 宏
1558// ---------------------------------------------------------------------------
1559
1560/// Diesel 风格强类型 AST 宏(与 `sql_string!` / `query!` 并存)。
1561///
1562/// # 设计
1563///
1564/// 接收 `table { col1: Type, col2: Type, ... }` 声明,生成:
1565/// 1. 一个 `table` 模块
1566/// 2. 每列对应一个零大小标记类型(如 `table::id`)
1567/// 3. 实现 `TypedColumn` trait,把列名 + Rust 类型提升到类型系统
1568///
1569/// 这样,`typed_query!(SELECT id FROM users WHERE name = ?)` 在编译期就能:
1570/// - 校验 `id` / `name` 列是否存在于 `users` 表声明中
1571/// - 校验 `?` 参数的 Rust 类型与列声明的类型一致
1572///
1573/// # 用法
1574///
1575/// ```ignore
1576/// use sz_orm_macros::typed_query;
1577///
1578/// // 1. 声明表 schema(编译期生成 column 标记类型)
1579/// typed_query! {
1580///     table users {
1581///         id: i64,
1582///         name: String,
1583///         email: String,
1584///         age: i32,
1585///     }
1586/// }
1587///
1588/// // 2. 编译期校验 SELECT:列名必须存在于 users 表
1589/// let sql = typed_query!(SELECT id, name FROM users WHERE age > ?);
1590/// // ❌ 编译错误:unknown column 'foo' in table 'users'
1591/// // let sql = typed_query!(SELECT foo FROM users);
1592/// ```
1593#[proc_macro]
1594pub fn typed_query(input: TokenStream) -> TokenStream {
1595    let tokens: Vec<TokenTree> = input.into_iter().collect();
1596
1597    // 分支 1:table 声明
1598    if tokens.iter().any(|t| {
1599        if let TokenTree::Ident(id) = t {
1600            id.to_string() == "table"
1601        } else {
1602            false
1603        }
1604    }) {
1605        return parse_table_decl(&tokens);
1606    }
1607
1608    // 分支 2:SELECT 表达式
1609    if tokens.iter().any(|t| {
1610        if let TokenTree::Ident(id) = t {
1611            id.to_string().eq_ignore_ascii_case("SELECT")
1612        } else {
1613            false
1614        }
1615    }) {
1616        return parse_typed_select(&tokens);
1617    }
1618
1619    compile_error(
1620        Span::call_site(),
1621        "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
1622    )
1623}
1624
1625/// 解析 `table name { col: Type, ... }` 声明
1626fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
1627    // 期望格式:table <ident> { <ident> : <ident> [, ...] }
1628    let mut idx = 0;
1629
1630    // 跳过 'table' 关键字
1631    if idx >= tokens.len() {
1632        return compile_error(Span::call_site(), "expected table name after 'table'");
1633    }
1634    if let TokenTree::Ident(id) = &tokens[idx] {
1635        if id.to_string() != "table" {
1636            return compile_error(id.span(), "expected 'table' keyword");
1637        }
1638    }
1639    idx += 1;
1640
1641    // 表名
1642    let table_name = if idx < tokens.len() {
1643        if let TokenTree::Ident(id) = &tokens[idx] {
1644            id.to_string()
1645        } else {
1646            return compile_error(tokens[idx].span(), "expected table name identifier");
1647        }
1648    } else {
1649        return compile_error(Span::call_site(), "expected table name");
1650    };
1651    idx += 1;
1652
1653    // 表体({} 内)
1654    let body_group = if idx < tokens.len() {
1655        if let TokenTree::Group(g) = &tokens[idx] {
1656            if g.delimiter() != Delimiter::Brace {
1657                return compile_error(g.span(), "expected '{' after table name");
1658            }
1659            g.clone()
1660        } else {
1661            return compile_error(tokens[idx].span(), "expected '{' after table name");
1662        }
1663    } else {
1664        return compile_error(Span::call_site(), "expected table body in '{ }'");
1665    };
1666
1667    // 解析列声明
1668    let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
1669    let columns = match parse_column_list(&body_tokens) {
1670        Ok(c) => c,
1671        Err(e) => return compile_error(Span::call_site(), &e),
1672    };
1673
1674    // 使用 quote! 构建类型安全的 TokenStream
1675    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1676    let table_name_lit = table_name.as_str();
1677
1678    // 为每列构建标记类型 + trait 实现
1679    let col_impls: Vec<TokenStream2> = columns
1680        .iter()
1681        .map(|(col_name, col_type)| {
1682            let col_ident =
1683                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1684            let col_name_lit = col_name.as_str();
1685            // 解析类型字符串为 TokenStream(quote! 会处理)
1686            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1687            quote! {
1688                #[derive(Debug, Clone, Copy)]
1689                pub struct #col_ident;
1690                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1691                    const NAME: &'static str = #col_name_lit;
1692                    type Table = table;
1693                    type RustType = #rust_type;
1694                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1695                }
1696            }
1697        })
1698        .collect();
1699
1700    // schema 常量条目
1701    let schema_entries: Vec<TokenStream2> = columns
1702        .iter()
1703        .map(|(n, t)| {
1704            let n_lit = n.as_str();
1705            let t_lit = t.as_str();
1706            quote! { (#n_lit, #t_lit) }
1707        })
1708        .collect();
1709
1710    let schema_const_ident = proc_macro2::Ident::new(
1711        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1712        Span::call_site().into(),
1713    );
1714
1715    let expanded = quote! {
1716        pub mod #table_ident {
1717            use super::*;
1718            pub struct table;
1719            impl ::sz_orm_core::typed::TypedTable for table {
1720                const NAME: &'static str = #table_name_lit;
1721            }
1722            #(#col_impls)*
1723        }
1724        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1725    };
1726
1727    expanded.into()
1728}
1729
1730/// 解析列声明列表:`col: Type, col2: Type2, ...`
1731fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
1732    let mut cols = Vec::new();
1733    let mut i = 0;
1734    while i < tokens.len() {
1735        // 列名
1736        let col_name = if let TokenTree::Ident(id) = &tokens[i] {
1737            id.to_string()
1738        } else {
1739            return Err(format!("expected column name at position {}", i));
1740        };
1741        i += 1;
1742
1743        // 冒号
1744        if i >= tokens.len() {
1745            return Err(format!("expected ':' after column '{}'", col_name));
1746        }
1747        if let TokenTree::Punct(p) = &tokens[i] {
1748            if p.as_char() != ':' {
1749                return Err(format!("expected ':' after column '{}'", col_name));
1750            }
1751        } else {
1752            return Err(format!("expected ':' after column '{}'", col_name));
1753        }
1754        i += 1;
1755
1756        // 类型(可能是 ident 或 path,如 String / i64 / Option<i64>)
1757        // 简化处理:收集直到遇到 ',' 或末尾
1758        let mut type_str = String::new();
1759        let mut depth = 0;
1760        while i < tokens.len() {
1761            match &tokens[i] {
1762                TokenTree::Punct(p) => {
1763                    if p.as_char() == ',' && depth == 0 {
1764                        i += 1;
1765                        break;
1766                    } else if p.as_char() == '<' || p.as_char() == '(' {
1767                        depth += 1;
1768                        type_str.push(p.as_char());
1769                    } else if p.as_char() == '>' || p.as_char() == ')' {
1770                        depth -= 1;
1771                        type_str.push(p.as_char());
1772                    } else {
1773                        type_str.push(p.as_char());
1774                    }
1775                }
1776                TokenTree::Ident(id) => {
1777                    if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
1778                    {
1779                        type_str.push(' ');
1780                    }
1781                    type_str.push_str(&id.to_string());
1782                }
1783                _ => {}
1784            }
1785            i += 1;
1786        }
1787
1788        cols.push((col_name, type_str.trim().to_string()));
1789    }
1790    Ok(cols)
1791}
1792
1793/// 解析 `SELECT col1, col2 FROM table WHERE col = ?` 表达式
1794///
1795/// 校验列名是否在表 schema 中(通过编译期常量查找)。
1796fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
1797    // 收集所有 ident 与 literal,构造 SQL 字符串
1798    let mut sql_parts: Vec<String> = Vec::new();
1799    let mut table_name: Option<String> = None;
1800    let mut in_from = false;
1801
1802    for (i, t) in tokens.iter().enumerate() {
1803        match t {
1804            TokenTree::Ident(id) => {
1805                let s = id.to_string();
1806                if s.eq_ignore_ascii_case("SELECT") {
1807                    sql_parts.push("SELECT".to_string());
1808                } else if s.eq_ignore_ascii_case("FROM") {
1809                    in_from = true;
1810                    sql_parts.push("FROM".to_string());
1811                } else if s.eq_ignore_ascii_case("WHERE")
1812                    || s.eq_ignore_ascii_case("AND")
1813                    || s.eq_ignore_ascii_case("OR")
1814                    || s.eq_ignore_ascii_case("LIMIT")
1815                    || s.eq_ignore_ascii_case("OFFSET")
1816                    || s.eq_ignore_ascii_case("ORDER")
1817                    || s.eq_ignore_ascii_case("BY")
1818                    || s.eq_ignore_ascii_case("GROUP")
1819                    || s.eq_ignore_ascii_case("HAVING")
1820                    || s.eq_ignore_ascii_case("JOIN")
1821                    || s.eq_ignore_ascii_case("INNER")
1822                    || s.eq_ignore_ascii_case("LEFT")
1823                    || s.eq_ignore_ascii_case("RIGHT")
1824                    || s.eq_ignore_ascii_case("ON")
1825                    || s.eq_ignore_ascii_case("AS")
1826                    || s.eq_ignore_ascii_case("ASC")
1827                    || s.eq_ignore_ascii_case("DESC")
1828                    || s.eq_ignore_ascii_case("DISTINCT")
1829                    || s.eq_ignore_ascii_case("NOT")
1830                    || s.eq_ignore_ascii_case("NULL")
1831                    || s.eq_ignore_ascii_case("IN")
1832                    || s.eq_ignore_ascii_case("BETWEEN")
1833                    || s.eq_ignore_ascii_case("LIKE")
1834                    || s.eq_ignore_ascii_case("IS")
1835                {
1836                    sql_parts.push(s.to_uppercase());
1837                } else if in_from && table_name.is_none() {
1838                    // FROM 后第一个 ident 是表名
1839                    table_name = Some(s.clone());
1840                    sql_parts.push(s.clone());
1841                } else {
1842                    sql_parts.push(s.clone());
1843                }
1844            }
1845            TokenTree::Literal(lit) => {
1846                sql_parts.push(lit.to_string());
1847            }
1848            TokenTree::Punct(p) => {
1849                let c = p.as_char();
1850                // SQL 中常见标点:, ; * ? = > < ( ) . 等
1851                let part = if c == ',' {
1852                    ",".to_string()
1853                } else if c == '?' {
1854                    "?".to_string()
1855                } else if c == '*' {
1856                    "*".to_string()
1857                } else if c == '=' {
1858                    "=".to_string()
1859                } else if c == '>' {
1860                    ">".to_string()
1861                } else if c == '<' {
1862                    "<".to_string()
1863                } else if c == '.' {
1864                    ".".to_string()
1865                } else if c == ';' {
1866                    ";".to_string()
1867                } else {
1868                    c.to_string()
1869                };
1870                sql_parts.push(part);
1871            }
1872            TokenTree::Group(g) => {
1873                // 处理 group(如 (1, 2, 3))
1874                let inner: String = g.stream().to_string();
1875                let delim = match g.delimiter() {
1876                    Delimiter::Parenthesis => "(",
1877                    Delimiter::Brace => "{",
1878                    Delimiter::Bracket => "[",
1879                    Delimiter::None => "",
1880                };
1881                let close = match g.delimiter() {
1882                    Delimiter::Parenthesis => ")",
1883                    Delimiter::Brace => "}",
1884                    Delimiter::Bracket => "]",
1885                    Delimiter::None => "",
1886                };
1887                sql_parts.push(format!("{}{}{}", delim, inner, close));
1888            }
1889        }
1890        // 单空格分隔(去重多个空格由 trim 处理)
1891        let _ = i;
1892    }
1893
1894    let sql = sql_parts
1895        .join(" ")
1896        .replace(", ", ",")
1897        .replace(" ,", ",")
1898        .replace("= ", "=")
1899        .replace(" =", "=")
1900        .replace("> ", ">")
1901        .replace(" >", ">")
1902        .replace("< ", "<")
1903        .replace(" <", "<")
1904        .replace("  ", " ");
1905
1906    // 验证 SQL 语法
1907    if let Err(e) = validate_sql_content(&sql, None) {
1908        return compile_error(
1909            Span::call_site(),
1910            &format!("typed_query! SQL validation failed: {}", e),
1911        );
1912    }
1913
1914    // 生成 SQL 字符串字面量
1915    let mut ts = TokenStream::new();
1916    let lit = Literal::string(&sql);
1917    ts.extend([TokenTree::Literal(lit)]);
1918    ts
1919}
1920
1921// ---------------------------------------------------------------------------
1922// schema! — Compile-time SQL schema generator
1923// ---------------------------------------------------------------------------
1924
1925/// Compile-time SQL schema generator.
1926///
1927/// Parses a SQL `CREATE TABLE` statement and generates typed table declarations
1928/// equivalent to `typed_query! { table ... }`.
1929///
1930/// # Syntax
1931///
1932/// ```ignore
1933/// use sz_orm_macros::schema;
1934///
1935/// schema! {
1936///     "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL, email TEXT)"
1937/// }
1938/// ```
1939///
1940/// 生成与以下手动声明等价的代码:
1941/// ```ignore
1942/// typed_query! {
1943///     table users {
1944///         id: i64,
1945///         name: String,
1946///         email: Option<String>,
1947///     }
1948/// }
1949/// ```
1950#[proc_macro]
1951/// 类型化裸 SQL 查询宏(SQLx `query_as!` 风格)。
1952///
1953/// 用法:`query_as!(RecordType, "SELECT col1, col2 FROM table WHERE id = ?")`
1954///
1955/// 生成 `sz_orm_core::queryable::QueryAs::<RecordType>::new("SELECT ...")`。
1956/// 在 `db-verify` feature + `SZ_ORM_QUERY_VERIFY=1` 环境下会连真 DB
1957/// 执行 EXPLAIN 验证 SQL 合法性。
1958///
1959/// **运行时列名验证**(P0-2):`QueryAs::fetch_all` 会比对 DB 返回的列名
1960/// 与 `RecordType::row_desc()`(由 `#[derive(FromQueryResult)]` 自动生成)。
1961/// 若 SQL SELECT 列不在 struct 字段中,返回 `DbError::QueryError`。
1962///
1963/// # 示例
1964///
1965/// ```ignore
1966/// #[derive(FromQueryResult)]
1967/// struct User { id: i64, name: String }
1968///
1969/// let q = query_as!(User, "SELECT id, name FROM users WHERE id = 1");
1970/// let users: Vec<User> = q.fetch_all(&mut conn).await?;
1971/// ```
1972pub fn query_as(input: TokenStream) -> TokenStream {
1973    let mut tokens = input.into_iter().peekable();
1974
1975    // 解析记录类型(第一个标识符/路径,如 User 或 crate::User)
1976    let mut record_type = String::new();
1977    loop {
1978        match tokens.next() {
1979            Some(TokenTree::Ident(ident)) => {
1980                record_type.push_str(&ident.to_string());
1981            }
1982            Some(TokenTree::Punct(p)) if p.as_char() == ':' => {
1983                // 处理 :: 路径分隔符
1984                record_type.push_str("::");
1985                // 跳过第二个 :
1986                if let Some(TokenTree::Punct(p2)) = tokens.peek() {
1987                    if p2.as_char() == ':' {
1988                        let _ = tokens.next();
1989                    }
1990                }
1991            }
1992            Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
1993            Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
1994            Some(other) => {
1995                return compile_error(
1996                    other.span(),
1997                    "query_as! 第一个参数必须是记录类型,如 query_as!(User, \"SELECT ...\")",
1998                );
1999            }
2000            None => {
2001                return compile_error(
2002                    Span::call_site(),
2003                    "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2004                );
2005            }
2006        }
2007    }
2008
2009    // 解析 SQL 字符串字面量
2010    let sql_raw = match tokens.next() {
2011        Some(TokenTree::Literal(lit)) => lit.to_string(),
2012        Some(other) => {
2013            return compile_error(other.span(), "query_as! 第二个参数必须是 SQL 字符串字面量");
2014        }
2015        None => {
2016            return compile_error(
2017                Span::call_site(),
2018                "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2019            );
2020        }
2021    };
2022
2023    let sql_content = match strip_string_literal(&sql_raw) {
2024        Some(s) => s,
2025        None => {
2026            return compile_error(Span::call_site(), "query_as! 的 SQL 参数必须是字符串字面量");
2027        }
2028    };
2029
2030    // 语法验证
2031    if let Err(err_msg) = validate_sql_content(sql_content, None) {
2032        return compile_error(Span::call_site(), &err_msg);
2033    }
2034
2035    // db-verify 验证
2036    #[cfg(feature = "db-verify")]
2037    let verify_cols: Option<Vec<(String, String)>> = {
2038        match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
2039            // 模式 1:连真 DB 执行 EXPLAIN 验证,并获取 SELECT 列的实际类型
2040            Some("1") => match verify_with_real_db(sql_content) {
2041                Ok(cols) => Some(cols),
2042                Err(err) => {
2043                    return compile_error(
2044                        Span::call_site(),
2045                        &format!("query_as! real DB verification failed: {}", err),
2046                    )
2047                }
2048            },
2049            // 模式 cache:从离线缓存文件查找(无需 DB,适合 CI)
2050            Some("cache") => {
2051                if let Err(err) = verify_with_cache(sql_content) {
2052                    return compile_error(
2053                        Span::call_site(),
2054                        &format!("query_as! offline cache verification failed: {}", err),
2055                    );
2056                }
2057                None
2058            }
2059            _ => None,
2060        }
2061    };
2062    #[cfg(not(feature = "db-verify"))]
2063    let _verify_cols: Option<Vec<(String, String)>> = None;
2064
2065    // 生成 QueryAs::<T>::new("...")
2066    // db-verify 通过时,附加编译期类型验证块(P0-2):
2067    // const 上下文中将 DB 实际列类型与结构体 __sz_orm_column_types() 对比,
2068    // 不匹配则 const panic → 编译失败。
2069    let escaped = sql_content.escape_default();
2070    let base = format!(
2071        "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
2072        record_type, escaped
2073    );
2074    #[cfg(feature = "db-verify")]
2075    let output = match &verify_cols {
2076        Some(cols) if !cols.is_empty() => {
2077            gen_compile_time_type_check(&record_type, sql_content, cols, &base)
2078        }
2079        _ => base,
2080    };
2081    #[cfg(not(feature = "db-verify"))]
2082    let output = base;
2083    output
2084        .parse()
2085        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query_as output"))
2086}
2087
2088#[proc_macro]
2089pub fn schema(input: TokenStream) -> TokenStream {
2090    let mut tokens = input.into_iter().peekable();
2091
2092    // 解析 SQL 字符串字面量
2093    let sql_raw = match tokens.next() {
2094        Some(TokenTree::Literal(lit)) => lit.to_string(),
2095        Some(other) => {
2096            return compile_error(
2097                other.span(),
2098                "Expected a string literal as the argument to schema!",
2099            );
2100        }
2101        None => {
2102            return compile_error(
2103                Span::call_site(),
2104                "Expected a string literal argument to schema!",
2105            );
2106        }
2107    };
2108
2109    let sql = match strip_string_literal(&sql_raw) {
2110        Some(s) => s,
2111        None => {
2112            return compile_error(
2113                Span::call_site(),
2114                "schema! requires a string literal argument",
2115            );
2116        }
2117    };
2118
2119    // 解析 CREATE TABLE
2120    let (table_name, columns) = match parse_create_table(sql) {
2121        Ok(v) => v,
2122        Err(e) => return compile_error(Span::call_site(), &e),
2123    };
2124
2125    // 生成代码(与 parse_table_decl 一致)
2126    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
2127    let table_name_lit = table_name.as_str();
2128
2129    let col_impls: Vec<TokenStream2> = columns
2130        .iter()
2131        .map(|(col_name, col_type)| {
2132            let col_ident =
2133                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
2134            let col_name_lit = col_name.as_str();
2135            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
2136            quote! {
2137                #[derive(Debug, Clone, Copy)]
2138                pub struct #col_ident;
2139                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
2140                    const NAME: &'static str = #col_name_lit;
2141                    type Table = table;
2142                    type RustType = #rust_type;
2143                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
2144                }
2145            }
2146        })
2147        .collect();
2148
2149    let schema_entries: Vec<TokenStream2> = columns
2150        .iter()
2151        .map(|(n, t)| {
2152            let n_lit = n.as_str();
2153            let t_lit = t.as_str();
2154            quote! { (#n_lit, #t_lit) }
2155        })
2156        .collect();
2157
2158    let schema_const_ident = proc_macro2::Ident::new(
2159        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
2160        Span::call_site().into(),
2161    );
2162
2163    let expanded = quote! {
2164        pub mod #table_ident {
2165            use super::*;
2166            pub struct table;
2167            impl ::sz_orm_core::typed::TypedTable for table {
2168                const NAME: &'static str = #table_name_lit;
2169            }
2170            #(#col_impls)*
2171        }
2172        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
2173    };
2174
2175    expanded.into()
2176}
2177
2178/// 解析 SQL `CREATE TABLE` 语句,返回 (表名, Vec<(列名, Rust 类型字符串)>)。
2179///
2180/// 支持以下语法:
2181/// - `CREATE TABLE [IF NOT EXISTS] <name> ( ... )`
2182/// - 表名/列名可带反引号、双引号或无引号
2183/// - 跳过 PRIMARY KEY / FOREIGN KEY / CONSTRAINT / UNIQUE / INDEX / KEY 约束行
2184/// - 列定义按顶层逗号分隔(嵌套括号如 DECIMAL(10,2) 不拆分)
2185fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
2186    let trimmed = sql.trim();
2187    let upper = trimmed.to_uppercase();
2188
2189    // 必须以 CREATE TABLE 开头
2190    if !upper.starts_with("CREATE TABLE") {
2191        return Err("schema! expects a CREATE TABLE statement".to_string());
2192    }
2193
2194    // 跳过 "CREATE TABLE"
2195    let mut rest = &trimmed["CREATE TABLE".len()..];
2196
2197    // 跳过可选的 "IF NOT EXISTS"
2198    let rest_upper = rest.trim_start().to_uppercase();
2199    if rest_upper.starts_with("IF NOT EXISTS") {
2200        rest = &rest.trim_start()["IF NOT EXISTS".len()..];
2201    }
2202
2203    rest = rest.trim_start();
2204
2205    // 解析表名(可能带反引号、双引号或无引号)
2206    let (table_name, after_name) = parse_identifier(rest)?;
2207    let rest = after_name.trim_start();
2208
2209    // 找到列定义起始的 '(' 与匹配的最后一个 ')'
2210    let paren_start = rest
2211        .find('(')
2212        .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
2213    let paren_end = rest
2214        .rfind(')')
2215        .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
2216    if paren_end <= paren_start {
2217        return Err("CREATE TABLE has malformed parentheses".to_string());
2218    }
2219
2220    let cols_str = &rest[paren_start + 1..paren_end];
2221
2222    // 按顶层逗号分隔列定义(注意嵌套括号,如 DECIMAL(10,2))
2223    let col_defs = split_top_level_commas(cols_str);
2224
2225    let mut columns = Vec::new();
2226    for def in col_defs {
2227        let def = def.trim();
2228        if def.is_empty() {
2229            continue;
2230        }
2231
2232        // 跳过约束定义行
2233        let def_upper = def.to_uppercase();
2234        if def_upper.starts_with("PRIMARY KEY")
2235            || def_upper.starts_with("FOREIGN KEY")
2236            || def_upper.starts_with("CONSTRAINT")
2237            || def_upper.starts_with("UNIQUE")
2238            || def_upper.starts_with("INDEX")
2239            || def_upper.starts_with("KEY")
2240        {
2241            continue;
2242        }
2243
2244        // 解析列名
2245        let (col_name, after_col) = parse_identifier(def)?;
2246        let rest = after_col.trim_start();
2247
2248        // 解析类型(取第一个 token,去掉括号参数)
2249        let (sql_type, after_type) = parse_type_token(rest)?;
2250        let rest = after_type.trim();
2251
2252        // 判断 nullability:NOT NULL 或 PRIMARY KEY 隐含 NOT NULL
2253        let rest_upper = rest.to_uppercase();
2254        let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
2255        let rust_type = sql_type_to_rust(&sql_type, !not_null);
2256
2257        columns.push((col_name, rust_type));
2258    }
2259
2260    Ok((table_name, columns))
2261}
2262
2263/// 解析标识符:支持反引号、双引号或无引号。
2264/// 返回 (标识符, 剩余字符串)。
2265fn parse_identifier(s: &str) -> Result<(String, &str), String> {
2266    let s = s.trim_start();
2267    if s.is_empty() {
2268        return Err("expected identifier".to_string());
2269    }
2270
2271    let bytes = s.as_bytes();
2272    match bytes[0] {
2273        b'`' => {
2274            let end = s[1..]
2275                .find('`')
2276                .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
2277            let ident = s[1..1 + end].to_string();
2278            Ok((ident, &s[1 + end + 1..]))
2279        }
2280        b'"' => {
2281            let end = s[1..]
2282                .find('"')
2283                .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
2284            let ident = s[1..1 + end].to_string();
2285            Ok((ident, &s[1 + end + 1..]))
2286        }
2287        _ => {
2288            let end = s
2289                .find(|c: char| !c.is_alphanumeric() && c != '_')
2290                .unwrap_or(s.len());
2291            if end == 0 {
2292                return Err(format!("invalid identifier: '{}'", s));
2293            }
2294            let ident = s[..end].to_string();
2295            Ok((ident, &s[end..]))
2296        }
2297    }
2298}
2299
2300/// 解析类型 token:取第一个标识符,可选跟随括号参数(如 VARCHAR(255) → VARCHAR)。
2301/// 返回 (类型名, 剩余字符串)。
2302fn parse_type_token(s: &str) -> Result<(String, &str), String> {
2303    let s = s.trim_start();
2304    if s.is_empty() {
2305        return Err("expected column type".to_string());
2306    }
2307
2308    let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
2309    if end == 0 {
2310        return Err(format!("invalid type: '{}'", s));
2311    }
2312    let type_name = s[..end].to_string();
2313    let mut rest = &s[end..];
2314
2315    // 跳过可选的括号参数,如 (255) 或 (10,2)
2316    rest = rest.trim_start();
2317    if rest.starts_with('(') {
2318        let close = rest
2319            .find(')')
2320            .ok_or_else(|| "unterminated type parameter list".to_string())?;
2321        rest = &rest[close + 1..];
2322    }
2323
2324    Ok((type_name, rest))
2325}
2326
2327/// 按顶层逗号分隔字符串(不进入嵌套括号)。
2328fn split_top_level_commas(s: &str) -> Vec<String> {
2329    let mut parts = Vec::new();
2330    let mut depth: i32 = 0;
2331    let mut current = String::new();
2332
2333    for ch in s.chars() {
2334        match ch {
2335            '(' => {
2336                depth += 1;
2337                current.push(ch);
2338            }
2339            ')' => {
2340                depth -= 1;
2341                current.push(ch);
2342            }
2343            ',' if depth == 0 => {
2344                parts.push(std::mem::take(&mut current));
2345            }
2346            _ => {
2347                current.push(ch);
2348            }
2349        }
2350    }
2351
2352    if !current.trim().is_empty() {
2353        parts.push(current);
2354    }
2355
2356    parts
2357}
2358
2359/// 将 SQL 类型映射为 Rust 类型字符串。
2360///
2361/// 匹配规则:取类型名第一个 token(去掉括号参数),不区分大小写匹配。
2362/// 未识别的类型默认映射为 `String`。若 `nullable == true`,用 `Option<T>` 包裹。
2363fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
2364    let upper = sql_type.to_uppercase();
2365    let rust = match upper.as_str() {
2366        // 8 字节整数
2367        "BIGINT" | "INT8" => "i64",
2368        // 4 字节整数(INT/INTEGER/INT4/SERIAL)
2369        "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
2370        // 2 字节整数
2371        "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
2372        // 1 字节整数
2373        "TINYINT" => "i8",
2374        // 浮点(4 字节)
2375        "FLOAT" | "REAL" | "FLOAT4" => "f32",
2376        // 浮点(8 字节)/ 定点数
2377        "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
2378        // 布尔
2379        "BOOLEAN" | "BOOL" => "bool",
2380        // 二进制(与 schema_gen::sql_type_to_rust 保持一致)
2381        "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
2382        // 字符串/日期/JSON/UUID(统一映射到 String,运行时再解析)
2383        "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
2384        | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
2385        _ => "String",
2386    };
2387
2388    if nullable {
2389        format!("Option<{}>", rust)
2390    } else {
2391        rust.to_string()
2392    }
2393}
2394
2395// ---------------------------------------------------------------------------
2396// `#[derive(Schema)]` — auto-generate table structure from a struct
2397// ---------------------------------------------------------------------------
2398
2399/// 派生宏:自动从 Rust 结构体生成表结构信息。
2400///
2401/// 解析 `#[table(name = "...")]` 和 `#[column(...)]` 属性,
2402/// 生成 `Schema` trait 实现,便于在运行时反射表名与列信息。
2403///
2404/// # 支持的属性
2405///
2406/// - `#[table(name = "users")]` — 指定表名(默认使用结构体名的蛇形形式)
2407/// - `#[column(name = "user_id")]` — 指定列名(默认使用字段名)
2408/// - `#[column(type = "VARCHAR(255)")]` — 指定 SQL 类型
2409/// - `#[column(primary_key)]` — 标记主键
2410/// - `#[column(nullable)]` — 显式标记允许 NULL
2411/// - `#[column(skip)]` — 跳过此字段,不生成 schema 条目
2412/// - `#[column(default = "0")]` — 标记字段有默认值
2413///
2414/// # 类型推断
2415///
2416/// 字段的 Rust 类型会自动映射为 SQL 类型:
2417/// - `i64`/`u64` → `BIGINT`
2418/// - `i32`/`u32` → `INTEGER`
2419/// - `String` → `TEXT`
2420/// - `f64` → `DOUBLE`
2421/// - `bool` → `BOOLEAN`
2422/// - `Vec<u8>` → `BLOB`
2423/// - `Option<T>` → 与 `T` 相同,但标记为 nullable
2424#[proc_macro_derive(Schema, attributes(table, column))]
2425pub fn derive_schema(input: TokenStream) -> TokenStream {
2426    let input = parse_macro_input!(input as syn::DeriveInput);
2427    derive::derive_schema_impl(input).into()
2428}
2429
2430// ---------------------------------------------------------------------------
2431// `#[derive(Builder)]` — auto-generate builder pattern code
2432// ---------------------------------------------------------------------------
2433
2434/// 派生宏:自动生成构造器模式代码。
2435///
2436/// 为目标结构体生成一个 `XxxBuilder` 类型,包含:
2437/// - `new()` 构造空 builder
2438/// - 每个字段的 setter 方法
2439/// - `build()` 方法返回 `Result<T, String>`
2440///
2441/// # 支持的属性
2442///
2443/// - `#[builder(skip)]` — 跳过此字段(不生成 setter,使用 Default)
2444/// - `#[builder(default = expr)]` — 指定默认值表达式
2445///
2446/// # 示例
2447///
2448/// ```ignore
2449/// use sz_orm_macros::Builder;
2450///
2451/// #[derive(Builder)]
2452/// struct User {
2453///     id: i64,
2454///     name: String,
2455/// }
2456///
2457/// let user = User::builder()
2458///     .id(1)
2459///     .name("Alice".to_string())
2460///     .build()
2461///     .unwrap();
2462/// ```
2463#[proc_macro_derive(Builder, attributes(builder))]
2464pub fn derive_builder(input: TokenStream) -> TokenStream {
2465    let input = parse_macro_input!(input as syn::DeriveInput);
2466    derive::derive_builder_impl(input).into()
2467}
2468
2469// ---------------------------------------------------------------------------
2470// `#[derive(Entity)]` — auto-generate `impl Model for Struct`
2471// ---------------------------------------------------------------------------
2472
2473/// 派生宏:自动生成 `sz_orm_core::Model` trait 实现。
2474///
2475/// 要求结构体恰好有一个 `#[column(primary_key)]` 字段,
2476/// 该字段的类型即为 `Model::PrimaryKey`。
2477///
2478/// # 支持的属性
2479///
2480/// - `#[table(name = "...")]` — 指定表名,默认蛇形结构体名
2481/// - `#[column(primary_key)]` — 标记主键字段(必需,恰好一个)
2482/// - `#[column(name = "...")]` — 覆盖主键列名(默认与字段名相同)
2483///
2484/// # 示例
2485///
2486/// ```ignore
2487/// use sz_orm_macros::Entity;
2488///
2489/// #[derive(Entity)]
2490/// #[table(name = "users")]
2491/// struct User {
2492///     #[column(primary_key)]
2493///     id: i64,
2494///     name: String,
2495/// }
2496///
2497/// assert_eq!(User::table_name(), "users");
2498/// assert_eq!(User::pk_name(), "id");
2499/// ```
2500#[proc_macro_derive(Entity, attributes(table, column))]
2501pub fn derive_entity(input: TokenStream) -> TokenStream {
2502    let input = parse_macro_input!(input as syn::DeriveInput);
2503    derive::derive_entity_impl(input).into()
2504}
2505
2506// ---------------------------------------------------------------------------
2507// `#[derive(FromQueryResult)]` — auto-generate `impl FromQueryResult for Struct`
2508// ---------------------------------------------------------------------------
2509
2510/// 派生宏:自动生成 `sz_orm_core::FromQueryResult` trait 实现。
2511///
2512/// 从查询结果行(`HashMap<String, Value>`)反序列化为结构体实例。
2513/// `Option<T>` 字段在列缺失或值为 NULL 时自动返回 `None`。
2514///
2515/// # 支持的属性
2516///
2517/// - `#[column(name = "...")]` — 覆盖列名映射(默认使用字段名)
2518///
2519/// # 示例
2520///
2521/// ```ignore
2522/// use sz_orm_macros::FromQueryResult;
2523///
2524/// #[derive(FromQueryResult)]
2525/// struct UserRow {
2526///     id: i64,
2527///     name: String,
2528///     #[column(name = "user_email")]
2529///     email: Option<String>,
2530/// }
2531/// ```
2532#[proc_macro_derive(FromQueryResult, attributes(column))]
2533pub fn derive_from_query_result(input: TokenStream) -> TokenStream {
2534    let input = parse_macro_input!(input as syn::DeriveInput);
2535    derive::derive_from_query_result_impl(input).into()
2536}
2537
2538// ---------------------------------------------------------------------------
2539// `#[derive(ColumnEnum)]` — auto-generate column name enum (P2-2)
2540// ---------------------------------------------------------------------------
2541
2542/// 派生宏:从结构体字段自动生成 `<StructName>Column` 列名枚举(P2-2)。
2543///
2544/// 每个字段生成一个变体(snake_case → CamelCase),通过 `ColumnTrait::as_str()`
2545/// 返回数据库列名;`#[column(name = "...")]` 可覆盖列名(与 FromQueryResult 一致)。
2546/// 同时实现 `std::fmt::Display`。
2547///
2548/// # 示例
2549///
2550/// ```rust,ignore
2551/// use sz_orm_macros::ColumnEnum;
2552/// use sz_orm_core::ColumnTrait;
2553///
2554/// #[derive(ColumnEnum)]
2555/// struct User {
2556///     id: i64,
2557///     #[column(name = "user_name")]
2558///     name: String,
2559/// }
2560///
2561/// assert_eq!(UserColumn::Id.as_str(), "id");
2562/// assert_eq!(UserColumn::Name.as_str(), "user_name");
2563/// assert_eq!(UserColumn::Id.to_string(), "id");
2564/// ```
2565#[proc_macro_derive(ColumnEnum, attributes(column))]
2566pub fn derive_column_enum(input: TokenStream) -> TokenStream {
2567    let input = parse_macro_input!(input as syn::DeriveInput);
2568    derive::derive_column_enum_impl(input).into()
2569}
2570
2571// ---------------------------------------------------------------------------
2572// `#[derive(FromRow)]` — auto-generate `impl FromRow for Struct`
2573// ---------------------------------------------------------------------------
2574
2575/// 派生宏:自动生成 `sz_orm_core::queryable::FromRow` trait 实现。
2576///
2577/// 从 `HashMap<String, Value>` 按列名反序列化为结构体实例。
2578/// 与 `FromQueryResult` 的区别在于错误类型为 `QueryError`(含列信息),
2579/// 适合需要精确错误定位的底层场景。
2580///
2581/// # 支持的属性
2582///
2583/// - `#[column(name = "...")]` — 覆盖列名映射(默认使用字段名)
2584///
2585/// # 示例
2586///
2587/// ```ignore
2588/// use sz_orm_macros::FromRow;
2589///
2590/// #[derive(FromRow)]
2591/// struct User {
2592///     id: i64,
2593///     name: String,
2594///     #[column(name = "user_email")]
2595///     email: Option<String>,
2596/// }
2597/// ```
2598#[proc_macro_derive(FromRow, attributes(column))]
2599pub fn derive_from_row(input: TokenStream) -> TokenStream {
2600    let input = parse_macro_input!(input as syn::DeriveInput);
2601    derive::derive_from_row_impl(input).into()
2602}
2603
2604// ---------------------------------------------------------------------------
2605// `#[derive(SqlType)]` — auto-generate `impl FromQueryResult + to_value()` for enums
2606// ---------------------------------------------------------------------------
2607
2608/// 派生宏:为 Rust 枚举自动生成 `sz_orm_core::FromQueryResult` trait 实现
2609/// 和 `to_value()` 方法。
2610///
2611/// 这是 sz-orm 对 SQLx `#[derive(Type)]` 的等效实现:
2612/// 让自定义枚举可以直接用于查询结果的字段映射和查询参数的绑定。
2613///
2614/// # 支持的属性
2615///
2616/// - `#[sql_type(rename_all = "snake_case")]` — 控制变体名的序列化格式
2617///   (snake_case / SCREAMING_SNAKE_CASE / camelCase / PascalCase / lowercase / UPPERCASE)
2618/// - `#[sql_type(rename = "...")]`(变体级)— 覆盖单个变体的序列化名
2619///
2620/// # 示例
2621///
2622/// ```ignore
2623/// use sz_orm_macros::SqlType;
2624///
2625/// #[derive(SqlType)]
2626/// enum Status {
2627///     Active,    // → "active"
2628///     Inactive,  // → "inactive"
2629/// }
2630///
2631/// let v = Status::Active.to_value();  // Value::String("active")
2632/// ```
2633#[proc_macro_derive(SqlType, attributes(sql_type))]
2634pub fn derive_sql_type(input: TokenStream) -> TokenStream {
2635    let input = parse_macro_input!(input as syn::DeriveInput);
2636    derive::derive_sql_type_impl(input).into()
2637}
2638
2639// ---------------------------------------------------------------------------
2640// `#[derive(Relation)]` — auto-generate `impl ModelExt` with relations()
2641// ---------------------------------------------------------------------------
2642
2643/// 派生宏:自动生成 `sz_orm_core::model::ModelExt` trait 实现,
2644/// 填充 `relations()` 映射,消除手写关系样板代码。
2645///
2646/// # 支持的属性
2647///
2648/// - `#[relation(has_many = "orders", fk = "user_id", pk = "id")]`
2649/// - `#[relation(belongs_to = "users", fk = "user_id", pk = "id")]`
2650/// - `#[relation(has_one = "profile", fk = "user_id", pk = "id")]`
2651/// - `#[relation(belongs_to_many = "roles", junction = "user_roles", fk = "user_id", other_key = "role_id", target = "roles", target_pk = "id")]`
2652/// - `#[relation(morph_many = "comments", morph_type = "commentable_type", morph_id = "commentable_id", morph_type_value = "Post")]`
2653/// - `#[relation(morph_to, morph_type = "commentable_type", morph_id = "commentable_id")]`
2654///
2655/// # 示例
2656///
2657/// ```ignore
2658/// use sz_orm_macros::{Entity, Relation};
2659///
2660/// #[derive(Entity, Relation)]
2661/// #[table(name = "users")]
2662/// struct User {
2663///     #[column(primary_key)]
2664///     id: i64,
2665/// }
2666///
2667/// // 自动生成:
2668/// // impl ModelExt for User {
2669/// //     fn relations() -> HashMap<&str, Relation> {
2670/// //         // 包含 #[relation] 定义的关系
2671/// //     }
2672/// // }
2673/// ```
2674#[proc_macro_derive(Relation, attributes(relation, table, column))]
2675pub fn derive_relation(input: TokenStream) -> TokenStream {
2676    let input = parse_macro_input!(input as syn::DeriveInput);
2677    derive::derive_relation_impl(input).into()
2678}
2679
2680// ---------------------------------------------------------------------------
2681// Unit tests — cover helper functions used by both macros
2682// ---------------------------------------------------------------------------
2683
2684#[cfg(test)]
2685mod tests {
2686    use super::*;
2687
2688    // ---- strip_string_literal ----
2689
2690    #[test]
2691    fn test_strip_plain_double_quoted() {
2692        assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
2693    }
2694
2695    #[test]
2696    fn test_strip_raw_double_hash() {
2697        assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
2698    }
2699
2700    #[test]
2701    fn test_strip_raw_double_no_hash() {
2702        assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
2703    }
2704
2705    #[test]
2706    fn test_strip_byte_string() {
2707        assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
2708        assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
2709    }
2710
2711    #[test]
2712    fn test_strip_non_string_returns_none() {
2713        assert_eq!(strip_string_literal("123"), None);
2714        assert_eq!(strip_string_literal("foo"), None);
2715    }
2716
2717    // ---- validate_sql_content ----
2718
2719    #[test]
2720    fn test_validate_select_with_from_ok() {
2721        assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
2722    }
2723
2724    #[test]
2725    fn test_validate_select_missing_from_fails() {
2726        assert!(validate_sql_content("SELECT * users", None).is_err());
2727    }
2728
2729    #[test]
2730    fn test_validate_insert_missing_into_fails() {
2731        assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
2732        assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
2733    }
2734
2735    #[test]
2736    fn test_validate_update_missing_set_fails() {
2737        assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
2738        assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
2739    }
2740
2741    #[test]
2742    fn test_validate_delete_missing_from_fails() {
2743        assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
2744        assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
2745    }
2746
2747    #[test]
2748    fn test_validate_empty_sql_fails() {
2749        assert!(validate_sql_content("", None).is_err());
2750        assert!(validate_sql_content("   ", None).is_err());
2751    }
2752
2753    // ---- balanced parens ----
2754
2755    #[test]
2756    fn test_validate_balanced_parens_ok() {
2757        assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
2758    }
2759
2760    #[test]
2761    fn test_validate_balanced_parens_unbalanced() {
2762        assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
2763        assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
2764    }
2765
2766    // ---- injection patterns ----
2767
2768    #[test]
2769    fn test_validate_no_injection_clean() {
2770        assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
2771    }
2772
2773    #[test]
2774    fn test_validate_no_injection_drop_table() {
2775        assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
2776    }
2777
2778    #[test]
2779    fn test_validate_no_injection_or_1_1() {
2780        // 编译期 SQL 已剥离外层引号,检测模式不再依赖引号字符。
2781        // "' OR '1'='1" 因引号分隔不再匹配 "or 1=1",故不再检测;
2782        // 但不含引号分隔的 "OR 1=1" 仍可被检测。
2783        assert!(validate_no_injection("' OR 1=1").is_err());
2784        assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
2785    }
2786
2787    #[test]
2788    fn test_validate_no_injection_drop_database() {
2789        assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
2790    }
2791
2792    #[test]
2793    fn test_validate_no_injection_information_schema() {
2794        assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
2795    }
2796
2797    #[test]
2798    fn test_validate_no_injection_xp_cmdshell() {
2799        assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
2800    }
2801
2802    #[test]
2803    fn test_validate_no_injection_union_select() {
2804        assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
2805    }
2806
2807    #[test]
2808    fn test_validate_no_injection_comment_dashes() {
2809        assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
2810    }
2811
2812    #[test]
2813    fn test_validate_no_injection_block_comment() {
2814        assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
2815    }
2816
2817    // ---- string literal closure ----
2818
2819    #[test]
2820    fn test_validate_string_literals_closed_ok() {
2821        assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
2822        assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
2823    }
2824
2825    #[test]
2826    fn test_validate_string_literals_closed_unclosed_single() {
2827        assert!(validate_string_literals_closed("'hello").is_err());
2828    }
2829
2830    #[test]
2831    fn test_validate_string_literals_closed_unclosed_double() {
2832        assert!(validate_string_literals_closed(r#""hello"#).is_err());
2833    }
2834
2835    // ---- param count check ----
2836
2837    #[test]
2838    fn test_validate_param_count_match() {
2839        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
2840        assert!(
2841            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
2842        );
2843    }
2844
2845    #[test]
2846    fn test_validate_param_count_mismatch() {
2847        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
2848        assert!(
2849            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
2850        );
2851    }
2852
2853    // ---- db-verify feature: detect_db_kind ----
2854
2855    #[cfg(feature = "db-verify")]
2856    #[test]
2857    fn test_detect_db_kind_mysql() {
2858        assert_eq!(
2859            detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
2860            DbKind::MySql
2861        );
2862    }
2863
2864    #[cfg(feature = "db-verify")]
2865    #[test]
2866    fn test_detect_db_kind_postgres() {
2867        assert_eq!(
2868            detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
2869            DbKind::Postgres
2870        );
2871        assert_eq!(
2872            detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
2873            DbKind::Postgres
2874        );
2875    }
2876
2877    #[cfg(feature = "db-verify")]
2878    #[test]
2879    fn test_detect_db_kind_sqlite() {
2880        assert_eq!(
2881            detect_db_kind("sqlite://path/to/db.db").unwrap(),
2882            DbKind::Sqlite
2883        );
2884        assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
2885    }
2886
2887    #[cfg(feature = "db-verify")]
2888    #[test]
2889    fn test_detect_db_kind_oracle() {
2890        assert_eq!(
2891            detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
2892            DbKind::Oracle
2893        );
2894        assert_eq!(
2895            detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
2896            DbKind::Oracle
2897        );
2898    }
2899
2900    #[cfg(feature = "db-verify")]
2901    #[test]
2902    fn test_detect_db_kind_sqlserver() {
2903        assert_eq!(
2904            detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
2905            DbKind::SqlServer
2906        );
2907        assert_eq!(
2908            detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
2909            DbKind::SqlServer
2910        );
2911        assert_eq!(
2912            detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
2913            DbKind::SqlServer
2914        );
2915    }
2916
2917    #[cfg(feature = "db-verify")]
2918    #[test]
2919    fn test_detect_db_kind_unsupported() {
2920        assert!(detect_db_kind("redis://user:pass@host/db").is_err());
2921        assert!(detect_db_kind("not-a-url").is_err());
2922    }
2923
2924    #[cfg(feature = "db-verify")]
2925    #[test]
2926    fn test_parse_oracle_dsn_basic() {
2927        let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
2928        let p = parse_oracle_dsn(dsn).unwrap();
2929        assert_eq!(p.user, "sys");
2930        assert_eq!(p.password, "test123");
2931        assert_eq!(p.host, "127.0.0.1");
2932        assert_eq!(p.port, 1521);
2933        assert_eq!(p.service, "freepdb1.FALSE");
2934        assert!(p.sysdba);
2935    }
2936
2937    #[cfg(feature = "db-verify")]
2938    #[test]
2939    fn test_parse_oracle_dsn_default_port() {
2940        // 无端口号时默认 1521
2941        let dsn = "oracle://sys:test123@127.0.0.1/FREE";
2942        let p = parse_oracle_dsn(dsn).unwrap();
2943        assert_eq!(p.port, 1521);
2944        assert_eq!(p.service, "FREE");
2945        assert!(!p.sysdba);
2946    }
2947
2948    #[cfg(feature = "db-verify")]
2949    #[test]
2950    fn test_parse_sqlserver_dsn_basic() {
2951        let dsn =
2952            "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
2953        let p = parse_sqlserver_dsn(dsn).unwrap();
2954        assert_eq!(p.user, "test");
2955        assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
2956        assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
2957        assert_eq!(p.port, 22527);
2958        assert_eq!(p.database, "test");
2959    }
2960
2961    #[cfg(feature = "db-verify")]
2962    #[test]
2963    fn test_parse_sqlserver_dsn_default_port() {
2964        let dsn = "mssql://user:pass@host/db";
2965        let p = parse_sqlserver_dsn(dsn).unwrap();
2966        assert_eq!(p.port, 1433);
2967        assert_eq!(p.database, "db");
2968    }
2969
2970    // ---- schema! 宏 parse_create_table 测试 ----
2971
2972    #[test]
2973    fn test_parse_create_table_basic() {
2974        let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
2975        let (table, cols) = parse_create_table(sql).unwrap();
2976        assert_eq!(table, "users");
2977        assert_eq!(
2978            cols,
2979            vec![
2980                ("id".to_string(), "i32".to_string()),
2981                ("name".to_string(), "String".to_string())
2982            ]
2983        );
2984    }
2985
2986    #[test]
2987    fn test_parse_create_table_with_if_not_exists() {
2988        let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
2989        let (table, cols) = parse_create_table(sql).unwrap();
2990        assert_eq!(table, "orders");
2991        assert_eq!(
2992            cols,
2993            vec![
2994                ("id".to_string(), "i64".to_string()),
2995                ("total".to_string(), "f64".to_string())
2996            ]
2997        );
2998    }
2999
3000    #[test]
3001    fn test_parse_create_table_nullable() {
3002        let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
3003        let (_, cols) = parse_create_table(sql).unwrap();
3004        assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
3005        assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
3006    }
3007
3008    #[test]
3009    fn test_parse_create_table_skip_constraints() {
3010        let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
3011        let (_, cols) = parse_create_table(sql).unwrap();
3012        assert_eq!(cols.len(), 2);
3013        assert_eq!(cols[0].0, "id");
3014        assert_eq!(cols[1].0, "name");
3015    }
3016
3017    #[test]
3018    fn test_parse_create_table_varchar_with_len() {
3019        let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
3020        let (_, cols) = parse_create_table(sql).unwrap();
3021        assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
3022        assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
3023    }
3024
3025    #[test]
3026    fn test_sql_type_to_rust_mappings() {
3027        // 整数(按字节宽度严格映射,与 SQL 标准一致)
3028        assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
3029        assert_eq!(sql_type_to_rust("INT8", false), "i64");
3030        assert_eq!(sql_type_to_rust("INT", false), "i32");
3031        assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
3032        assert_eq!(sql_type_to_rust("INT4", false), "i32");
3033        assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
3034        assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
3035        assert_eq!(sql_type_to_rust("INT2", false), "i16");
3036        assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
3037        assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
3038        // 浮点
3039        assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
3040        assert_eq!(sql_type_to_rust("REAL", false), "f32");
3041        assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
3042        assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
3043        assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
3044        assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
3045        assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
3046        assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
3047        // 布尔
3048        assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
3049        assert_eq!(sql_type_to_rust("BOOL", false), "bool");
3050        // 字符串
3051        assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
3052        assert_eq!(sql_type_to_rust("TEXT", false), "String");
3053        assert_eq!(sql_type_to_rust("CHAR", false), "String");
3054        assert_eq!(sql_type_to_rust("UUID", false), "String");
3055        assert_eq!(sql_type_to_rust("DATE", false), "String");
3056        assert_eq!(sql_type_to_rust("DATETIME", false), "String");
3057        assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
3058        assert_eq!(sql_type_to_rust("JSON", false), "String");
3059        assert_eq!(sql_type_to_rust("JSONB", false), "String");
3060        // 二进制
3061        assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
3062        assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
3063        assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
3064        assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
3065        // nullable
3066        assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
3067        assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
3068        assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
3069        assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
3070        // unknown
3071        assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
3072    }
3073
3074    #[test]
3075    fn test_parse_create_table_error_no_create() {
3076        assert!(parse_create_table("SELECT * FROM users").is_err());
3077    }
3078
3079    #[test]
3080    fn test_parse_create_table_error_no_parens() {
3081        assert!(parse_create_table("CREATE TABLE foo").is_err());
3082    }
3083
3084    // -----------------------------------------------------------------------
3085    // Gap 1 测试:列名/表名提取
3086    // -----------------------------------------------------------------------
3087
3088    #[cfg(feature = "db-verify")]
3089    #[test]
3090    fn test_extract_tables_simple() {
3091        let tables = extract_tables("SELECT id, name FROM users WHERE id = ?");
3092        assert!(tables.contains(&"users".to_string()));
3093    }
3094
3095    #[cfg(feature = "db-verify")]
3096    #[test]
3097    fn test_extract_tables_multiple() {
3098        let tables = extract_tables(
3099            "SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id WHERE u.id = ?",
3100        );
3101        assert!(tables.contains(&"users".to_string()));
3102        assert!(tables.contains(&"orders".to_string()));
3103    }
3104
3105    #[cfg(feature = "db-verify")]
3106    #[test]
3107    fn test_extract_columns_select_and_where() {
3108        let cols =
3109            extract_columns("SELECT id, name FROM users WHERE email = ? ORDER BY created_at");
3110        // id, name from SELECT; email from WHERE; created_at from ORDER BY
3111        assert!(cols.contains(&"id".to_string()));
3112        assert!(cols.contains(&"name".to_string()));
3113        assert!(cols.contains(&"email".to_string()));
3114        assert!(cols.contains(&"created_at".to_string()));
3115    }
3116
3117    #[cfg(feature = "db-verify")]
3118    #[test]
3119    fn test_extract_columns_skips_keywords() {
3120        let cols = extract_columns("SELECT COUNT(id), name FROM users WHERE status = ?");
3121        // COUNT is a function, should be skipped
3122        assert!(!cols.contains(&"count".to_string()));
3123        assert!(cols.contains(&"id".to_string()));
3124        assert!(cols.contains(&"name".to_string()));
3125        assert!(cols.contains(&"status".to_string()));
3126    }
3127
3128    #[cfg(feature = "db-verify")]
3129    #[test]
3130    fn test_is_sql_function() {
3131        assert!(is_sql_function("COUNT"));
3132        assert!(is_sql_function("now"));
3133        assert!(is_sql_function("COALESCE"));
3134        assert!(!is_sql_function("name"));
3135        assert!(!is_sql_function("user_id"));
3136    }
3137}