1#![allow(linker_messages)]
35extern crate proc_macro;
57
58use proc_macro::{Delimiter, Group, Ident, Literal, Punct, Spacing, Span, TokenStream, TokenTree};
59
60use proc_macro2::TokenStream as TokenStream2;
62use quote::quote;
63use syn::parse_macro_input;
64
65mod derive;
67
68#[proc_macro]
88pub fn sql_string(input: TokenStream) -> TokenStream {
89 let mut tokens = input.into_iter().peekable();
90
91 let sql = match tokens.next() {
93 Some(TokenTree::Literal(lit)) => lit.to_string(),
94 Some(other) => {
95 return compile_error(
96 other.span(),
97 "Expected a string literal as the first argument to sql_string!",
98 );
99 }
100 None => {
101 return compile_error(
102 Span::call_site(),
103 "Expected a string literal argument to sql_string!",
104 );
105 }
106 };
107
108 let sql_content = if sql.starts_with("r#\"") {
110 &sql[3..sql.len() - 2]
111 } else if sql.starts_with("r\"") {
112 &sql[2..sql.len() - 1]
113 } else if sql.starts_with('"') {
114 &sql[1..sql.len() - 1]
115 } else if sql.starts_with("b\"") || sql.starts_with("b\'") {
116 &sql[2..sql.len() - 1]
117 } else {
118 return compile_error(
119 Span::call_site(),
120 "sql_string! requires a string literal argument",
121 );
122 };
123
124 let mut expected_params = None;
126 if tokens.peek().is_some() {
127 match tokens.next() {
129 Some(TokenTree::Punct(p)) if p.as_char() == ';' => {}
130 Some(other) => {
131 return compile_error(
132 other.span(),
133 "Expected `;` before param count, e.g. sql_string!(\"...\"; params: 2)",
134 );
135 }
136 None => {}
137 }
138
139 match tokens.next() {
141 Some(TokenTree::Ident(id)) if id.to_string() == "params" => {}
142 Some(other) => {
143 return compile_error(
144 other.span(),
145 "Expected `params:` keyword, e.g. sql_string!(\"...\"; params: 2)",
146 );
147 }
148 None => {
149 return compile_error(Span::call_site(), "Expected param count after `;`");
150 }
151 }
152
153 match tokens.next() {
155 Some(TokenTree::Punct(p)) if p.as_char() == ':' => {}
156 Some(other) => {
157 return compile_error(
158 other.span(),
159 "Expected `:` after `params`, e.g. sql_string!(\"...\"; params: 2)",
160 );
161 }
162 None => {
163 return compile_error(Span::call_site(), "Expected param count after `params`");
164 }
165 }
166
167 match tokens.next() {
169 Some(TokenTree::Literal(lit)) => {
170 let num_str = lit.to_string();
171 if let Ok(n) = num_str.parse::<usize>() {
172 expected_params = Some(n);
173 } else {
174 return compile_error(
175 lit.span(),
176 "Expected a positive integer for param count",
177 );
178 }
179 }
180 Some(other) => {
181 return compile_error(
182 other.span(),
183 "Expected a number after `params:`, e.g. sql_string!(\"...\"; params: 2)",
184 );
185 }
186 None => {
187 return compile_error(Span::call_site(), "Expected a number after `params:`");
188 }
189 }
190 }
191
192 if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
194 return compile_error(Span::call_site(), &err_msg);
195 }
196
197 let output = format!("\"{}\"", sql_content.escape_default());
199 output
200 .parse()
201 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
202}
203
204fn validate_sql_content(sql: &str, expected_params: Option<usize>) -> Result<(), String> {
209 let trimmed = sql.trim();
210 if trimmed.is_empty() {
211 return Err("SQL statement is empty".to_string());
212 }
213
214 validate_balanced_parens(trimmed)?;
215 validate_string_literals_closed(trimmed)?;
216 validate_no_injection(trimmed)?;
217
218 let sql_upper = trimmed.to_uppercase();
220 if sql_upper.starts_with("SELECT") {
221 if !sql_upper.contains("FROM") {
222 return Err("SELECT statement missing FROM clause".to_string());
223 }
224 } else if sql_upper.starts_with("INSERT") {
225 if !sql_upper.contains("INTO") {
226 return Err("INSERT statement missing INTO clause".to_string());
227 }
228 if !sql_upper.contains("VALUES") {
229 return Err("INSERT statement missing VALUES clause".to_string());
230 }
231 } else if sql_upper.starts_with("UPDATE") {
232 if !sql_upper.contains("SET") {
233 return Err("UPDATE statement missing SET clause".to_string());
234 }
235 } else if sql_upper.starts_with("DELETE") && !sql_upper.contains("FROM") {
236 return Err("DELETE statement missing FROM clause".to_string());
237 }
238
239 if let Some(expected) = expected_params {
241 let actual = sql.chars().filter(|&c| c == '?').count();
242 if actual != expected {
243 return Err(format!(
244 "Parameter count mismatch: expected {} parameters, found {}",
245 expected, actual
246 ));
247 }
248 }
249
250 Ok(())
251}
252
253fn validate_balanced_parens(sql: &str) -> Result<(), String> {
254 let mut depth: i32 = 0;
255 for (i, ch) in sql.char_indices() {
256 match ch {
257 '(' => depth += 1,
258 ')' => {
259 depth -= 1;
260 if depth < 0 {
261 return Err(format!(
262 "Unbalanced parentheses: unexpected ')' at position {}",
263 i
264 ));
265 }
266 }
267 _ => {}
268 }
269 }
270 if depth != 0 {
271 return Err(format!("Unbalanced parentheses: {} unclosed '('", depth));
272 }
273 Ok(())
274}
275
276fn validate_string_literals_closed(sql: &str) -> Result<(), String> {
277 let mut in_single = false;
278 let mut in_double = false;
279 let mut prev = '\0';
280
281 for ch in sql.chars() {
282 if prev == '\\' {
283 prev = ch;
284 continue;
285 }
286
287 match ch {
288 '\'' if !in_double => in_single = !in_single,
289 '"' if !in_single => in_double = !in_double,
290 _ => {}
291 }
292 prev = ch;
293 }
294
295 if in_single {
296 return Err("Unclosed single-quoted string literal".to_string());
297 }
298 if in_double {
299 return Err("Unclosed double-quoted string literal".to_string());
300 }
301
302 Ok(())
303}
304
305fn validate_no_injection(sql: &str) -> Result<(), String> {
306 let sql_lower = sql.to_lowercase();
307
308 let injection_patterns: &[&str] = &[
311 "drop table",
313 "drop database",
314 "; drop",
315 "or 1=1",
317 "or 1 = 1",
318 "union select",
319 "union all select",
320 "--",
322 "/*",
323 "*/",
324 "xp_cmdshell",
326 "sp_executesql",
327 "exec(",
328 "execute(",
329 "information_schema",
331 "sys.tables",
332 "sys.columns",
333 ];
334
335 for pattern in injection_patterns {
336 if sql_lower.contains(pattern) {
337 return Err(format!("潜在的 SQL 注入模式被检测到: '{}'", pattern));
338 }
339 }
340
341 Ok(())
342}
343
344#[proc_macro]
376pub fn query(input: TokenStream) -> TokenStream {
377 let mut tokens = input.into_iter().peekable();
378
379 let sql = match tokens.next() {
381 Some(TokenTree::Literal(lit)) => lit.to_string(),
382 Some(other) => {
383 return compile_error(
384 other.span(),
385 "Expected a string literal as the first argument to query!",
386 );
387 }
388 None => {
389 return compile_error(
390 Span::call_site(),
391 "Expected a string literal argument to query!",
392 );
393 }
394 };
395
396 let sql_content = match strip_string_literal(&sql) {
397 Some(s) => s,
398 None => {
399 return compile_error(
400 Span::call_site(),
401 "query! requires a string literal argument",
402 );
403 }
404 };
405
406 if let Err(err_msg) = validate_sql_content(sql_content, None) {
408 return compile_error(Span::call_site(), &err_msg);
409 }
410
411 #[cfg(feature = "db-verify")]
413 {
414 if std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() == Some("1") {
415 if let Err(err) = verify_with_real_db(sql_content) {
416 return compile_error(
417 Span::call_site(),
418 &format!("query! real DB verification failed: {}", err),
419 );
420 }
421 }
422 }
423
424 let output = format!("\"{}\"", sql_content.escape_default());
426 output
427 .parse()
428 .unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
429}
430
431fn strip_string_literal(raw: &str) -> Option<&str> {
434 if raw.starts_with("r#\"") {
435 Some(&raw[3..raw.len() - 2])
436 } else if raw.starts_with("r\"") {
437 Some(&raw[2..raw.len() - 1])
438 } else if raw.starts_with('"') {
439 Some(&raw[1..raw.len() - 1])
440 } else if raw.starts_with("b\"") || raw.starts_with("b\'") {
441 Some(&raw[2..raw.len() - 1])
442 } else {
443 None
444 }
445}
446
447#[cfg(feature = "db-verify")]
452fn verify_with_real_db(sql: &str) -> Result<(), String> {
453 let dsn = std::env::var("DATABASE_URL")
454 .map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
455
456 let db_kind =
457 detect_db_kind(&dsn).map_err(|e| format!("Failed to detect DB kind from DSN: {}", e))?;
458
459 let explain_sql = match db_kind {
460 DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql),
461 DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
462 };
463
464 let rt = tokio::runtime::Runtime::new()
465 .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
466
467 rt.block_on(async {
468 match db_kind {
469 DbKind::MySql => verify_mysql(&dsn, &explain_sql).await,
470 DbKind::Postgres => verify_postgres(&dsn, &explain_sql).await,
471 DbKind::Sqlite => verify_sqlite(&dsn, &explain_sql).await,
472 }
473 })
474}
475
476#[cfg(feature = "db-verify")]
477#[derive(Debug, Clone, Copy, PartialEq, Eq)]
478enum DbKind {
479 MySql,
480 Postgres,
481 Sqlite,
482}
483
484#[cfg(feature = "db-verify")]
485fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
486 let lower = dsn.to_lowercase();
487 if lower.starts_with("mysql://") {
488 Ok(DbKind::MySql)
489 } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
490 Ok(DbKind::Postgres)
491 } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
492 Ok(DbKind::Sqlite)
493 } else {
494 Err(format!("Unsupported DSN scheme: {}", dsn))
495 }
496}
497
498#[cfg(feature = "db-verify")]
499async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
500 let pool = sqlx::MySqlPool::connect(dsn)
501 .await
502 .map_err(|e| format!("MySQL connect failed: {}", e))?;
503 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
504 .execute(&pool)
505 .await
506 .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
507 Ok(())
508}
509
510#[cfg(feature = "db-verify")]
511async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
512 let pool = sqlx::PgPool::connect(dsn)
513 .await
514 .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
515 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
516 .execute(&pool)
517 .await
518 .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
519 Ok(())
520}
521
522#[cfg(feature = "db-verify")]
523async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
524 let pool = sqlx::SqlitePool::connect(dsn)
525 .await
526 .map_err(|e| format!("SQLite connect failed: {}", e))?;
527 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
528 .execute(&pool)
529 .await
530 .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
531 Ok(())
532}
533
534fn compile_error(span: Span, msg: &str) -> TokenStream {
540 let mut ts = TokenStream::new();
542 ts.extend([
543 TokenTree::Ident(Ident::new("compile_error", span)),
544 TokenTree::Punct(Punct::new('!', Spacing::Alone)),
545 TokenTree::Group(Group::new(
546 Delimiter::Parenthesis,
547 TokenStream::from(TokenTree::Literal(Literal::string(msg))),
548 )),
549 ]);
550 ts
551}
552
553#[proc_macro]
591pub fn typed_query(input: TokenStream) -> TokenStream {
592 let tokens: Vec<TokenTree> = input.into_iter().collect();
593
594 if tokens.iter().any(|t| {
596 if let TokenTree::Ident(id) = t {
597 id.to_string() == "table"
598 } else {
599 false
600 }
601 }) {
602 return parse_table_decl(&tokens);
603 }
604
605 if tokens.iter().any(|t| {
607 if let TokenTree::Ident(id) = t {
608 id.to_string().eq_ignore_ascii_case("SELECT")
609 } else {
610 false
611 }
612 }) {
613 return parse_typed_select(&tokens);
614 }
615
616 compile_error(
617 Span::call_site(),
618 "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
619 )
620}
621
622fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
624 let mut idx = 0;
626
627 if idx >= tokens.len() {
629 return compile_error(Span::call_site(), "expected table name after 'table'");
630 }
631 if let TokenTree::Ident(id) = &tokens[idx] {
632 if id.to_string() != "table" {
633 return compile_error(id.span(), "expected 'table' keyword");
634 }
635 }
636 idx += 1;
637
638 let table_name = if idx < tokens.len() {
640 if let TokenTree::Ident(id) = &tokens[idx] {
641 id.to_string()
642 } else {
643 return compile_error(tokens[idx].span(), "expected table name identifier");
644 }
645 } else {
646 return compile_error(Span::call_site(), "expected table name");
647 };
648 idx += 1;
649
650 let body_group = if idx < tokens.len() {
652 if let TokenTree::Group(g) = &tokens[idx] {
653 if g.delimiter() != Delimiter::Brace {
654 return compile_error(g.span(), "expected '{' after table name");
655 }
656 g.clone()
657 } else {
658 return compile_error(tokens[idx].span(), "expected '{' after table name");
659 }
660 } else {
661 return compile_error(Span::call_site(), "expected table body in '{ }'");
662 };
663
664 let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
666 let columns = match parse_column_list(&body_tokens) {
667 Ok(c) => c,
668 Err(e) => return compile_error(Span::call_site(), &e),
669 };
670
671 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
673 let table_name_lit = table_name.as_str();
674
675 let col_impls: Vec<TokenStream2> = columns
677 .iter()
678 .map(|(col_name, col_type)| {
679 let col_ident =
680 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
681 let col_name_lit = col_name.as_str();
682 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
684 quote! {
685 #[derive(Debug, Clone, Copy)]
686 pub struct #col_ident;
687 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
688 const NAME: &'static str = #col_name_lit;
689 type Table = table;
690 type RustType = #rust_type;
691 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
692 }
693 }
694 })
695 .collect();
696
697 let schema_entries: Vec<TokenStream2> = columns
699 .iter()
700 .map(|(n, t)| {
701 let n_lit = n.as_str();
702 let t_lit = t.as_str();
703 quote! { (#n_lit, #t_lit) }
704 })
705 .collect();
706
707 let schema_const_ident = proc_macro2::Ident::new(
708 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
709 Span::call_site().into(),
710 );
711
712 let expanded = quote! {
713 pub mod #table_ident {
714 use super::*;
715 pub struct table;
716 impl ::sz_orm_core::typed::TypedTable for table {
717 const NAME: &'static str = #table_name_lit;
718 }
719 #(#col_impls)*
720 }
721 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
722 };
723
724 expanded.into()
725}
726
727fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
729 let mut cols = Vec::new();
730 let mut i = 0;
731 while i < tokens.len() {
732 let col_name = if let TokenTree::Ident(id) = &tokens[i] {
734 id.to_string()
735 } else {
736 return Err(format!("expected column name at position {}", i));
737 };
738 i += 1;
739
740 if i >= tokens.len() {
742 return Err(format!("expected ':' after column '{}'", col_name));
743 }
744 if let TokenTree::Punct(p) = &tokens[i] {
745 if p.as_char() != ':' {
746 return Err(format!("expected ':' after column '{}'", col_name));
747 }
748 } else {
749 return Err(format!("expected ':' after column '{}'", col_name));
750 }
751 i += 1;
752
753 let mut type_str = String::new();
756 let mut depth = 0;
757 while i < tokens.len() {
758 match &tokens[i] {
759 TokenTree::Punct(p) => {
760 if p.as_char() == ',' && depth == 0 {
761 i += 1;
762 break;
763 } else if p.as_char() == '<' || p.as_char() == '(' {
764 depth += 1;
765 type_str.push(p.as_char());
766 } else if p.as_char() == '>' || p.as_char() == ')' {
767 depth -= 1;
768 type_str.push(p.as_char());
769 } else {
770 type_str.push(p.as_char());
771 }
772 }
773 TokenTree::Ident(id) => {
774 if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
775 {
776 type_str.push(' ');
777 }
778 type_str.push_str(&id.to_string());
779 }
780 _ => {}
781 }
782 i += 1;
783 }
784
785 cols.push((col_name, type_str.trim().to_string()));
786 }
787 Ok(cols)
788}
789
790fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
794 let mut sql_parts: Vec<String> = Vec::new();
796 let mut table_name: Option<String> = None;
797 let mut in_from = false;
798
799 for (i, t) in tokens.iter().enumerate() {
800 match t {
801 TokenTree::Ident(id) => {
802 let s = id.to_string();
803 if s.eq_ignore_ascii_case("SELECT") {
804 sql_parts.push("SELECT".to_string());
805 } else if s.eq_ignore_ascii_case("FROM") {
806 in_from = true;
807 sql_parts.push("FROM".to_string());
808 } else if s.eq_ignore_ascii_case("WHERE")
809 || s.eq_ignore_ascii_case("AND")
810 || s.eq_ignore_ascii_case("OR")
811 || s.eq_ignore_ascii_case("LIMIT")
812 || s.eq_ignore_ascii_case("OFFSET")
813 || s.eq_ignore_ascii_case("ORDER")
814 || s.eq_ignore_ascii_case("BY")
815 || s.eq_ignore_ascii_case("GROUP")
816 || s.eq_ignore_ascii_case("HAVING")
817 || s.eq_ignore_ascii_case("JOIN")
818 || s.eq_ignore_ascii_case("INNER")
819 || s.eq_ignore_ascii_case("LEFT")
820 || s.eq_ignore_ascii_case("RIGHT")
821 || s.eq_ignore_ascii_case("ON")
822 || s.eq_ignore_ascii_case("AS")
823 || s.eq_ignore_ascii_case("ASC")
824 || s.eq_ignore_ascii_case("DESC")
825 || s.eq_ignore_ascii_case("DISTINCT")
826 || s.eq_ignore_ascii_case("NOT")
827 || s.eq_ignore_ascii_case("NULL")
828 || s.eq_ignore_ascii_case("IN")
829 || s.eq_ignore_ascii_case("BETWEEN")
830 || s.eq_ignore_ascii_case("LIKE")
831 || s.eq_ignore_ascii_case("IS")
832 {
833 sql_parts.push(s.to_uppercase());
834 } else if in_from && table_name.is_none() {
835 table_name = Some(s.clone());
837 sql_parts.push(s.clone());
838 } else {
839 sql_parts.push(s.clone());
840 }
841 }
842 TokenTree::Literal(lit) => {
843 sql_parts.push(lit.to_string());
844 }
845 TokenTree::Punct(p) => {
846 let c = p.as_char();
847 let part = if c == ',' {
849 ",".to_string()
850 } else if c == '?' {
851 "?".to_string()
852 } else if c == '*' {
853 "*".to_string()
854 } else if c == '=' {
855 "=".to_string()
856 } else if c == '>' {
857 ">".to_string()
858 } else if c == '<' {
859 "<".to_string()
860 } else if c == '.' {
861 ".".to_string()
862 } else if c == ';' {
863 ";".to_string()
864 } else {
865 c.to_string()
866 };
867 sql_parts.push(part);
868 }
869 TokenTree::Group(g) => {
870 let inner: String = g.stream().to_string();
872 let delim = match g.delimiter() {
873 Delimiter::Parenthesis => "(",
874 Delimiter::Brace => "{",
875 Delimiter::Bracket => "[",
876 Delimiter::None => "",
877 };
878 let close = match g.delimiter() {
879 Delimiter::Parenthesis => ")",
880 Delimiter::Brace => "}",
881 Delimiter::Bracket => "]",
882 Delimiter::None => "",
883 };
884 sql_parts.push(format!("{}{}{}", delim, inner, close));
885 }
886 }
887 let _ = i;
889 }
890
891 let sql = sql_parts
892 .join(" ")
893 .replace(", ", ",")
894 .replace(" ,", ",")
895 .replace("= ", "=")
896 .replace(" =", "=")
897 .replace("> ", ">")
898 .replace(" >", ">")
899 .replace("< ", "<")
900 .replace(" <", "<")
901 .replace(" ", " ");
902
903 if let Err(e) = validate_sql_content(&sql, None) {
905 return compile_error(
906 Span::call_site(),
907 &format!("typed_query! SQL validation failed: {}", e),
908 );
909 }
910
911 let mut ts = TokenStream::new();
913 let lit = Literal::string(&sql);
914 ts.extend([TokenTree::Literal(lit)]);
915 ts
916}
917
918#[proc_macro]
948pub fn schema(input: TokenStream) -> TokenStream {
949 let mut tokens = input.into_iter().peekable();
950
951 let sql_raw = match tokens.next() {
953 Some(TokenTree::Literal(lit)) => lit.to_string(),
954 Some(other) => {
955 return compile_error(
956 other.span(),
957 "Expected a string literal as the argument to schema!",
958 );
959 }
960 None => {
961 return compile_error(
962 Span::call_site(),
963 "Expected a string literal argument to schema!",
964 );
965 }
966 };
967
968 let sql = match strip_string_literal(&sql_raw) {
969 Some(s) => s,
970 None => {
971 return compile_error(
972 Span::call_site(),
973 "schema! requires a string literal argument",
974 );
975 }
976 };
977
978 let (table_name, columns) = match parse_create_table(sql) {
980 Ok(v) => v,
981 Err(e) => return compile_error(Span::call_site(), &e),
982 };
983
984 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
986 let table_name_lit = table_name.as_str();
987
988 let col_impls: Vec<TokenStream2> = columns
989 .iter()
990 .map(|(col_name, col_type)| {
991 let col_ident =
992 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
993 let col_name_lit = col_name.as_str();
994 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
995 quote! {
996 #[derive(Debug, Clone, Copy)]
997 pub struct #col_ident;
998 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
999 const NAME: &'static str = #col_name_lit;
1000 type Table = table;
1001 type RustType = #rust_type;
1002 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1003 }
1004 }
1005 })
1006 .collect();
1007
1008 let schema_entries: Vec<TokenStream2> = columns
1009 .iter()
1010 .map(|(n, t)| {
1011 let n_lit = n.as_str();
1012 let t_lit = t.as_str();
1013 quote! { (#n_lit, #t_lit) }
1014 })
1015 .collect();
1016
1017 let schema_const_ident = proc_macro2::Ident::new(
1018 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1019 Span::call_site().into(),
1020 );
1021
1022 let expanded = quote! {
1023 pub mod #table_ident {
1024 use super::*;
1025 pub struct table;
1026 impl ::sz_orm_core::typed::TypedTable for table {
1027 const NAME: &'static str = #table_name_lit;
1028 }
1029 #(#col_impls)*
1030 }
1031 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1032 };
1033
1034 expanded.into()
1035}
1036
1037fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
1045 let trimmed = sql.trim();
1046 let upper = trimmed.to_uppercase();
1047
1048 if !upper.starts_with("CREATE TABLE") {
1050 return Err("schema! expects a CREATE TABLE statement".to_string());
1051 }
1052
1053 let mut rest = &trimmed["CREATE TABLE".len()..];
1055
1056 let rest_upper = rest.trim_start().to_uppercase();
1058 if rest_upper.starts_with("IF NOT EXISTS") {
1059 rest = &rest.trim_start()["IF NOT EXISTS".len()..];
1060 }
1061
1062 rest = rest.trim_start();
1063
1064 let (table_name, after_name) = parse_identifier(rest)?;
1066 let rest = after_name.trim_start();
1067
1068 let paren_start = rest
1070 .find('(')
1071 .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
1072 let paren_end = rest
1073 .rfind(')')
1074 .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
1075 if paren_end <= paren_start {
1076 return Err("CREATE TABLE has malformed parentheses".to_string());
1077 }
1078
1079 let cols_str = &rest[paren_start + 1..paren_end];
1080
1081 let col_defs = split_top_level_commas(cols_str);
1083
1084 let mut columns = Vec::new();
1085 for def in col_defs {
1086 let def = def.trim();
1087 if def.is_empty() {
1088 continue;
1089 }
1090
1091 let def_upper = def.to_uppercase();
1093 if def_upper.starts_with("PRIMARY KEY")
1094 || def_upper.starts_with("FOREIGN KEY")
1095 || def_upper.starts_with("CONSTRAINT")
1096 || def_upper.starts_with("UNIQUE")
1097 || def_upper.starts_with("INDEX")
1098 || def_upper.starts_with("KEY")
1099 {
1100 continue;
1101 }
1102
1103 let (col_name, after_col) = parse_identifier(def)?;
1105 let rest = after_col.trim_start();
1106
1107 let (sql_type, after_type) = parse_type_token(rest)?;
1109 let rest = after_type.trim();
1110
1111 let rest_upper = rest.to_uppercase();
1113 let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
1114 let rust_type = sql_type_to_rust(&sql_type, !not_null);
1115
1116 columns.push((col_name, rust_type));
1117 }
1118
1119 Ok((table_name, columns))
1120}
1121
1122fn parse_identifier(s: &str) -> Result<(String, &str), String> {
1125 let s = s.trim_start();
1126 if s.is_empty() {
1127 return Err("expected identifier".to_string());
1128 }
1129
1130 let bytes = s.as_bytes();
1131 match bytes[0] {
1132 b'`' => {
1133 let end = s[1..]
1134 .find('`')
1135 .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
1136 let ident = s[1..1 + end].to_string();
1137 Ok((ident, &s[1 + end + 1..]))
1138 }
1139 b'"' => {
1140 let end = s[1..]
1141 .find('"')
1142 .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
1143 let ident = s[1..1 + end].to_string();
1144 Ok((ident, &s[1 + end + 1..]))
1145 }
1146 _ => {
1147 let end = s
1148 .find(|c: char| !c.is_alphanumeric() && c != '_')
1149 .unwrap_or(s.len());
1150 if end == 0 {
1151 return Err(format!("invalid identifier: '{}'", s));
1152 }
1153 let ident = s[..end].to_string();
1154 Ok((ident, &s[end..]))
1155 }
1156 }
1157}
1158
1159fn parse_type_token(s: &str) -> Result<(String, &str), String> {
1162 let s = s.trim_start();
1163 if s.is_empty() {
1164 return Err("expected column type".to_string());
1165 }
1166
1167 let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
1168 if end == 0 {
1169 return Err(format!("invalid type: '{}'", s));
1170 }
1171 let type_name = s[..end].to_string();
1172 let mut rest = &s[end..];
1173
1174 rest = rest.trim_start();
1176 if rest.starts_with('(') {
1177 let close = rest
1178 .find(')')
1179 .ok_or_else(|| "unterminated type parameter list".to_string())?;
1180 rest = &rest[close + 1..];
1181 }
1182
1183 Ok((type_name, rest))
1184}
1185
1186fn split_top_level_commas(s: &str) -> Vec<String> {
1188 let mut parts = Vec::new();
1189 let mut depth: i32 = 0;
1190 let mut current = String::new();
1191
1192 for ch in s.chars() {
1193 match ch {
1194 '(' => {
1195 depth += 1;
1196 current.push(ch);
1197 }
1198 ')' => {
1199 depth -= 1;
1200 current.push(ch);
1201 }
1202 ',' if depth == 0 => {
1203 parts.push(std::mem::take(&mut current));
1204 }
1205 _ => {
1206 current.push(ch);
1207 }
1208 }
1209 }
1210
1211 if !current.trim().is_empty() {
1212 parts.push(current);
1213 }
1214
1215 parts
1216}
1217
1218fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
1223 let upper = sql_type.to_uppercase();
1224 let rust = match upper.as_str() {
1225 "BIGINT" | "INT8" => "i64",
1227 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
1229 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
1231 "TINYINT" => "i8",
1233 "FLOAT" | "REAL" | "FLOAT4" => "f32",
1235 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
1237 "BOOLEAN" | "BOOL" => "bool",
1239 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
1241 "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
1243 | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
1244 _ => "String",
1245 };
1246
1247 if nullable {
1248 format!("Option<{}>", rust)
1249 } else {
1250 rust.to_string()
1251 }
1252}
1253
1254#[proc_macro_derive(Schema, attributes(table, column))]
1284pub fn derive_schema(input: TokenStream) -> TokenStream {
1285 let input = parse_macro_input!(input as syn::DeriveInput);
1286 derive::derive_schema_impl(input).into()
1287}
1288
1289#[proc_macro_derive(Builder, attributes(builder))]
1323pub fn derive_builder(input: TokenStream) -> TokenStream {
1324 let input = parse_macro_input!(input as syn::DeriveInput);
1325 derive::derive_builder_impl(input).into()
1326}
1327
1328#[cfg(test)]
1333mod tests {
1334 use super::*;
1335
1336 #[test]
1339 fn test_strip_plain_double_quoted() {
1340 assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
1341 }
1342
1343 #[test]
1344 fn test_strip_raw_double_hash() {
1345 assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
1346 }
1347
1348 #[test]
1349 fn test_strip_raw_double_no_hash() {
1350 assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
1351 }
1352
1353 #[test]
1354 fn test_strip_byte_string() {
1355 assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
1356 assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
1357 }
1358
1359 #[test]
1360 fn test_strip_non_string_returns_none() {
1361 assert_eq!(strip_string_literal("123"), None);
1362 assert_eq!(strip_string_literal("foo"), None);
1363 }
1364
1365 #[test]
1368 fn test_validate_select_with_from_ok() {
1369 assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
1370 }
1371
1372 #[test]
1373 fn test_validate_select_missing_from_fails() {
1374 assert!(validate_sql_content("SELECT * users", None).is_err());
1375 }
1376
1377 #[test]
1378 fn test_validate_insert_missing_into_fails() {
1379 assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
1380 assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
1381 }
1382
1383 #[test]
1384 fn test_validate_update_missing_set_fails() {
1385 assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
1386 assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
1387 }
1388
1389 #[test]
1390 fn test_validate_delete_missing_from_fails() {
1391 assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
1392 assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
1393 }
1394
1395 #[test]
1396 fn test_validate_empty_sql_fails() {
1397 assert!(validate_sql_content("", None).is_err());
1398 assert!(validate_sql_content(" ", None).is_err());
1399 }
1400
1401 #[test]
1404 fn test_validate_balanced_parens_ok() {
1405 assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
1406 }
1407
1408 #[test]
1409 fn test_validate_balanced_parens_unbalanced() {
1410 assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
1411 assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
1412 }
1413
1414 #[test]
1417 fn test_validate_no_injection_clean() {
1418 assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
1419 }
1420
1421 #[test]
1422 fn test_validate_no_injection_drop_table() {
1423 assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
1424 }
1425
1426 #[test]
1427 fn test_validate_no_injection_or_1_1() {
1428 assert!(validate_no_injection("' OR 1=1").is_err());
1432 assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
1433 }
1434
1435 #[test]
1436 fn test_validate_no_injection_drop_database() {
1437 assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
1438 }
1439
1440 #[test]
1441 fn test_validate_no_injection_information_schema() {
1442 assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
1443 }
1444
1445 #[test]
1446 fn test_validate_no_injection_xp_cmdshell() {
1447 assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
1448 }
1449
1450 #[test]
1451 fn test_validate_no_injection_union_select() {
1452 assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
1453 }
1454
1455 #[test]
1456 fn test_validate_no_injection_comment_dashes() {
1457 assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
1458 }
1459
1460 #[test]
1461 fn test_validate_no_injection_block_comment() {
1462 assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
1463 }
1464
1465 #[test]
1468 fn test_validate_string_literals_closed_ok() {
1469 assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
1470 assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
1471 }
1472
1473 #[test]
1474 fn test_validate_string_literals_closed_unclosed_single() {
1475 assert!(validate_string_literals_closed("'hello").is_err());
1476 }
1477
1478 #[test]
1479 fn test_validate_string_literals_closed_unclosed_double() {
1480 assert!(validate_string_literals_closed(r#""hello"#).is_err());
1481 }
1482
1483 #[test]
1486 fn test_validate_param_count_match() {
1487 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
1488 assert!(
1489 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
1490 );
1491 }
1492
1493 #[test]
1494 fn test_validate_param_count_mismatch() {
1495 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
1496 assert!(
1497 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
1498 );
1499 }
1500
1501 #[cfg(feature = "db-verify")]
1504 #[test]
1505 fn test_detect_db_kind_mysql() {
1506 assert_eq!(
1507 detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
1508 DbKind::MySql
1509 );
1510 }
1511
1512 #[cfg(feature = "db-verify")]
1513 #[test]
1514 fn test_detect_db_kind_postgres() {
1515 assert_eq!(
1516 detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
1517 DbKind::Postgres
1518 );
1519 assert_eq!(
1520 detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
1521 DbKind::Postgres
1522 );
1523 }
1524
1525 #[cfg(feature = "db-verify")]
1526 #[test]
1527 fn test_detect_db_kind_sqlite() {
1528 assert_eq!(
1529 detect_db_kind("sqlite://path/to/db.db").unwrap(),
1530 DbKind::Sqlite
1531 );
1532 assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
1533 }
1534
1535 #[cfg(feature = "db-verify")]
1536 #[test]
1537 fn test_detect_db_kind_unsupported() {
1538 assert!(detect_db_kind("oracle://user:pass@host/db").is_err());
1539 assert!(detect_db_kind("not-a-url").is_err());
1540 }
1541
1542 #[test]
1545 fn test_parse_create_table_basic() {
1546 let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
1547 let (table, cols) = parse_create_table(sql).unwrap();
1548 assert_eq!(table, "users");
1549 assert_eq!(
1550 cols,
1551 vec![
1552 ("id".to_string(), "i32".to_string()),
1553 ("name".to_string(), "String".to_string())
1554 ]
1555 );
1556 }
1557
1558 #[test]
1559 fn test_parse_create_table_with_if_not_exists() {
1560 let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
1561 let (table, cols) = parse_create_table(sql).unwrap();
1562 assert_eq!(table, "orders");
1563 assert_eq!(
1564 cols,
1565 vec![
1566 ("id".to_string(), "i64".to_string()),
1567 ("total".to_string(), "f64".to_string())
1568 ]
1569 );
1570 }
1571
1572 #[test]
1573 fn test_parse_create_table_nullable() {
1574 let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
1575 let (_, cols) = parse_create_table(sql).unwrap();
1576 assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
1577 assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
1578 }
1579
1580 #[test]
1581 fn test_parse_create_table_skip_constraints() {
1582 let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
1583 let (_, cols) = parse_create_table(sql).unwrap();
1584 assert_eq!(cols.len(), 2);
1585 assert_eq!(cols[0].0, "id");
1586 assert_eq!(cols[1].0, "name");
1587 }
1588
1589 #[test]
1590 fn test_parse_create_table_varchar_with_len() {
1591 let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
1592 let (_, cols) = parse_create_table(sql).unwrap();
1593 assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
1594 assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
1595 }
1596
1597 #[test]
1598 fn test_sql_type_to_rust_mappings() {
1599 assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
1601 assert_eq!(sql_type_to_rust("INT8", false), "i64");
1602 assert_eq!(sql_type_to_rust("INT", false), "i32");
1603 assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
1604 assert_eq!(sql_type_to_rust("INT4", false), "i32");
1605 assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
1606 assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
1607 assert_eq!(sql_type_to_rust("INT2", false), "i16");
1608 assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
1609 assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
1610 assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
1612 assert_eq!(sql_type_to_rust("REAL", false), "f32");
1613 assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
1614 assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
1615 assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
1616 assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
1617 assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
1618 assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
1619 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
1621 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
1622 assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
1624 assert_eq!(sql_type_to_rust("TEXT", false), "String");
1625 assert_eq!(sql_type_to_rust("CHAR", false), "String");
1626 assert_eq!(sql_type_to_rust("UUID", false), "String");
1627 assert_eq!(sql_type_to_rust("DATE", false), "String");
1628 assert_eq!(sql_type_to_rust("DATETIME", false), "String");
1629 assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
1630 assert_eq!(sql_type_to_rust("JSON", false), "String");
1631 assert_eq!(sql_type_to_rust("JSONB", false), "String");
1632 assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
1634 assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
1635 assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
1636 assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
1637 assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
1639 assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
1640 assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
1641 assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
1642 assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
1644 }
1645
1646 #[test]
1647 fn test_parse_create_table_error_no_create() {
1648 assert!(parse_create_table("SELECT * FROM users").is_err());
1649 }
1650
1651 #[test]
1652 fn test_parse_create_table_error_no_parens() {
1653 assert!(parse_create_table("CREATE TABLE foo").is_err());
1654 }
1655}