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!(
1588 "{}/{}@{}:{}/{}",
1589 parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
1590 );
1591 if parsed.sysdba {
1592 conn_str.push_str(" AS SYSDBA");
1593 }
1594 let full_script = format!(
1596 "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
1597 EXPLAIN PLAN FOR {};\n\
1598 SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
1599 EXIT;\n",
1600 explain_sql
1601 );
1602 let output = std::process::Command::new("sqlplus")
1603 .args(["-S", "-L", &conn_str])
1604 .stdin(std::process::Stdio::piped())
1605 .stdout(std::process::Stdio::piped())
1606 .stderr(std::process::Stdio::piped())
1607 .spawn()
1608 .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
1609 use std::io::Write;
1610 let mut child = output;
1611 if let Some(mut stdin) = child.stdin.take() {
1612 stdin
1613 .write_all(full_script.as_bytes())
1614 .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
1615 }
1616 let out = child
1617 .wait_with_output()
1618 .map_err(|e| format!("sqlplus wait failed: {}", e))?;
1619 let stdout = String::from_utf8_lossy(&out.stdout);
1620 let stderr = String::from_utf8_lossy(&out.stderr);
1621 if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
1622 return Err(format!(
1623 "Oracle EXPLAIN failed: stdout={} stderr={}",
1624 stdout.trim(),
1625 stderr.trim()
1626 ));
1627 }
1628 Ok(())
1629}
1630
1631#[cfg(feature = "db-verify")]
1636fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
1637 let parsed = parse_sqlserver_dsn(dsn)?;
1638 let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
1640 let out = std::process::Command::new("sqlcmd")
1641 .args([
1642 "-S",
1643 &format!("{},{}", parsed.host, parsed.port),
1644 "-U",
1645 &parsed.user,
1646 "-P",
1647 &parsed.password,
1648 "-d",
1649 &parsed.database,
1650 "-Q",
1651 &query,
1652 "-h",
1653 "-1",
1654 "-W",
1655 ])
1656 .output()
1657 .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
1658 let stdout = String::from_utf8_lossy(&out.stdout);
1659 let stderr = String::from_utf8_lossy(&out.stderr);
1660 if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
1661 return Err(format!(
1662 "SQL Server SHOWPLAN failed: stdout={} stderr={}",
1663 stdout.trim(),
1664 stderr.trim()
1665 ));
1666 }
1667 Ok(())
1668}
1669
1670#[cfg(feature = "db-verify")]
1672struct OracleDsn {
1673 user: String,
1674 password: String,
1675 host: String,
1676 port: u16,
1677 service: String,
1678 sysdba: bool,
1679}
1680
1681#[cfg(feature = "db-verify")]
1683fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
1684 let raw = dsn
1685 .strip_prefix("oracle://")
1686 .or_else(|| dsn.strip_prefix("oracle:"))
1687 .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
1688 let (auth_host_service, query) = match raw.find('?') {
1690 Some(idx) => (&raw[..idx], &raw[idx + 1..]),
1691 None => (raw, ""),
1692 };
1693 let sysdba = query
1694 .split('&')
1695 .any(|p| p == "sysdba=1" || p == "sysdba=true");
1696 let at = auth_host_service
1698 .find('@')
1699 .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
1700 let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
1701 let colon = user_pass
1702 .find(':')
1703 .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
1704 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1705 let (host_port, service) = match host_port_service.rfind('/') {
1706 Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
1707 None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
1708 };
1709 let (host, port) = match host_port.find(':') {
1710 Some(idx) => (
1711 &host_port[..idx],
1712 host_port[idx + 1..]
1713 .parse::<u16>()
1714 .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
1715 ),
1716 None => (host_port, 1521u16),
1717 };
1718 Ok(OracleDsn {
1719 user: user.to_string(),
1720 password: password.to_string(),
1721 host: host.to_string(),
1722 port,
1723 service: service.to_string(),
1724 sysdba,
1725 })
1726}
1727
1728#[cfg(feature = "db-verify")]
1730struct SqlServerDsn {
1731 user: String,
1732 password: String,
1733 host: String,
1734 port: u16,
1735 database: String,
1736}
1737
1738#[cfg(feature = "db-verify")]
1740fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
1741 let raw = dsn
1742 .strip_prefix("sqlserver://")
1743 .or_else(|| dsn.strip_prefix("mssql://"))
1744 .or_else(|| dsn.strip_prefix("tds://"))
1745 .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
1746 let at = raw
1747 .find('@')
1748 .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
1749 let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
1750 let colon = user_pass
1751 .find(':')
1752 .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
1753 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1754 let (host_port, database) = match host_port_db.rfind('/') {
1755 Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
1756 None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
1757 };
1758 let (host, port) = match host_port.find(':') {
1759 Some(idx) => (
1760 &host_port[..idx],
1761 host_port[idx + 1..]
1762 .parse::<u16>()
1763 .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
1764 ),
1765 None => (host_port, 1433u16),
1766 };
1767 Ok(SqlServerDsn {
1768 user: user.to_string(),
1769 password: password.to_string(),
1770 host: host.to_string(),
1771 port,
1772 database: database.to_string(),
1773 })
1774}
1775
1776fn compile_error(span: Span, msg: &str) -> TokenStream {
1782 let mut ts = TokenStream::new();
1784 ts.extend([
1785 TokenTree::Ident(Ident::new("compile_error", span)),
1786 TokenTree::Punct(Punct::new('!', Spacing::Alone)),
1787 TokenTree::Group(Group::new(
1788 Delimiter::Parenthesis,
1789 TokenStream::from(TokenTree::Literal(Literal::string(msg))),
1790 )),
1791 ]);
1792 ts
1793}
1794
1795#[proc_macro]
1833pub fn typed_query(input: TokenStream) -> TokenStream {
1834 let tokens: Vec<TokenTree> = input.into_iter().collect();
1835
1836 if tokens.iter().any(|t| {
1838 if let TokenTree::Ident(id) = t {
1839 id.to_string() == "table"
1840 } else {
1841 false
1842 }
1843 }) {
1844 return parse_table_decl(&tokens);
1845 }
1846
1847 if tokens.iter().any(|t| {
1849 if let TokenTree::Ident(id) = t {
1850 id.to_string().eq_ignore_ascii_case("SELECT")
1851 } else {
1852 false
1853 }
1854 }) {
1855 return parse_typed_select(&tokens);
1856 }
1857
1858 compile_error(
1859 Span::call_site(),
1860 "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
1861 )
1862}
1863
1864fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
1866 let mut idx = 0;
1868
1869 if idx >= tokens.len() {
1871 return compile_error(Span::call_site(), "expected table name after 'table'");
1872 }
1873 if let TokenTree::Ident(id) = &tokens[idx] {
1874 if id.to_string() != "table" {
1875 return compile_error(id.span(), "expected 'table' keyword");
1876 }
1877 }
1878 idx += 1;
1879
1880 let table_name = if idx < tokens.len() {
1882 if let TokenTree::Ident(id) = &tokens[idx] {
1883 id.to_string()
1884 } else {
1885 return compile_error(tokens[idx].span(), "expected table name identifier");
1886 }
1887 } else {
1888 return compile_error(Span::call_site(), "expected table name");
1889 };
1890 idx += 1;
1891
1892 let body_group = if idx < tokens.len() {
1894 if let TokenTree::Group(g) = &tokens[idx] {
1895 if g.delimiter() != Delimiter::Brace {
1896 return compile_error(g.span(), "expected '{' after table name");
1897 }
1898 g.clone()
1899 } else {
1900 return compile_error(tokens[idx].span(), "expected '{' after table name");
1901 }
1902 } else {
1903 return compile_error(Span::call_site(), "expected table body in '{ }'");
1904 };
1905
1906 let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
1908 let columns = match parse_column_list(&body_tokens) {
1909 Ok(c) => c,
1910 Err(e) => return compile_error(Span::call_site(), &e),
1911 };
1912
1913 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1915 let table_name_lit = table_name.as_str();
1916
1917 let col_impls: Vec<TokenStream2> = columns
1919 .iter()
1920 .map(|(col_name, col_type)| {
1921 let col_ident =
1922 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1923 let col_name_lit = col_name.as_str();
1924 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1926 quote! {
1927 #[derive(Debug, Clone, Copy)]
1928 pub struct #col_ident;
1929 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1930 const NAME: &'static str = #col_name_lit;
1931 type Table = table;
1932 type RustType = #rust_type;
1933 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1934 }
1935 }
1936 })
1937 .collect();
1938
1939 let schema_entries: Vec<TokenStream2> = columns
1941 .iter()
1942 .map(|(n, t)| {
1943 let n_lit = n.as_str();
1944 let t_lit = t.as_str();
1945 quote! { (#n_lit, #t_lit) }
1946 })
1947 .collect();
1948
1949 let schema_const_ident = proc_macro2::Ident::new(
1950 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1951 Span::call_site().into(),
1952 );
1953
1954 let expanded = quote! {
1955 pub mod #table_ident {
1956 use super::*;
1957 pub struct table;
1958 impl ::sz_orm_core::typed::TypedTable for table {
1959 const NAME: &'static str = #table_name_lit;
1960 }
1961 #(#col_impls)*
1962 }
1963 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1964 };
1965
1966 expanded.into()
1967}
1968
1969fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
1971 let mut cols = Vec::new();
1972 let mut i = 0;
1973 while i < tokens.len() {
1974 let col_name = if let TokenTree::Ident(id) = &tokens[i] {
1976 id.to_string()
1977 } else {
1978 return Err(format!("expected column name at position {}", i));
1979 };
1980 i += 1;
1981
1982 if i >= tokens.len() {
1984 return Err(format!("expected ':' after column '{}'", col_name));
1985 }
1986 if let TokenTree::Punct(p) = &tokens[i] {
1987 if p.as_char() != ':' {
1988 return Err(format!("expected ':' after column '{}'", col_name));
1989 }
1990 } else {
1991 return Err(format!("expected ':' after column '{}'", col_name));
1992 }
1993 i += 1;
1994
1995 let mut type_str = String::new();
1998 let mut depth = 0;
1999 while i < tokens.len() {
2000 match &tokens[i] {
2001 TokenTree::Punct(p) => {
2002 if p.as_char() == ',' && depth == 0 {
2003 i += 1;
2004 break;
2005 } else if p.as_char() == '<' || p.as_char() == '(' {
2006 depth += 1;
2007 type_str.push(p.as_char());
2008 } else if p.as_char() == '>' || p.as_char() == ')' {
2009 depth -= 1;
2010 type_str.push(p.as_char());
2011 } else {
2012 type_str.push(p.as_char());
2013 }
2014 }
2015 TokenTree::Ident(id) => {
2016 if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
2017 {
2018 type_str.push(' ');
2019 }
2020 type_str.push_str(&id.to_string());
2021 }
2022 _ => {}
2023 }
2024 i += 1;
2025 }
2026
2027 cols.push((col_name, type_str.trim().to_string()));
2028 }
2029 Ok(cols)
2030}
2031
2032fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
2036 let mut sql_parts: Vec<String> = Vec::new();
2038 let mut table_name: Option<String> = None;
2039 let mut in_from = false;
2040
2041 for (i, t) in tokens.iter().enumerate() {
2042 match t {
2043 TokenTree::Ident(id) => {
2044 let s = id.to_string();
2045 if s.eq_ignore_ascii_case("SELECT") {
2046 sql_parts.push("SELECT".to_string());
2047 } else if s.eq_ignore_ascii_case("FROM") {
2048 in_from = true;
2049 sql_parts.push("FROM".to_string());
2050 } else if s.eq_ignore_ascii_case("WHERE")
2051 || s.eq_ignore_ascii_case("AND")
2052 || s.eq_ignore_ascii_case("OR")
2053 || s.eq_ignore_ascii_case("LIMIT")
2054 || s.eq_ignore_ascii_case("OFFSET")
2055 || s.eq_ignore_ascii_case("ORDER")
2056 || s.eq_ignore_ascii_case("BY")
2057 || s.eq_ignore_ascii_case("GROUP")
2058 || s.eq_ignore_ascii_case("HAVING")
2059 || s.eq_ignore_ascii_case("JOIN")
2060 || s.eq_ignore_ascii_case("INNER")
2061 || s.eq_ignore_ascii_case("LEFT")
2062 || s.eq_ignore_ascii_case("RIGHT")
2063 || s.eq_ignore_ascii_case("ON")
2064 || s.eq_ignore_ascii_case("AS")
2065 || s.eq_ignore_ascii_case("ASC")
2066 || s.eq_ignore_ascii_case("DESC")
2067 || s.eq_ignore_ascii_case("DISTINCT")
2068 || s.eq_ignore_ascii_case("NOT")
2069 || s.eq_ignore_ascii_case("NULL")
2070 || s.eq_ignore_ascii_case("IN")
2071 || s.eq_ignore_ascii_case("BETWEEN")
2072 || s.eq_ignore_ascii_case("LIKE")
2073 || s.eq_ignore_ascii_case("IS")
2074 {
2075 sql_parts.push(s.to_uppercase());
2076 } else if in_from && table_name.is_none() {
2077 table_name = Some(s.clone());
2079 sql_parts.push(s.clone());
2080 } else {
2081 sql_parts.push(s.clone());
2082 }
2083 }
2084 TokenTree::Literal(lit) => {
2085 sql_parts.push(lit.to_string());
2086 }
2087 TokenTree::Punct(p) => {
2088 let c = p.as_char();
2089 let part = if c == ',' {
2091 ",".to_string()
2092 } else if c == '?' {
2093 "?".to_string()
2094 } else if c == '*' {
2095 "*".to_string()
2096 } else if c == '=' {
2097 "=".to_string()
2098 } else if c == '>' {
2099 ">".to_string()
2100 } else if c == '<' {
2101 "<".to_string()
2102 } else if c == '.' {
2103 ".".to_string()
2104 } else if c == ';' {
2105 ";".to_string()
2106 } else {
2107 c.to_string()
2108 };
2109 sql_parts.push(part);
2110 }
2111 TokenTree::Group(g) => {
2112 let inner: String = g.stream().to_string();
2114 let delim = match g.delimiter() {
2115 Delimiter::Parenthesis => "(",
2116 Delimiter::Brace => "{",
2117 Delimiter::Bracket => "[",
2118 Delimiter::None => "",
2119 };
2120 let close = match g.delimiter() {
2121 Delimiter::Parenthesis => ")",
2122 Delimiter::Brace => "}",
2123 Delimiter::Bracket => "]",
2124 Delimiter::None => "",
2125 };
2126 sql_parts.push(format!("{}{}{}", delim, inner, close));
2127 }
2128 }
2129 let _ = i;
2131 }
2132
2133 let sql = sql_parts
2134 .join(" ")
2135 .replace(", ", ",")
2136 .replace(" ,", ",")
2137 .replace("= ", "=")
2138 .replace(" =", "=")
2139 .replace("> ", ">")
2140 .replace(" >", ">")
2141 .replace("< ", "<")
2142 .replace(" <", "<")
2143 .replace(" ", " ");
2144
2145 if let Err(e) = validate_sql_content(&sql, None) {
2147 return compile_error(
2148 Span::call_site(),
2149 &format!("typed_query! SQL validation failed: {}", e),
2150 );
2151 }
2152
2153 let mut ts = TokenStream::new();
2155 let lit = Literal::string(&sql);
2156 ts.extend([TokenTree::Literal(lit)]);
2157 ts
2158}
2159
2160#[proc_macro]
2190pub fn query_as(input: TokenStream) -> TokenStream {
2212 let mut tokens = input.into_iter().peekable();
2213
2214 let mut record_type = String::new();
2216 loop {
2217 match tokens.next() {
2218 Some(TokenTree::Ident(ident)) => {
2219 record_type.push_str(&ident.to_string());
2220 }
2221 Some(TokenTree::Punct(p)) if p.as_char() == ':' => {
2222 record_type.push_str("::");
2224 if let Some(TokenTree::Punct(p2)) = tokens.peek() {
2226 if p2.as_char() == ':' {
2227 let _ = tokens.next();
2228 }
2229 }
2230 }
2231 Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
2232 Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
2233 Some(other) => {
2234 return compile_error(
2235 other.span(),
2236 "query_as! 第一个参数必须是记录类型,如 query_as!(User, \"SELECT ...\")",
2237 );
2238 }
2239 None => {
2240 return compile_error(
2241 Span::call_site(),
2242 "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2243 );
2244 }
2245 }
2246 }
2247
2248 let sql_raw = match tokens.next() {
2250 Some(TokenTree::Literal(lit)) => lit.to_string(),
2251 Some(other) => {
2252 return compile_error(other.span(), "query_as! 第二个参数必须是 SQL 字符串字面量");
2253 }
2254 None => {
2255 return compile_error(
2256 Span::call_site(),
2257 "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2258 );
2259 }
2260 };
2261
2262 let sql_content = match strip_string_literal(&sql_raw) {
2263 Some(s) => s,
2264 None => {
2265 return compile_error(Span::call_site(), "query_as! 的 SQL 参数必须是字符串字面量");
2266 }
2267 };
2268
2269 if let Err(err_msg) = validate_sql_content(sql_content, None) {
2271 return compile_error(Span::call_site(), &err_msg);
2272 }
2273
2274 #[cfg(feature = "db-verify")]
2276 let verify_cols: Option<Vec<(String, String)>> = {
2277 match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
2278 Some("1") => match verify_with_real_db(sql_content) {
2280 Ok(cols) => Some(cols),
2281 Err(err) => {
2282 return compile_error(
2283 Span::call_site(),
2284 &format!("query_as! real DB verification failed: {}", err),
2285 )
2286 }
2287 },
2288 Some("cache") => {
2290 if let Err(err) = verify_with_cache(sql_content) {
2291 return compile_error(
2292 Span::call_site(),
2293 &format!("query_as! offline cache verification failed: {}", err),
2294 );
2295 }
2296 None
2297 }
2298 _ => None,
2299 }
2300 };
2301 #[cfg(not(feature = "db-verify"))]
2302 let _verify_cols: Option<Vec<(String, String)>> = None;
2303
2304 let escaped = sql_content.escape_default();
2309 let base = format!(
2310 "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
2311 record_type, escaped
2312 );
2313 #[cfg(feature = "db-verify")]
2314 let output = match &verify_cols {
2315 Some(cols) if !cols.is_empty() => {
2316 gen_compile_time_type_check(&record_type, sql_content, cols, &base)
2317 }
2318 _ => base,
2319 };
2320 #[cfg(not(feature = "db-verify"))]
2321 let output = base;
2322 output
2323 .parse()
2324 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query_as output"))
2325}
2326
2327#[proc_macro]
2328pub fn schema(input: TokenStream) -> TokenStream {
2329 let mut tokens = input.into_iter().peekable();
2330
2331 let sql_raw = match tokens.next() {
2333 Some(TokenTree::Literal(lit)) => lit.to_string(),
2334 Some(other) => {
2335 return compile_error(
2336 other.span(),
2337 "Expected a string literal as the argument to schema!",
2338 );
2339 }
2340 None => {
2341 return compile_error(
2342 Span::call_site(),
2343 "Expected a string literal argument to schema!",
2344 );
2345 }
2346 };
2347
2348 let sql = match strip_string_literal(&sql_raw) {
2349 Some(s) => s,
2350 None => {
2351 return compile_error(
2352 Span::call_site(),
2353 "schema! requires a string literal argument",
2354 );
2355 }
2356 };
2357
2358 let (table_name, columns) = match parse_create_table(sql) {
2360 Ok(v) => v,
2361 Err(e) => return compile_error(Span::call_site(), &e),
2362 };
2363
2364 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
2366 let table_name_lit = table_name.as_str();
2367
2368 let col_impls: Vec<TokenStream2> = columns
2369 .iter()
2370 .map(|(col_name, col_type)| {
2371 let col_ident =
2372 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
2373 let col_name_lit = col_name.as_str();
2374 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
2375 quote! {
2376 #[derive(Debug, Clone, Copy)]
2377 pub struct #col_ident;
2378 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
2379 const NAME: &'static str = #col_name_lit;
2380 type Table = table;
2381 type RustType = #rust_type;
2382 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
2383 }
2384 }
2385 })
2386 .collect();
2387
2388 let schema_entries: Vec<TokenStream2> = columns
2389 .iter()
2390 .map(|(n, t)| {
2391 let n_lit = n.as_str();
2392 let t_lit = t.as_str();
2393 quote! { (#n_lit, #t_lit) }
2394 })
2395 .collect();
2396
2397 let schema_const_ident = proc_macro2::Ident::new(
2398 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
2399 Span::call_site().into(),
2400 );
2401
2402 let expanded = quote! {
2403 pub mod #table_ident {
2404 use super::*;
2405 pub struct table;
2406 impl ::sz_orm_core::typed::TypedTable for table {
2407 const NAME: &'static str = #table_name_lit;
2408 }
2409 #(#col_impls)*
2410 }
2411 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
2412 };
2413
2414 expanded.into()
2415}
2416
2417fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
2425 let trimmed = sql.trim();
2426 let upper = trimmed.to_uppercase();
2427
2428 if !upper.starts_with("CREATE TABLE") {
2430 return Err("schema! expects a CREATE TABLE statement".to_string());
2431 }
2432
2433 let mut rest = &trimmed["CREATE TABLE".len()..];
2435
2436 let rest_upper = rest.trim_start().to_uppercase();
2438 if rest_upper.starts_with("IF NOT EXISTS") {
2439 rest = &rest.trim_start()["IF NOT EXISTS".len()..];
2440 }
2441
2442 rest = rest.trim_start();
2443
2444 let (table_name, after_name) = parse_identifier(rest)?;
2446 let rest = after_name.trim_start();
2447
2448 let paren_start = rest
2450 .find('(')
2451 .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
2452 let paren_end = rest
2453 .rfind(')')
2454 .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
2455 if paren_end <= paren_start {
2456 return Err("CREATE TABLE has malformed parentheses".to_string());
2457 }
2458
2459 let cols_str = &rest[paren_start + 1..paren_end];
2460
2461 let col_defs = split_top_level_commas(cols_str);
2463
2464 let mut columns = Vec::new();
2465 for def in col_defs {
2466 let def = def.trim();
2467 if def.is_empty() {
2468 continue;
2469 }
2470
2471 let def_upper = def.to_uppercase();
2473 if def_upper.starts_with("PRIMARY KEY")
2474 || def_upper.starts_with("FOREIGN KEY")
2475 || def_upper.starts_with("CONSTRAINT")
2476 || def_upper.starts_with("UNIQUE")
2477 || def_upper.starts_with("INDEX")
2478 || def_upper.starts_with("KEY")
2479 {
2480 continue;
2481 }
2482
2483 let (col_name, after_col) = parse_identifier(def)?;
2485 let rest = after_col.trim_start();
2486
2487 let (sql_type, after_type) = parse_type_token(rest)?;
2489 let rest = after_type.trim();
2490
2491 let rest_upper = rest.to_uppercase();
2493 let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
2494 let rust_type = sql_type_to_rust(&sql_type, !not_null);
2495
2496 columns.push((col_name, rust_type));
2497 }
2498
2499 Ok((table_name, columns))
2500}
2501
2502fn parse_identifier(s: &str) -> Result<(String, &str), String> {
2505 let s = s.trim_start();
2506 if s.is_empty() {
2507 return Err("expected identifier".to_string());
2508 }
2509
2510 let bytes = s.as_bytes();
2511 match bytes[0] {
2512 b'`' => {
2513 let end = s[1..]
2514 .find('`')
2515 .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
2516 let ident = s[1..1 + end].to_string();
2517 Ok((ident, &s[1 + end + 1..]))
2518 }
2519 b'"' => {
2520 let end = s[1..]
2521 .find('"')
2522 .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
2523 let ident = s[1..1 + end].to_string();
2524 Ok((ident, &s[1 + end + 1..]))
2525 }
2526 _ => {
2527 let end = s
2528 .find(|c: char| !c.is_alphanumeric() && c != '_')
2529 .unwrap_or(s.len());
2530 if end == 0 {
2531 return Err(format!("invalid identifier: '{}'", s));
2532 }
2533 let ident = s[..end].to_string();
2534 Ok((ident, &s[end..]))
2535 }
2536 }
2537}
2538
2539fn parse_type_token(s: &str) -> Result<(String, &str), String> {
2542 let s = s.trim_start();
2543 if s.is_empty() {
2544 return Err("expected column type".to_string());
2545 }
2546
2547 let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
2548 if end == 0 {
2549 return Err(format!("invalid type: '{}'", s));
2550 }
2551 let type_name = s[..end].to_string();
2552 let mut rest = &s[end..];
2553
2554 rest = rest.trim_start();
2556 if rest.starts_with('(') {
2557 let close = rest
2558 .find(')')
2559 .ok_or_else(|| "unterminated type parameter list".to_string())?;
2560 rest = &rest[close + 1..];
2561 }
2562
2563 Ok((type_name, rest))
2564}
2565
2566fn split_top_level_commas(s: &str) -> Vec<String> {
2568 let mut parts = Vec::new();
2569 let mut depth: i32 = 0;
2570 let mut current = String::new();
2571
2572 for ch in s.chars() {
2573 match ch {
2574 '(' => {
2575 depth += 1;
2576 current.push(ch);
2577 }
2578 ')' => {
2579 depth -= 1;
2580 current.push(ch);
2581 }
2582 ',' if depth == 0 => {
2583 parts.push(std::mem::take(&mut current));
2584 }
2585 _ => {
2586 current.push(ch);
2587 }
2588 }
2589 }
2590
2591 if !current.trim().is_empty() {
2592 parts.push(current);
2593 }
2594
2595 parts
2596}
2597
2598fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
2603 let upper = sql_type.to_uppercase();
2604 let rust = match upper.as_str() {
2605 "BIGINT" | "INT8" => "i64",
2607 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
2609 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
2611 "TINYINT" => "i8",
2613 "FLOAT" | "REAL" | "FLOAT4" => "f32",
2615 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
2617 "BOOLEAN" | "BOOL" => "bool",
2619 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
2621 "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
2623 | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
2624 _ => "String",
2625 };
2626
2627 if nullable {
2628 format!("Option<{}>", rust)
2629 } else {
2630 rust.to_string()
2631 }
2632}
2633
2634#[proc_macro_derive(Schema, attributes(table, column))]
2664pub fn derive_schema(input: TokenStream) -> TokenStream {
2665 let input = parse_macro_input!(input as syn::DeriveInput);
2666 derive::derive_schema_impl(input).into()
2667}
2668
2669#[proc_macro_derive(GraphQLModel, attributes(table, column))]
2698pub fn derive_graphql_model(input: TokenStream) -> TokenStream {
2699 let input = parse_macro_input!(input as syn::DeriveInput);
2700 derive::derive_graphql_model_impl(input).into()
2701}
2702
2703#[proc_macro_derive(Builder, attributes(builder))]
2737pub fn derive_builder(input: TokenStream) -> TokenStream {
2738 let input = parse_macro_input!(input as syn::DeriveInput);
2739 derive::derive_builder_impl(input).into()
2740}
2741
2742#[proc_macro_derive(Entity, attributes(table, column))]
2774pub fn derive_entity(input: TokenStream) -> TokenStream {
2775 let input = parse_macro_input!(input as syn::DeriveInput);
2776 derive::derive_entity_impl(input).into()
2777}
2778
2779#[proc_macro_derive(FromQueryResult, attributes(column))]
2806pub fn derive_from_query_result(input: TokenStream) -> TokenStream {
2807 let input = parse_macro_input!(input as syn::DeriveInput);
2808 derive::derive_from_query_result_impl(input).into()
2809}
2810
2811#[proc_macro_derive(ColumnEnum, attributes(column))]
2839pub fn derive_column_enum(input: TokenStream) -> TokenStream {
2840 let input = parse_macro_input!(input as syn::DeriveInput);
2841 derive::derive_column_enum_impl(input).into()
2842}
2843
2844#[proc_macro_derive(FromRow, attributes(column))]
2872pub fn derive_from_row(input: TokenStream) -> TokenStream {
2873 let input = parse_macro_input!(input as syn::DeriveInput);
2874 derive::derive_from_row_impl(input).into()
2875}
2876
2877#[proc_macro_derive(SqlType, attributes(sql_type))]
2907pub fn derive_sql_type(input: TokenStream) -> TokenStream {
2908 let input = parse_macro_input!(input as syn::DeriveInput);
2909 derive::derive_sql_type_impl(input).into()
2910}
2911
2912#[proc_macro_derive(Relation, attributes(relation, table, column))]
2948pub fn derive_relation(input: TokenStream) -> TokenStream {
2949 let input = parse_macro_input!(input as syn::DeriveInput);
2950 derive::derive_relation_impl(input).into()
2951}
2952
2953#[proc_macro_derive(RelationTrait, attributes(relation, table, column))]
2970pub fn derive_relation_trait(input: TokenStream) -> TokenStream {
2971 let input = parse_macro_input!(input as syn::DeriveInput);
2972 derive::derive_relation_trait_impl(input).into()
2973}
2974
2975#[cfg(feature = "data-validation")]
3005#[proc_macro_derive(Validate, attributes(validate))]
3006pub fn derive_validate(input: TokenStream) -> TokenStream {
3007 crate::derive_validate::derive_validate_impl(input)
3008}
3009
3010#[cfg(feature = "governance-derive")]
3050#[proc_macro_derive(Governed, attributes(pii, mask))]
3051pub fn derive_governed(input: TokenStream) -> TokenStream {
3052 let input = parse_macro_input!(input as syn::DeriveInput);
3053 let name = &input.ident;
3054
3055 const VALID_STRATEGIES: [&str; 4] = ["hash", "partial", "replace", "encrypt"];
3056
3057 let mut pii_fields: Vec<(String, String)> = Vec::new();
3058 let mut errors: Vec<syn::Error> = Vec::new();
3059
3060 if let syn::Data::Struct(data) = &input.data {
3061 for field in &data.fields {
3062 let Some(field_name) = field.ident.as_ref().map(|i| i.to_string()) else {
3063 continue;
3064 };
3065 let is_pii = field.attrs.iter().any(|a| a.path().is_ident("pii"));
3066
3067 let mut mask_strategy: Option<String> = None;
3069 for attr in &field.attrs {
3070 if !attr.path().is_ident("mask") {
3071 continue;
3072 }
3073 let _ = attr.parse_nested_meta(|meta| {
3074 if meta.path.is_ident("strategy") {
3075 let lit: syn::LitStr = meta.value()?.parse()?;
3076 mask_strategy = Some(lit.value());
3077 Ok(())
3078 } else {
3079 Err(meta.error("unsupported #[mask] attribute, only 'strategy' is allowed"))
3080 }
3081 });
3082 }
3083
3084 if is_pii {
3085 match mask_strategy {
3086 Some(strategy) => {
3087 if !VALID_STRATEGIES.contains(&strategy.as_str()) {
3088 errors.push(syn::Error::new_spanned(
3089 field,
3090 format!(
3091 "invalid #[mask(strategy = \"{strategy}\")]: allowed strategies are {:?}",
3092 VALID_STRATEGIES
3093 ),
3094 ));
3095 } else {
3096 pii_fields.push((field_name, strategy));
3097 }
3098 }
3099 None => errors.push(syn::Error::new_spanned(
3100 field,
3101 "#[pii] field must declare #[mask(strategy = \"...\")]",
3102 )),
3103 }
3104 }
3105 }
3106 }
3107
3108 if !errors.is_empty() {
3109 let err_tokens: proc_macro2::TokenStream =
3111 errors.iter().map(|e| e.to_compile_error()).collect();
3112 return err_tokens.into();
3113 }
3114
3115 let entries = pii_fields.iter().map(|(f, s)| {
3116 let f = f.as_str();
3117 let s = s.as_str();
3118 quote::quote!((#f, #s))
3119 });
3120
3121 quote::quote! {
3122 impl ::sz_orm_core::governance::GovernedModel for #name {
3123 fn pii_fields() -> Vec<(&'static str, &'static str)> {
3124 vec![#(#entries),*]
3125 }
3126 }
3127 }
3128 .into()
3129}
3130
3131#[cfg(feature = "n1-lint")]
3157#[proc_macro_attribute]
3158pub fn detect_n_plus_one(_attr: TokenStream, item: TokenStream) -> TokenStream {
3159 let item_fn = parse_macro_input!(item as syn::ItemFn);
3160 let findings = sz_orm_n1_lint::analyze_fn(&item_fn);
3161 for f in &findings {
3162 eprintln!(
3163 "warning: [sz-orm-n1-lint] {} at line {}: {}",
3164 f.pattern.as_str(),
3165 f.line,
3166 f.message
3167 );
3168 }
3169 quote::quote!(#item_fn).into()
3170}
3171
3172#[cfg(test)]
3177mod tests {
3178 use super::*;
3179
3180 #[test]
3183 fn test_strip_plain_double_quoted() {
3184 assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
3185 }
3186
3187 #[test]
3188 fn test_strip_raw_double_hash() {
3189 assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
3190 }
3191
3192 #[test]
3193 fn test_strip_raw_double_no_hash() {
3194 assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
3195 }
3196
3197 #[test]
3198 fn test_strip_byte_string() {
3199 assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
3200 assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
3201 }
3202
3203 #[test]
3204 fn test_strip_non_string_returns_none() {
3205 assert_eq!(strip_string_literal("123"), None);
3206 assert_eq!(strip_string_literal("foo"), None);
3207 }
3208
3209 #[test]
3212 fn test_validate_select_with_from_ok() {
3213 assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
3214 }
3215
3216 #[test]
3217 fn test_validate_select_missing_from_fails() {
3218 assert!(validate_sql_content("SELECT * users", None).is_err());
3219 }
3220
3221 #[test]
3222 fn test_validate_insert_missing_into_fails() {
3223 assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
3224 assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
3225 }
3226
3227 #[test]
3228 fn test_validate_update_missing_set_fails() {
3229 assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
3230 assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
3231 }
3232
3233 #[test]
3234 fn test_validate_delete_missing_from_fails() {
3235 assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
3236 assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
3237 }
3238
3239 #[test]
3240 fn test_validate_empty_sql_fails() {
3241 assert!(validate_sql_content("", None).is_err());
3242 assert!(validate_sql_content(" ", None).is_err());
3243 }
3244
3245 #[test]
3248 fn test_validate_balanced_parens_ok() {
3249 assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
3250 }
3251
3252 #[test]
3253 fn test_validate_balanced_parens_unbalanced() {
3254 assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
3255 assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
3256 }
3257
3258 #[test]
3261 fn test_validate_no_injection_clean() {
3262 assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
3263 }
3264
3265 #[test]
3266 fn test_validate_no_injection_drop_table() {
3267 assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
3268 }
3269
3270 #[test]
3271 fn test_validate_no_injection_or_1_1() {
3272 assert!(validate_no_injection("' OR 1=1").is_err());
3276 assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
3277 }
3278
3279 #[test]
3280 fn test_validate_no_injection_drop_database() {
3281 assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
3282 }
3283
3284 #[test]
3285 fn test_validate_no_injection_information_schema() {
3286 assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
3287 }
3288
3289 #[test]
3290 fn test_validate_no_injection_xp_cmdshell() {
3291 assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
3292 }
3293
3294 #[test]
3295 fn test_validate_no_injection_union_select() {
3296 assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
3297 }
3298
3299 #[test]
3300 fn test_validate_no_injection_comment_dashes() {
3301 assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
3302 }
3303
3304 #[test]
3305 fn test_validate_no_injection_block_comment() {
3306 assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
3307 }
3308
3309 #[test]
3312 fn test_validate_string_literals_closed_ok() {
3313 assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
3314 assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
3315 }
3316
3317 #[test]
3318 fn test_validate_string_literals_closed_unclosed_single() {
3319 assert!(validate_string_literals_closed("'hello").is_err());
3320 }
3321
3322 #[test]
3323 fn test_validate_string_literals_closed_unclosed_double() {
3324 assert!(validate_string_literals_closed(r#""hello"#).is_err());
3325 }
3326
3327 #[test]
3330 fn test_validate_param_count_match() {
3331 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
3332 assert!(
3333 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
3334 );
3335 }
3336
3337 #[test]
3338 fn test_validate_param_count_mismatch() {
3339 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
3340 assert!(
3341 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
3342 );
3343 }
3344
3345 #[cfg(feature = "db-verify")]
3348 #[test]
3349 fn test_detect_db_kind_mysql() {
3350 assert_eq!(
3351 detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
3352 DbKind::MySql
3353 );
3354 }
3355
3356 #[cfg(feature = "db-verify")]
3357 #[test]
3358 fn test_detect_db_kind_postgres() {
3359 assert_eq!(
3360 detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
3361 DbKind::Postgres
3362 );
3363 assert_eq!(
3364 detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
3365 DbKind::Postgres
3366 );
3367 }
3368
3369 #[cfg(feature = "db-verify")]
3370 #[test]
3371 fn test_detect_db_kind_sqlite() {
3372 assert_eq!(
3373 detect_db_kind("sqlite://path/to/db.db").unwrap(),
3374 DbKind::Sqlite
3375 );
3376 assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
3377 }
3378
3379 #[cfg(feature = "db-verify")]
3380 #[test]
3381 fn test_detect_db_kind_oracle() {
3382 assert_eq!(
3383 detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
3384 DbKind::Oracle
3385 );
3386 assert_eq!(
3387 detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
3388 DbKind::Oracle
3389 );
3390 }
3391
3392 #[cfg(feature = "db-verify")]
3393 #[test]
3394 fn test_detect_db_kind_sqlserver() {
3395 assert_eq!(
3396 detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
3397 DbKind::SqlServer
3398 );
3399 assert_eq!(
3400 detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
3401 DbKind::SqlServer
3402 );
3403 assert_eq!(
3404 detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
3405 DbKind::SqlServer
3406 );
3407 }
3408
3409 #[cfg(feature = "db-verify")]
3410 #[test]
3411 fn test_detect_db_kind_unsupported() {
3412 assert!(detect_db_kind("redis://user:pass@host/db").is_err());
3413 assert!(detect_db_kind("not-a-url").is_err());
3414 }
3415
3416 #[cfg(feature = "db-verify")]
3417 #[test]
3418 fn test_parse_oracle_dsn_basic() {
3419 let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
3420 let p = parse_oracle_dsn(dsn).unwrap();
3421 assert_eq!(p.user, "sys");
3422 assert_eq!(p.password, "test123");
3423 assert_eq!(p.host, "127.0.0.1");
3424 assert_eq!(p.port, 1521);
3425 assert_eq!(p.service, "freepdb1.FALSE");
3426 assert!(p.sysdba);
3427 }
3428
3429 #[cfg(feature = "db-verify")]
3430 #[test]
3431 fn test_parse_oracle_dsn_default_port() {
3432 let dsn = "oracle://sys:test123@127.0.0.1/FREE";
3434 let p = parse_oracle_dsn(dsn).unwrap();
3435 assert_eq!(p.port, 1521);
3436 assert_eq!(p.service, "FREE");
3437 assert!(!p.sysdba);
3438 }
3439
3440 #[cfg(feature = "db-verify")]
3441 #[test]
3442 fn test_parse_sqlserver_dsn_basic() {
3443 let dsn =
3444 "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
3445 let p = parse_sqlserver_dsn(dsn).unwrap();
3446 assert_eq!(p.user, "test");
3447 assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
3448 assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
3449 assert_eq!(p.port, 22527);
3450 assert_eq!(p.database, "test");
3451 }
3452
3453 #[cfg(feature = "db-verify")]
3454 #[test]
3455 fn test_parse_sqlserver_dsn_default_port() {
3456 let dsn = "mssql://user:pass@host/db";
3457 let p = parse_sqlserver_dsn(dsn).unwrap();
3458 assert_eq!(p.port, 1433);
3459 assert_eq!(p.database, "db");
3460 }
3461
3462 #[test]
3465 fn test_parse_create_table_basic() {
3466 let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
3467 let (table, cols) = parse_create_table(sql).unwrap();
3468 assert_eq!(table, "users");
3469 assert_eq!(
3470 cols,
3471 vec![
3472 ("id".to_string(), "i32".to_string()),
3473 ("name".to_string(), "String".to_string())
3474 ]
3475 );
3476 }
3477
3478 #[test]
3479 fn test_parse_create_table_with_if_not_exists() {
3480 let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
3481 let (table, cols) = parse_create_table(sql).unwrap();
3482 assert_eq!(table, "orders");
3483 assert_eq!(
3484 cols,
3485 vec![
3486 ("id".to_string(), "i64".to_string()),
3487 ("total".to_string(), "f64".to_string())
3488 ]
3489 );
3490 }
3491
3492 #[test]
3493 fn test_parse_create_table_nullable() {
3494 let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
3495 let (_, cols) = parse_create_table(sql).unwrap();
3496 assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
3497 assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
3498 }
3499
3500 #[test]
3501 fn test_parse_create_table_skip_constraints() {
3502 let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
3503 let (_, cols) = parse_create_table(sql).unwrap();
3504 assert_eq!(cols.len(), 2);
3505 assert_eq!(cols[0].0, "id");
3506 assert_eq!(cols[1].0, "name");
3507 }
3508
3509 #[test]
3510 fn test_parse_create_table_varchar_with_len() {
3511 let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
3512 let (_, cols) = parse_create_table(sql).unwrap();
3513 assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
3514 assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
3515 }
3516
3517 #[test]
3518 fn test_sql_type_to_rust_mappings() {
3519 assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
3521 assert_eq!(sql_type_to_rust("INT8", false), "i64");
3522 assert_eq!(sql_type_to_rust("INT", false), "i32");
3523 assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
3524 assert_eq!(sql_type_to_rust("INT4", false), "i32");
3525 assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
3526 assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
3527 assert_eq!(sql_type_to_rust("INT2", false), "i16");
3528 assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
3529 assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
3530 assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
3532 assert_eq!(sql_type_to_rust("REAL", false), "f32");
3533 assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
3534 assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
3535 assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
3536 assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
3537 assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
3538 assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
3539 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
3541 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
3542 assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
3544 assert_eq!(sql_type_to_rust("TEXT", false), "String");
3545 assert_eq!(sql_type_to_rust("CHAR", false), "String");
3546 assert_eq!(sql_type_to_rust("UUID", false), "String");
3547 assert_eq!(sql_type_to_rust("DATE", false), "String");
3548 assert_eq!(sql_type_to_rust("DATETIME", false), "String");
3549 assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
3550 assert_eq!(sql_type_to_rust("JSON", false), "String");
3551 assert_eq!(sql_type_to_rust("JSONB", false), "String");
3552 assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
3554 assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
3555 assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
3556 assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
3557 assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
3559 assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
3560 assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
3561 assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
3562 assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
3564 }
3565
3566 #[test]
3567 fn test_parse_create_table_error_no_create() {
3568 assert!(parse_create_table("SELECT * FROM users").is_err());
3569 }
3570
3571 #[test]
3572 fn test_parse_create_table_error_no_parens() {
3573 assert!(parse_create_table("CREATE TABLE foo").is_err());
3574 }
3575
3576 #[cfg(feature = "db-verify")]
3581 #[test]
3582 fn test_extract_tables_simple() {
3583 let tables = extract_tables("SELECT id, name FROM users WHERE id = ?");
3584 assert!(tables.contains(&"users".to_string()));
3585 }
3586
3587 #[cfg(feature = "db-verify")]
3588 #[test]
3589 fn test_extract_tables_multiple() {
3590 let tables = extract_tables(
3591 "SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id WHERE u.id = ?",
3592 );
3593 assert!(tables.contains(&"users".to_string()));
3594 assert!(tables.contains(&"orders".to_string()));
3595 }
3596
3597 #[cfg(feature = "db-verify")]
3598 #[test]
3599 fn test_extract_columns_select_and_where() {
3600 let cols =
3601 extract_columns("SELECT id, name FROM users WHERE email = ? ORDER BY created_at");
3602 assert!(cols.contains(&"id".to_string()));
3604 assert!(cols.contains(&"name".to_string()));
3605 assert!(cols.contains(&"email".to_string()));
3606 assert!(cols.contains(&"created_at".to_string()));
3607 }
3608
3609 #[cfg(feature = "db-verify")]
3610 #[test]
3611 fn test_extract_columns_skips_keywords() {
3612 let cols = extract_columns("SELECT COUNT(id), name FROM users WHERE status = ?");
3613 assert!(!cols.contains(&"count".to_string()));
3615 assert!(cols.contains(&"id".to_string()));
3616 assert!(cols.contains(&"name".to_string()));
3617 assert!(cols.contains(&"status".to_string()));
3618 }
3619
3620 #[cfg(feature = "db-verify")]
3621 #[test]
3622 fn test_is_sql_function() {
3623 assert!(is_sql_function("COUNT"));
3624 assert!(is_sql_function("now"));
3625 assert!(is_sql_function("COALESCE"));
3626 assert!(!is_sql_function("name"));
3627 assert!(!is_sql_function("user_id"));
3628 }
3629}