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 sql_no_placeholders = replace_placeholders_with_null(sql);
462
463 let explain_sql = match db_kind {
465 DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql_no_placeholders),
466 DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql_no_placeholders),
467 DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
469 DbKind::SqlServer => sql_no_placeholders,
471 };
472
473 if matches!(db_kind, DbKind::MySql | DbKind::Postgres | DbKind::Sqlite) {
475 let rt = tokio::runtime::Runtime::new()
476 .map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
477 return rt.block_on(async {
478 match db_kind {
479 DbKind::MySql => verify_mysql(&dsn, &explain_sql).await,
480 DbKind::Postgres => verify_postgres(&dsn, &explain_sql).await,
481 DbKind::Sqlite => verify_sqlite(&dsn, &explain_sql).await,
482 _ => unreachable!(),
483 }
484 });
485 }
486
487 match db_kind {
489 DbKind::Oracle => verify_oracle(&dsn, &explain_sql),
490 DbKind::SqlServer => verify_sqlserver(&dsn, &explain_sql),
491 _ => unreachable!(),
492 }
493}
494
495#[cfg(feature = "db-verify")]
496#[derive(Debug, Clone, Copy, PartialEq, Eq)]
497enum DbKind {
498 MySql,
499 Postgres,
500 Sqlite,
501 Oracle,
502 SqlServer,
503}
504
505#[cfg(feature = "db-verify")]
510fn replace_placeholders_with_null(sql: &str) -> String {
511 let mut result = String::with_capacity(sql.len() + 16);
512 let mut in_single_quote = false;
513 let mut in_double_quote = false;
514 let mut prev = '\0';
515
516 for ch in sql.chars() {
517 if prev == '\\' {
518 result.push(ch);
520 prev = ch;
521 continue;
522 }
523 match ch {
524 '\'' if !in_double_quote => in_single_quote = !in_single_quote,
525 '"' if !in_single_quote => in_double_quote = !in_double_quote,
526 '?' if !in_single_quote && !in_double_quote => {
527 result.push_str("NULL");
528 prev = ch;
529 continue;
530 }
531 _ => {}
532 }
533 result.push(ch);
534 prev = ch;
535 }
536 result
537}
538
539#[cfg(feature = "db-verify")]
540fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
541 let lower = dsn.to_lowercase();
542 if lower.starts_with("mysql://") {
543 Ok(DbKind::MySql)
544 } else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
545 Ok(DbKind::Postgres)
546 } else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
547 Ok(DbKind::Sqlite)
548 } else if lower.starts_with("oracle://") || lower.starts_with("oracle:") {
549 Ok(DbKind::Oracle)
550 } else if lower.starts_with("sqlserver://")
551 || lower.starts_with("mssql://")
552 || lower.starts_with("tds://")
553 {
554 Ok(DbKind::SqlServer)
555 } else {
556 Err(format!("Unsupported DSN scheme: {}", dsn))
557 }
558}
559
560#[cfg(feature = "db-verify")]
561async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
562 let pool = sqlx::MySqlPool::connect(dsn)
563 .await
564 .map_err(|e| format!("MySQL connect failed: {}", e))?;
565 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
566 .execute(&pool)
567 .await
568 .map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
569 Ok(())
570}
571
572#[cfg(feature = "db-verify")]
573async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
574 let pool = sqlx::PgPool::connect(dsn)
575 .await
576 .map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
577 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
578 .execute(&pool)
579 .await
580 .map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
581 Ok(())
582}
583
584#[cfg(feature = "db-verify")]
585async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
586 let pool = sqlx::SqlitePool::connect(dsn)
587 .await
588 .map_err(|e| format!("SQLite connect failed: {}", e))?;
589 sqlx::query(sqlx::AssertSqlSafe(explain_sql))
590 .execute(&pool)
591 .await
592 .map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
593 Ok(())
594}
595
596#[cfg(feature = "db-verify")]
601fn verify_oracle(dsn: &str, explain_sql: &str) -> Result<(), String> {
602 let parsed = parse_oracle_dsn(dsn)?;
603 let mut conn_str = format!(
605 "{}/{}@{}:{}/{}",
606 parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
607 );
608 if parsed.sysdba {
609 conn_str.push_str(" AS SYSDBA");
610 }
611 let full_script = format!(
613 "SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
614 EXPLAIN PLAN FOR {};\n\
615 SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
616 EXIT;\n",
617 explain_sql
618 );
619 let output = std::process::Command::new("sqlplus")
620 .args(["-S", "-L", &conn_str])
621 .stdin(std::process::Stdio::piped())
622 .stdout(std::process::Stdio::piped())
623 .stderr(std::process::Stdio::piped())
624 .spawn()
625 .map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
626 use std::io::Write;
627 let mut child = output;
628 if let Some(mut stdin) = child.stdin.take() {
629 stdin
630 .write_all(full_script.as_bytes())
631 .map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
632 }
633 let out = child
634 .wait_with_output()
635 .map_err(|e| format!("sqlplus wait failed: {}", e))?;
636 let stdout = String::from_utf8_lossy(&out.stdout);
637 let stderr = String::from_utf8_lossy(&out.stderr);
638 if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
639 return Err(format!(
640 "Oracle EXPLAIN failed: stdout={} stderr={}",
641 stdout.trim(),
642 stderr.trim()
643 ));
644 }
645 Ok(())
646}
647
648#[cfg(feature = "db-verify")]
653fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
654 let parsed = parse_sqlserver_dsn(dsn)?;
655 let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
657 let out = std::process::Command::new("sqlcmd")
658 .args([
659 "-S",
660 &format!("{},{}", parsed.host, parsed.port),
661 "-U",
662 &parsed.user,
663 "-P",
664 &parsed.password,
665 "-d",
666 &parsed.database,
667 "-Q",
668 &query,
669 "-h",
670 "-1",
671 "-W",
672 ])
673 .output()
674 .map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
675 let stdout = String::from_utf8_lossy(&out.stdout);
676 let stderr = String::from_utf8_lossy(&out.stderr);
677 if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
678 return Err(format!(
679 "SQL Server SHOWPLAN failed: stdout={} stderr={}",
680 stdout.trim(),
681 stderr.trim()
682 ));
683 }
684 Ok(())
685}
686
687#[cfg(feature = "db-verify")]
689struct OracleDsn {
690 user: String,
691 password: String,
692 host: String,
693 port: u16,
694 service: String,
695 sysdba: bool,
696}
697
698#[cfg(feature = "db-verify")]
700fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
701 let raw = dsn
702 .strip_prefix("oracle://")
703 .or_else(|| dsn.strip_prefix("oracle:"))
704 .ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
705 let (auth_host_service, query) = match raw.find('?') {
707 Some(idx) => (&raw[..idx], &raw[idx + 1..]),
708 None => (raw, ""),
709 };
710 let sysdba = query.split('&').any(|p| p == "sysdba=1" || p == "sysdba=true");
711 let at = auth_host_service
713 .find('@')
714 .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
715 let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
716 let colon = user_pass
717 .find(':')
718 .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
719 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
720 let (host_port, service) = match host_port_service.rfind('/') {
721 Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
722 None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
723 };
724 let (host, port) = match host_port.find(':') {
725 Some(idx) => (
726 &host_port[..idx],
727 host_port[idx + 1..]
728 .parse::<u16>()
729 .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
730 ),
731 None => (host_port, 1521u16),
732 };
733 Ok(OracleDsn {
734 user: user.to_string(),
735 password: password.to_string(),
736 host: host.to_string(),
737 port,
738 service: service.to_string(),
739 sysdba,
740 })
741}
742
743#[cfg(feature = "db-verify")]
745struct SqlServerDsn {
746 user: String,
747 password: String,
748 host: String,
749 port: u16,
750 database: String,
751}
752
753#[cfg(feature = "db-verify")]
755fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
756 let raw = dsn
757 .strip_prefix("sqlserver://")
758 .or_else(|| dsn.strip_prefix("mssql://"))
759 .or_else(|| dsn.strip_prefix("tds://"))
760 .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
761 let at = raw
762 .find('@')
763 .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
764 let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
765 let colon = user_pass
766 .find(':')
767 .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
768 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
769 let (host_port, database) = match host_port_db.rfind('/') {
770 Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
771 None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
772 };
773 let (host, port) = match host_port.find(':') {
774 Some(idx) => (
775 &host_port[..idx],
776 host_port[idx + 1..]
777 .parse::<u16>()
778 .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
779 ),
780 None => (host_port, 1433u16),
781 };
782 Ok(SqlServerDsn {
783 user: user.to_string(),
784 password: password.to_string(),
785 host: host.to_string(),
786 port,
787 database: database.to_string(),
788 })
789}
790
791fn compile_error(span: Span, msg: &str) -> TokenStream {
797 let mut ts = TokenStream::new();
799 ts.extend([
800 TokenTree::Ident(Ident::new("compile_error", span)),
801 TokenTree::Punct(Punct::new('!', Spacing::Alone)),
802 TokenTree::Group(Group::new(
803 Delimiter::Parenthesis,
804 TokenStream::from(TokenTree::Literal(Literal::string(msg))),
805 )),
806 ]);
807 ts
808}
809
810#[proc_macro]
848pub fn typed_query(input: TokenStream) -> TokenStream {
849 let tokens: Vec<TokenTree> = input.into_iter().collect();
850
851 if tokens.iter().any(|t| {
853 if let TokenTree::Ident(id) = t {
854 id.to_string() == "table"
855 } else {
856 false
857 }
858 }) {
859 return parse_table_decl(&tokens);
860 }
861
862 if tokens.iter().any(|t| {
864 if let TokenTree::Ident(id) = t {
865 id.to_string().eq_ignore_ascii_case("SELECT")
866 } else {
867 false
868 }
869 }) {
870 return parse_typed_select(&tokens);
871 }
872
873 compile_error(
874 Span::call_site(),
875 "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
876 )
877}
878
879fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
881 let mut idx = 0;
883
884 if idx >= tokens.len() {
886 return compile_error(Span::call_site(), "expected table name after 'table'");
887 }
888 if let TokenTree::Ident(id) = &tokens[idx] {
889 if id.to_string() != "table" {
890 return compile_error(id.span(), "expected 'table' keyword");
891 }
892 }
893 idx += 1;
894
895 let table_name = if idx < tokens.len() {
897 if let TokenTree::Ident(id) = &tokens[idx] {
898 id.to_string()
899 } else {
900 return compile_error(tokens[idx].span(), "expected table name identifier");
901 }
902 } else {
903 return compile_error(Span::call_site(), "expected table name");
904 };
905 idx += 1;
906
907 let body_group = if idx < tokens.len() {
909 if let TokenTree::Group(g) = &tokens[idx] {
910 if g.delimiter() != Delimiter::Brace {
911 return compile_error(g.span(), "expected '{' after table name");
912 }
913 g.clone()
914 } else {
915 return compile_error(tokens[idx].span(), "expected '{' after table name");
916 }
917 } else {
918 return compile_error(Span::call_site(), "expected table body in '{ }'");
919 };
920
921 let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
923 let columns = match parse_column_list(&body_tokens) {
924 Ok(c) => c,
925 Err(e) => return compile_error(Span::call_site(), &e),
926 };
927
928 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
930 let table_name_lit = table_name.as_str();
931
932 let col_impls: Vec<TokenStream2> = columns
934 .iter()
935 .map(|(col_name, col_type)| {
936 let col_ident =
937 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
938 let col_name_lit = col_name.as_str();
939 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
941 quote! {
942 #[derive(Debug, Clone, Copy)]
943 pub struct #col_ident;
944 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
945 const NAME: &'static str = #col_name_lit;
946 type Table = table;
947 type RustType = #rust_type;
948 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
949 }
950 }
951 })
952 .collect();
953
954 let schema_entries: Vec<TokenStream2> = columns
956 .iter()
957 .map(|(n, t)| {
958 let n_lit = n.as_str();
959 let t_lit = t.as_str();
960 quote! { (#n_lit, #t_lit) }
961 })
962 .collect();
963
964 let schema_const_ident = proc_macro2::Ident::new(
965 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
966 Span::call_site().into(),
967 );
968
969 let expanded = quote! {
970 pub mod #table_ident {
971 use super::*;
972 pub struct table;
973 impl ::sz_orm_core::typed::TypedTable for table {
974 const NAME: &'static str = #table_name_lit;
975 }
976 #(#col_impls)*
977 }
978 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
979 };
980
981 expanded.into()
982}
983
984fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
986 let mut cols = Vec::new();
987 let mut i = 0;
988 while i < tokens.len() {
989 let col_name = if let TokenTree::Ident(id) = &tokens[i] {
991 id.to_string()
992 } else {
993 return Err(format!("expected column name at position {}", i));
994 };
995 i += 1;
996
997 if i >= tokens.len() {
999 return Err(format!("expected ':' after column '{}'", col_name));
1000 }
1001 if let TokenTree::Punct(p) = &tokens[i] {
1002 if p.as_char() != ':' {
1003 return Err(format!("expected ':' after column '{}'", col_name));
1004 }
1005 } else {
1006 return Err(format!("expected ':' after column '{}'", col_name));
1007 }
1008 i += 1;
1009
1010 let mut type_str = String::new();
1013 let mut depth = 0;
1014 while i < tokens.len() {
1015 match &tokens[i] {
1016 TokenTree::Punct(p) => {
1017 if p.as_char() == ',' && depth == 0 {
1018 i += 1;
1019 break;
1020 } else if p.as_char() == '<' || p.as_char() == '(' {
1021 depth += 1;
1022 type_str.push(p.as_char());
1023 } else if p.as_char() == '>' || p.as_char() == ')' {
1024 depth -= 1;
1025 type_str.push(p.as_char());
1026 } else {
1027 type_str.push(p.as_char());
1028 }
1029 }
1030 TokenTree::Ident(id) => {
1031 if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
1032 {
1033 type_str.push(' ');
1034 }
1035 type_str.push_str(&id.to_string());
1036 }
1037 _ => {}
1038 }
1039 i += 1;
1040 }
1041
1042 cols.push((col_name, type_str.trim().to_string()));
1043 }
1044 Ok(cols)
1045}
1046
1047fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
1051 let mut sql_parts: Vec<String> = Vec::new();
1053 let mut table_name: Option<String> = None;
1054 let mut in_from = false;
1055
1056 for (i, t) in tokens.iter().enumerate() {
1057 match t {
1058 TokenTree::Ident(id) => {
1059 let s = id.to_string();
1060 if s.eq_ignore_ascii_case("SELECT") {
1061 sql_parts.push("SELECT".to_string());
1062 } else if s.eq_ignore_ascii_case("FROM") {
1063 in_from = true;
1064 sql_parts.push("FROM".to_string());
1065 } else if s.eq_ignore_ascii_case("WHERE")
1066 || s.eq_ignore_ascii_case("AND")
1067 || s.eq_ignore_ascii_case("OR")
1068 || s.eq_ignore_ascii_case("LIMIT")
1069 || s.eq_ignore_ascii_case("OFFSET")
1070 || s.eq_ignore_ascii_case("ORDER")
1071 || s.eq_ignore_ascii_case("BY")
1072 || s.eq_ignore_ascii_case("GROUP")
1073 || s.eq_ignore_ascii_case("HAVING")
1074 || s.eq_ignore_ascii_case("JOIN")
1075 || s.eq_ignore_ascii_case("INNER")
1076 || s.eq_ignore_ascii_case("LEFT")
1077 || s.eq_ignore_ascii_case("RIGHT")
1078 || s.eq_ignore_ascii_case("ON")
1079 || s.eq_ignore_ascii_case("AS")
1080 || s.eq_ignore_ascii_case("ASC")
1081 || s.eq_ignore_ascii_case("DESC")
1082 || s.eq_ignore_ascii_case("DISTINCT")
1083 || s.eq_ignore_ascii_case("NOT")
1084 || s.eq_ignore_ascii_case("NULL")
1085 || s.eq_ignore_ascii_case("IN")
1086 || s.eq_ignore_ascii_case("BETWEEN")
1087 || s.eq_ignore_ascii_case("LIKE")
1088 || s.eq_ignore_ascii_case("IS")
1089 {
1090 sql_parts.push(s.to_uppercase());
1091 } else if in_from && table_name.is_none() {
1092 table_name = Some(s.clone());
1094 sql_parts.push(s.clone());
1095 } else {
1096 sql_parts.push(s.clone());
1097 }
1098 }
1099 TokenTree::Literal(lit) => {
1100 sql_parts.push(lit.to_string());
1101 }
1102 TokenTree::Punct(p) => {
1103 let c = p.as_char();
1104 let part = if c == ',' {
1106 ",".to_string()
1107 } else if c == '?' {
1108 "?".to_string()
1109 } else if c == '*' {
1110 "*".to_string()
1111 } else if c == '=' {
1112 "=".to_string()
1113 } else if c == '>' {
1114 ">".to_string()
1115 } else if c == '<' {
1116 "<".to_string()
1117 } else if c == '.' {
1118 ".".to_string()
1119 } else if c == ';' {
1120 ";".to_string()
1121 } else {
1122 c.to_string()
1123 };
1124 sql_parts.push(part);
1125 }
1126 TokenTree::Group(g) => {
1127 let inner: String = g.stream().to_string();
1129 let delim = match g.delimiter() {
1130 Delimiter::Parenthesis => "(",
1131 Delimiter::Brace => "{",
1132 Delimiter::Bracket => "[",
1133 Delimiter::None => "",
1134 };
1135 let close = match g.delimiter() {
1136 Delimiter::Parenthesis => ")",
1137 Delimiter::Brace => "}",
1138 Delimiter::Bracket => "]",
1139 Delimiter::None => "",
1140 };
1141 sql_parts.push(format!("{}{}{}", delim, inner, close));
1142 }
1143 }
1144 let _ = i;
1146 }
1147
1148 let sql = sql_parts
1149 .join(" ")
1150 .replace(", ", ",")
1151 .replace(" ,", ",")
1152 .replace("= ", "=")
1153 .replace(" =", "=")
1154 .replace("> ", ">")
1155 .replace(" >", ">")
1156 .replace("< ", "<")
1157 .replace(" <", "<")
1158 .replace(" ", " ");
1159
1160 if let Err(e) = validate_sql_content(&sql, None) {
1162 return compile_error(
1163 Span::call_site(),
1164 &format!("typed_query! SQL validation failed: {}", e),
1165 );
1166 }
1167
1168 let mut ts = TokenStream::new();
1170 let lit = Literal::string(&sql);
1171 ts.extend([TokenTree::Literal(lit)]);
1172 ts
1173}
1174
1175#[proc_macro]
1205pub fn schema(input: TokenStream) -> TokenStream {
1206 let mut tokens = input.into_iter().peekable();
1207
1208 let sql_raw = match tokens.next() {
1210 Some(TokenTree::Literal(lit)) => lit.to_string(),
1211 Some(other) => {
1212 return compile_error(
1213 other.span(),
1214 "Expected a string literal as the argument to schema!",
1215 );
1216 }
1217 None => {
1218 return compile_error(
1219 Span::call_site(),
1220 "Expected a string literal argument to schema!",
1221 );
1222 }
1223 };
1224
1225 let sql = match strip_string_literal(&sql_raw) {
1226 Some(s) => s,
1227 None => {
1228 return compile_error(
1229 Span::call_site(),
1230 "schema! requires a string literal argument",
1231 );
1232 }
1233 };
1234
1235 let (table_name, columns) = match parse_create_table(sql) {
1237 Ok(v) => v,
1238 Err(e) => return compile_error(Span::call_site(), &e),
1239 };
1240
1241 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1243 let table_name_lit = table_name.as_str();
1244
1245 let col_impls: Vec<TokenStream2> = columns
1246 .iter()
1247 .map(|(col_name, col_type)| {
1248 let col_ident =
1249 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1250 let col_name_lit = col_name.as_str();
1251 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1252 quote! {
1253 #[derive(Debug, Clone, Copy)]
1254 pub struct #col_ident;
1255 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1256 const NAME: &'static str = #col_name_lit;
1257 type Table = table;
1258 type RustType = #rust_type;
1259 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1260 }
1261 }
1262 })
1263 .collect();
1264
1265 let schema_entries: Vec<TokenStream2> = columns
1266 .iter()
1267 .map(|(n, t)| {
1268 let n_lit = n.as_str();
1269 let t_lit = t.as_str();
1270 quote! { (#n_lit, #t_lit) }
1271 })
1272 .collect();
1273
1274 let schema_const_ident = proc_macro2::Ident::new(
1275 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1276 Span::call_site().into(),
1277 );
1278
1279 let expanded = quote! {
1280 pub mod #table_ident {
1281 use super::*;
1282 pub struct table;
1283 impl ::sz_orm_core::typed::TypedTable for table {
1284 const NAME: &'static str = #table_name_lit;
1285 }
1286 #(#col_impls)*
1287 }
1288 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1289 };
1290
1291 expanded.into()
1292}
1293
1294fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
1302 let trimmed = sql.trim();
1303 let upper = trimmed.to_uppercase();
1304
1305 if !upper.starts_with("CREATE TABLE") {
1307 return Err("schema! expects a CREATE TABLE statement".to_string());
1308 }
1309
1310 let mut rest = &trimmed["CREATE TABLE".len()..];
1312
1313 let rest_upper = rest.trim_start().to_uppercase();
1315 if rest_upper.starts_with("IF NOT EXISTS") {
1316 rest = &rest.trim_start()["IF NOT EXISTS".len()..];
1317 }
1318
1319 rest = rest.trim_start();
1320
1321 let (table_name, after_name) = parse_identifier(rest)?;
1323 let rest = after_name.trim_start();
1324
1325 let paren_start = rest
1327 .find('(')
1328 .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
1329 let paren_end = rest
1330 .rfind(')')
1331 .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
1332 if paren_end <= paren_start {
1333 return Err("CREATE TABLE has malformed parentheses".to_string());
1334 }
1335
1336 let cols_str = &rest[paren_start + 1..paren_end];
1337
1338 let col_defs = split_top_level_commas(cols_str);
1340
1341 let mut columns = Vec::new();
1342 for def in col_defs {
1343 let def = def.trim();
1344 if def.is_empty() {
1345 continue;
1346 }
1347
1348 let def_upper = def.to_uppercase();
1350 if def_upper.starts_with("PRIMARY KEY")
1351 || def_upper.starts_with("FOREIGN KEY")
1352 || def_upper.starts_with("CONSTRAINT")
1353 || def_upper.starts_with("UNIQUE")
1354 || def_upper.starts_with("INDEX")
1355 || def_upper.starts_with("KEY")
1356 {
1357 continue;
1358 }
1359
1360 let (col_name, after_col) = parse_identifier(def)?;
1362 let rest = after_col.trim_start();
1363
1364 let (sql_type, after_type) = parse_type_token(rest)?;
1366 let rest = after_type.trim();
1367
1368 let rest_upper = rest.to_uppercase();
1370 let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
1371 let rust_type = sql_type_to_rust(&sql_type, !not_null);
1372
1373 columns.push((col_name, rust_type));
1374 }
1375
1376 Ok((table_name, columns))
1377}
1378
1379fn parse_identifier(s: &str) -> Result<(String, &str), String> {
1382 let s = s.trim_start();
1383 if s.is_empty() {
1384 return Err("expected identifier".to_string());
1385 }
1386
1387 let bytes = s.as_bytes();
1388 match bytes[0] {
1389 b'`' => {
1390 let end = s[1..]
1391 .find('`')
1392 .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
1393 let ident = s[1..1 + end].to_string();
1394 Ok((ident, &s[1 + end + 1..]))
1395 }
1396 b'"' => {
1397 let end = s[1..]
1398 .find('"')
1399 .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
1400 let ident = s[1..1 + end].to_string();
1401 Ok((ident, &s[1 + end + 1..]))
1402 }
1403 _ => {
1404 let end = s
1405 .find(|c: char| !c.is_alphanumeric() && c != '_')
1406 .unwrap_or(s.len());
1407 if end == 0 {
1408 return Err(format!("invalid identifier: '{}'", s));
1409 }
1410 let ident = s[..end].to_string();
1411 Ok((ident, &s[end..]))
1412 }
1413 }
1414}
1415
1416fn parse_type_token(s: &str) -> Result<(String, &str), String> {
1419 let s = s.trim_start();
1420 if s.is_empty() {
1421 return Err("expected column type".to_string());
1422 }
1423
1424 let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
1425 if end == 0 {
1426 return Err(format!("invalid type: '{}'", s));
1427 }
1428 let type_name = s[..end].to_string();
1429 let mut rest = &s[end..];
1430
1431 rest = rest.trim_start();
1433 if rest.starts_with('(') {
1434 let close = rest
1435 .find(')')
1436 .ok_or_else(|| "unterminated type parameter list".to_string())?;
1437 rest = &rest[close + 1..];
1438 }
1439
1440 Ok((type_name, rest))
1441}
1442
1443fn split_top_level_commas(s: &str) -> Vec<String> {
1445 let mut parts = Vec::new();
1446 let mut depth: i32 = 0;
1447 let mut current = String::new();
1448
1449 for ch in s.chars() {
1450 match ch {
1451 '(' => {
1452 depth += 1;
1453 current.push(ch);
1454 }
1455 ')' => {
1456 depth -= 1;
1457 current.push(ch);
1458 }
1459 ',' if depth == 0 => {
1460 parts.push(std::mem::take(&mut current));
1461 }
1462 _ => {
1463 current.push(ch);
1464 }
1465 }
1466 }
1467
1468 if !current.trim().is_empty() {
1469 parts.push(current);
1470 }
1471
1472 parts
1473}
1474
1475fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
1480 let upper = sql_type.to_uppercase();
1481 let rust = match upper.as_str() {
1482 "BIGINT" | "INT8" => "i64",
1484 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
1486 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
1488 "TINYINT" => "i8",
1490 "FLOAT" | "REAL" | "FLOAT4" => "f32",
1492 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
1494 "BOOLEAN" | "BOOL" => "bool",
1496 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
1498 "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
1500 | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
1501 _ => "String",
1502 };
1503
1504 if nullable {
1505 format!("Option<{}>", rust)
1506 } else {
1507 rust.to_string()
1508 }
1509}
1510
1511#[proc_macro_derive(Schema, attributes(table, column))]
1541pub fn derive_schema(input: TokenStream) -> TokenStream {
1542 let input = parse_macro_input!(input as syn::DeriveInput);
1543 derive::derive_schema_impl(input).into()
1544}
1545
1546#[proc_macro_derive(Builder, attributes(builder))]
1580pub fn derive_builder(input: TokenStream) -> TokenStream {
1581 let input = parse_macro_input!(input as syn::DeriveInput);
1582 derive::derive_builder_impl(input).into()
1583}
1584
1585#[cfg(test)]
1590mod tests {
1591 use super::*;
1592
1593 #[test]
1596 fn test_strip_plain_double_quoted() {
1597 assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
1598 }
1599
1600 #[test]
1601 fn test_strip_raw_double_hash() {
1602 assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
1603 }
1604
1605 #[test]
1606 fn test_strip_raw_double_no_hash() {
1607 assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
1608 }
1609
1610 #[test]
1611 fn test_strip_byte_string() {
1612 assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
1613 assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
1614 }
1615
1616 #[test]
1617 fn test_strip_non_string_returns_none() {
1618 assert_eq!(strip_string_literal("123"), None);
1619 assert_eq!(strip_string_literal("foo"), None);
1620 }
1621
1622 #[test]
1625 fn test_validate_select_with_from_ok() {
1626 assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
1627 }
1628
1629 #[test]
1630 fn test_validate_select_missing_from_fails() {
1631 assert!(validate_sql_content("SELECT * users", None).is_err());
1632 }
1633
1634 #[test]
1635 fn test_validate_insert_missing_into_fails() {
1636 assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
1637 assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
1638 }
1639
1640 #[test]
1641 fn test_validate_update_missing_set_fails() {
1642 assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
1643 assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
1644 }
1645
1646 #[test]
1647 fn test_validate_delete_missing_from_fails() {
1648 assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
1649 assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
1650 }
1651
1652 #[test]
1653 fn test_validate_empty_sql_fails() {
1654 assert!(validate_sql_content("", None).is_err());
1655 assert!(validate_sql_content(" ", None).is_err());
1656 }
1657
1658 #[test]
1661 fn test_validate_balanced_parens_ok() {
1662 assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
1663 }
1664
1665 #[test]
1666 fn test_validate_balanced_parens_unbalanced() {
1667 assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
1668 assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
1669 }
1670
1671 #[test]
1674 fn test_validate_no_injection_clean() {
1675 assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
1676 }
1677
1678 #[test]
1679 fn test_validate_no_injection_drop_table() {
1680 assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
1681 }
1682
1683 #[test]
1684 fn test_validate_no_injection_or_1_1() {
1685 assert!(validate_no_injection("' OR 1=1").is_err());
1689 assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
1690 }
1691
1692 #[test]
1693 fn test_validate_no_injection_drop_database() {
1694 assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
1695 }
1696
1697 #[test]
1698 fn test_validate_no_injection_information_schema() {
1699 assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
1700 }
1701
1702 #[test]
1703 fn test_validate_no_injection_xp_cmdshell() {
1704 assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
1705 }
1706
1707 #[test]
1708 fn test_validate_no_injection_union_select() {
1709 assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
1710 }
1711
1712 #[test]
1713 fn test_validate_no_injection_comment_dashes() {
1714 assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
1715 }
1716
1717 #[test]
1718 fn test_validate_no_injection_block_comment() {
1719 assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
1720 }
1721
1722 #[test]
1725 fn test_validate_string_literals_closed_ok() {
1726 assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
1727 assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
1728 }
1729
1730 #[test]
1731 fn test_validate_string_literals_closed_unclosed_single() {
1732 assert!(validate_string_literals_closed("'hello").is_err());
1733 }
1734
1735 #[test]
1736 fn test_validate_string_literals_closed_unclosed_double() {
1737 assert!(validate_string_literals_closed(r#""hello"#).is_err());
1738 }
1739
1740 #[test]
1743 fn test_validate_param_count_match() {
1744 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
1745 assert!(
1746 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
1747 );
1748 }
1749
1750 #[test]
1751 fn test_validate_param_count_mismatch() {
1752 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
1753 assert!(
1754 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
1755 );
1756 }
1757
1758 #[cfg(feature = "db-verify")]
1761 #[test]
1762 fn test_detect_db_kind_mysql() {
1763 assert_eq!(
1764 detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
1765 DbKind::MySql
1766 );
1767 }
1768
1769 #[cfg(feature = "db-verify")]
1770 #[test]
1771 fn test_detect_db_kind_postgres() {
1772 assert_eq!(
1773 detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
1774 DbKind::Postgres
1775 );
1776 assert_eq!(
1777 detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
1778 DbKind::Postgres
1779 );
1780 }
1781
1782 #[cfg(feature = "db-verify")]
1783 #[test]
1784 fn test_detect_db_kind_sqlite() {
1785 assert_eq!(
1786 detect_db_kind("sqlite://path/to/db.db").unwrap(),
1787 DbKind::Sqlite
1788 );
1789 assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
1790 }
1791
1792 #[cfg(feature = "db-verify")]
1793 #[test]
1794 fn test_detect_db_kind_oracle() {
1795 assert_eq!(
1796 detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
1797 DbKind::Oracle
1798 );
1799 assert_eq!(
1800 detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
1801 DbKind::Oracle
1802 );
1803 }
1804
1805 #[cfg(feature = "db-verify")]
1806 #[test]
1807 fn test_detect_db_kind_sqlserver() {
1808 assert_eq!(
1809 detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
1810 DbKind::SqlServer
1811 );
1812 assert_eq!(
1813 detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
1814 DbKind::SqlServer
1815 );
1816 assert_eq!(
1817 detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
1818 DbKind::SqlServer
1819 );
1820 }
1821
1822 #[cfg(feature = "db-verify")]
1823 #[test]
1824 fn test_detect_db_kind_unsupported() {
1825 assert!(detect_db_kind("redis://user:pass@host/db").is_err());
1826 assert!(detect_db_kind("not-a-url").is_err());
1827 }
1828
1829 #[cfg(feature = "db-verify")]
1830 #[test]
1831 fn test_parse_oracle_dsn_basic() {
1832 let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
1833 let p = parse_oracle_dsn(dsn).unwrap();
1834 assert_eq!(p.user, "sys");
1835 assert_eq!(p.password, "test123");
1836 assert_eq!(p.host, "127.0.0.1");
1837 assert_eq!(p.port, 1521);
1838 assert_eq!(p.service, "freepdb1.FALSE");
1839 assert!(p.sysdba);
1840 }
1841
1842 #[cfg(feature = "db-verify")]
1843 #[test]
1844 fn test_parse_oracle_dsn_default_port() {
1845 let dsn = "oracle://sys:test123@127.0.0.1/FREE";
1847 let p = parse_oracle_dsn(dsn).unwrap();
1848 assert_eq!(p.port, 1521);
1849 assert_eq!(p.service, "FREE");
1850 assert!(!p.sysdba);
1851 }
1852
1853 #[cfg(feature = "db-verify")]
1854 #[test]
1855 fn test_parse_sqlserver_dsn_basic() {
1856 let dsn = "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
1857 let p = parse_sqlserver_dsn(dsn).unwrap();
1858 assert_eq!(p.user, "test");
1859 assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
1860 assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
1861 assert_eq!(p.port, 22527);
1862 assert_eq!(p.database, "test");
1863 }
1864
1865 #[cfg(feature = "db-verify")]
1866 #[test]
1867 fn test_parse_sqlserver_dsn_default_port() {
1868 let dsn = "mssql://user:pass@host/db";
1869 let p = parse_sqlserver_dsn(dsn).unwrap();
1870 assert_eq!(p.port, 1433);
1871 assert_eq!(p.database, "db");
1872 }
1873
1874 #[test]
1877 fn test_parse_create_table_basic() {
1878 let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
1879 let (table, cols) = parse_create_table(sql).unwrap();
1880 assert_eq!(table, "users");
1881 assert_eq!(
1882 cols,
1883 vec![
1884 ("id".to_string(), "i32".to_string()),
1885 ("name".to_string(), "String".to_string())
1886 ]
1887 );
1888 }
1889
1890 #[test]
1891 fn test_parse_create_table_with_if_not_exists() {
1892 let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
1893 let (table, cols) = parse_create_table(sql).unwrap();
1894 assert_eq!(table, "orders");
1895 assert_eq!(
1896 cols,
1897 vec![
1898 ("id".to_string(), "i64".to_string()),
1899 ("total".to_string(), "f64".to_string())
1900 ]
1901 );
1902 }
1903
1904 #[test]
1905 fn test_parse_create_table_nullable() {
1906 let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
1907 let (_, cols) = parse_create_table(sql).unwrap();
1908 assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
1909 assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
1910 }
1911
1912 #[test]
1913 fn test_parse_create_table_skip_constraints() {
1914 let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
1915 let (_, cols) = parse_create_table(sql).unwrap();
1916 assert_eq!(cols.len(), 2);
1917 assert_eq!(cols[0].0, "id");
1918 assert_eq!(cols[1].0, "name");
1919 }
1920
1921 #[test]
1922 fn test_parse_create_table_varchar_with_len() {
1923 let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
1924 let (_, cols) = parse_create_table(sql).unwrap();
1925 assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
1926 assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
1927 }
1928
1929 #[test]
1930 fn test_sql_type_to_rust_mappings() {
1931 assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
1933 assert_eq!(sql_type_to_rust("INT8", false), "i64");
1934 assert_eq!(sql_type_to_rust("INT", false), "i32");
1935 assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
1936 assert_eq!(sql_type_to_rust("INT4", false), "i32");
1937 assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
1938 assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
1939 assert_eq!(sql_type_to_rust("INT2", false), "i16");
1940 assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
1941 assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
1942 assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
1944 assert_eq!(sql_type_to_rust("REAL", false), "f32");
1945 assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
1946 assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
1947 assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
1948 assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
1949 assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
1950 assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
1951 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
1953 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
1954 assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
1956 assert_eq!(sql_type_to_rust("TEXT", false), "String");
1957 assert_eq!(sql_type_to_rust("CHAR", false), "String");
1958 assert_eq!(sql_type_to_rust("UUID", false), "String");
1959 assert_eq!(sql_type_to_rust("DATE", false), "String");
1960 assert_eq!(sql_type_to_rust("DATETIME", false), "String");
1961 assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
1962 assert_eq!(sql_type_to_rust("JSON", false), "String");
1963 assert_eq!(sql_type_to_rust("JSONB", false), "String");
1964 assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
1966 assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
1967 assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
1968 assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
1969 assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
1971 assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
1972 assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
1973 assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
1974 assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
1976 }
1977
1978 #[test]
1979 fn test_parse_create_table_error_no_create() {
1980 assert!(parse_create_table("SELECT * FROM users").is_err());
1981 }
1982
1983 #[test]
1984 fn test_parse_create_table_error_no_parens() {
1985 assert!(parse_create_table("CREATE TABLE foo").is_err());
1986 }
1987}