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#[proc_macro]
91pub fn sql_string(input: TokenStream) -> TokenStream {
92 let mut tokens = input.into_iter().peekable();
93
94 let sql = match tokens.next() {
96 Some(TokenTree::Literal(lit)) => lit.to_string(),
97 Some(other) => {
98 return compile_error(
99 other.span(),
100 "Expected a string literal as the first argument to sql_string!",
101 );
102 }
103 None => {
104 return compile_error(
105 Span::call_site(),
106 "Expected a string literal argument to sql_string!",
107 );
108 }
109 };
110
111 let sql_content = if sql.starts_with("r#\"") {
113 &sql[3..sql.len() - 2]
114 } else if sql.starts_with("r\"") {
115 &sql[2..sql.len() - 1]
116 } else if sql.starts_with('"') {
117 &sql[1..sql.len() - 1]
118 } else if sql.starts_with("b\"") || sql.starts_with("b\'") {
119 &sql[2..sql.len() - 1]
120 } else {
121 return compile_error(
122 Span::call_site(),
123 "sql_string! requires a string literal argument",
124 );
125 };
126
127 let mut expected_params = None;
129 if tokens.peek().is_some() {
130 match tokens.next() {
132 Some(TokenTree::Punct(p)) if p.as_char() == ';' => {}
133 Some(other) => {
134 return compile_error(
135 other.span(),
136 "Expected `;` before param count, e.g. sql_string!(\"...\"; params: 2)",
137 );
138 }
139 None => {}
140 }
141
142 match tokens.next() {
144 Some(TokenTree::Ident(id)) if id.to_string() == "params" => {}
145 Some(other) => {
146 return compile_error(
147 other.span(),
148 "Expected `params:` keyword, e.g. sql_string!(\"...\"; params: 2)",
149 );
150 }
151 None => {
152 return compile_error(Span::call_site(), "Expected param count after `;`");
153 }
154 }
155
156 match tokens.next() {
158 Some(TokenTree::Punct(p)) if p.as_char() == ':' => {}
159 Some(other) => {
160 return compile_error(
161 other.span(),
162 "Expected `:` after `params`, e.g. sql_string!(\"...\"; params: 2)",
163 );
164 }
165 None => {
166 return compile_error(Span::call_site(), "Expected param count after `params`");
167 }
168 }
169
170 match tokens.next() {
172 Some(TokenTree::Literal(lit)) => {
173 let num_str = lit.to_string();
174 if let Ok(n) = num_str.parse::<usize>() {
175 expected_params = Some(n);
176 } else {
177 return compile_error(
178 lit.span(),
179 "Expected a positive integer for param count",
180 );
181 }
182 }
183 Some(other) => {
184 return compile_error(
185 other.span(),
186 "Expected a number after `params:`, e.g. sql_string!(\"...\"; params: 2)",
187 );
188 }
189 None => {
190 return compile_error(Span::call_site(), "Expected a number after `params:`");
191 }
192 }
193 }
194
195 if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
197 return compile_error(Span::call_site(), &err_msg);
198 }
199
200 let output = format!("\"{}\"", sql_content.escape_default());
202 output
203 .parse()
204 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
205}
206
207fn validate_sql_content(sql: &str, expected_params: Option<usize>) -> Result<(), String> {
212 let trimmed = sql.trim();
213 if trimmed.is_empty() {
214 return Err("SQL statement is empty".to_string());
215 }
216
217 validate_balanced_parens(trimmed)?;
218 validate_string_literals_closed(trimmed)?;
219 validate_no_injection(trimmed)?;
220
221 let sql_upper = trimmed.to_uppercase();
223 if sql_upper.starts_with("SELECT") {
224 if !sql_upper.contains("FROM") {
225 return Err("SELECT statement missing FROM clause".to_string());
226 }
227 } else if sql_upper.starts_with("INSERT") {
228 if !sql_upper.contains("INTO") {
229 return Err("INSERT statement missing INTO clause".to_string());
230 }
231 if !sql_upper.contains("VALUES") {
232 return Err("INSERT statement missing VALUES clause".to_string());
233 }
234 } else if sql_upper.starts_with("UPDATE") {
235 if !sql_upper.contains("SET") {
236 return Err("UPDATE statement missing SET clause".to_string());
237 }
238 } else if sql_upper.starts_with("DELETE") && !sql_upper.contains("FROM") {
239 return Err("DELETE statement missing FROM clause".to_string());
240 }
241
242 if let Some(expected) = expected_params {
244 let actual = sql.chars().filter(|&c| c == '?').count();
245 if actual != expected {
246 return Err(format!(
247 "Parameter count mismatch: expected {} parameters, found {}",
248 expected, actual
249 ));
250 }
251 }
252
253 Ok(())
254}
255
256fn validate_balanced_parens(sql: &str) -> Result<(), String> {
257 let mut depth: i32 = 0;
258 for (i, ch) in sql.char_indices() {
259 match ch {
260 '(' => depth += 1,
261 ')' => {
262 depth -= 1;
263 if depth < 0 {
264 return Err(format!(
265 "Unbalanced parentheses: unexpected ')' at position {}",
266 i
267 ));
268 }
269 }
270 _ => {}
271 }
272 }
273 if depth != 0 {
274 return Err(format!("Unbalanced parentheses: {} unclosed '('", depth));
275 }
276 Ok(())
277}
278
279fn validate_string_literals_closed(sql: &str) -> Result<(), String> {
280 let mut in_single = false;
281 let mut in_double = false;
282 let mut prev = '\0';
283
284 for ch in sql.chars() {
285 if prev == '\\' {
286 prev = ch;
287 continue;
288 }
289
290 match ch {
291 '\'' if !in_double => in_single = !in_single,
292 '"' if !in_single => in_double = !in_double,
293 _ => {}
294 }
295 prev = ch;
296 }
297
298 if in_single {
299 return Err("Unclosed single-quoted string literal".to_string());
300 }
301 if in_double {
302 return Err("Unclosed double-quoted string literal".to_string());
303 }
304
305 Ok(())
306}
307
308fn validate_no_injection(sql: &str) -> Result<(), String> {
309 let sql_lower = sql.to_lowercase();
310
311 let injection_patterns: &[&str] = &[
314 "drop table",
316 "drop database",
317 "; drop",
318 "or 1=1",
320 "or 1 = 1",
321 "union select",
322 "union all select",
323 "--",
325 "/*",
326 "*/",
327 "xp_cmdshell",
329 "sp_executesql",
330 "exec(",
331 "execute(",
332 "information_schema",
334 "sys.tables",
335 "sys.columns",
336 ];
337
338 for pattern in injection_patterns {
339 if sql_lower.contains(pattern) {
340 return Err(format!("潜在的 SQL 注入模式被检测到: '{}'", pattern));
341 }
342 }
343
344 Ok(())
345}
346
347#[proc_macro]
381pub fn query(input: TokenStream) -> TokenStream {
382 let mut tokens = input.into_iter().peekable();
383
384 let type_param: Option<TokenStream2> = match tokens.peek() {
387 Some(TokenTree::Ident(_)) | Some(TokenTree::Punct(_)) => {
388 let mut ty_tokens = Vec::new();
390 while let Some(tok) = tokens.peek() {
391 match tok {
392 TokenTree::Punct(p) if p.as_char() == ',' => break,
393 TokenTree::Punct(p) if p.as_char() == ':' => {
394 ty_tokens.push(tokens.next().unwrap());
395 if let Some(TokenTree::Punct(p2)) = tokens.peek() {
397 if p2.as_char() == ':' {
398 ty_tokens.push(tokens.next().unwrap());
399 }
400 }
401 }
402 _ => ty_tokens.push(tokens.next().unwrap()),
403 }
404 }
405 match tokens.peek() {
407 Some(TokenTree::Punct(p)) if p.as_char() == ',' => {
408 tokens.next(); let ts: proc_macro::TokenStream = ty_tokens.into_iter().collect();
410 Some(TokenStream2::from(ts))
411 }
412 _ => None, }
414 }
415 _ => None,
416 };
417
418 let sql = match tokens.next() {
420 Some(TokenTree::Literal(lit)) => lit.to_string(),
421 Some(other) => {
422 return compile_error(
423 other.span(),
424 if type_param.is_some() {
425 "query!(T, \"SQL\"): expected a string literal as the second argument"
426 } else {
427 "Expected a string literal as the first argument to query!"
428 },
429 );
430 }
431 None => {
432 return compile_error(
433 Span::call_site(),
434 if type_param.is_some() {
435 "query!(T, \"SQL\"): missing SQL string argument"
436 } else {
437 "Expected a string literal argument to query!"
438 },
439 );
440 }
441 };
442
443 let sql_content = match strip_string_literal(&sql) {
444 Some(s) => s,
445 None => {
446 return compile_error(
447 Span::call_site(),
448 "query! requires a string literal argument",
449 );
450 }
451 };
452
453 if let Err(err_msg) = validate_sql_content(sql_content, None) {
455 return compile_error(Span::call_site(), &err_msg);
456 }
457
458 #[cfg(feature = "db-verify")]
460 let verify_cols: Option<Vec<(String, String)>> = {
461 match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
462 Some("1") => match verify_with_real_db(sql_content) {
464 Ok(cols) => Some(cols),
465 Err(err) => {
466 return compile_error(
467 Span::call_site(),
468 &format!("query! real DB verification failed: {}", err),
469 )
470 }
471 },
472 Some("cache") => {
474 if let Err(err) = verify_with_cache(sql_content) {
475 return compile_error(
476 Span::call_site(),
477 &format!("query! offline cache verification failed: {}", err),
478 );
479 }
480 None
481 }
482 _ => None,
483 }
484 };
485 #[cfg(not(feature = "db-verify"))]
486 let _verify_cols: Option<Vec<(String, String)>> = None;
487
488 let escaped = sql_content.escape_default().to_string();
490 let base = if let Some(ref ty) = type_param {
491 format!(
493 "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
494 ty, escaped
495 )
496 } else {
497 format!("::sz_orm_core::queryable::Query::new(\"{}\")", escaped)
499 };
500 #[cfg(feature = "db-verify")]
502 let output = match (&verify_cols, &type_param) {
503 (Some(cols), Some(ty)) if !cols.is_empty() => {
504 gen_compile_time_type_check(&ty.to_string(), sql_content, cols, &base)
505 }
506 _ => base,
507 };
508 #[cfg(not(feature = "db-verify"))]
509 let output = base;
510 output
511 .parse()
512 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query! output"))
513}
514
515fn strip_string_literal(raw: &str) -> Option<&str> {
518 if raw.starts_with("r#\"") {
519 Some(&raw[3..raw.len() - 2])
520 } else if raw.starts_with("r\"") {
521 Some(&raw[2..raw.len() - 1])
522 } else if raw.starts_with('"') {
523 Some(&raw[1..raw.len() - 1])
524 } else if raw.starts_with("b\"") || raw.starts_with("b\'") {
525 Some(&raw[2..raw.len() - 1])
526 } else {
527 None
528 }
529}
530
531#[cfg(feature = "db-verify")]
536fn verify_with_real_db(sql: &str) -> Result<Vec<(String, String)>, String> {
537 let dsn = std::env::var("DATABASE_URL")
538 .map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
539
540 let db_kind =
541 detect_db_kind(&dsn).map_err(|e| format!("Failed to detect DB kind from DSN: {}", e))?;
542
543 let sql_no_placeholders = replace_placeholders_with_null(sql);
546
547 let explain_sql = match db_kind {
549 DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql_no_placeholders),
550 DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql_no_placeholders),
551 DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
553 DbKind::SqlServer => sql_no_placeholders,
555 };
556
557 if matches!(db_kind, DbKind::MySql | DbKind::Postgres | DbKind::Sqlite) {
559 let rt = tokio::runtime::Runtime::new()
560 .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
561 return rt.block_on(async {
562 if let DbKind::MySql = db_kind {
564 verify_mysql(&dsn, &explain_sql).await?;
565 } else if let DbKind::Postgres = db_kind {
566 verify_postgres(&dsn, &explain_sql).await?;
567 } else {
568 verify_sqlite(&dsn, &explain_sql).await?;
570 }
571 verify_columns(&dsn, db_kind, sql).await?;
573 fetch_column_types(&dsn, db_kind, sql).await
576 });
577 }
578
579 if let DbKind::Oracle = db_kind {
581 verify_oracle(&dsn, &explain_sql).map(|_| Vec::new())
582 } else {
583 verify_sqlserver(&dsn, &explain_sql).map(|_| Vec::new())
585 }
586}
587
588#[cfg(feature = "db-verify")]
601fn verify_with_cache(sql: &str) -> Result<(), String> {
602 let cache_path = std::env::var("SZ_ORM_SQLX_CACHE").map_err(|_| {
603 "SZ_ORM_SQLX_CACHE not set. \
604 Set it to the path of a JSON file containing verified SQL statements, \
605 e.g. SZ_ORM_SQLX_CACHE=.sz-orm/query-cache.json"
606 .to_string()
607 })?;
608
609 let cache_content = std::fs::read_to_string(&cache_path).map_err(|e| {
610 format!(
611 "Failed to read cache file '{}': {}. \
612 Run `cargo sz-orm prepare` or build with SZ_ORM_QUERY_VERIFY=1 to generate it.",
613 cache_path, e
614 )
615 })?;
616
617 let verified: Vec<String> = serde_json::from_str(&cache_content).unwrap_or_else(|_| {
619 cache_content
620 .lines()
621 .map(|l| l.trim().to_string())
622 .filter(|l| !l.is_empty() && !l.starts_with('#'))
623 .collect()
624 });
625
626 if verified.iter().any(|v| v.trim() == sql.trim()) {
627 Ok(())
628 } else {
629 Err(format!(
630 "SQL not found in offline cache ({} entries): \"{}\". \
631 Add it to the cache by running with SZ_ORM_QUERY_VERIFY=1 first.",
632 verified.len(),
633 truncate_sql(sql, 80)
634 ))
635 }
636}
637
638#[cfg(feature = "db-verify")]
640fn truncate_sql(sql: &str, max: usize) -> String {
641 if sql.len() <= max {
642 sql.to_string()
643 } else {
644 format!("{}...", &sql[..max])
645 }
646}
647
648#[cfg(feature = "db-verify")]
649#[derive(Debug, Clone, Copy, PartialEq, Eq)]
650enum DbKind {
651 MySql,
652 Postgres,
653 Sqlite,
654 Oracle,
655 SqlServer,
656}
657
658#[cfg(feature = "db-verify")]
663fn replace_placeholders_with_null(sql: &str) -> String {
664 let mut result = String::with_capacity(sql.len() + 16);
665 let mut in_single_quote = false;
666 let mut in_double_quote = false;
667 let mut prev = '\0';
668
669 for ch in sql.chars() {
670 if prev == '\\' {
671 result.push(ch);
673 prev = ch;
674 continue;
675 }
676 match ch {
677 '\'' if !in_double_quote => in_single_quote = !in_single_quote,
678 '"' if !in_single_quote => in_double_quote = !in_double_quote,
679 '?' if !in_single_quote && !in_double_quote => {
680 result.push_str("NULL");
681 prev = ch;
682 continue;
683 }
684 _ => {}
685 }
686 result.push(ch);
687 prev = ch;
688 }
689 result
690}
691
692#[cfg(feature = "db-verify")]
693fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
694 let lower = dsn.to_lowercase();
695 if lower.starts_with("mysql://") {
696 Ok(DbKind::MySql)
697 } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
698 Ok(DbKind::Postgres)
699 } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
700 Ok(DbKind::Sqlite)
701 } else if lower.starts_with("oracle://") || lower.starts_with("oracle:") {
702 Ok(DbKind::Oracle)
703 } else if lower.starts_with("sqlserver://")
704 || lower.starts_with("mssql://")
705 || lower.starts_with("tds://")
706 {
707 Ok(DbKind::SqlServer)
708 } else {
709 Err(format!("Unsupported DSN scheme: {}", dsn))
710 }
711}
712
713#[cfg(feature = "db-verify")]
714async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
715 let pool = sqlx::MySqlPool::connect(dsn)
716 .await
717 .map_err(|e| format!("MySQL connect failed: {}", e))?;
718 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
719 .execute(&pool)
720 .await
721 .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
722 Ok(())
723}
724
725#[cfg(feature = "db-verify")]
726async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
727 let pool = sqlx::PgPool::connect(dsn)
728 .await
729 .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
730 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
731 .execute(&pool)
732 .await
733 .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
734 Ok(())
735}
736
737#[cfg(feature = "db-verify")]
738async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
739 let pool = sqlx::SqlitePool::connect(dsn)
740 .await
741 .map_err(|e| format!("SQLite connect failed: {}", e))?;
742 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
743 .execute(&pool)
744 .await
745 .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
746 Ok(())
747}
748
749#[cfg(feature = "db-verify")]
764async fn verify_columns(dsn: &str, db_kind: DbKind, sql: &str) -> Result<(), String> {
765 if matches!(db_kind, DbKind::Sqlite | DbKind::Oracle | DbKind::SqlServer) {
767 return Ok(());
768 }
769
770 let tables = extract_tables(sql);
771 let columns = extract_columns(sql);
772
773 if tables.is_empty() || columns.is_empty() {
774 return Ok(());
775 }
776
777 match db_kind {
778 DbKind::MySql => verify_columns_mysql(dsn, &tables, &columns, sql).await,
779 DbKind::Postgres => verify_columns_postgres(dsn, &tables, &columns, sql).await,
780 _ => Ok(()),
781 }
782}
783
784#[cfg(feature = "db-verify")]
786fn extract_tables(sql: &str) -> Vec<String> {
787 let mut tables = Vec::new();
788 let upper = sql.to_uppercase();
789
790 let from_idx = match upper.find("FROM") {
792 Some(i) => i,
793 None => return tables,
794 };
795
796 let end_patterns = ["WHERE", "ORDER", "GROUP", "LIMIT", "HAVING", "UNION"];
797 let end_idx = end_patterns
798 .iter()
799 .filter_map(|p| {
800 let mut search_start = 0;
802 while let Some(i) = upper[search_start..].find(*p) {
803 let abs_i = search_start + i;
804 let before = upper[..abs_i].chars().last().unwrap_or(' ');
805 let after = upper[abs_i + p.len()..].chars().next().unwrap_or(' ');
806 if !before.is_alphanumeric()
807 && !after.is_alphanumeric()
808 && before != '_'
809 && after != '_'
810 {
811 return Some(abs_i);
812 }
813 search_start = abs_i + p.len();
814 }
815 None
816 })
817 .filter(|&i| i > from_idx)
818 .min()
819 .unwrap_or(sql.len());
820
821 let from_clause = &sql[from_idx + 4..end_idx];
822
823 let join_split = {
825 let lower = from_clause.to_lowercase();
826 let mut result = String::with_capacity(from_clause.len());
827 let mut i = 0;
828 let bytes = from_clause.as_bytes();
829 let lower_bytes = lower.as_bytes();
830 while i < bytes.len() {
831 let mut matched = false;
832 for join_kw in &[
833 " join ",
834 " inner join ",
835 " left join ",
836 " right join ",
837 " left outer join ",
838 " right outer join ",
839 " cross join ",
840 " full join ",
841 " full outer join ",
842 ] {
843 let kw = join_kw.as_bytes();
844 if i + kw.len() <= bytes.len() && &lower_bytes[i..i + kw.len()] == kw {
845 result.push(',');
846 i += kw.len();
847 matched = true;
848 break;
849 }
850 }
851 if !matched {
852 result.push(bytes[i] as char);
853 i += 1;
854 }
855 }
856 result
857 };
858 let parts: Vec<&str> = join_split.split([',', '\n']).collect();
859
860 for part in parts {
861 let part = part.trim();
862 if part.is_empty() {
863 continue;
864 }
865 let table_word = part
867 .split_whitespace()
868 .next()
869 .unwrap_or(part)
870 .trim_end_matches([',', ';']);
871 let clean = table_word.trim_matches(|c| c == '`' || c == '"');
873 if !clean.is_empty()
874 && !matches!(
875 clean.to_uppercase().as_str(),
876 "INNER"
877 | "LEFT"
878 | "RIGHT"
879 | "OUTER"
880 | "CROSS"
881 | "FULL"
882 | "NATURAL"
883 | "ON"
884 | "USING"
885 | "AS"
886 )
887 {
888 tables.push(clean.to_lowercase());
889 }
890 }
891
892 tables
893}
894
895#[cfg(feature = "db-verify")]
897fn extract_columns(sql: &str) -> Vec<String> {
898 let mut columns = Vec::new();
899 let upper = sql.to_uppercase();
900
901 let mut collect_from_segment = |segment: &str| {
903 let keywords = [
906 "SELECT",
907 "FROM",
908 "WHERE",
909 "AND",
910 "OR",
911 "NOT",
912 "IN",
913 "IS",
914 "NULL",
915 "LIKE",
916 "BETWEEN",
917 "AS",
918 "ON",
919 "JOIN",
920 "INNER",
921 "LEFT",
922 "RIGHT",
923 "OUTER",
924 "CROSS",
925 "FULL",
926 "NATURAL",
927 "ORDER",
928 "BY",
929 "GROUP",
930 "HAVING",
931 "LIMIT",
932 "OFFSET",
933 "ASC",
934 "DESC",
935 "DISTINCT",
936 "COUNT",
937 "SUM",
938 "AVG",
939 "MIN",
940 "MAX",
941 "CASE",
942 "WHEN",
943 "THEN",
944 "ELSE",
945 "END",
946 "COALESCE",
947 "NULLIF",
948 "CAST",
949 "TRUE",
950 "FALSE",
951 "INSERT",
952 "INTO",
953 "VALUES",
954 "UPDATE",
955 "SET",
956 "DELETE",
957 "CREATE",
958 "TABLE",
959 "INDEX",
960 "IF",
961 "EXISTS",
962 "PRIMARY",
963 "KEY",
964 "REFERENCES",
965 "FOREIGN",
966 ];
967
968 for word in segment.split(|c: char| !c.is_alphanumeric() && c != '_') {
969 if word.is_empty() || word.len() < 2 {
970 continue;
971 }
972 let w = word.to_uppercase();
973 if keywords.contains(&w.as_str()) {
974 continue;
975 }
976 if word.chars().all(|c| c.is_ascii_digit()) {
978 continue;
979 }
980 let pos = segment.find(word).unwrap_or(0);
983 if pos > 0 && segment.chars().nth(pos - 1) == Some('.') {
984 continue;
985 }
986 if word == "*" {
988 continue;
989 }
990 let lower = word.to_lowercase();
991 if !columns.contains(&lower) {
992 columns.push(lower);
993 }
994 }
995 };
996
997 if let Some(from_idx) = upper.find("FROM") {
999 if let Some(sel_idx) = upper.find("SELECT") {
1000 let sel_segment = &sql[sel_idx + 6..from_idx];
1001 collect_from_segment(sel_segment);
1002 }
1003 }
1004
1005 if let Some(where_idx) = upper.find("WHERE") {
1007 let end_idx = ["ORDER", "GROUP", "LIMIT", "HAVING", "UNION"]
1008 .iter()
1009 .filter_map(|p| upper.find(p))
1010 .filter(|&i| i > where_idx)
1011 .min()
1012 .unwrap_or(sql.len());
1013 collect_from_segment(&sql[where_idx + 5..end_idx]);
1014 }
1015
1016 if let Some(order_idx) = upper.find("ORDER BY") {
1018 let end_idx = ["GROUP", "LIMIT", "HAVING", "UNION"]
1019 .iter()
1020 .filter_map(|p| upper.find(p))
1021 .filter(|&i| i > order_idx)
1022 .min()
1023 .unwrap_or(sql.len());
1024 collect_from_segment(&sql[order_idx + 8..end_idx]);
1025 }
1026
1027 columns
1028}
1029
1030#[cfg(feature = "db-verify")]
1031async fn verify_columns_mysql(
1032 dsn: &str,
1033 tables: &[String],
1034 columns: &[String],
1035 sql: &str,
1036) -> Result<(), String> {
1037 let pool = sqlx::MySqlPool::connect(dsn)
1038 .await
1039 .map_err(|e| format!("MySQL connect failed: {}", e))?;
1040
1041 for col in columns {
1042 let rows = sqlx::query(
1044 "SELECT TABLE_NAME, COLUMN_NAME FROM INFORMATION_SCHEMA.COLUMNS \
1045 WHERE TABLE_SCHEMA = DATABASE() AND COLUMN_NAME = ?",
1046 )
1047 .bind(col)
1048 .fetch_all(&pool)
1049 .await
1050 .map_err(|e| format!("MySQL column lookup failed for '{}': {}", col, e))?;
1051
1052 if rows.is_empty() {
1053 if is_sql_function(col) {
1055 continue;
1056 }
1057 return Err(format!(
1058 "query! column verification failed: column '{}' not found in any table of the current database. \
1059 SQL: {}",
1060 col,
1061 truncate_sql(sql, 120)
1062 ));
1063 }
1064
1065 if !tables.is_empty() {
1067 let found_in_table = rows.iter().any(|row| {
1068 let table_name: String = row.get("TABLE_NAME");
1069 tables.iter().any(|t| t == &table_name.to_lowercase())
1070 });
1071 if !found_in_table {
1072 let available: Vec<String> = rows.iter().map(|r| r.get("TABLE_NAME")).collect();
1073 return Err(format!(
1074 "query! column verification failed: column '{}' exists but not in FROM table(s) {:?}. \
1075 Found in: {:?}. SQL: {}",
1076 col,
1077 tables,
1078 available,
1079 truncate_sql(sql, 120)
1080 ));
1081 }
1082 }
1083 }
1084
1085 Ok(())
1086}
1087
1088#[cfg(feature = "db-verify")]
1089async fn verify_columns_postgres(
1090 dsn: &str,
1091 tables: &[String],
1092 columns: &[String],
1093 sql: &str,
1094) -> Result<(), String> {
1095 let pool = sqlx::PgPool::connect(dsn)
1096 .await
1097 .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
1098
1099 for col in columns {
1100 let rows = sqlx::query(
1101 "SELECT TABLE_NAME, COLUMN_NAME FROM INFORMATION_SCHEMA.COLUMNS \
1102 WHERE TABLE_CATALOG = CURRENT_CATALOG AND COLUMN_NAME = $1",
1103 )
1104 .bind(col)
1105 .fetch_all(&pool)
1106 .await
1107 .map_err(|e| format!("PostgreSQL column lookup failed for '{}': {}", col, e))?;
1108
1109 if rows.is_empty() && !is_sql_function(col) {
1110 return Err(format!(
1111 "query! column verification failed: column '{}' not found in any table of the current database. \
1112 SQL: {}",
1113 col,
1114 truncate_sql(sql, 120)
1115 ));
1116 }
1117
1118 if !tables.is_empty() && !rows.is_empty() {
1119 let found_in_table = rows.iter().any(|row| {
1120 let table_name: String = row.get("TABLE_NAME");
1121 tables.iter().any(|t| t == &table_name.to_lowercase())
1122 });
1123 if !found_in_table {
1124 let available: Vec<String> = rows.iter().map(|r| r.get("TABLE_NAME")).collect();
1125 return Err(format!(
1126 "query! column verification failed: column '{}' exists but not in FROM table(s) {:?}. \
1127 Found in: {:?}. SQL: {}",
1128 col, tables, available,
1129 truncate_sql(sql, 120)
1130 ));
1131 }
1132 }
1133 }
1134
1135 Ok(())
1136}
1137
1138#[cfg(feature = "db-verify")]
1150async fn fetch_column_types(
1151 dsn: &str,
1152 db_kind: DbKind,
1153 sql: &str,
1154) -> Result<Vec<(String, String)>, String> {
1155 if !matches!(db_kind, DbKind::MySql | DbKind::Postgres) {
1156 return Ok(Vec::new());
1157 }
1158
1159 let tables = extract_tables(sql);
1160 let columns = extract_columns(sql);
1161 if tables.is_empty() || columns.is_empty() {
1162 return Ok(Vec::new());
1163 }
1164
1165 match db_kind {
1166 DbKind::MySql => fetch_column_types_mysql(dsn, &tables, &columns).await,
1167 DbKind::Postgres => fetch_column_types_postgres(dsn, &tables, &columns).await,
1168 _ => Ok(Vec::new()),
1169 }
1170}
1171
1172#[cfg(feature = "db-verify")]
1173async fn fetch_column_types_mysql(
1174 dsn: &str,
1175 tables: &[String],
1176 columns: &[String],
1177) -> Result<Vec<(String, String)>, String> {
1178 let pool = sqlx::MySqlPool::connect(dsn)
1179 .await
1180 .map_err(|e| format!("MySQL connect failed for type fetch: {}", e))?;
1181
1182 let mut result = Vec::new();
1183 for col in columns {
1184 let rows = sqlx::query(
1185 "SELECT TABLE_NAME, COLUMN_NAME, DATA_TYPE \
1186 FROM INFORMATION_SCHEMA.COLUMNS \
1187 WHERE TABLE_SCHEMA = DATABASE() AND COLUMN_NAME = ?",
1188 )
1189 .bind(col)
1190 .fetch_all(&pool)
1191 .await
1192 .map_err(|e| format!("MySQL type lookup failed for '{}': {}", col, e))?;
1193
1194 let ty = rows
1196 .iter()
1197 .find(|row| {
1198 let tn: String = row.get("TABLE_NAME");
1199 tables.iter().any(|t| t == &tn.to_lowercase())
1200 })
1201 .and_then(|r| r.try_get::<String, _>("DATA_TYPE").ok());
1202 if let Some(ty) = ty {
1203 result.push((col.to_lowercase(), ty));
1204 }
1205 }
1206 Ok(result)
1207}
1208
1209#[cfg(feature = "db-verify")]
1210async fn fetch_column_types_postgres(
1211 dsn: &str,
1212 tables: &[String],
1213 columns: &[String],
1214) -> Result<Vec<(String, String)>, String> {
1215 let pool = sqlx::PgPool::connect(dsn)
1216 .await
1217 .map_err(|e| format!("PostgreSQL connect failed for type fetch: {}", e))?;
1218
1219 let mut result = Vec::new();
1220 for col in columns {
1221 let rows = sqlx::query(
1222 "SELECT TABLE_NAME, COLUMN_NAME, udt_name \
1223 FROM INFORMATION_SCHEMA.COLUMNS \
1224 WHERE TABLE_CATALOG = CURRENT_CATALOG AND COLUMN_NAME = $1",
1225 )
1226 .bind(col)
1227 .fetch_all(&pool)
1228 .await
1229 .map_err(|e| format!("PostgreSQL type lookup failed for '{}': {}", col, e))?;
1230
1231 let ty = rows
1232 .iter()
1233 .find(|row| {
1234 let tn: String = row.get("TABLE_NAME");
1235 tables.iter().any(|t| t == &tn.to_lowercase())
1236 })
1237 .and_then(|r| r.try_get::<String, _>("udt_name").ok());
1238 if let Some(ty) = ty {
1239 result.push((col.to_lowercase(), ty));
1240 }
1241 }
1242 Ok(result)
1243}
1244
1245#[cfg(feature = "db-verify")]
1256fn gen_compile_time_type_check(
1257 record_type: &str,
1258 sql: &str,
1259 cols: &[(String, String)],
1260 query_expr: &str,
1261) -> String {
1262 let n = cols.len();
1263 let sql_esc = sql.escape_default().to_string();
1264 let mut checks = String::new();
1265 checks.push_str(&format!(
1266 "if exp.len() != {} {{ panic!(\"sz-orm compile-time type check failed for `{}`: SELECT returns {} columns but struct field count differs\"); }}",
1267 n, sql_esc, n
1268 ));
1269 for (i, (name, ty)) in cols.iter().enumerate() {
1270 let name_esc = name.escape_default().to_string();
1271 let ty_esc = ty.escape_default().to_string();
1272 checks.push_str(&format!(
1273 "if !::sz_orm_core::__sz_orm_const_str_eq(exp[{}].0, \"{}\") {{ panic!(\"sz-orm compile-time type check failed for `{}`: SELECT column #{} `{}` not found in struct fields\"); }}",
1274 i, name_esc, sql_esc, i, name_esc
1275 ));
1276 checks.push_str(&format!(
1277 "if !::sz_orm_core::__sz_orm_const_types_compatible(\"{}\", exp[{}].1) {{ panic!(\"sz-orm compile-time type check failed for `{}`: column `{}` type mismatch (db type `{}` not compatible with struct field type)\"); }}",
1278 ty_esc, i, sql_esc, name_esc, ty_esc
1279 ));
1280 }
1281 format!(
1282 "{{ const _: () = {{ let exp = <{}>::__sz_orm_column_types(); {} }}; {} }}",
1283 record_type, checks, query_expr
1284 )
1285}
1286
1287#[cfg(feature = "db-verify")]
1289fn is_sql_function(name: &str) -> bool {
1290 matches!(
1291 name.to_uppercase().as_str(),
1292 "NOW"
1293 | "CURRENT_TIMESTAMP"
1294 | "CURRENT_DATE"
1295 | "CURRENT_TIME"
1296 | "COUNT"
1297 | "SUM"
1298 | "AVG"
1299 | "MIN"
1300 | "MAX"
1301 | "COALESCE"
1302 | "NULLIF"
1303 | "CAST"
1304 | "CONVERT"
1305 | "IFNULL"
1306 | "NVL"
1307 | "UPPER"
1308 | "LOWER"
1309 | "LENGTH"
1310 | "TRIM"
1311 | "SUBSTRING"
1312 | "CONCAT"
1313 | "REPLACE"
1314 | "ROUND"
1315 | "CEIL"
1316 | "FLOOR"
1317 | "ABS"
1318 | "MOD"
1319 | "POWER"
1320 | "SQRT"
1321 | "LOG"
1322 | "EXP"
1323 | "DATE"
1324 | "YEAR"
1325 | "MONTH"
1326 | "DAY"
1327 | "HOUR"
1328 | "MINUTE"
1329 | "SECOND"
1330 | "NOW()"
1331 | "UUID"
1332 | "RANDOM"
1333 | "MD5"
1334 | "TRUE"
1335 | "FALSE"
1336 | "NULL"
1337 )
1338}
1339
1340#[cfg(feature = "db-verify")]
1345fn verify_oracle(dsn: &str, explain_sql: &str) -> Result<(), String> {
1346 let parsed = parse_oracle_dsn(dsn)?;
1347 let mut conn_str = format!(
1349 "{}/{}@{}:{}/{}",
1350 parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
1351 );
1352 if parsed.sysdba {
1353 conn_str.push_str(" AS SYSDBA");
1354 }
1355 let full_script = format!(
1357 "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
1358 EXPLAIN PLAN FOR {};\n\
1359 SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
1360 EXIT;\n",
1361 explain_sql
1362 );
1363 let output = std::process::Command::new("sqlplus")
1364 .args(["-S", "-L", &conn_str])
1365 .stdin(std::process::Stdio::piped())
1366 .stdout(std::process::Stdio::piped())
1367 .stderr(std::process::Stdio::piped())
1368 .spawn()
1369 .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
1370 use std::io::Write;
1371 let mut child = output;
1372 if let Some(mut stdin) = child.stdin.take() {
1373 stdin
1374 .write_all(full_script.as_bytes())
1375 .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
1376 }
1377 let out = child
1378 .wait_with_output()
1379 .map_err(|e| format!("sqlplus wait failed: {}", e))?;
1380 let stdout = String::from_utf8_lossy(&out.stdout);
1381 let stderr = String::from_utf8_lossy(&out.stderr);
1382 if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
1383 return Err(format!(
1384 "Oracle EXPLAIN failed: stdout={} stderr={}",
1385 stdout.trim(),
1386 stderr.trim()
1387 ));
1388 }
1389 Ok(())
1390}
1391
1392#[cfg(feature = "db-verify")]
1397fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
1398 let parsed = parse_sqlserver_dsn(dsn)?;
1399 let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
1401 let out = std::process::Command::new("sqlcmd")
1402 .args([
1403 "-S",
1404 &format!("{},{}", parsed.host, parsed.port),
1405 "-U",
1406 &parsed.user,
1407 "-P",
1408 &parsed.password,
1409 "-d",
1410 &parsed.database,
1411 "-Q",
1412 &query,
1413 "-h",
1414 "-1",
1415 "-W",
1416 ])
1417 .output()
1418 .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
1419 let stdout = String::from_utf8_lossy(&out.stdout);
1420 let stderr = String::from_utf8_lossy(&out.stderr);
1421 if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
1422 return Err(format!(
1423 "SQL Server SHOWPLAN failed: stdout={} stderr={}",
1424 stdout.trim(),
1425 stderr.trim()
1426 ));
1427 }
1428 Ok(())
1429}
1430
1431#[cfg(feature = "db-verify")]
1433struct OracleDsn {
1434 user: String,
1435 password: String,
1436 host: String,
1437 port: u16,
1438 service: String,
1439 sysdba: bool,
1440}
1441
1442#[cfg(feature = "db-verify")]
1444fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
1445 let raw = dsn
1446 .strip_prefix("oracle://")
1447 .or_else(|| dsn.strip_prefix("oracle:"))
1448 .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
1449 let (auth_host_service, query) = match raw.find('?') {
1451 Some(idx) => (&raw[..idx], &raw[idx + 1..]),
1452 None => (raw, ""),
1453 };
1454 let sysdba = query
1455 .split('&')
1456 .any(|p| p == "sysdba=1" || p == "sysdba=true");
1457 let at = auth_host_service
1459 .find('@')
1460 .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
1461 let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
1462 let colon = user_pass
1463 .find(':')
1464 .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
1465 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1466 let (host_port, service) = match host_port_service.rfind('/') {
1467 Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
1468 None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
1469 };
1470 let (host, port) = match host_port.find(':') {
1471 Some(idx) => (
1472 &host_port[..idx],
1473 host_port[idx + 1..]
1474 .parse::<u16>()
1475 .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
1476 ),
1477 None => (host_port, 1521u16),
1478 };
1479 Ok(OracleDsn {
1480 user: user.to_string(),
1481 password: password.to_string(),
1482 host: host.to_string(),
1483 port,
1484 service: service.to_string(),
1485 sysdba,
1486 })
1487}
1488
1489#[cfg(feature = "db-verify")]
1491struct SqlServerDsn {
1492 user: String,
1493 password: String,
1494 host: String,
1495 port: u16,
1496 database: String,
1497}
1498
1499#[cfg(feature = "db-verify")]
1501fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
1502 let raw = dsn
1503 .strip_prefix("sqlserver://")
1504 .or_else(|| dsn.strip_prefix("mssql://"))
1505 .or_else(|| dsn.strip_prefix("tds://"))
1506 .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
1507 let at = raw
1508 .find('@')
1509 .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
1510 let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
1511 let colon = user_pass
1512 .find(':')
1513 .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
1514 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
1515 let (host_port, database) = match host_port_db.rfind('/') {
1516 Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
1517 None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
1518 };
1519 let (host, port) = match host_port.find(':') {
1520 Some(idx) => (
1521 &host_port[..idx],
1522 host_port[idx + 1..]
1523 .parse::<u16>()
1524 .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
1525 ),
1526 None => (host_port, 1433u16),
1527 };
1528 Ok(SqlServerDsn {
1529 user: user.to_string(),
1530 password: password.to_string(),
1531 host: host.to_string(),
1532 port,
1533 database: database.to_string(),
1534 })
1535}
1536
1537fn compile_error(span: Span, msg: &str) -> TokenStream {
1543 let mut ts = TokenStream::new();
1545 ts.extend([
1546 TokenTree::Ident(Ident::new("compile_error", span)),
1547 TokenTree::Punct(Punct::new('!', Spacing::Alone)),
1548 TokenTree::Group(Group::new(
1549 Delimiter::Parenthesis,
1550 TokenStream::from(TokenTree::Literal(Literal::string(msg))),
1551 )),
1552 ]);
1553 ts
1554}
1555
1556#[proc_macro]
1594pub fn typed_query(input: TokenStream) -> TokenStream {
1595 let tokens: Vec<TokenTree> = input.into_iter().collect();
1596
1597 if tokens.iter().any(|t| {
1599 if let TokenTree::Ident(id) = t {
1600 id.to_string() == "table"
1601 } else {
1602 false
1603 }
1604 }) {
1605 return parse_table_decl(&tokens);
1606 }
1607
1608 if tokens.iter().any(|t| {
1610 if let TokenTree::Ident(id) = t {
1611 id.to_string().eq_ignore_ascii_case("SELECT")
1612 } else {
1613 false
1614 }
1615 }) {
1616 return parse_typed_select(&tokens);
1617 }
1618
1619 compile_error(
1620 Span::call_site(),
1621 "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
1622 )
1623}
1624
1625fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
1627 let mut idx = 0;
1629
1630 if idx >= tokens.len() {
1632 return compile_error(Span::call_site(), "expected table name after 'table'");
1633 }
1634 if let TokenTree::Ident(id) = &tokens[idx] {
1635 if id.to_string() != "table" {
1636 return compile_error(id.span(), "expected 'table' keyword");
1637 }
1638 }
1639 idx += 1;
1640
1641 let table_name = if idx < tokens.len() {
1643 if let TokenTree::Ident(id) = &tokens[idx] {
1644 id.to_string()
1645 } else {
1646 return compile_error(tokens[idx].span(), "expected table name identifier");
1647 }
1648 } else {
1649 return compile_error(Span::call_site(), "expected table name");
1650 };
1651 idx += 1;
1652
1653 let body_group = if idx < tokens.len() {
1655 if let TokenTree::Group(g) = &tokens[idx] {
1656 if g.delimiter() != Delimiter::Brace {
1657 return compile_error(g.span(), "expected '{' after table name");
1658 }
1659 g.clone()
1660 } else {
1661 return compile_error(tokens[idx].span(), "expected '{' after table name");
1662 }
1663 } else {
1664 return compile_error(Span::call_site(), "expected table body in '{ }'");
1665 };
1666
1667 let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
1669 let columns = match parse_column_list(&body_tokens) {
1670 Ok(c) => c,
1671 Err(e) => return compile_error(Span::call_site(), &e),
1672 };
1673
1674 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1676 let table_name_lit = table_name.as_str();
1677
1678 let col_impls: Vec<TokenStream2> = columns
1680 .iter()
1681 .map(|(col_name, col_type)| {
1682 let col_ident =
1683 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1684 let col_name_lit = col_name.as_str();
1685 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1687 quote! {
1688 #[derive(Debug, Clone, Copy)]
1689 pub struct #col_ident;
1690 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1691 const NAME: &'static str = #col_name_lit;
1692 type Table = table;
1693 type RustType = #rust_type;
1694 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1695 }
1696 }
1697 })
1698 .collect();
1699
1700 let schema_entries: Vec<TokenStream2> = columns
1702 .iter()
1703 .map(|(n, t)| {
1704 let n_lit = n.as_str();
1705 let t_lit = t.as_str();
1706 quote! { (#n_lit, #t_lit) }
1707 })
1708 .collect();
1709
1710 let schema_const_ident = proc_macro2::Ident::new(
1711 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1712 Span::call_site().into(),
1713 );
1714
1715 let expanded = quote! {
1716 pub mod #table_ident {
1717 use super::*;
1718 pub struct table;
1719 impl ::sz_orm_core::typed::TypedTable for table {
1720 const NAME: &'static str = #table_name_lit;
1721 }
1722 #(#col_impls)*
1723 }
1724 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1725 };
1726
1727 expanded.into()
1728}
1729
1730fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
1732 let mut cols = Vec::new();
1733 let mut i = 0;
1734 while i < tokens.len() {
1735 let col_name = if let TokenTree::Ident(id) = &tokens[i] {
1737 id.to_string()
1738 } else {
1739 return Err(format!("expected column name at position {}", i));
1740 };
1741 i += 1;
1742
1743 if i >= tokens.len() {
1745 return Err(format!("expected ':' after column '{}'", col_name));
1746 }
1747 if let TokenTree::Punct(p) = &tokens[i] {
1748 if p.as_char() != ':' {
1749 return Err(format!("expected ':' after column '{}'", col_name));
1750 }
1751 } else {
1752 return Err(format!("expected ':' after column '{}'", col_name));
1753 }
1754 i += 1;
1755
1756 let mut type_str = String::new();
1759 let mut depth = 0;
1760 while i < tokens.len() {
1761 match &tokens[i] {
1762 TokenTree::Punct(p) => {
1763 if p.as_char() == ',' && depth == 0 {
1764 i += 1;
1765 break;
1766 } else if p.as_char() == '<' || p.as_char() == '(' {
1767 depth += 1;
1768 type_str.push(p.as_char());
1769 } else if p.as_char() == '>' || p.as_char() == ')' {
1770 depth -= 1;
1771 type_str.push(p.as_char());
1772 } else {
1773 type_str.push(p.as_char());
1774 }
1775 }
1776 TokenTree::Ident(id) => {
1777 if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
1778 {
1779 type_str.push(' ');
1780 }
1781 type_str.push_str(&id.to_string());
1782 }
1783 _ => {}
1784 }
1785 i += 1;
1786 }
1787
1788 cols.push((col_name, type_str.trim().to_string()));
1789 }
1790 Ok(cols)
1791}
1792
1793fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
1797 let mut sql_parts: Vec<String> = Vec::new();
1799 let mut table_name: Option<String> = None;
1800 let mut in_from = false;
1801
1802 for (i, t) in tokens.iter().enumerate() {
1803 match t {
1804 TokenTree::Ident(id) => {
1805 let s = id.to_string();
1806 if s.eq_ignore_ascii_case("SELECT") {
1807 sql_parts.push("SELECT".to_string());
1808 } else if s.eq_ignore_ascii_case("FROM") {
1809 in_from = true;
1810 sql_parts.push("FROM".to_string());
1811 } else if s.eq_ignore_ascii_case("WHERE")
1812 || s.eq_ignore_ascii_case("AND")
1813 || s.eq_ignore_ascii_case("OR")
1814 || s.eq_ignore_ascii_case("LIMIT")
1815 || s.eq_ignore_ascii_case("OFFSET")
1816 || s.eq_ignore_ascii_case("ORDER")
1817 || s.eq_ignore_ascii_case("BY")
1818 || s.eq_ignore_ascii_case("GROUP")
1819 || s.eq_ignore_ascii_case("HAVING")
1820 || s.eq_ignore_ascii_case("JOIN")
1821 || s.eq_ignore_ascii_case("INNER")
1822 || s.eq_ignore_ascii_case("LEFT")
1823 || s.eq_ignore_ascii_case("RIGHT")
1824 || s.eq_ignore_ascii_case("ON")
1825 || s.eq_ignore_ascii_case("AS")
1826 || s.eq_ignore_ascii_case("ASC")
1827 || s.eq_ignore_ascii_case("DESC")
1828 || s.eq_ignore_ascii_case("DISTINCT")
1829 || s.eq_ignore_ascii_case("NOT")
1830 || s.eq_ignore_ascii_case("NULL")
1831 || s.eq_ignore_ascii_case("IN")
1832 || s.eq_ignore_ascii_case("BETWEEN")
1833 || s.eq_ignore_ascii_case("LIKE")
1834 || s.eq_ignore_ascii_case("IS")
1835 {
1836 sql_parts.push(s.to_uppercase());
1837 } else if in_from && table_name.is_none() {
1838 table_name = Some(s.clone());
1840 sql_parts.push(s.clone());
1841 } else {
1842 sql_parts.push(s.clone());
1843 }
1844 }
1845 TokenTree::Literal(lit) => {
1846 sql_parts.push(lit.to_string());
1847 }
1848 TokenTree::Punct(p) => {
1849 let c = p.as_char();
1850 let part = if c == ',' {
1852 ",".to_string()
1853 } else if c == '?' {
1854 "?".to_string()
1855 } else if c == '*' {
1856 "*".to_string()
1857 } else if c == '=' {
1858 "=".to_string()
1859 } else if c == '>' {
1860 ">".to_string()
1861 } else if c == '<' {
1862 "<".to_string()
1863 } else if c == '.' {
1864 ".".to_string()
1865 } else if c == ';' {
1866 ";".to_string()
1867 } else {
1868 c.to_string()
1869 };
1870 sql_parts.push(part);
1871 }
1872 TokenTree::Group(g) => {
1873 let inner: String = g.stream().to_string();
1875 let delim = match g.delimiter() {
1876 Delimiter::Parenthesis => "(",
1877 Delimiter::Brace => "{",
1878 Delimiter::Bracket => "[",
1879 Delimiter::None => "",
1880 };
1881 let close = match g.delimiter() {
1882 Delimiter::Parenthesis => ")",
1883 Delimiter::Brace => "}",
1884 Delimiter::Bracket => "]",
1885 Delimiter::None => "",
1886 };
1887 sql_parts.push(format!("{}{}{}", delim, inner, close));
1888 }
1889 }
1890 let _ = i;
1892 }
1893
1894 let sql = sql_parts
1895 .join(" ")
1896 .replace(", ", ",")
1897 .replace(" ,", ",")
1898 .replace("= ", "=")
1899 .replace(" =", "=")
1900 .replace("> ", ">")
1901 .replace(" >", ">")
1902 .replace("< ", "<")
1903 .replace(" <", "<")
1904 .replace(" ", " ");
1905
1906 if let Err(e) = validate_sql_content(&sql, None) {
1908 return compile_error(
1909 Span::call_site(),
1910 &format!("typed_query! SQL validation failed: {}", e),
1911 );
1912 }
1913
1914 let mut ts = TokenStream::new();
1916 let lit = Literal::string(&sql);
1917 ts.extend([TokenTree::Literal(lit)]);
1918 ts
1919}
1920
1921#[proc_macro]
1951pub fn query_as(input: TokenStream) -> TokenStream {
1973 let mut tokens = input.into_iter().peekable();
1974
1975 let mut record_type = String::new();
1977 loop {
1978 match tokens.next() {
1979 Some(TokenTree::Ident(ident)) => {
1980 record_type.push_str(&ident.to_string());
1981 }
1982 Some(TokenTree::Punct(p)) if p.as_char() == ':' => {
1983 record_type.push_str("::");
1985 if let Some(TokenTree::Punct(p2)) = tokens.peek() {
1987 if p2.as_char() == ':' {
1988 let _ = tokens.next();
1989 }
1990 }
1991 }
1992 Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
1993 Some(TokenTree::Punct(p)) if p.as_char() == ',' => break,
1994 Some(other) => {
1995 return compile_error(
1996 other.span(),
1997 "query_as! 第一个参数必须是记录类型,如 query_as!(User, \"SELECT ...\")",
1998 );
1999 }
2000 None => {
2001 return compile_error(
2002 Span::call_site(),
2003 "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2004 );
2005 }
2006 }
2007 }
2008
2009 let sql_raw = match tokens.next() {
2011 Some(TokenTree::Literal(lit)) => lit.to_string(),
2012 Some(other) => {
2013 return compile_error(other.span(), "query_as! 第二个参数必须是 SQL 字符串字面量");
2014 }
2015 None => {
2016 return compile_error(
2017 Span::call_site(),
2018 "query_as! 需要两个参数:query_as!(RecordType, \"SELECT ...\")",
2019 );
2020 }
2021 };
2022
2023 let sql_content = match strip_string_literal(&sql_raw) {
2024 Some(s) => s,
2025 None => {
2026 return compile_error(Span::call_site(), "query_as! 的 SQL 参数必须是字符串字面量");
2027 }
2028 };
2029
2030 if let Err(err_msg) = validate_sql_content(sql_content, None) {
2032 return compile_error(Span::call_site(), &err_msg);
2033 }
2034
2035 #[cfg(feature = "db-verify")]
2037 let verify_cols: Option<Vec<(String, String)>> = {
2038 match std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() {
2039 Some("1") => match verify_with_real_db(sql_content) {
2041 Ok(cols) => Some(cols),
2042 Err(err) => {
2043 return compile_error(
2044 Span::call_site(),
2045 &format!("query_as! real DB verification failed: {}", err),
2046 )
2047 }
2048 },
2049 Some("cache") => {
2051 if let Err(err) = verify_with_cache(sql_content) {
2052 return compile_error(
2053 Span::call_site(),
2054 &format!("query_as! offline cache verification failed: {}", err),
2055 );
2056 }
2057 None
2058 }
2059 _ => None,
2060 }
2061 };
2062 #[cfg(not(feature = "db-verify"))]
2063 let _verify_cols: Option<Vec<(String, String)>> = None;
2064
2065 let escaped = sql_content.escape_default();
2070 let base = format!(
2071 "::sz_orm_core::queryable::QueryAs::<{}>::new(\"{}\")",
2072 record_type, escaped
2073 );
2074 #[cfg(feature = "db-verify")]
2075 let output = match &verify_cols {
2076 Some(cols) if !cols.is_empty() => {
2077 gen_compile_time_type_check(&record_type, sql_content, cols, &base)
2078 }
2079 _ => base,
2080 };
2081 #[cfg(not(feature = "db-verify"))]
2082 let output = base;
2083 output
2084 .parse()
2085 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate query_as output"))
2086}
2087
2088#[proc_macro]
2089pub fn schema(input: TokenStream) -> TokenStream {
2090 let mut tokens = input.into_iter().peekable();
2091
2092 let sql_raw = match tokens.next() {
2094 Some(TokenTree::Literal(lit)) => lit.to_string(),
2095 Some(other) => {
2096 return compile_error(
2097 other.span(),
2098 "Expected a string literal as the argument to schema!",
2099 );
2100 }
2101 None => {
2102 return compile_error(
2103 Span::call_site(),
2104 "Expected a string literal argument to schema!",
2105 );
2106 }
2107 };
2108
2109 let sql = match strip_string_literal(&sql_raw) {
2110 Some(s) => s,
2111 None => {
2112 return compile_error(
2113 Span::call_site(),
2114 "schema! requires a string literal argument",
2115 );
2116 }
2117 };
2118
2119 let (table_name, columns) = match parse_create_table(sql) {
2121 Ok(v) => v,
2122 Err(e) => return compile_error(Span::call_site(), &e),
2123 };
2124
2125 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
2127 let table_name_lit = table_name.as_str();
2128
2129 let col_impls: Vec<TokenStream2> = columns
2130 .iter()
2131 .map(|(col_name, col_type)| {
2132 let col_ident =
2133 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
2134 let col_name_lit = col_name.as_str();
2135 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
2136 quote! {
2137 #[derive(Debug, Clone, Copy)]
2138 pub struct #col_ident;
2139 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
2140 const NAME: &'static str = #col_name_lit;
2141 type Table = table;
2142 type RustType = #rust_type;
2143 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
2144 }
2145 }
2146 })
2147 .collect();
2148
2149 let schema_entries: Vec<TokenStream2> = columns
2150 .iter()
2151 .map(|(n, t)| {
2152 let n_lit = n.as_str();
2153 let t_lit = t.as_str();
2154 quote! { (#n_lit, #t_lit) }
2155 })
2156 .collect();
2157
2158 let schema_const_ident = proc_macro2::Ident::new(
2159 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
2160 Span::call_site().into(),
2161 );
2162
2163 let expanded = quote! {
2164 pub mod #table_ident {
2165 use super::*;
2166 pub struct table;
2167 impl ::sz_orm_core::typed::TypedTable for table {
2168 const NAME: &'static str = #table_name_lit;
2169 }
2170 #(#col_impls)*
2171 }
2172 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
2173 };
2174
2175 expanded.into()
2176}
2177
2178fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
2186 let trimmed = sql.trim();
2187 let upper = trimmed.to_uppercase();
2188
2189 if !upper.starts_with("CREATE TABLE") {
2191 return Err("schema! expects a CREATE TABLE statement".to_string());
2192 }
2193
2194 let mut rest = &trimmed["CREATE TABLE".len()..];
2196
2197 let rest_upper = rest.trim_start().to_uppercase();
2199 if rest_upper.starts_with("IF NOT EXISTS") {
2200 rest = &rest.trim_start()["IF NOT EXISTS".len()..];
2201 }
2202
2203 rest = rest.trim_start();
2204
2205 let (table_name, after_name) = parse_identifier(rest)?;
2207 let rest = after_name.trim_start();
2208
2209 let paren_start = rest
2211 .find('(')
2212 .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
2213 let paren_end = rest
2214 .rfind(')')
2215 .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
2216 if paren_end <= paren_start {
2217 return Err("CREATE TABLE has malformed parentheses".to_string());
2218 }
2219
2220 let cols_str = &rest[paren_start + 1..paren_end];
2221
2222 let col_defs = split_top_level_commas(cols_str);
2224
2225 let mut columns = Vec::new();
2226 for def in col_defs {
2227 let def = def.trim();
2228 if def.is_empty() {
2229 continue;
2230 }
2231
2232 let def_upper = def.to_uppercase();
2234 if def_upper.starts_with("PRIMARY KEY")
2235 || def_upper.starts_with("FOREIGN KEY")
2236 || def_upper.starts_with("CONSTRAINT")
2237 || def_upper.starts_with("UNIQUE")
2238 || def_upper.starts_with("INDEX")
2239 || def_upper.starts_with("KEY")
2240 {
2241 continue;
2242 }
2243
2244 let (col_name, after_col) = parse_identifier(def)?;
2246 let rest = after_col.trim_start();
2247
2248 let (sql_type, after_type) = parse_type_token(rest)?;
2250 let rest = after_type.trim();
2251
2252 let rest_upper = rest.to_uppercase();
2254 let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
2255 let rust_type = sql_type_to_rust(&sql_type, !not_null);
2256
2257 columns.push((col_name, rust_type));
2258 }
2259
2260 Ok((table_name, columns))
2261}
2262
2263fn parse_identifier(s: &str) -> Result<(String, &str), String> {
2266 let s = s.trim_start();
2267 if s.is_empty() {
2268 return Err("expected identifier".to_string());
2269 }
2270
2271 let bytes = s.as_bytes();
2272 match bytes[0] {
2273 b'`' => {
2274 let end = s[1..]
2275 .find('`')
2276 .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
2277 let ident = s[1..1 + end].to_string();
2278 Ok((ident, &s[1 + end + 1..]))
2279 }
2280 b'"' => {
2281 let end = s[1..]
2282 .find('"')
2283 .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
2284 let ident = s[1..1 + end].to_string();
2285 Ok((ident, &s[1 + end + 1..]))
2286 }
2287 _ => {
2288 let end = s
2289 .find(|c: char| !c.is_alphanumeric() && c != '_')
2290 .unwrap_or(s.len());
2291 if end == 0 {
2292 return Err(format!("invalid identifier: '{}'", s));
2293 }
2294 let ident = s[..end].to_string();
2295 Ok((ident, &s[end..]))
2296 }
2297 }
2298}
2299
2300fn parse_type_token(s: &str) -> Result<(String, &str), String> {
2303 let s = s.trim_start();
2304 if s.is_empty() {
2305 return Err("expected column type".to_string());
2306 }
2307
2308 let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
2309 if end == 0 {
2310 return Err(format!("invalid type: '{}'", s));
2311 }
2312 let type_name = s[..end].to_string();
2313 let mut rest = &s[end..];
2314
2315 rest = rest.trim_start();
2317 if rest.starts_with('(') {
2318 let close = rest
2319 .find(')')
2320 .ok_or_else(|| "unterminated type parameter list".to_string())?;
2321 rest = &rest[close + 1..];
2322 }
2323
2324 Ok((type_name, rest))
2325}
2326
2327fn split_top_level_commas(s: &str) -> Vec<String> {
2329 let mut parts = Vec::new();
2330 let mut depth: i32 = 0;
2331 let mut current = String::new();
2332
2333 for ch in s.chars() {
2334 match ch {
2335 '(' => {
2336 depth += 1;
2337 current.push(ch);
2338 }
2339 ')' => {
2340 depth -= 1;
2341 current.push(ch);
2342 }
2343 ',' if depth == 0 => {
2344 parts.push(std::mem::take(&mut current));
2345 }
2346 _ => {
2347 current.push(ch);
2348 }
2349 }
2350 }
2351
2352 if !current.trim().is_empty() {
2353 parts.push(current);
2354 }
2355
2356 parts
2357}
2358
2359fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
2364 let upper = sql_type.to_uppercase();
2365 let rust = match upper.as_str() {
2366 "BIGINT" | "INT8" => "i64",
2368 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
2370 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
2372 "TINYINT" => "i8",
2374 "FLOAT" | "REAL" | "FLOAT4" => "f32",
2376 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
2378 "BOOLEAN" | "BOOL" => "bool",
2380 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
2382 "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
2384 | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
2385 _ => "String",
2386 };
2387
2388 if nullable {
2389 format!("Option<{}>", rust)
2390 } else {
2391 rust.to_string()
2392 }
2393}
2394
2395#[proc_macro_derive(Schema, attributes(table, column))]
2425pub fn derive_schema(input: TokenStream) -> TokenStream {
2426 let input = parse_macro_input!(input as syn::DeriveInput);
2427 derive::derive_schema_impl(input).into()
2428}
2429
2430#[proc_macro_derive(Builder, attributes(builder))]
2464pub fn derive_builder(input: TokenStream) -> TokenStream {
2465 let input = parse_macro_input!(input as syn::DeriveInput);
2466 derive::derive_builder_impl(input).into()
2467}
2468
2469#[proc_macro_derive(Entity, attributes(table, column))]
2501pub fn derive_entity(input: TokenStream) -> TokenStream {
2502 let input = parse_macro_input!(input as syn::DeriveInput);
2503 derive::derive_entity_impl(input).into()
2504}
2505
2506#[proc_macro_derive(FromQueryResult, attributes(column))]
2533pub fn derive_from_query_result(input: TokenStream) -> TokenStream {
2534 let input = parse_macro_input!(input as syn::DeriveInput);
2535 derive::derive_from_query_result_impl(input).into()
2536}
2537
2538#[proc_macro_derive(ColumnEnum, attributes(column))]
2566pub fn derive_column_enum(input: TokenStream) -> TokenStream {
2567 let input = parse_macro_input!(input as syn::DeriveInput);
2568 derive::derive_column_enum_impl(input).into()
2569}
2570
2571#[proc_macro_derive(FromRow, attributes(column))]
2599pub fn derive_from_row(input: TokenStream) -> TokenStream {
2600 let input = parse_macro_input!(input as syn::DeriveInput);
2601 derive::derive_from_row_impl(input).into()
2602}
2603
2604#[proc_macro_derive(SqlType, attributes(sql_type))]
2634pub fn derive_sql_type(input: TokenStream) -> TokenStream {
2635 let input = parse_macro_input!(input as syn::DeriveInput);
2636 derive::derive_sql_type_impl(input).into()
2637}
2638
2639#[proc_macro_derive(Relation, attributes(relation, table, column))]
2675pub fn derive_relation(input: TokenStream) -> TokenStream {
2676 let input = parse_macro_input!(input as syn::DeriveInput);
2677 derive::derive_relation_impl(input).into()
2678}
2679
2680#[cfg(test)]
2685mod tests {
2686 use super::*;
2687
2688 #[test]
2691 fn test_strip_plain_double_quoted() {
2692 assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
2693 }
2694
2695 #[test]
2696 fn test_strip_raw_double_hash() {
2697 assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
2698 }
2699
2700 #[test]
2701 fn test_strip_raw_double_no_hash() {
2702 assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
2703 }
2704
2705 #[test]
2706 fn test_strip_byte_string() {
2707 assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
2708 assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
2709 }
2710
2711 #[test]
2712 fn test_strip_non_string_returns_none() {
2713 assert_eq!(strip_string_literal("123"), None);
2714 assert_eq!(strip_string_literal("foo"), None);
2715 }
2716
2717 #[test]
2720 fn test_validate_select_with_from_ok() {
2721 assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
2722 }
2723
2724 #[test]
2725 fn test_validate_select_missing_from_fails() {
2726 assert!(validate_sql_content("SELECT * users", None).is_err());
2727 }
2728
2729 #[test]
2730 fn test_validate_insert_missing_into_fails() {
2731 assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
2732 assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
2733 }
2734
2735 #[test]
2736 fn test_validate_update_missing_set_fails() {
2737 assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
2738 assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
2739 }
2740
2741 #[test]
2742 fn test_validate_delete_missing_from_fails() {
2743 assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
2744 assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
2745 }
2746
2747 #[test]
2748 fn test_validate_empty_sql_fails() {
2749 assert!(validate_sql_content("", None).is_err());
2750 assert!(validate_sql_content(" ", None).is_err());
2751 }
2752
2753 #[test]
2756 fn test_validate_balanced_parens_ok() {
2757 assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
2758 }
2759
2760 #[test]
2761 fn test_validate_balanced_parens_unbalanced() {
2762 assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
2763 assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
2764 }
2765
2766 #[test]
2769 fn test_validate_no_injection_clean() {
2770 assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
2771 }
2772
2773 #[test]
2774 fn test_validate_no_injection_drop_table() {
2775 assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
2776 }
2777
2778 #[test]
2779 fn test_validate_no_injection_or_1_1() {
2780 assert!(validate_no_injection("' OR 1=1").is_err());
2784 assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
2785 }
2786
2787 #[test]
2788 fn test_validate_no_injection_drop_database() {
2789 assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
2790 }
2791
2792 #[test]
2793 fn test_validate_no_injection_information_schema() {
2794 assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
2795 }
2796
2797 #[test]
2798 fn test_validate_no_injection_xp_cmdshell() {
2799 assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
2800 }
2801
2802 #[test]
2803 fn test_validate_no_injection_union_select() {
2804 assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
2805 }
2806
2807 #[test]
2808 fn test_validate_no_injection_comment_dashes() {
2809 assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
2810 }
2811
2812 #[test]
2813 fn test_validate_no_injection_block_comment() {
2814 assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
2815 }
2816
2817 #[test]
2820 fn test_validate_string_literals_closed_ok() {
2821 assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
2822 assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
2823 }
2824
2825 #[test]
2826 fn test_validate_string_literals_closed_unclosed_single() {
2827 assert!(validate_string_literals_closed("'hello").is_err());
2828 }
2829
2830 #[test]
2831 fn test_validate_string_literals_closed_unclosed_double() {
2832 assert!(validate_string_literals_closed(r#""hello"#).is_err());
2833 }
2834
2835 #[test]
2838 fn test_validate_param_count_match() {
2839 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
2840 assert!(
2841 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
2842 );
2843 }
2844
2845 #[test]
2846 fn test_validate_param_count_mismatch() {
2847 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
2848 assert!(
2849 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
2850 );
2851 }
2852
2853 #[cfg(feature = "db-verify")]
2856 #[test]
2857 fn test_detect_db_kind_mysql() {
2858 assert_eq!(
2859 detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
2860 DbKind::MySql
2861 );
2862 }
2863
2864 #[cfg(feature = "db-verify")]
2865 #[test]
2866 fn test_detect_db_kind_postgres() {
2867 assert_eq!(
2868 detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
2869 DbKind::Postgres
2870 );
2871 assert_eq!(
2872 detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
2873 DbKind::Postgres
2874 );
2875 }
2876
2877 #[cfg(feature = "db-verify")]
2878 #[test]
2879 fn test_detect_db_kind_sqlite() {
2880 assert_eq!(
2881 detect_db_kind("sqlite://path/to/db.db").unwrap(),
2882 DbKind::Sqlite
2883 );
2884 assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
2885 }
2886
2887 #[cfg(feature = "db-verify")]
2888 #[test]
2889 fn test_detect_db_kind_oracle() {
2890 assert_eq!(
2891 detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
2892 DbKind::Oracle
2893 );
2894 assert_eq!(
2895 detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
2896 DbKind::Oracle
2897 );
2898 }
2899
2900 #[cfg(feature = "db-verify")]
2901 #[test]
2902 fn test_detect_db_kind_sqlserver() {
2903 assert_eq!(
2904 detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
2905 DbKind::SqlServer
2906 );
2907 assert_eq!(
2908 detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
2909 DbKind::SqlServer
2910 );
2911 assert_eq!(
2912 detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
2913 DbKind::SqlServer
2914 );
2915 }
2916
2917 #[cfg(feature = "db-verify")]
2918 #[test]
2919 fn test_detect_db_kind_unsupported() {
2920 assert!(detect_db_kind("redis://user:pass@host/db").is_err());
2921 assert!(detect_db_kind("not-a-url").is_err());
2922 }
2923
2924 #[cfg(feature = "db-verify")]
2925 #[test]
2926 fn test_parse_oracle_dsn_basic() {
2927 let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
2928 let p = parse_oracle_dsn(dsn).unwrap();
2929 assert_eq!(p.user, "sys");
2930 assert_eq!(p.password, "test123");
2931 assert_eq!(p.host, "127.0.0.1");
2932 assert_eq!(p.port, 1521);
2933 assert_eq!(p.service, "freepdb1.FALSE");
2934 assert!(p.sysdba);
2935 }
2936
2937 #[cfg(feature = "db-verify")]
2938 #[test]
2939 fn test_parse_oracle_dsn_default_port() {
2940 let dsn = "oracle://sys:test123@127.0.0.1/FREE";
2942 let p = parse_oracle_dsn(dsn).unwrap();
2943 assert_eq!(p.port, 1521);
2944 assert_eq!(p.service, "FREE");
2945 assert!(!p.sysdba);
2946 }
2947
2948 #[cfg(feature = "db-verify")]
2949 #[test]
2950 fn test_parse_sqlserver_dsn_basic() {
2951 let dsn =
2952 "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
2953 let p = parse_sqlserver_dsn(dsn).unwrap();
2954 assert_eq!(p.user, "test");
2955 assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
2956 assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
2957 assert_eq!(p.port, 22527);
2958 assert_eq!(p.database, "test");
2959 }
2960
2961 #[cfg(feature = "db-verify")]
2962 #[test]
2963 fn test_parse_sqlserver_dsn_default_port() {
2964 let dsn = "mssql://user:pass@host/db";
2965 let p = parse_sqlserver_dsn(dsn).unwrap();
2966 assert_eq!(p.port, 1433);
2967 assert_eq!(p.database, "db");
2968 }
2969
2970 #[test]
2973 fn test_parse_create_table_basic() {
2974 let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
2975 let (table, cols) = parse_create_table(sql).unwrap();
2976 assert_eq!(table, "users");
2977 assert_eq!(
2978 cols,
2979 vec![
2980 ("id".to_string(), "i32".to_string()),
2981 ("name".to_string(), "String".to_string())
2982 ]
2983 );
2984 }
2985
2986 #[test]
2987 fn test_parse_create_table_with_if_not_exists() {
2988 let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
2989 let (table, cols) = parse_create_table(sql).unwrap();
2990 assert_eq!(table, "orders");
2991 assert_eq!(
2992 cols,
2993 vec![
2994 ("id".to_string(), "i64".to_string()),
2995 ("total".to_string(), "f64".to_string())
2996 ]
2997 );
2998 }
2999
3000 #[test]
3001 fn test_parse_create_table_nullable() {
3002 let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
3003 let (_, cols) = parse_create_table(sql).unwrap();
3004 assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
3005 assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
3006 }
3007
3008 #[test]
3009 fn test_parse_create_table_skip_constraints() {
3010 let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
3011 let (_, cols) = parse_create_table(sql).unwrap();
3012 assert_eq!(cols.len(), 2);
3013 assert_eq!(cols[0].0, "id");
3014 assert_eq!(cols[1].0, "name");
3015 }
3016
3017 #[test]
3018 fn test_parse_create_table_varchar_with_len() {
3019 let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
3020 let (_, cols) = parse_create_table(sql).unwrap();
3021 assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
3022 assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
3023 }
3024
3025 #[test]
3026 fn test_sql_type_to_rust_mappings() {
3027 assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
3029 assert_eq!(sql_type_to_rust("INT8", false), "i64");
3030 assert_eq!(sql_type_to_rust("INT", false), "i32");
3031 assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
3032 assert_eq!(sql_type_to_rust("INT4", false), "i32");
3033 assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
3034 assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
3035 assert_eq!(sql_type_to_rust("INT2", false), "i16");
3036 assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
3037 assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
3038 assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
3040 assert_eq!(sql_type_to_rust("REAL", false), "f32");
3041 assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
3042 assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
3043 assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
3044 assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
3045 assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
3046 assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
3047 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
3049 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
3050 assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
3052 assert_eq!(sql_type_to_rust("TEXT", false), "String");
3053 assert_eq!(sql_type_to_rust("CHAR", false), "String");
3054 assert_eq!(sql_type_to_rust("UUID", false), "String");
3055 assert_eq!(sql_type_to_rust("DATE", false), "String");
3056 assert_eq!(sql_type_to_rust("DATETIME", false), "String");
3057 assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
3058 assert_eq!(sql_type_to_rust("JSON", false), "String");
3059 assert_eq!(sql_type_to_rust("JSONB", false), "String");
3060 assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
3062 assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
3063 assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
3064 assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
3065 assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
3067 assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
3068 assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
3069 assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
3070 assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
3072 }
3073
3074 #[test]
3075 fn test_parse_create_table_error_no_create() {
3076 assert!(parse_create_table("SELECT * FROM users").is_err());
3077 }
3078
3079 #[test]
3080 fn test_parse_create_table_error_no_parens() {
3081 assert!(parse_create_table("CREATE TABLE foo").is_err());
3082 }
3083
3084 #[cfg(feature = "db-verify")]
3089 #[test]
3090 fn test_extract_tables_simple() {
3091 let tables = extract_tables("SELECT id, name FROM users WHERE id = ?");
3092 assert!(tables.contains(&"users".to_string()));
3093 }
3094
3095 #[cfg(feature = "db-verify")]
3096 #[test]
3097 fn test_extract_tables_multiple() {
3098 let tables = extract_tables(
3099 "SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id WHERE u.id = ?",
3100 );
3101 assert!(tables.contains(&"users".to_string()));
3102 assert!(tables.contains(&"orders".to_string()));
3103 }
3104
3105 #[cfg(feature = "db-verify")]
3106 #[test]
3107 fn test_extract_columns_select_and_where() {
3108 let cols =
3109 extract_columns("SELECT id, name FROM users WHERE email = ? ORDER BY created_at");
3110 assert!(cols.contains(&"id".to_string()));
3112 assert!(cols.contains(&"name".to_string()));
3113 assert!(cols.contains(&"email".to_string()));
3114 assert!(cols.contains(&"created_at".to_string()));
3115 }
3116
3117 #[cfg(feature = "db-verify")]
3118 #[test]
3119 fn test_extract_columns_skips_keywords() {
3120 let cols = extract_columns("SELECT COUNT(id), name FROM users WHERE status = ?");
3121 assert!(!cols.contains(&"count".to_string()));
3123 assert!(cols.contains(&"id".to_string()));
3124 assert!(cols.contains(&"name".to_string()));
3125 assert!(cols.contains(&"status".to_string()));
3126 }
3127
3128 #[cfg(feature = "db-verify")]
3129 #[test]
3130 fn test_is_sql_function() {
3131 assert!(is_sql_function("COUNT"));
3132 assert!(is_sql_function("now"));
3133 assert!(is_sql_function("COALESCE"));
3134 assert!(!is_sql_function("name"));
3135 assert!(!is_sql_function("user_id"));
3136 }
3137}