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