1#![allow(linker_messages)]
35extern 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
63use proc_macro2::TokenStream as TokenStream2;
65use quote::quote;
66use syn::parse_macro_input;
67
68mod derive;
70
71#[cfg(feature = "data-validation")]
73mod derive_validate;
74
75#[cfg(any(feature = "typed-dsl", feature = "custom-diagnostic"))]
77mod diagnostic;
78
79#[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#[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#[proc_macro]
178pub fn sql_string(input: TokenStream) -> TokenStream {
179 let mut tokens = input.into_iter().peekable();
180
181 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 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 let mut expected_params = None;
216 if tokens.peek().is_some() {
217 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 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 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 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 if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
284 return compile_error(Span::call_site(), &err_msg);
285 }
286
287 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
294fn 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 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 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 let injection_patterns: &[&str] = &[
401 "drop table",
403 "drop database",
404 "; drop",
405 "or 1=1",
407 "or 1 = 1",
408 "union select",
409 "union all select",
410 "--",
412 "/*",
413 "*/",
414 "xp_cmdshell",
416 "sp_executesql",
417 "exec(",
418 "execute(",
419 "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#[proc_macro]
468pub fn query(input: TokenStream) -> TokenStream {
469 let mut tokens = input.into_iter().peekable();
470
471 let type_param: Option<TokenStream2> = match tokens.peek() {
474 Some(TokenTree::Ident(_)) | Some(TokenTree::Punct(_)) => {
475 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 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 match tokens.peek() {
494 Some(TokenTree::Punct(p)) if p.as_char() == ',' => {
495 tokens.next(); let ts: proc_macro::TokenStream = ty_tokens.into_iter().collect();
497 Some(TokenStream2::from(ts))
498 }
499 _ => None, }
501 }
502 _ => None,
503 };
504
505 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 if let Err(err_msg) = validate_sql_content(sql_content, None) {
542 return compile_error(Span::call_site(), &err_msg);
543 }
544
545 #[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 Some("1") => match verify_with_real_db(sql_content) {
551 Ok(cols) => {
552 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 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 let escaped = sql_content.escape_default().to_string();
583 let base = if let Some(ref ty) = type_param {
584 format!(
586 "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
587 ty, escaped
588 )
589 } else {
590 format!("::sz_orm_core::queryable::Query::new(\"{}\")", escaped)
592 };
593 #[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
608fn 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#[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 let sql_no_placeholders = replace_placeholders_with_null(sql);
639
640 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 DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
646 DbKind::SqlServer => sql_no_placeholders,
648 };
649
650 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 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 verify_sqlite(&dsn, &explain_sql).await?;
663 }
664 verify_columns(&dsn, db_kind, sql).await?;
666 fetch_column_types(&dsn, db_kind, sql).await
669 });
670 }
671
672 if let DbKind::Oracle = db_kind {
674 verify_oracle(&dsn, &explain_sql).map(|_| Vec::new())
675 } else {
676 verify_sqlserver(&dsn, &explain_sql).map(|_| Vec::new())
678 }
679}
680
681#[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#[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 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#[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#[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#[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#[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 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#[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#[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 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#[cfg(feature = "db-verify")]
998async fn verify_columns(dsn: &str, db_kind: DbKind, sql: &str) -> Result<(), String> {
999 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#[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 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 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 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 let table_word = part
1101 .split_whitespace()
1102 .next()
1103 .unwrap_or(part)
1104 .trim_end_matches([',', ';']);
1105 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#[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 let mut collect_from_segment = |segment: &str| {
1137 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 if word.chars().all(|c| c.is_ascii_digit()) {
1212 continue;
1213 }
1214 let pos = segment.find(word).unwrap_or(0);
1217 if pos > 0 && segment.chars().nth(pos - 1) == Some('.') {
1218 continue;
1219 }
1220 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 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 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 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 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 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 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 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#[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 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 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#[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#[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#[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#[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#[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#[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 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 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#[cfg(feature = "db-verify")]
1727struct SqlServerDsn {
1728 user: String,
1729 password: String,
1730 host: String,
1731 port: u16,
1732 database: String,
1733}
1734
1735#[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
1773fn compile_error(span: Span, msg: &str) -> TokenStream {
1779 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#[proc_macro]
1830pub fn typed_query(input: TokenStream) -> TokenStream {
1831 let tokens: Vec<TokenTree> = input.into_iter().collect();
1832
1833 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 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
1861fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
1863 let mut idx = 0;
1865
1866 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 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 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 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 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1912 let table_name_lit = table_name.as_str();
1913
1914 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 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 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
1966fn 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 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 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 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
2029fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
2033 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 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 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 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 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 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 let mut ts = TokenStream::new();
2152 let lit = Literal::string(&sql);
2153 ts.extend([TokenTree::Literal(lit)]);
2154 ts
2155}
2156
2157#[proc_macro]
2187pub fn query_as(input: TokenStream) -> TokenStream {
2209 let mut tokens = input.into_iter().peekable();
2210
2211 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 record_type.push_str("::");
2221 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 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 if let Err(err_msg) = validate_sql_content(sql_content, None) {
2268 return compile_error(Span::call_site(), &err_msg);
2269 }
2270
2271 #[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 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 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 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 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 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 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
2414fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
2422 let trimmed = sql.trim();
2423 let upper = trimmed.to_uppercase();
2424
2425 if !upper.starts_with("CREATE TABLE") {
2427 return Err("schema! expects a CREATE TABLE statement".to_string());
2428 }
2429
2430 let mut rest = &trimmed["CREATE TABLE".len()..];
2432
2433 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 let (table_name, after_name) = parse_identifier(rest)?;
2443 let rest = after_name.trim_start();
2444
2445 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 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 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 let (col_name, after_col) = parse_identifier(def)?;
2482 let rest = after_col.trim_start();
2483
2484 let (sql_type, after_type) = parse_type_token(rest)?;
2486 let rest = after_type.trim();
2487
2488 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
2499fn 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
2536fn 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 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
2563fn 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
2595fn 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 "BIGINT" | "INT8" => "i64",
2604 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
2606 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
2608 "TINYINT" => "i8",
2610 "FLOAT" | "REAL" | "FLOAT4" => "f32",
2612 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
2614 "BOOLEAN" | "BOOL" => "bool",
2616 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
2618 "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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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 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 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#[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#[cfg(test)]
3174mod tests {
3175 use super::*;
3176
3177 #[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 #[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 #[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 #[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 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 #[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 #[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 #[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 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 #[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 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 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 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
3538 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
3539 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 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 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 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 #[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 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 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}