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    let mut conn_str = format!(
1587        "{}/{}@{}:{}/{}",
1588        parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
1589    );
1590    if parsed.sysdba {
1591        conn_str.push_str(" AS SYSDBA");
1592    }
1593    let full_script = format!(
1594        "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
1595         EXPLAIN PLAN FOR {};\n\
1596         SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
1597         EXIT;\n",
1598        explain_sql
1599    );
1600    let connect_script = format!("connect {}\n{}", conn_str, full_script);
1601    let output = std::process::Command::new("sqlplus")
1602        .args(["-S", "-L", "/nolog"])
1603        .stdin(std::process::Stdio::piped())
1604        .stdout(std::process::Stdio::piped())
1605        .stderr(std::process::Stdio::piped())
1606        .spawn()
1607        .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
1608    use std::io::Write;
1609    let mut child = output;
1610    if let Some(mut stdin) = child.stdin.take() {
1611        stdin
1612            .write_all(connect_script.as_bytes())
1613            .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
1614    }
1615    let out = child
1616        .wait_with_output()
1617        .map_err(|e| format!("sqlplus wait failed: {}", e))?;
1618    let stdout = String::from_utf8_lossy(&out.stdout);
1619    let stderr = String::from_utf8_lossy(&out.stderr);
1620    if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
1621        return Err(format!(
1622            "Oracle EXPLAIN failed: stdout={} stderr={}",
1623            stdout.trim(),
1624            stderr.trim()
1625        ));
1626    }
1627    Ok(())
1628}
1629
1630/// SQL Server 编译期验证:通过 sqlcmd 命令行工具执行 SET SHOWPLAN_TEXT ON
1631///
1632/// DSN 格式:`sqlserver://user:pass@host:port/db`
1633/// 例如:`sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test`
1634#[cfg(feature = "db-verify")]
1635fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
1636    let parsed = parse_sqlserver_dsn(dsn)?;
1637    let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
1638    let out = std::process::Command::new("sqlcmd")
1639        .args([
1640            "-S",
1641            &format!("{},{}", parsed.host, parsed.port),
1642            "-U",
1643            &parsed.user,
1644            "-d",
1645            &parsed.database,
1646            "-Q",
1647            &query,
1648            "-h",
1649            "-1",
1650            "-W",
1651        ])
1652        .env("SQLCMDPASSWORD", &parsed.password)
1653        .output()
1654        .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
1655    let stdout = String::from_utf8_lossy(&out.stdout);
1656    let stderr = String::from_utf8_lossy(&out.stderr);
1657    if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
1658        return Err(format!(
1659            "SQL Server SHOWPLAN failed: stdout={} stderr={}",
1660            stdout.trim(),
1661            stderr.trim()
1662        ));
1663    }
1664    Ok(())
1665}
1666
1667/// Oracle DSN 解析结果
1668#[cfg(feature = "db-verify")]
1669struct OracleDsn {
1670    user: String,
1671    password: String,
1672    host: String,
1673    port: u16,
1674    service: String,
1675    sysdba: bool,
1676}
1677
1678/// 解析 oracle://user:pass@host:port/service?sysdba=1
1679#[cfg(feature = "db-verify")]
1680fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
1681    let raw = dsn
1682        .strip_prefix("oracle://")
1683        .or_else(|| dsn.strip_prefix("oracle:"))
1684        .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
1685    // 分离 query
1686    let (auth_host_service, query) = match raw.find('?') {
1687        Some(idx) => (&raw[..idx], &raw[idx + 1..]),
1688        None => (raw, ""),
1689    };
1690    let sysdba = query
1691        .split('&')
1692        .any(|p| p == "sysdba=1" || p == "sysdba=true");
1693    // user:pass@host:port/service
1694    let at = auth_host_service
1695        .find('@')
1696        .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
1697    let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
1698    let colon = user_pass
1699        .find(':')
1700        .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
1701    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1702    let (host_port, service) = match host_port_service.rfind('/') {
1703        Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
1704        None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
1705    };
1706    let (host, port) = match host_port.find(':') {
1707        Some(idx) => (
1708            &host_port[..idx],
1709            host_port[idx + 1..]
1710                .parse::<u16>()
1711                .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
1712        ),
1713        None => (host_port, 1521u16),
1714    };
1715    Ok(OracleDsn {
1716        user: user.to_string(),
1717        password: password.to_string(),
1718        host: host.to_string(),
1719        port,
1720        service: service.to_string(),
1721        sysdba,
1722    })
1723}
1724
1725/// SQL Server DSN 解析结果
1726#[cfg(feature = "db-verify")]
1727struct SqlServerDsn {
1728    user: String,
1729    password: String,
1730    host: String,
1731    port: u16,
1732    database: String,
1733}
1734
1735/// 解析 sqlserver://user:pass@host:port/db
1736#[cfg(feature = "db-verify")]
1737fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
1738    let raw = dsn
1739        .strip_prefix("sqlserver://")
1740        .or_else(|| dsn.strip_prefix("mssql://"))
1741        .or_else(|| dsn.strip_prefix("tds://"))
1742        .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
1743    let at = raw
1744        .find('@')
1745        .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
1746    let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
1747    let colon = user_pass
1748        .find(':')
1749        .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
1750    let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1751    let (host_port, database) = match host_port_db.rfind('/') {
1752        Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
1753        None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
1754    };
1755    let (host, port) = match host_port.find(':') {
1756        Some(idx) => (
1757            &host_port[..idx],
1758            host_port[idx + 1..]
1759                .parse::<u16>()
1760                .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
1761        ),
1762        None => (host_port, 1433u16),
1763    };
1764    Ok(SqlServerDsn {
1765        user: user.to_string(),
1766        password: password.to_string(),
1767        host: host.to_string(),
1768        port,
1769        database: database.to_string(),
1770    })
1771}
1772
1773// ---------------------------------------------------------------------------
1774// Helpers
1775// ---------------------------------------------------------------------------
1776
1777/// Create a compile_error! token stream
1778fn compile_error(span: Span, msg: &str) -> TokenStream {
1779    // emit: compile_error!("msg")
1780    let mut ts = TokenStream::new();
1781    ts.extend([
1782        TokenTree::Ident(Ident::new("compile_error", span)),
1783        TokenTree::Punct(Punct::new('!', Spacing::Alone)),
1784        TokenTree::Group(Group::new(
1785            Delimiter::Parenthesis,
1786            TokenStream::from(TokenTree::Literal(Literal::string(msg))),
1787        )),
1788    ]);
1789    ts
1790}
1791
1792// ---------------------------------------------------------------------------
1793// typed_query! — Diesel 风格强类型 AST 宏
1794// ---------------------------------------------------------------------------
1795
1796/// Diesel 风格强类型 AST 宏(与 `sql_string!` / `query!` 并存)。
1797///
1798/// # 设计
1799///
1800/// 接收 `table { col1: Type, col2: Type, ... }` 声明,生成:
1801/// 1. 一个 `table` 模块
1802/// 2. 每列对应一个零大小标记类型(如 `table::id`)
1803/// 3. 实现 `TypedColumn` trait,把列名 + Rust 类型提升到类型系统
1804///
1805/// 这样,`typed_query!(SELECT id FROM users WHERE name = ?)` 在编译期就能:
1806/// - 校验 `id` / `name` 列是否存在于 `users` 表声明中
1807/// - 校验 `?` 参数的 Rust 类型与列声明的类型一致
1808///
1809/// # 用法
1810///
1811/// ```ignore
1812/// use sz_orm_macros::typed_query;
1813///
1814/// // 1. 声明表 schema(编译期生成 column 标记类型)
1815/// typed_query! {
1816///     table users {
1817///         id: i64,
1818///         name: String,
1819///         email: String,
1820///         age: i32,
1821///     }
1822/// }
1823///
1824/// // 2. 编译期校验 SELECT:列名必须存在于 users 表
1825/// let sql = typed_query!(SELECT id, name FROM users WHERE age > ?);
1826/// // ❌ 编译错误:unknown column 'foo' in table 'users'
1827/// // let sql = typed_query!(SELECT foo FROM users);
1828/// ```
1829#[proc_macro]
1830pub fn typed_query(input: TokenStream) -> TokenStream {
1831    let tokens: Vec<TokenTree> = input.into_iter().collect();
1832
1833    // 分支 1:table 声明
1834    if tokens.iter().any(|t| {
1835        if let TokenTree::Ident(id) = t {
1836            id.to_string() == "table"
1837        } else {
1838            false
1839        }
1840    }) {
1841        return parse_table_decl(&tokens);
1842    }
1843
1844    // 分支 2:SELECT 表达式
1845    if tokens.iter().any(|t| {
1846        if let TokenTree::Ident(id) = t {
1847            id.to_string().eq_ignore_ascii_case("SELECT")
1848        } else {
1849            false
1850        }
1851    }) {
1852        return parse_typed_select(&tokens);
1853    }
1854
1855    compile_error(
1856        Span::call_site(),
1857        "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
1858    )
1859}
1860
1861/// 解析 `table name { col: Type, ... }` 声明
1862fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
1863    // 期望格式:table <ident> { <ident> : <ident> [, ...] }
1864    let mut idx = 0;
1865
1866    // 跳过 'table' 关键字
1867    if idx >= tokens.len() {
1868        return compile_error(Span::call_site(), "expected table name after 'table'");
1869    }
1870    if let TokenTree::Ident(id) = &tokens[idx] {
1871        if id.to_string() != "table" {
1872            return compile_error(id.span(), "expected 'table' keyword");
1873        }
1874    }
1875    idx += 1;
1876
1877    // 表名
1878    let table_name = if idx < tokens.len() {
1879        if let TokenTree::Ident(id) = &tokens[idx] {
1880            id.to_string()
1881        } else {
1882            return compile_error(tokens[idx].span(), "expected table name identifier");
1883        }
1884    } else {
1885        return compile_error(Span::call_site(), "expected table name");
1886    };
1887    idx += 1;
1888
1889    // 表体({} 内)
1890    let body_group = if idx < tokens.len() {
1891        if let TokenTree::Group(g) = &tokens[idx] {
1892            if g.delimiter() != Delimiter::Brace {
1893                return compile_error(g.span(), "expected '{' after table name");
1894            }
1895            g.clone()
1896        } else {
1897            return compile_error(tokens[idx].span(), "expected '{' after table name");
1898        }
1899    } else {
1900        return compile_error(Span::call_site(), "expected table body in '{ }'");
1901    };
1902
1903    // 解析列声明
1904    let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
1905    let columns = match parse_column_list(&body_tokens) {
1906        Ok(c) => c,
1907        Err(e) => return compile_error(Span::call_site(), &e),
1908    };
1909
1910    // 使用 quote! 构建类型安全的 TokenStream
1911    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1912    let table_name_lit = table_name.as_str();
1913
1914    // 为每列构建标记类型 + trait 实现
1915    let col_impls: Vec<TokenStream2> = columns
1916        .iter()
1917        .map(|(col_name, col_type)| {
1918            let col_ident =
1919                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1920            let col_name_lit = col_name.as_str();
1921            // 解析类型字符串为 TokenStream(quote! 会处理)
1922            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1923            quote! {
1924                #[derive(Debug, Clone, Copy)]
1925                pub struct #col_ident;
1926                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1927                    const NAME: &'static str = #col_name_lit;
1928                    type Table = table;
1929                    type RustType = #rust_type;
1930                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1931                }
1932            }
1933        })
1934        .collect();
1935
1936    // schema 常量条目
1937    let schema_entries: Vec<TokenStream2> = columns
1938        .iter()
1939        .map(|(n, t)| {
1940            let n_lit = n.as_str();
1941            let t_lit = t.as_str();
1942            quote! { (#n_lit, #t_lit) }
1943        })
1944        .collect();
1945
1946    let schema_const_ident = proc_macro2::Ident::new(
1947        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1948        Span::call_site().into(),
1949    );
1950
1951    let expanded = quote! {
1952        pub mod #table_ident {
1953            use super::*;
1954            pub struct table;
1955            impl ::sz_orm_core::typed::TypedTable for table {
1956                const NAME: &'static str = #table_name_lit;
1957            }
1958            #(#col_impls)*
1959        }
1960        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1961    };
1962
1963    expanded.into()
1964}
1965
1966/// 解析列声明列表:`col: Type, col2: Type2, ...`
1967fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
1968    let mut cols = Vec::new();
1969    let mut i = 0;
1970    while i < tokens.len() {
1971        // 列名
1972        let col_name = if let TokenTree::Ident(id) = &tokens[i] {
1973            id.to_string()
1974        } else {
1975            return Err(format!("expected column name at position {}", i));
1976        };
1977        i += 1;
1978
1979        // 冒号
1980        if i >= tokens.len() {
1981            return Err(format!("expected ':' after column '{}'", col_name));
1982        }
1983        if let TokenTree::Punct(p) = &tokens[i] {
1984            if p.as_char() != ':' {
1985                return Err(format!("expected ':' after column '{}'", col_name));
1986            }
1987        } else {
1988            return Err(format!("expected ':' after column '{}'", col_name));
1989        }
1990        i += 1;
1991
1992        // 类型(可能是 ident 或 path,如 String / i64 / Option<i64>)
1993        // 简化处理:收集直到遇到 ',' 或末尾
1994        let mut type_str = String::new();
1995        let mut depth = 0;
1996        while i < tokens.len() {
1997            match &tokens[i] {
1998                TokenTree::Punct(p) => {
1999                    if p.as_char() == ',' && depth == 0 {
2000                        i += 1;
2001                        break;
2002                    } else if p.as_char() == '<' || p.as_char() == '(' {
2003                        depth += 1;
2004                        type_str.push(p.as_char());
2005                    } else if p.as_char() == '>' || p.as_char() == ')' {
2006                        depth -= 1;
2007                        type_str.push(p.as_char());
2008                    } else {
2009                        type_str.push(p.as_char());
2010                    }
2011                }
2012                TokenTree::Ident(id) => {
2013                    if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
2014                    {
2015                        type_str.push(' ');
2016                    }
2017                    type_str.push_str(&id.to_string());
2018                }
2019                _ => {}
2020            }
2021            i += 1;
2022        }
2023
2024        cols.push((col_name, type_str.trim().to_string()));
2025    }
2026    Ok(cols)
2027}
2028
2029/// 解析 `SELECT col1, col2 FROM table WHERE col = ?` 表达式
2030///
2031/// 校验列名是否在表 schema 中(通过编译期常量查找)。
2032fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
2033    // 收集所有 ident 与 literal,构造 SQL 字符串
2034    let mut sql_parts: Vec<String> = Vec::new();
2035    let mut table_name: Option<String> = None;
2036    let mut in_from = false;
2037
2038    for (i, t) in tokens.iter().enumerate() {
2039        match t {
2040            TokenTree::Ident(id) => {
2041                let s = id.to_string();
2042                if s.eq_ignore_ascii_case("SELECT") {
2043                    sql_parts.push("SELECT".to_string());
2044                } else if s.eq_ignore_ascii_case("FROM") {
2045                    in_from = true;
2046                    sql_parts.push("FROM".to_string());
2047                } else if s.eq_ignore_ascii_case("WHERE")
2048                    || s.eq_ignore_ascii_case("AND")
2049                    || s.eq_ignore_ascii_case("OR")
2050                    || s.eq_ignore_ascii_case("LIMIT")
2051                    || s.eq_ignore_ascii_case("OFFSET")
2052                    || s.eq_ignore_ascii_case("ORDER")
2053                    || s.eq_ignore_ascii_case("BY")
2054                    || s.eq_ignore_ascii_case("GROUP")
2055                    || s.eq_ignore_ascii_case("HAVING")
2056                    || s.eq_ignore_ascii_case("JOIN")
2057                    || s.eq_ignore_ascii_case("INNER")
2058                    || s.eq_ignore_ascii_case("LEFT")
2059                    || s.eq_ignore_ascii_case("RIGHT")
2060                    || s.eq_ignore_ascii_case("ON")
2061                    || s.eq_ignore_ascii_case("AS")
2062                    || s.eq_ignore_ascii_case("ASC")
2063                    || s.eq_ignore_ascii_case("DESC")
2064                    || s.eq_ignore_ascii_case("DISTINCT")
2065                    || s.eq_ignore_ascii_case("NOT")
2066                    || s.eq_ignore_ascii_case("NULL")
2067                    || s.eq_ignore_ascii_case("IN")
2068                    || s.eq_ignore_ascii_case("BETWEEN")
2069                    || s.eq_ignore_ascii_case("LIKE")
2070                    || s.eq_ignore_ascii_case("IS")
2071                {
2072                    sql_parts.push(s.to_uppercase());
2073                } else if in_from && table_name.is_none() {
2074                    // FROM 后第一个 ident 是表名
2075                    table_name = Some(s.clone());
2076                    sql_parts.push(s.clone());
2077                } else {
2078                    sql_parts.push(s.clone());
2079                }
2080            }
2081            TokenTree::Literal(lit) => {
2082                sql_parts.push(lit.to_string());
2083            }
2084            TokenTree::Punct(p) => {
2085                let c = p.as_char();
2086                // SQL 中常见标点:, ; * ? = > < ( ) . 等
2087                let part = if c == ',' {
2088                    ",".to_string()
2089                } else if c == '?' {
2090                    "?".to_string()
2091                } else if c == '*' {
2092                    "*".to_string()
2093                } else if c == '=' {
2094                    "=".to_string()
2095                } else if c == '>' {
2096                    ">".to_string()
2097                } else if c == '<' {
2098                    "<".to_string()
2099                } else if c == '.' {
2100                    ".".to_string()
2101                } else if c == ';' {
2102                    ";".to_string()
2103                } else {
2104                    c.to_string()
2105                };
2106                sql_parts.push(part);
2107            }
2108            TokenTree::Group(g) => {
2109                // 处理 group(如 (1, 2, 3))
2110                let inner: String = g.stream().to_string();
2111                let delim = match g.delimiter() {
2112                    Delimiter::Parenthesis => "(",
2113                    Delimiter::Brace => "{",
2114                    Delimiter::Bracket => "[",
2115                    Delimiter::None => "",
2116                };
2117                let close = match g.delimiter() {
2118                    Delimiter::Parenthesis => ")",
2119                    Delimiter::Brace => "}",
2120                    Delimiter::Bracket => "]",
2121                    Delimiter::None => "",
2122                };
2123                sql_parts.push(format!("{}{}{}", delim, inner, close));
2124            }
2125        }
2126        // 单空格分隔(去重多个空格由 trim 处理)
2127        let _ = i;
2128    }
2129
2130    let sql = sql_parts
2131        .join(" ")
2132        .replace(", ", ",")
2133        .replace(" ,", ",")
2134        .replace("= ", "=")
2135        .replace(" =", "=")
2136        .replace("> ", ">")
2137        .replace(" >", ">")
2138        .replace("< ", "<")
2139        .replace(" <", "<")
2140        .replace("  ", " ");
2141
2142    // 验证 SQL 语法
2143    if let Err(e) = validate_sql_content(&sql, None) {
2144        return compile_error(
2145            Span::call_site(),
2146            &format!("typed_query! SQL validation failed: {}", e),
2147        );
2148    }
2149
2150    // 生成 SQL 字符串字面量
2151    let mut ts = TokenStream::new();
2152    let lit = Literal::string(&sql);
2153    ts.extend([TokenTree::Literal(lit)]);
2154    ts
2155}
2156
2157// ---------------------------------------------------------------------------
2158// schema! — Compile-time SQL schema generator
2159// ---------------------------------------------------------------------------
2160
2161/// Compile-time SQL schema generator.
2162///
2163/// Parses a SQL `CREATE TABLE` statement and generates typed table declarations
2164/// equivalent to `typed_query! { table ... }`.
2165///
2166/// # Syntax
2167///
2168/// ```ignore
2169/// use sz_orm_macros::schema;
2170///
2171/// schema! {
2172///     "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL, email TEXT)"
2173/// }
2174/// ```
2175///
2176/// 生成与以下手动声明等价的代码:
2177/// ```ignore
2178/// typed_query! {
2179///     table users {
2180///         id: i64,
2181///         name: String,
2182///         email: Option<String>,
2183///     }
2184/// }
2185/// ```
2186#[proc_macro]
2187/// 类型化裸 SQL 查询宏(SQLx `query_as!` 风格)。
2188///
2189/// 用法:`query_as!(RecordType, "SELECT col1, col2 FROM table WHERE id = ?")`
2190///
2191/// 生成 `sz_orm_core::queryable::QueryAs::<RecordType>::new("SELECT ...")`。
2192/// 在 `db-verify` feature + `SZ_ORM_QUERY_VERIFY=1` 环境下会连真 DB
2193/// 执行 EXPLAIN 验证 SQL 合法性。
2194///
2195/// **运行时列名验证**(P0-2):`QueryAs::fetch_all` 会比对 DB 返回的列名
2196/// 与 `RecordType::row_desc()`(由 `#[derive(FromQueryResult)]` 自动生成)。
2197/// 若 SQL SELECT 列不在 struct 字段中,返回 `DbError::QueryError`。
2198///
2199/// # 示例
2200///
2201/// ```ignore
2202/// #[derive(FromQueryResult)]
2203/// struct User { id: i64, name: String }
2204///
2205/// let q = query_as!(User, "SELECT id, name FROM users WHERE id = 1");
2206/// let users: Vec<User> = q.fetch_all(&mut conn).await?;
2207/// ```
2208pub fn query_as(input: TokenStream) -> TokenStream {
2209    let mut tokens = input.into_iter().peekable();
2210
2211    // 解析记录类型(第一个标识符/路径,如 User 或 crate::User)
2212    let mut record_type = String::new();
2213    loop {
2214        match tokens.next() {
2215            Some(TokenTree::Ident(ident)) => {
2216                record_type.push_str(&ident.to_string());
2217            }
2218            Some(TokenTree::Punct(p)) if p.as_char() == ':' => {
2219                // 处理 :: 路径分隔符
2220                record_type.push_str("::");
2221                // 跳过第二个 :
2222                if let Some(TokenTree::Punct(p2)) = tokens.peek() {
2223                    if p2.as_char() == ':' {
2224                        let _ = tokens.next();
2225                    }
2226                }
2227            }
2228            Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
2229            Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
2230            Some(other) => {
2231                return compile_error(
2232                    other.span(),
2233                    "query_as! 第一个参数必须是记录类型,如 query_as!(User, \"SELECT ...\")",
2234                );
2235            }
2236            None => {
2237                return compile_error(
2238                    Span::call_site(),
2239                    "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2240                );
2241            }
2242        }
2243    }
2244
2245    // 解析 SQL 字符串字面量
2246    let sql_raw = match tokens.next() {
2247        Some(TokenTree::Literal(lit)) => lit.to_string(),
2248        Some(other) => {
2249            return compile_error(other.span(), "query_as! 第二个参数必须是 SQL 字符串字面量");
2250        }
2251        None => {
2252            return compile_error(
2253                Span::call_site(),
2254                "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2255            );
2256        }
2257    };
2258
2259    let sql_content = match strip_string_literal(&sql_raw) {
2260        Some(s) => s,
2261        None => {
2262            return compile_error(Span::call_site(), "query_as! 的 SQL 参数必须是字符串字面量");
2263        }
2264    };
2265
2266    // 语法验证
2267    if let Err(err_msg) = validate_sql_content(sql_content, None) {
2268        return compile_error(Span::call_site(), &err_msg);
2269    }
2270
2271    // db-verify 验证
2272    #[cfg(feature = "db-verify")]
2273    let verify_cols: Option<Vec<(String, String)>> = {
2274        match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
2275            // 模式 1:连真 DB 执行 EXPLAIN 验证,并获取 SELECT 列的实际类型
2276            Some("1") => match verify_with_real_db(sql_content) {
2277                Ok(cols) => Some(cols),
2278                Err(err) => {
2279                    return compile_error(
2280                        Span::call_site(),
2281                        &format!("query_as! real DB verification failed: {}", err),
2282                    )
2283                }
2284            },
2285            // 模式 cache:从离线缓存文件查找(无需 DB,适合 CI)
2286            Some("cache") => {
2287                if let Err(err) = verify_with_cache(sql_content) {
2288                    return compile_error(
2289                        Span::call_site(),
2290                        &format!("query_as! offline cache verification failed: {}", err),
2291                    );
2292                }
2293                None
2294            }
2295            _ => None,
2296        }
2297    };
2298    #[cfg(not(feature = "db-verify"))]
2299    let _verify_cols: Option<Vec<(String, String)>> = None;
2300
2301    // 生成 QueryAs::<T>::new("...")
2302    // db-verify 通过时,附加编译期类型验证块(P0-2):
2303    // const 上下文中将 DB 实际列类型与结构体 __sz_orm_column_types() 对比,
2304    // 不匹配则 const panic → 编译失败。
2305    let escaped = sql_content.escape_default();
2306    let base = format!(
2307        "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
2308        record_type, escaped
2309    );
2310    #[cfg(feature = "db-verify")]
2311    let output = match &verify_cols {
2312        Some(cols) if !cols.is_empty() => {
2313            gen_compile_time_type_check(&record_type, sql_content, cols, &base)
2314        }
2315        _ => base,
2316    };
2317    #[cfg(not(feature = "db-verify"))]
2318    let output = base;
2319    output
2320        .parse()
2321        .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query_as output"))
2322}
2323
2324#[proc_macro]
2325pub fn schema(input: TokenStream) -> TokenStream {
2326    let mut tokens = input.into_iter().peekable();
2327
2328    // 解析 SQL 字符串字面量
2329    let sql_raw = match tokens.next() {
2330        Some(TokenTree::Literal(lit)) => lit.to_string(),
2331        Some(other) => {
2332            return compile_error(
2333                other.span(),
2334                "Expected a string literal as the argument to schema!",
2335            );
2336        }
2337        None => {
2338            return compile_error(
2339                Span::call_site(),
2340                "Expected a string literal argument to schema!",
2341            );
2342        }
2343    };
2344
2345    let sql = match strip_string_literal(&sql_raw) {
2346        Some(s) => s,
2347        None => {
2348            return compile_error(
2349                Span::call_site(),
2350                "schema! requires a string literal argument",
2351            );
2352        }
2353    };
2354
2355    // 解析 CREATE TABLE
2356    let (table_name, columns) = match parse_create_table(sql) {
2357        Ok(v) => v,
2358        Err(e) => return compile_error(Span::call_site(), &e),
2359    };
2360
2361    // 生成代码(与 parse_table_decl 一致)
2362    let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
2363    let table_name_lit = table_name.as_str();
2364
2365    let col_impls: Vec<TokenStream2> = columns
2366        .iter()
2367        .map(|(col_name, col_type)| {
2368            let col_ident =
2369                proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
2370            let col_name_lit = col_name.as_str();
2371            let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
2372            quote! {
2373                #[derive(Debug, Clone, Copy)]
2374                pub struct #col_ident;
2375                impl ::sz_orm_core::typed::TypedColumn for #col_ident {
2376                    const NAME: &'static str = #col_name_lit;
2377                    type Table = table;
2378                    type RustType = #rust_type;
2379                    type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
2380                }
2381            }
2382        })
2383        .collect();
2384
2385    let schema_entries: Vec<TokenStream2> = columns
2386        .iter()
2387        .map(|(n, t)| {
2388            let n_lit = n.as_str();
2389            let t_lit = t.as_str();
2390            quote! { (#n_lit, #t_lit) }
2391        })
2392        .collect();
2393
2394    let schema_const_ident = proc_macro2::Ident::new(
2395        &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
2396        Span::call_site().into(),
2397    );
2398
2399    let expanded = quote! {
2400        pub mod #table_ident {
2401            use super::*;
2402            pub struct table;
2403            impl ::sz_orm_core::typed::TypedTable for table {
2404                const NAME: &'static str = #table_name_lit;
2405            }
2406            #(#col_impls)*
2407        }
2408        const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
2409    };
2410
2411    expanded.into()
2412}
2413
2414/// 解析 SQL `CREATE TABLE` 语句,返回 (表名, Vec<(列名, Rust 类型字符串)>)。
2415///
2416/// 支持以下语法:
2417/// - `CREATE TABLE [IF NOT EXISTS] <name> ( ... )`
2418/// - 表名/列名可带反引号、双引号或无引号
2419/// - 跳过 PRIMARY KEY / FOREIGN KEY / CONSTRAINT / UNIQUE / INDEX / KEY 约束行
2420/// - 列定义按顶层逗号分隔(嵌套括号如 DECIMAL(10,2) 不拆分)
2421fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
2422    let trimmed = sql.trim();
2423    let upper = trimmed.to_uppercase();
2424
2425    // 必须以 CREATE TABLE 开头
2426    if !upper.starts_with("CREATE TABLE") {
2427        return Err("schema! expects a CREATE TABLE statement".to_string());
2428    }
2429
2430    // 跳过 "CREATE TABLE"
2431    let mut rest = &trimmed["CREATE TABLE".len()..];
2432
2433    // 跳过可选的 "IF NOT EXISTS"
2434    let rest_upper = rest.trim_start().to_uppercase();
2435    if rest_upper.starts_with("IF NOT EXISTS") {
2436        rest = &rest.trim_start()["IF NOT EXISTS".len()..];
2437    }
2438
2439    rest = rest.trim_start();
2440
2441    // 解析表名(可能带反引号、双引号或无引号)
2442    let (table_name, after_name) = parse_identifier(rest)?;
2443    let rest = after_name.trim_start();
2444
2445    // 找到列定义起始的 '(' 与匹配的最后一个 ')'
2446    let paren_start = rest
2447        .find('(')
2448        .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
2449    let paren_end = rest
2450        .rfind(')')
2451        .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
2452    if paren_end <= paren_start {
2453        return Err("CREATE TABLE has malformed parentheses".to_string());
2454    }
2455
2456    let cols_str = &rest[paren_start + 1..paren_end];
2457
2458    // 按顶层逗号分隔列定义(注意嵌套括号,如 DECIMAL(10,2))
2459    let col_defs = split_top_level_commas(cols_str);
2460
2461    let mut columns = Vec::new();
2462    for def in col_defs {
2463        let def = def.trim();
2464        if def.is_empty() {
2465            continue;
2466        }
2467
2468        // 跳过约束定义行
2469        let def_upper = def.to_uppercase();
2470        if def_upper.starts_with("PRIMARY KEY")
2471            || def_upper.starts_with("FOREIGN KEY")
2472            || def_upper.starts_with("CONSTRAINT")
2473            || def_upper.starts_with("UNIQUE")
2474            || def_upper.starts_with("INDEX")
2475            || def_upper.starts_with("KEY")
2476        {
2477            continue;
2478        }
2479
2480        // 解析列名
2481        let (col_name, after_col) = parse_identifier(def)?;
2482        let rest = after_col.trim_start();
2483
2484        // 解析类型(取第一个 token,去掉括号参数)
2485        let (sql_type, after_type) = parse_type_token(rest)?;
2486        let rest = after_type.trim();
2487
2488        // 判断 nullability:NOT NULL 或 PRIMARY KEY 隐含 NOT NULL
2489        let rest_upper = rest.to_uppercase();
2490        let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
2491        let rust_type = sql_type_to_rust(&sql_type, !not_null);
2492
2493        columns.push((col_name, rust_type));
2494    }
2495
2496    Ok((table_name, columns))
2497}
2498
2499/// 解析标识符:支持反引号、双引号或无引号。
2500/// 返回 (标识符, 剩余字符串)。
2501fn parse_identifier(s: &str) -> Result<(String, &str), String> {
2502    let s = s.trim_start();
2503    if s.is_empty() {
2504        return Err("expected identifier".to_string());
2505    }
2506
2507    let bytes = s.as_bytes();
2508    match bytes[0] {
2509        b'`' => {
2510            let end = s[1..]
2511                .find('`')
2512                .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
2513            let ident = s[1..1 + end].to_string();
2514            Ok((ident, &s[1 + end + 1..]))
2515        }
2516        b'"' => {
2517            let end = s[1..]
2518                .find('"')
2519                .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
2520            let ident = s[1..1 + end].to_string();
2521            Ok((ident, &s[1 + end + 1..]))
2522        }
2523        _ => {
2524            let end = s
2525                .find(|c: char| !c.is_alphanumeric() && c != '_')
2526                .unwrap_or(s.len());
2527            if end == 0 {
2528                return Err(format!("invalid identifier: '{}'", s));
2529            }
2530            let ident = s[..end].to_string();
2531            Ok((ident, &s[end..]))
2532        }
2533    }
2534}
2535
2536/// 解析类型 token:取第一个标识符,可选跟随括号参数(如 VARCHAR(255) → VARCHAR)。
2537/// 返回 (类型名, 剩余字符串)。
2538fn parse_type_token(s: &str) -> Result<(String, &str), String> {
2539    let s = s.trim_start();
2540    if s.is_empty() {
2541        return Err("expected column type".to_string());
2542    }
2543
2544    let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
2545    if end == 0 {
2546        return Err(format!("invalid type: '{}'", s));
2547    }
2548    let type_name = s[..end].to_string();
2549    let mut rest = &s[end..];
2550
2551    // 跳过可选的括号参数,如 (255) 或 (10,2)
2552    rest = rest.trim_start();
2553    if rest.starts_with('(') {
2554        let close = rest
2555            .find(')')
2556            .ok_or_else(|| "unterminated type parameter list".to_string())?;
2557        rest = &rest[close + 1..];
2558    }
2559
2560    Ok((type_name, rest))
2561}
2562
2563/// 按顶层逗号分隔字符串(不进入嵌套括号)。
2564fn split_top_level_commas(s: &str) -> Vec<String> {
2565    let mut parts = Vec::new();
2566    let mut depth: i32 = 0;
2567    let mut current = String::new();
2568
2569    for ch in s.chars() {
2570        match ch {
2571            '(' => {
2572                depth += 1;
2573                current.push(ch);
2574            }
2575            ')' => {
2576                depth -= 1;
2577                current.push(ch);
2578            }
2579            ',' if depth == 0 => {
2580                parts.push(std::mem::take(&mut current));
2581            }
2582            _ => {
2583                current.push(ch);
2584            }
2585        }
2586    }
2587
2588    if !current.trim().is_empty() {
2589        parts.push(current);
2590    }
2591
2592    parts
2593}
2594
2595/// 将 SQL 类型映射为 Rust 类型字符串。
2596///
2597/// 匹配规则:取类型名第一个 token(去掉括号参数),不区分大小写匹配。
2598/// 未识别的类型默认映射为 `String`。若 `nullable == true`,用 `Option<T>` 包裹。
2599fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
2600    let upper = sql_type.to_uppercase();
2601    let rust = match upper.as_str() {
2602        // 8 字节整数
2603        "BIGINT" | "INT8" => "i64",
2604        // 4 字节整数(INT/INTEGER/INT4/SERIAL)
2605        "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
2606        // 2 字节整数
2607        "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
2608        // 1 字节整数
2609        "TINYINT" => "i8",
2610        // 浮点(4 字节)
2611        "FLOAT" | "REAL" | "FLOAT4" => "f32",
2612        // 浮点(8 字节)/ 定点数
2613        "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
2614        // 布尔
2615        "BOOLEAN" | "BOOL" => "bool",
2616        // 二进制(与 schema_gen::sql_type_to_rust 保持一致)
2617        "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
2618        // 字符串/日期/JSON/UUID(统一映射到 String,运行时再解析)
2619        "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
2620        | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
2621        _ => "String",
2622    };
2623
2624    if nullable {
2625        format!("Option<{}>", rust)
2626    } else {
2627        rust.to_string()
2628    }
2629}
2630
2631// ---------------------------------------------------------------------------
2632// `#[derive(Schema)]` — auto-generate table structure from a struct
2633// ---------------------------------------------------------------------------
2634
2635/// 派生宏:自动从 Rust 结构体生成表结构信息。
2636///
2637/// 解析 `#[table(name = "...")]` 和 `#[column(...)]` 属性,
2638/// 生成 `Schema` trait 实现,便于在运行时反射表名与列信息。
2639///
2640/// # 支持的属性
2641///
2642/// - `#[table(name = "users")]` — 指定表名(默认使用结构体名的蛇形形式)
2643/// - `#[column(name = "user_id")]` — 指定列名(默认使用字段名)
2644/// - `#[column(type = "VARCHAR(255)")]` — 指定 SQL 类型
2645/// - `#[column(primary_key)]` — 标记主键
2646/// - `#[column(nullable)]` — 显式标记允许 NULL
2647/// - `#[column(skip)]` — 跳过此字段,不生成 schema 条目
2648/// - `#[column(default = "0")]` — 标记字段有默认值
2649///
2650/// # 类型推断
2651///
2652/// 字段的 Rust 类型会自动映射为 SQL 类型:
2653/// - `i64`/`u64` → `BIGINT`
2654/// - `i32`/`u32` → `INTEGER`
2655/// - `String` → `TEXT`
2656/// - `f64` → `DOUBLE`
2657/// - `bool` → `BOOLEAN`
2658/// - `Vec<u8>` → `BLOB`
2659/// - `Option<T>` → 与 `T` 相同,但标记为 nullable
2660#[proc_macro_derive(Schema, attributes(table, column))]
2661pub fn derive_schema(input: TokenStream) -> TokenStream {
2662    let input = parse_macro_input!(input as syn::DeriveInput);
2663    derive::derive_schema_impl(input).into()
2664}
2665
2666// ---------------------------------------------------------------------------
2667// `#[derive(GraphQLModel)]` — auto-generate `impl GraphQLModelInfo`
2668// ---------------------------------------------------------------------------
2669
2670/// 派生宏:自动生成 `sz_orm_graphql::schema_gen::GraphQLModelInfo` 实现。
2671///
2672/// 从 `#[derive(GraphQLModel)]` 结构体提取字段元数据(字段名 + Rust 类型 + 可空性),
2673/// 供 `SchemaGenerator::from_model` 使用。零运行时开销。
2674///
2675/// # 支持的属性
2676///
2677/// - `#[table(name = "users")]` — 指定表名(默认使用结构体名的 snake_case)
2678/// - `#[column(skip)]` — 跳过此字段
2679/// - `#[column(name = "custom_name")]` — 指定列名
2680///
2681/// # 示例
2682///
2683/// ```ignore
2684/// use sz_orm_macros::GraphQLModel;
2685///
2686/// #[derive(GraphQLModel)]
2687/// #[table(name = "users")]
2688/// struct User {
2689///     id: i64,
2690///     name: String,
2691///     email: Option<String>,
2692/// }
2693/// ```
2694#[proc_macro_derive(GraphQLModel, attributes(table, column))]
2695pub fn derive_graphql_model(input: TokenStream) -> TokenStream {
2696    let input = parse_macro_input!(input as syn::DeriveInput);
2697    derive::derive_graphql_model_impl(input).into()
2698}
2699
2700// ---------------------------------------------------------------------------
2701// `#[derive(Builder)]` — auto-generate builder pattern code
2702// ---------------------------------------------------------------------------
2703
2704/// 派生宏:自动生成构造器模式代码。
2705///
2706/// 为目标结构体生成一个 `XxxBuilder` 类型,包含:
2707/// - `new()` 构造空 builder
2708/// - 每个字段的 setter 方法
2709/// - `build()` 方法返回 `Result<T, String>`
2710///
2711/// # 支持的属性
2712///
2713/// - `#[builder(skip)]` — 跳过此字段(不生成 setter,使用 Default)
2714/// - `#[builder(default = expr)]` — 指定默认值表达式
2715///
2716/// # 示例
2717///
2718/// ```ignore
2719/// use sz_orm_macros::Builder;
2720///
2721/// #[derive(Builder)]
2722/// struct User {
2723///     id: i64,
2724///     name: String,
2725/// }
2726///
2727/// let user = User::builder()
2728///     .id(1)
2729///     .name("Alice".to_string())
2730///     .build()
2731///     .unwrap();
2732/// ```
2733#[proc_macro_derive(Builder, attributes(builder))]
2734pub fn derive_builder(input: TokenStream) -> TokenStream {
2735    let input = parse_macro_input!(input as syn::DeriveInput);
2736    derive::derive_builder_impl(input).into()
2737}
2738
2739// ---------------------------------------------------------------------------
2740// `#[derive(Entity)]` — auto-generate `impl Model for Struct`
2741// ---------------------------------------------------------------------------
2742
2743/// 派生宏:自动生成 `sz_orm_core::Model` trait 实现。
2744///
2745/// 要求结构体恰好有一个 `#[column(primary_key)]` 字段,
2746/// 该字段的类型即为 `Model::PrimaryKey`。
2747///
2748/// # 支持的属性
2749///
2750/// - `#[table(name = "...")]` — 指定表名,默认蛇形结构体名
2751/// - `#[column(primary_key)]` — 标记主键字段(必需,恰好一个)
2752/// - `#[column(name = "...")]` — 覆盖主键列名(默认与字段名相同)
2753///
2754/// # 示例
2755///
2756/// ```ignore
2757/// use sz_orm_macros::Entity;
2758///
2759/// #[derive(Entity)]
2760/// #[table(name = "users")]
2761/// struct User {
2762///     #[column(primary_key)]
2763///     id: i64,
2764///     name: String,
2765/// }
2766///
2767/// assert_eq!(User::table_name(), "users");
2768/// assert_eq!(User::pk_name(), "id");
2769/// ```
2770#[proc_macro_derive(Entity, attributes(table, column))]
2771pub fn derive_entity(input: TokenStream) -> TokenStream {
2772    let input = parse_macro_input!(input as syn::DeriveInput);
2773    derive::derive_entity_impl(input).into()
2774}
2775
2776// ---------------------------------------------------------------------------
2777// `#[derive(FromQueryResult)]` — auto-generate `impl FromQueryResult for Struct`
2778// ---------------------------------------------------------------------------
2779
2780/// 派生宏:自动生成 `sz_orm_core::FromQueryResult` trait 实现。
2781///
2782/// 从查询结果行(`HashMap<String, Value>`)反序列化为结构体实例。
2783/// `Option<T>` 字段在列缺失或值为 NULL 时自动返回 `None`。
2784///
2785/// # 支持的属性
2786///
2787/// - `#[column(name = "...")]` — 覆盖列名映射(默认使用字段名)
2788///
2789/// # 示例
2790///
2791/// ```ignore
2792/// use sz_orm_macros::FromQueryResult;
2793///
2794/// #[derive(FromQueryResult)]
2795/// struct UserRow {
2796///     id: i64,
2797///     name: String,
2798///     #[column(name = "user_email")]
2799///     email: Option<String>,
2800/// }
2801/// ```
2802#[proc_macro_derive(FromQueryResult, attributes(column))]
2803pub fn derive_from_query_result(input: TokenStream) -> TokenStream {
2804    let input = parse_macro_input!(input as syn::DeriveInput);
2805    derive::derive_from_query_result_impl(input).into()
2806}
2807
2808// ---------------------------------------------------------------------------
2809// `#[derive(ColumnEnum)]` — auto-generate column name enum (P2-2)
2810// ---------------------------------------------------------------------------
2811
2812/// 派生宏:从结构体字段自动生成 `<StructName>Column` 列名枚举(P2-2)。
2813///
2814/// 每个字段生成一个变体(snake_case → CamelCase),通过 `ColumnTrait::as_str()`
2815/// 返回数据库列名;`#[column(name = "...")]` 可覆盖列名(与 FromQueryResult 一致)。
2816/// 同时实现 `std::fmt::Display`。
2817///
2818/// # 示例
2819///
2820/// ```rust,ignore
2821/// use sz_orm_macros::ColumnEnum;
2822/// use sz_orm_core::ColumnTrait;
2823///
2824/// #[derive(ColumnEnum)]
2825/// struct User {
2826///     id: i64,
2827///     #[column(name = "user_name")]
2828///     name: String,
2829/// }
2830///
2831/// assert_eq!(UserColumn::Id.as_str(), "id");
2832/// assert_eq!(UserColumn::Name.as_str(), "user_name");
2833/// assert_eq!(UserColumn::Id.to_string(), "id");
2834/// ```
2835#[proc_macro_derive(ColumnEnum, attributes(column))]
2836pub fn derive_column_enum(input: TokenStream) -> TokenStream {
2837    let input = parse_macro_input!(input as syn::DeriveInput);
2838    derive::derive_column_enum_impl(input).into()
2839}
2840
2841// ---------------------------------------------------------------------------
2842// `#[derive(FromRow)]` — auto-generate `impl FromRow for Struct`
2843// ---------------------------------------------------------------------------
2844
2845/// 派生宏:自动生成 `sz_orm_core::queryable::FromRow` trait 实现。
2846///
2847/// 从 `HashMap<String, Value>` 按列名反序列化为结构体实例。
2848/// 与 `FromQueryResult` 的区别在于错误类型为 `QueryError`(含列信息),
2849/// 适合需要精确错误定位的底层场景。
2850///
2851/// # 支持的属性
2852///
2853/// - `#[column(name = "...")]` — 覆盖列名映射(默认使用字段名)
2854///
2855/// # 示例
2856///
2857/// ```ignore
2858/// use sz_orm_macros::FromRow;
2859///
2860/// #[derive(FromRow)]
2861/// struct User {
2862///     id: i64,
2863///     name: String,
2864///     #[column(name = "user_email")]
2865///     email: Option<String>,
2866/// }
2867/// ```
2868#[proc_macro_derive(FromRow, attributes(column))]
2869pub fn derive_from_row(input: TokenStream) -> TokenStream {
2870    let input = parse_macro_input!(input as syn::DeriveInput);
2871    derive::derive_from_row_impl(input).into()
2872}
2873
2874// ---------------------------------------------------------------------------
2875// `#[derive(SqlType)]` — auto-generate `impl FromQueryResult + to_value()` for enums
2876// ---------------------------------------------------------------------------
2877
2878/// 派生宏:为 Rust 枚举自动生成 `sz_orm_core::FromQueryResult` trait 实现
2879/// 和 `to_value()` 方法。
2880///
2881/// 这是 sz-orm 对 SQLx `#[derive(Type)]` 的等效实现:
2882/// 让自定义枚举可以直接用于查询结果的字段映射和查询参数的绑定。
2883///
2884/// # 支持的属性
2885///
2886/// - `#[sql_type(rename_all = "snake_case")]` — 控制变体名的序列化格式
2887///   (snake_case / SCREAMING_SNAKE_CASE / camelCase / PascalCase / lowercase / UPPERCASE)
2888/// - `#[sql_type(rename = "...")]`(变体级)— 覆盖单个变体的序列化名
2889///
2890/// # 示例
2891///
2892/// ```ignore
2893/// use sz_orm_macros::SqlType;
2894///
2895/// #[derive(SqlType)]
2896/// enum Status {
2897///     Active,    // → "active"
2898///     Inactive,  // → "inactive"
2899/// }
2900///
2901/// let v = Status::Active.to_value();  // Value::String("active")
2902/// ```
2903#[proc_macro_derive(SqlType, attributes(sql_type))]
2904pub fn derive_sql_type(input: TokenStream) -> TokenStream {
2905    let input = parse_macro_input!(input as syn::DeriveInput);
2906    derive::derive_sql_type_impl(input).into()
2907}
2908
2909// ---------------------------------------------------------------------------
2910// `#[derive(Relation)]` — auto-generate `impl ModelExt` with relations()
2911// ---------------------------------------------------------------------------
2912
2913/// 派生宏:自动生成 `sz_orm_core::model::ModelExt` trait 实现,
2914/// 填充 `relations()` 映射,消除手写关系样板代码。
2915///
2916/// # 支持的属性
2917///
2918/// - `#[relation(has_many = "orders", fk = "user_id", pk = "id")]`
2919/// - `#[relation(belongs_to = "users", fk = "user_id", pk = "id")]`
2920/// - `#[relation(has_one = "profile", fk = "user_id", pk = "id")]`
2921/// - `#[relation(belongs_to_many = "roles", junction = "user_roles", fk = "user_id", other_key = "role_id", target = "roles", target_pk = "id")]`
2922/// - `#[relation(morph_many = "comments", morph_type = "commentable_type", morph_id = "commentable_id", morph_type_value = "Post")]`
2923/// - `#[relation(morph_to, morph_type = "commentable_type", morph_id = "commentable_id")]`
2924///
2925/// # 示例
2926///
2927/// ```ignore
2928/// use sz_orm_macros::{Entity, Relation};
2929///
2930/// #[derive(Entity, Relation)]
2931/// #[table(name = "users")]
2932/// struct User {
2933///     #[column(primary_key)]
2934///     id: i64,
2935/// }
2936///
2937/// // 自动生成:
2938/// // impl ModelExt for User {
2939/// //     fn relations() -> HashMap<&str, Relation> {
2940/// //         // 包含 #[relation] 定义的关系
2941/// //     }
2942/// // }
2943/// ```
2944#[proc_macro_derive(Relation, attributes(relation, table, column))]
2945pub fn derive_relation(input: TokenStream) -> TokenStream {
2946    let input = parse_macro_input!(input as syn::DeriveInput);
2947    derive::derive_relation_impl(input).into()
2948}
2949
2950/// `#[derive(RelationTrait)]` — 自动生成 `RelationTrait` 实现(P-F-2, v2.1.0)
2951///
2952/// 从 `#[relation(...)]` 属性生成 `RelationDef` 静量表 + `impl RelationTrait`。
2953/// 与 `#[derive(Relation)]` 共享属性解析,但生成零分配的静态切片而非 `HashMap`。
2954///
2955/// # 示例
2956///
2957/// ```ignore
2958/// #[derive(RelationTrait)]
2959/// #[relation(has_many = "Order", fk = "user_id", pk = "id")]
2960/// struct User { id: i64, name: String }
2961///
2962/// // 自动生成:
2963/// // static RELATIONS: &[RelationDef] = &[RelationDef::new("Order", "users", "orders", "id", "user_id", HasMany)];
2964/// // impl RelationTrait for User { ... }
2965/// ```
2966#[proc_macro_derive(RelationTrait, attributes(relation, table, column))]
2967pub fn derive_relation_trait(input: TokenStream) -> TokenStream {
2968    let input = parse_macro_input!(input as syn::DeriveInput);
2969    derive::derive_relation_trait_impl(input).into()
2970}
2971
2972// ---------------------------------------------------------------------------
2973// v3.9.0:数据验证派生宏(data-validation feature 隔离)
2974// ---------------------------------------------------------------------------
2975
2976/// 派生 `Validate` trait 实现。
2977///
2978/// 支持的 `#[validate(...)]` 规则:
2979/// - `email` — 邮箱格式校验
2980/// - `required` — 非空校验
2981/// - `length(min=N, max=N)` — 长度范围校验
2982/// - `range(min=N, max=N)` — 数值范围校验
2983/// - `regex(pattern=r"...")` — 正则匹配校验
2984/// - `contains(value="...")` — 包含子串校验
2985/// - `does_not_contain(value="...")` — 不包含子串校验
2986/// - `if = "condition"` — 条件校验
2987///
2988/// # 示例
2989///
2990/// ```ignore
2991/// #[derive(Validate)]
2992/// struct User {
2993///     #[validate(email)]
2994///     email: String,
2995///     #[validate(length(min=2, max=50))]
2996///     name: String,
2997///     #[validate(range(min=0, max=150))]
2998///     age: i64,
2999/// }
3000/// ```
3001#[cfg(feature = "data-validation")]
3002#[proc_macro_derive(Validate, attributes(validate))]
3003pub fn derive_validate(input: TokenStream) -> TokenStream {
3004    crate::derive_validate::derive_validate_impl(input)
3005}
3006
3007// ---------------------------------------------------------------------------
3008// v4.3.0 M3-T3:`#[derive(Governed)]` — 编译期数据治理(PII 标注强制)
3009// 通过 `compile-governance` feature(sz-orm-core)→ `governance-derive`(本包)启用。
3010// ---------------------------------------------------------------------------
3011
3012/// 派生宏:为模型生成数据治理元数据,并**编译期强制** PII 标注合规。
3013///
3014/// # 支持属性
3015///
3016/// - `#[pii]` — 标记字段为个人敏感数据(PII)
3017/// - `#[mask(strategy = "...")]` — 声明脱敏策略(hash / partial / replace / encrypt)
3018///
3019/// # 编译期强制规则
3020///
3021/// 1. `#[pii]` 字段必须同时声明 `#[mask(strategy = "...")]`,否则 **编译失败**
3022/// 2. `mask` 策略必须在白名单内(hash/partial/replace/encrypt),否则 **编译失败**
3023///
3024/// # 示例
3025///
3026/// ```ignore
3027/// use sz_orm_core::governance::GovernedModel;
3028///
3029/// #[derive(Governed)]
3030/// struct User {
3031///     id: i64,
3032///     #[pii]
3033///     #[mask(strategy = "partial")]
3034///     email: String,
3035///     #[pii]
3036///     #[mask(strategy = "hash")]
3037///     phone: String,
3038///     name: String, // 非 PII,无需 mask
3039/// }
3040///
3041/// assert_eq!(
3042///     User::pii_fields(),
3043///     vec![("email", "partial"), ("phone", "hash")]
3044/// );
3045/// ```
3046#[cfg(feature = "governance-derive")]
3047#[proc_macro_derive(Governed, attributes(pii, mask))]
3048pub fn derive_governed(input: TokenStream) -> TokenStream {
3049    let input = parse_macro_input!(input as syn::DeriveInput);
3050    let name = &input.ident;
3051
3052    const VALID_STRATEGIES: [&str; 4] = ["hash", "partial", "replace", "encrypt"];
3053
3054    let mut pii_fields: Vec<(String, String)> = Vec::new();
3055    let mut errors: Vec<syn::Error> = Vec::new();
3056
3057    if let syn::Data::Struct(data) = &input.data {
3058        for field in &data.fields {
3059            let Some(field_name) = field.ident.as_ref().map(|i| i.to_string()) else {
3060                continue;
3061            };
3062            let is_pii = field.attrs.iter().any(|a| a.path().is_ident("pii"));
3063
3064            // 解析 #[mask(strategy = "xxx")]
3065            let mut mask_strategy: Option<String> = None;
3066            for attr in &field.attrs {
3067                if !attr.path().is_ident("mask") {
3068                    continue;
3069                }
3070                let _ = attr.parse_nested_meta(|meta| {
3071                    if meta.path.is_ident("strategy") {
3072                        let lit: syn::LitStr = meta.value()?.parse()?;
3073                        mask_strategy = Some(lit.value());
3074                        Ok(())
3075                    } else {
3076                        Err(meta.error("unsupported #[mask] attribute, only 'strategy' is allowed"))
3077                    }
3078                });
3079            }
3080
3081            if is_pii {
3082                match mask_strategy {
3083                    Some(strategy) => {
3084                        if !VALID_STRATEGIES.contains(&strategy.as_str()) {
3085                            errors.push(syn::Error::new_spanned(
3086                                field,
3087                                format!(
3088                                    "invalid #[mask(strategy = \"{strategy}\")]: allowed strategies are {:?}",
3089                                    VALID_STRATEGIES
3090                                ),
3091                            ));
3092                        } else {
3093                            pii_fields.push((field_name, strategy));
3094                        }
3095                    }
3096                    None => errors.push(syn::Error::new_spanned(
3097                        field,
3098                        "#[pii] field must declare #[mask(strategy = \"...\")]",
3099                    )),
3100                }
3101            }
3102        }
3103    }
3104
3105    if !errors.is_empty() {
3106        // 每条错误独立转 compile_error!,合并为一个 TokenStream
3107        let err_tokens: proc_macro2::TokenStream =
3108            errors.iter().map(|e| e.to_compile_error()).collect();
3109        return err_tokens.into();
3110    }
3111
3112    let entries = pii_fields.iter().map(|(f, s)| {
3113        let f = f.as_str();
3114        let s = s.as_str();
3115        quote::quote!((#f, #s))
3116    });
3117
3118    quote::quote! {
3119        impl ::sz_orm_core::governance::GovernedModel for #name {
3120            fn pii_fields() -> Vec<(&'static str, &'static str)> {
3121                vec![#(#entries),*]
3122            }
3123        }
3124    }
3125    .into()
3126}
3127
3128// ---------------------------------------------------------------------------
3129// v4.3.0 M2:`#[detect_n_plus_one]` — N+1 静态检测标注宏(n1-lint feature)
3130// 分析函数体 AST,检测循环(for/while)内查询调用,编译期输出警告。
3131// 分析逻辑复用 sz-orm-n1-lint(与 CLI 批量扫描共用,避免重复实现)。
3132// ---------------------------------------------------------------------------
3133
3134/// 标注宏:分析函数体,检测 N+1 查询模式并输出编译期警告(非阻断)。
3135///
3136/// 检测模式(保守白名单,避免误报):
3137/// - 循环体内 `find_by_*` / `where_eq` / `query` 等查询调用 → `query-in-loop` 警告
3138/// - 循环内条件分支中的查询调用 → `query-in-loop` 警告
3139///
3140/// 警告通过 `eprintln!` 输出(`warning: [sz-orm-n1-lint] ...`),不阻断编译;
3141/// 函数原样透传(零运行时影响)。
3142///
3143/// # 示例
3144///
3145/// ```ignore
3146/// #[detect_n_plus_one]
3147/// fn process_users(users: Vec<User>) {
3148///     for user in users {
3149///         let orders = Order::find_by_user(user.id); // ⚠️ 编译期警告:query-in-loop
3150///     }
3151/// }
3152/// ```
3153#[cfg(feature = "n1-lint")]
3154#[proc_macro_attribute]
3155pub fn detect_n_plus_one(_attr: TokenStream, item: TokenStream) -> TokenStream {
3156    let item_fn = parse_macro_input!(item as syn::ItemFn);
3157    let findings = sz_orm_n1_lint::analyze_fn(&item_fn);
3158    for f in &findings {
3159        eprintln!(
3160            "warning: [sz-orm-n1-lint] {} at line {}: {}",
3161            f.pattern.as_str(),
3162            f.line,
3163            f.message
3164        );
3165    }
3166    quote::quote!(#item_fn).into()
3167}
3168
3169// ---------------------------------------------------------------------------
3170// Unit tests — cover helper functions used by both macros
3171// ---------------------------------------------------------------------------
3172
3173#[cfg(test)]
3174mod tests {
3175    use super::*;
3176
3177    // ---- strip_string_literal ----
3178
3179    #[test]
3180    fn test_strip_plain_double_quoted() {
3181        assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
3182    }
3183
3184    #[test]
3185    fn test_strip_raw_double_hash() {
3186        assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
3187    }
3188
3189    #[test]
3190    fn test_strip_raw_double_no_hash() {
3191        assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
3192    }
3193
3194    #[test]
3195    fn test_strip_byte_string() {
3196        assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
3197        assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
3198    }
3199
3200    #[test]
3201    fn test_strip_non_string_returns_none() {
3202        assert_eq!(strip_string_literal("123"), None);
3203        assert_eq!(strip_string_literal("foo"), None);
3204    }
3205
3206    // ---- validate_sql_content ----
3207
3208    #[test]
3209    fn test_validate_select_with_from_ok() {
3210        assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
3211    }
3212
3213    #[test]
3214    fn test_validate_select_missing_from_fails() {
3215        assert!(validate_sql_content("SELECT * users", None).is_err());
3216    }
3217
3218    #[test]
3219    fn test_validate_insert_missing_into_fails() {
3220        assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
3221        assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
3222    }
3223
3224    #[test]
3225    fn test_validate_update_missing_set_fails() {
3226        assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
3227        assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
3228    }
3229
3230    #[test]
3231    fn test_validate_delete_missing_from_fails() {
3232        assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
3233        assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
3234    }
3235
3236    #[test]
3237    fn test_validate_empty_sql_fails() {
3238        assert!(validate_sql_content("", None).is_err());
3239        assert!(validate_sql_content("   ", None).is_err());
3240    }
3241
3242    // ---- balanced parens ----
3243
3244    #[test]
3245    fn test_validate_balanced_parens_ok() {
3246        assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
3247    }
3248
3249    #[test]
3250    fn test_validate_balanced_parens_unbalanced() {
3251        assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
3252        assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
3253    }
3254
3255    // ---- injection patterns ----
3256
3257    #[test]
3258    fn test_validate_no_injection_clean() {
3259        assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
3260    }
3261
3262    #[test]
3263    fn test_validate_no_injection_drop_table() {
3264        assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
3265    }
3266
3267    #[test]
3268    fn test_validate_no_injection_or_1_1() {
3269        // 编译期 SQL 已剥离外层引号,检测模式不再依赖引号字符。
3270        // "' OR '1'='1" 因引号分隔不再匹配 "or 1=1",故不再检测;
3271        // 但不含引号分隔的 "OR 1=1" 仍可被检测。
3272        assert!(validate_no_injection("' OR 1=1").is_err());
3273        assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
3274    }
3275
3276    #[test]
3277    fn test_validate_no_injection_drop_database() {
3278        assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
3279    }
3280
3281    #[test]
3282    fn test_validate_no_injection_information_schema() {
3283        assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
3284    }
3285
3286    #[test]
3287    fn test_validate_no_injection_xp_cmdshell() {
3288        assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
3289    }
3290
3291    #[test]
3292    fn test_validate_no_injection_union_select() {
3293        assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
3294    }
3295
3296    #[test]
3297    fn test_validate_no_injection_comment_dashes() {
3298        assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
3299    }
3300
3301    #[test]
3302    fn test_validate_no_injection_block_comment() {
3303        assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
3304    }
3305
3306    // ---- string literal closure ----
3307
3308    #[test]
3309    fn test_validate_string_literals_closed_ok() {
3310        assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
3311        assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
3312    }
3313
3314    #[test]
3315    fn test_validate_string_literals_closed_unclosed_single() {
3316        assert!(validate_string_literals_closed("'hello").is_err());
3317    }
3318
3319    #[test]
3320    fn test_validate_string_literals_closed_unclosed_double() {
3321        assert!(validate_string_literals_closed(r#""hello"#).is_err());
3322    }
3323
3324    // ---- param count check ----
3325
3326    #[test]
3327    fn test_validate_param_count_match() {
3328        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
3329        assert!(
3330            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
3331        );
3332    }
3333
3334    #[test]
3335    fn test_validate_param_count_mismatch() {
3336        assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
3337        assert!(
3338            validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
3339        );
3340    }
3341
3342    // ---- db-verify feature: detect_db_kind ----
3343
3344    #[cfg(feature = "db-verify")]
3345    #[test]
3346    fn test_detect_db_kind_mysql() {
3347        assert_eq!(
3348            detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
3349            DbKind::MySql
3350        );
3351    }
3352
3353    #[cfg(feature = "db-verify")]
3354    #[test]
3355    fn test_detect_db_kind_postgres() {
3356        assert_eq!(
3357            detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
3358            DbKind::Postgres
3359        );
3360        assert_eq!(
3361            detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
3362            DbKind::Postgres
3363        );
3364    }
3365
3366    #[cfg(feature = "db-verify")]
3367    #[test]
3368    fn test_detect_db_kind_sqlite() {
3369        assert_eq!(
3370            detect_db_kind("sqlite://path/to/db.db").unwrap(),
3371            DbKind::Sqlite
3372        );
3373        assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
3374    }
3375
3376    #[cfg(feature = "db-verify")]
3377    #[test]
3378    fn test_detect_db_kind_oracle() {
3379        assert_eq!(
3380            detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
3381            DbKind::Oracle
3382        );
3383        assert_eq!(
3384            detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
3385            DbKind::Oracle
3386        );
3387    }
3388
3389    #[cfg(feature = "db-verify")]
3390    #[test]
3391    fn test_detect_db_kind_sqlserver() {
3392        assert_eq!(
3393            detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
3394            DbKind::SqlServer
3395        );
3396        assert_eq!(
3397            detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
3398            DbKind::SqlServer
3399        );
3400        assert_eq!(
3401            detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
3402            DbKind::SqlServer
3403        );
3404    }
3405
3406    #[cfg(feature = "db-verify")]
3407    #[test]
3408    fn test_detect_db_kind_unsupported() {
3409        assert!(detect_db_kind("redis://user:pass@host/db").is_err());
3410        assert!(detect_db_kind("not-a-url").is_err());
3411    }
3412
3413    #[cfg(feature = "db-verify")]
3414    #[test]
3415    fn test_parse_oracle_dsn_basic() {
3416        let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
3417        let p = parse_oracle_dsn(dsn).unwrap();
3418        assert_eq!(p.user, "sys");
3419        assert_eq!(p.password, "test123");
3420        assert_eq!(p.host, "127.0.0.1");
3421        assert_eq!(p.port, 1521);
3422        assert_eq!(p.service, "freepdb1.FALSE");
3423        assert!(p.sysdba);
3424    }
3425
3426    #[cfg(feature = "db-verify")]
3427    #[test]
3428    fn test_parse_oracle_dsn_default_port() {
3429        // 无端口号时默认 1521
3430        let dsn = "oracle://sys:test123@127.0.0.1/FREE";
3431        let p = parse_oracle_dsn(dsn).unwrap();
3432        assert_eq!(p.port, 1521);
3433        assert_eq!(p.service, "FREE");
3434        assert!(!p.sysdba);
3435    }
3436
3437    #[cfg(feature = "db-verify")]
3438    #[test]
3439    fn test_parse_sqlserver_dsn_basic() {
3440        let dsn =
3441            "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
3442        let p = parse_sqlserver_dsn(dsn).unwrap();
3443        assert_eq!(p.user, "test");
3444        assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
3445        assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
3446        assert_eq!(p.port, 22527);
3447        assert_eq!(p.database, "test");
3448    }
3449
3450    #[cfg(feature = "db-verify")]
3451    #[test]
3452    fn test_parse_sqlserver_dsn_default_port() {
3453        let dsn = "mssql://user:pass@host/db";
3454        let p = parse_sqlserver_dsn(dsn).unwrap();
3455        assert_eq!(p.port, 1433);
3456        assert_eq!(p.database, "db");
3457    }
3458
3459    // ---- schema! 宏 parse_create_table 测试 ----
3460
3461    #[test]
3462    fn test_parse_create_table_basic() {
3463        let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
3464        let (table, cols) = parse_create_table(sql).unwrap();
3465        assert_eq!(table, "users");
3466        assert_eq!(
3467            cols,
3468            vec![
3469                ("id".to_string(), "i32".to_string()),
3470                ("name".to_string(), "String".to_string())
3471            ]
3472        );
3473    }
3474
3475    #[test]
3476    fn test_parse_create_table_with_if_not_exists() {
3477        let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
3478        let (table, cols) = parse_create_table(sql).unwrap();
3479        assert_eq!(table, "orders");
3480        assert_eq!(
3481            cols,
3482            vec![
3483                ("id".to_string(), "i64".to_string()),
3484                ("total".to_string(), "f64".to_string())
3485            ]
3486        );
3487    }
3488
3489    #[test]
3490    fn test_parse_create_table_nullable() {
3491        let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
3492        let (_, cols) = parse_create_table(sql).unwrap();
3493        assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
3494        assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
3495    }
3496
3497    #[test]
3498    fn test_parse_create_table_skip_constraints() {
3499        let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
3500        let (_, cols) = parse_create_table(sql).unwrap();
3501        assert_eq!(cols.len(), 2);
3502        assert_eq!(cols[0].0, "id");
3503        assert_eq!(cols[1].0, "name");
3504    }
3505
3506    #[test]
3507    fn test_parse_create_table_varchar_with_len() {
3508        let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
3509        let (_, cols) = parse_create_table(sql).unwrap();
3510        assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
3511        assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
3512    }
3513
3514    #[test]
3515    fn test_sql_type_to_rust_mappings() {
3516        // 整数(按字节宽度严格映射,与 SQL 标准一致)
3517        assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
3518        assert_eq!(sql_type_to_rust("INT8", false), "i64");
3519        assert_eq!(sql_type_to_rust("INT", false), "i32");
3520        assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
3521        assert_eq!(sql_type_to_rust("INT4", false), "i32");
3522        assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
3523        assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
3524        assert_eq!(sql_type_to_rust("INT2", false), "i16");
3525        assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
3526        assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
3527        // 浮点
3528        assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
3529        assert_eq!(sql_type_to_rust("REAL", false), "f32");
3530        assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
3531        assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
3532        assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
3533        assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
3534        assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
3535        assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
3536        // 布尔
3537        assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
3538        assert_eq!(sql_type_to_rust("BOOL", false), "bool");
3539        // 字符串
3540        assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
3541        assert_eq!(sql_type_to_rust("TEXT", false), "String");
3542        assert_eq!(sql_type_to_rust("CHAR", false), "String");
3543        assert_eq!(sql_type_to_rust("UUID", false), "String");
3544        assert_eq!(sql_type_to_rust("DATE", false), "String");
3545        assert_eq!(sql_type_to_rust("DATETIME", false), "String");
3546        assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
3547        assert_eq!(sql_type_to_rust("JSON", false), "String");
3548        assert_eq!(sql_type_to_rust("JSONB", false), "String");
3549        // 二进制
3550        assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
3551        assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
3552        assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
3553        assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
3554        // nullable
3555        assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
3556        assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
3557        assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
3558        assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
3559        // unknown
3560        assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
3561    }
3562
3563    #[test]
3564    fn test_parse_create_table_error_no_create() {
3565        assert!(parse_create_table("SELECT * FROM users").is_err());
3566    }
3567
3568    #[test]
3569    fn test_parse_create_table_error_no_parens() {
3570        assert!(parse_create_table("CREATE TABLE foo").is_err());
3571    }
3572
3573    // -----------------------------------------------------------------------
3574    // Gap 1 测试:列名/表名提取
3575    // -----------------------------------------------------------------------
3576
3577    #[cfg(feature = "db-verify")]
3578    #[test]
3579    fn test_extract_tables_simple() {
3580        let tables = extract_tables("SELECT id, name FROM users WHERE id = ?");
3581        assert!(tables.contains(&"users".to_string()));
3582    }
3583
3584    #[cfg(feature = "db-verify")]
3585    #[test]
3586    fn test_extract_tables_multiple() {
3587        let tables = extract_tables(
3588            "SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id WHERE u.id = ?",
3589        );
3590        assert!(tables.contains(&"users".to_string()));
3591        assert!(tables.contains(&"orders".to_string()));
3592    }
3593
3594    #[cfg(feature = "db-verify")]
3595    #[test]
3596    fn test_extract_columns_select_and_where() {
3597        let cols =
3598            extract_columns("SELECT id, name FROM users WHERE email = ? ORDER BY created_at");
3599        // id, name from SELECT; email from WHERE; created_at from ORDER BY
3600        assert!(cols.contains(&"id".to_string()));
3601        assert!(cols.contains(&"name".to_string()));
3602        assert!(cols.contains(&"email".to_string()));
3603        assert!(cols.contains(&"created_at".to_string()));
3604    }
3605
3606    #[cfg(feature = "db-verify")]
3607    #[test]
3608    fn test_extract_columns_skips_keywords() {
3609        let cols = extract_columns("SELECT COUNT(id), name FROM users WHERE status = ?");
3610        // COUNT is a function, should be skipped
3611        assert!(!cols.contains(&"count".to_string()));
3612        assert!(cols.contains(&"id".to_string()));
3613        assert!(cols.contains(&"name".to_string()));
3614        assert!(cols.contains(&"status".to_string()));
3615    }
3616
3617    #[cfg(feature = "db-verify")]
3618    #[test]
3619    fn test_is_sql_function() {
3620        assert!(is_sql_function("COUNT"));
3621        assert!(is_sql_function("now"));
3622        assert!(is_sql_function("COALESCE"));
3623        assert!(!is_sql_function("name"));
3624        assert!(!is_sql_function("user_id"));
3625    }
3626}