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
711 .split('&')
712 .any(|p| p == "sysdba=1" || p == "sysdba=true");
713 let at = auth_host_service
715 .find('@')
716 .ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
717 let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
718 let colon = user_pass
719 .find(':')
720 .ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
721 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
722 let (host_port, service) = match host_port_service.rfind('/') {
723 Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
724 None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
725 };
726 let (host, port) = match host_port.find(':') {
727 Some(idx) => (
728 &host_port[..idx],
729 host_port[idx + 1..]
730 .parse::<u16>()
731 .map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
732 ),
733 None => (host_port, 1521u16),
734 };
735 Ok(OracleDsn {
736 user: user.to_string(),
737 password: password.to_string(),
738 host: host.to_string(),
739 port,
740 service: service.to_string(),
741 sysdba,
742 })
743}
744
745#[cfg(feature = "db-verify")]
747struct SqlServerDsn {
748 user: String,
749 password: String,
750 host: String,
751 port: u16,
752 database: String,
753}
754
755#[cfg(feature = "db-verify")]
757fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
758 let raw = dsn
759 .strip_prefix("sqlserver://")
760 .or_else(|| dsn.strip_prefix("mssql://"))
761 .or_else(|| dsn.strip_prefix("tds://"))
762 .ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
763 let at = raw
764 .find('@')
765 .ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
766 let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
767 let colon = user_pass
768 .find(':')
769 .ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
770 let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
771 let (host_port, database) = match host_port_db.rfind('/') {
772 Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
773 None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
774 };
775 let (host, port) = match host_port.find(':') {
776 Some(idx) => (
777 &host_port[..idx],
778 host_port[idx + 1..]
779 .parse::<u16>()
780 .map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
781 ),
782 None => (host_port, 1433u16),
783 };
784 Ok(SqlServerDsn {
785 user: user.to_string(),
786 password: password.to_string(),
787 host: host.to_string(),
788 port,
789 database: database.to_string(),
790 })
791}
792
793fn compile_error(span: Span, msg: &str) -> TokenStream {
799 let mut ts = TokenStream::new();
801 ts.extend([
802 TokenTree::Ident(Ident::new("compile_error", span)),
803 TokenTree::Punct(Punct::new('!', Spacing::Alone)),
804 TokenTree::Group(Group::new(
805 Delimiter::Parenthesis,
806 TokenStream::from(TokenTree::Literal(Literal::string(msg))),
807 )),
808 ]);
809 ts
810}
811
812#[proc_macro]
850pub fn typed_query(input: TokenStream) -> TokenStream {
851 let tokens: Vec<TokenTree> = input.into_iter().collect();
852
853 if tokens.iter().any(|t| {
855 if let TokenTree::Ident(id) = t {
856 id.to_string() == "table"
857 } else {
858 false
859 }
860 }) {
861 return parse_table_decl(&tokens);
862 }
863
864 if tokens.iter().any(|t| {
866 if let TokenTree::Ident(id) = t {
867 id.to_string().eq_ignore_ascii_case("SELECT")
868 } else {
869 false
870 }
871 }) {
872 return parse_typed_select(&tokens);
873 }
874
875 compile_error(
876 Span::call_site(),
877 "typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
878 )
879}
880
881fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
883 let mut idx = 0;
885
886 if idx >= tokens.len() {
888 return compile_error(Span::call_site(), "expected table name after 'table'");
889 }
890 if let TokenTree::Ident(id) = &tokens[idx] {
891 if id.to_string() != "table" {
892 return compile_error(id.span(), "expected 'table' keyword");
893 }
894 }
895 idx += 1;
896
897 let table_name = if idx < tokens.len() {
899 if let TokenTree::Ident(id) = &tokens[idx] {
900 id.to_string()
901 } else {
902 return compile_error(tokens[idx].span(), "expected table name identifier");
903 }
904 } else {
905 return compile_error(Span::call_site(), "expected table name");
906 };
907 idx += 1;
908
909 let body_group = if idx < tokens.len() {
911 if let TokenTree::Group(g) = &tokens[idx] {
912 if g.delimiter() != Delimiter::Brace {
913 return compile_error(g.span(), "expected '{' after table name");
914 }
915 g.clone()
916 } else {
917 return compile_error(tokens[idx].span(), "expected '{' after table name");
918 }
919 } else {
920 return compile_error(Span::call_site(), "expected table body in '{ }'");
921 };
922
923 let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
925 let columns = match parse_column_list(&body_tokens) {
926 Ok(c) => c,
927 Err(e) => return compile_error(Span::call_site(), &e),
928 };
929
930 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
932 let table_name_lit = table_name.as_str();
933
934 let col_impls: Vec<TokenStream2> = columns
936 .iter()
937 .map(|(col_name, col_type)| {
938 let col_ident =
939 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
940 let col_name_lit = col_name.as_str();
941 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
943 quote! {
944 #[derive(Debug, Clone, Copy)]
945 pub struct #col_ident;
946 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
947 const NAME: &'static str = #col_name_lit;
948 type Table = table;
949 type RustType = #rust_type;
950 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
951 }
952 }
953 })
954 .collect();
955
956 let schema_entries: Vec<TokenStream2> = columns
958 .iter()
959 .map(|(n, t)| {
960 let n_lit = n.as_str();
961 let t_lit = t.as_str();
962 quote! { (#n_lit, #t_lit) }
963 })
964 .collect();
965
966 let schema_const_ident = proc_macro2::Ident::new(
967 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
968 Span::call_site().into(),
969 );
970
971 let expanded = quote! {
972 pub mod #table_ident {
973 use super::*;
974 pub struct table;
975 impl ::sz_orm_core::typed::TypedTable for table {
976 const NAME: &'static str = #table_name_lit;
977 }
978 #(#col_impls)*
979 }
980 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
981 };
982
983 expanded.into()
984}
985
986fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
988 let mut cols = Vec::new();
989 let mut i = 0;
990 while i < tokens.len() {
991 let col_name = if let TokenTree::Ident(id) = &tokens[i] {
993 id.to_string()
994 } else {
995 return Err(format!("expected column name at position {}", i));
996 };
997 i += 1;
998
999 if i >= tokens.len() {
1001 return Err(format!("expected ':' after column '{}'", col_name));
1002 }
1003 if let TokenTree::Punct(p) = &tokens[i] {
1004 if p.as_char() != ':' {
1005 return Err(format!("expected ':' after column '{}'", col_name));
1006 }
1007 } else {
1008 return Err(format!("expected ':' after column '{}'", col_name));
1009 }
1010 i += 1;
1011
1012 let mut type_str = String::new();
1015 let mut depth = 0;
1016 while i < tokens.len() {
1017 match &tokens[i] {
1018 TokenTree::Punct(p) => {
1019 if p.as_char() == ',' && depth == 0 {
1020 i += 1;
1021 break;
1022 } else if p.as_char() == '<' || p.as_char() == '(' {
1023 depth += 1;
1024 type_str.push(p.as_char());
1025 } else if p.as_char() == '>' || p.as_char() == ')' {
1026 depth -= 1;
1027 type_str.push(p.as_char());
1028 } else {
1029 type_str.push(p.as_char());
1030 }
1031 }
1032 TokenTree::Ident(id) => {
1033 if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
1034 {
1035 type_str.push(' ');
1036 }
1037 type_str.push_str(&id.to_string());
1038 }
1039 _ => {}
1040 }
1041 i += 1;
1042 }
1043
1044 cols.push((col_name, type_str.trim().to_string()));
1045 }
1046 Ok(cols)
1047}
1048
1049fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
1053 let mut sql_parts: Vec<String> = Vec::new();
1055 let mut table_name: Option<String> = None;
1056 let mut in_from = false;
1057
1058 for (i, t) in tokens.iter().enumerate() {
1059 match t {
1060 TokenTree::Ident(id) => {
1061 let s = id.to_string();
1062 if s.eq_ignore_ascii_case("SELECT") {
1063 sql_parts.push("SELECT".to_string());
1064 } else if s.eq_ignore_ascii_case("FROM") {
1065 in_from = true;
1066 sql_parts.push("FROM".to_string());
1067 } else if s.eq_ignore_ascii_case("WHERE")
1068 || s.eq_ignore_ascii_case("AND")
1069 || s.eq_ignore_ascii_case("OR")
1070 || s.eq_ignore_ascii_case("LIMIT")
1071 || s.eq_ignore_ascii_case("OFFSET")
1072 || s.eq_ignore_ascii_case("ORDER")
1073 || s.eq_ignore_ascii_case("BY")
1074 || s.eq_ignore_ascii_case("GROUP")
1075 || s.eq_ignore_ascii_case("HAVING")
1076 || s.eq_ignore_ascii_case("JOIN")
1077 || s.eq_ignore_ascii_case("INNER")
1078 || s.eq_ignore_ascii_case("LEFT")
1079 || s.eq_ignore_ascii_case("RIGHT")
1080 || s.eq_ignore_ascii_case("ON")
1081 || s.eq_ignore_ascii_case("AS")
1082 || s.eq_ignore_ascii_case("ASC")
1083 || s.eq_ignore_ascii_case("DESC")
1084 || s.eq_ignore_ascii_case("DISTINCT")
1085 || s.eq_ignore_ascii_case("NOT")
1086 || s.eq_ignore_ascii_case("NULL")
1087 || s.eq_ignore_ascii_case("IN")
1088 || s.eq_ignore_ascii_case("BETWEEN")
1089 || s.eq_ignore_ascii_case("LIKE")
1090 || s.eq_ignore_ascii_case("IS")
1091 {
1092 sql_parts.push(s.to_uppercase());
1093 } else if in_from && table_name.is_none() {
1094 table_name = Some(s.clone());
1096 sql_parts.push(s.clone());
1097 } else {
1098 sql_parts.push(s.clone());
1099 }
1100 }
1101 TokenTree::Literal(lit) => {
1102 sql_parts.push(lit.to_string());
1103 }
1104 TokenTree::Punct(p) => {
1105 let c = p.as_char();
1106 let part = 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 if c == ';' {
1122 ";".to_string()
1123 } else {
1124 c.to_string()
1125 };
1126 sql_parts.push(part);
1127 }
1128 TokenTree::Group(g) => {
1129 let inner: String = g.stream().to_string();
1131 let delim = match g.delimiter() {
1132 Delimiter::Parenthesis => "(",
1133 Delimiter::Brace => "{",
1134 Delimiter::Bracket => "[",
1135 Delimiter::None => "",
1136 };
1137 let close = match g.delimiter() {
1138 Delimiter::Parenthesis => ")",
1139 Delimiter::Brace => "}",
1140 Delimiter::Bracket => "]",
1141 Delimiter::None => "",
1142 };
1143 sql_parts.push(format!("{}{}{}", delim, inner, close));
1144 }
1145 }
1146 let _ = i;
1148 }
1149
1150 let sql = sql_parts
1151 .join(" ")
1152 .replace(", ", ",")
1153 .replace(" ,", ",")
1154 .replace("= ", "=")
1155 .replace(" =", "=")
1156 .replace("> ", ">")
1157 .replace(" >", ">")
1158 .replace("< ", "<")
1159 .replace(" <", "<")
1160 .replace(" ", " ");
1161
1162 if let Err(e) = validate_sql_content(&sql, None) {
1164 return compile_error(
1165 Span::call_site(),
1166 &format!("typed_query! SQL validation failed: {}", e),
1167 );
1168 }
1169
1170 let mut ts = TokenStream::new();
1172 let lit = Literal::string(&sql);
1173 ts.extend([TokenTree::Literal(lit)]);
1174 ts
1175}
1176
1177#[proc_macro]
1207pub fn schema(input: TokenStream) -> TokenStream {
1208 let mut tokens = input.into_iter().peekable();
1209
1210 let sql_raw = match tokens.next() {
1212 Some(TokenTree::Literal(lit)) => lit.to_string(),
1213 Some(other) => {
1214 return compile_error(
1215 other.span(),
1216 "Expected a string literal as the argument to schema!",
1217 );
1218 }
1219 None => {
1220 return compile_error(
1221 Span::call_site(),
1222 "Expected a string literal argument to schema!",
1223 );
1224 }
1225 };
1226
1227 let sql = match strip_string_literal(&sql_raw) {
1228 Some(s) => s,
1229 None => {
1230 return compile_error(
1231 Span::call_site(),
1232 "schema! requires a string literal argument",
1233 );
1234 }
1235 };
1236
1237 let (table_name, columns) = match parse_create_table(sql) {
1239 Ok(v) => v,
1240 Err(e) => return compile_error(Span::call_site(), &e),
1241 };
1242
1243 let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
1245 let table_name_lit = table_name.as_str();
1246
1247 let col_impls: Vec<TokenStream2> = columns
1248 .iter()
1249 .map(|(col_name, col_type)| {
1250 let col_ident =
1251 proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
1252 let col_name_lit = col_name.as_str();
1253 let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
1254 quote! {
1255 #[derive(Debug, Clone, Copy)]
1256 pub struct #col_ident;
1257 impl ::sz_orm_core::typed::TypedColumn for #col_ident {
1258 const NAME: &'static str = #col_name_lit;
1259 type Table = table;
1260 type RustType = #rust_type;
1261 type SqlType = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
1262 }
1263 }
1264 })
1265 .collect();
1266
1267 let schema_entries: Vec<TokenStream2> = columns
1268 .iter()
1269 .map(|(n, t)| {
1270 let n_lit = n.as_str();
1271 let t_lit = t.as_str();
1272 quote! { (#n_lit, #t_lit) }
1273 })
1274 .collect();
1275
1276 let schema_const_ident = proc_macro2::Ident::new(
1277 &format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
1278 Span::call_site().into(),
1279 );
1280
1281 let expanded = quote! {
1282 pub mod #table_ident {
1283 use super::*;
1284 pub struct table;
1285 impl ::sz_orm_core::typed::TypedTable for table {
1286 const NAME: &'static str = #table_name_lit;
1287 }
1288 #(#col_impls)*
1289 }
1290 const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
1291 };
1292
1293 expanded.into()
1294}
1295
1296fn parse_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
1304 let trimmed = sql.trim();
1305 let upper = trimmed.to_uppercase();
1306
1307 if !upper.starts_with("CREATE TABLE") {
1309 return Err("schema! expects a CREATE TABLE statement".to_string());
1310 }
1311
1312 let mut rest = &trimmed["CREATE TABLE".len()..];
1314
1315 let rest_upper = rest.trim_start().to_uppercase();
1317 if rest_upper.starts_with("IF NOT EXISTS") {
1318 rest = &rest.trim_start()["IF NOT EXISTS".len()..];
1319 }
1320
1321 rest = rest.trim_start();
1322
1323 let (table_name, after_name) = parse_identifier(rest)?;
1325 let rest = after_name.trim_start();
1326
1327 let paren_start = rest
1329 .find('(')
1330 .ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
1331 let paren_end = rest
1332 .rfind(')')
1333 .ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
1334 if paren_end <= paren_start {
1335 return Err("CREATE TABLE has malformed parentheses".to_string());
1336 }
1337
1338 let cols_str = &rest[paren_start + 1..paren_end];
1339
1340 let col_defs = split_top_level_commas(cols_str);
1342
1343 let mut columns = Vec::new();
1344 for def in col_defs {
1345 let def = def.trim();
1346 if def.is_empty() {
1347 continue;
1348 }
1349
1350 let def_upper = def.to_uppercase();
1352 if def_upper.starts_with("PRIMARY KEY")
1353 || def_upper.starts_with("FOREIGN KEY")
1354 || def_upper.starts_with("CONSTRAINT")
1355 || def_upper.starts_with("UNIQUE")
1356 || def_upper.starts_with("INDEX")
1357 || def_upper.starts_with("KEY")
1358 {
1359 continue;
1360 }
1361
1362 let (col_name, after_col) = parse_identifier(def)?;
1364 let rest = after_col.trim_start();
1365
1366 let (sql_type, after_type) = parse_type_token(rest)?;
1368 let rest = after_type.trim();
1369
1370 let rest_upper = rest.to_uppercase();
1372 let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
1373 let rust_type = sql_type_to_rust(&sql_type, !not_null);
1374
1375 columns.push((col_name, rust_type));
1376 }
1377
1378 Ok((table_name, columns))
1379}
1380
1381fn parse_identifier(s: &str) -> Result<(String, &str), String> {
1384 let s = s.trim_start();
1385 if s.is_empty() {
1386 return Err("expected identifier".to_string());
1387 }
1388
1389 let bytes = s.as_bytes();
1390 match bytes[0] {
1391 b'`' => {
1392 let end = s[1..]
1393 .find('`')
1394 .ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
1395 let ident = s[1..1 + end].to_string();
1396 Ok((ident, &s[1 + end + 1..]))
1397 }
1398 b'"' => {
1399 let end = s[1..]
1400 .find('"')
1401 .ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
1402 let ident = s[1..1 + end].to_string();
1403 Ok((ident, &s[1 + end + 1..]))
1404 }
1405 _ => {
1406 let end = s
1407 .find(|c: char| !c.is_alphanumeric() && c != '_')
1408 .unwrap_or(s.len());
1409 if end == 0 {
1410 return Err(format!("invalid identifier: '{}'", s));
1411 }
1412 let ident = s[..end].to_string();
1413 Ok((ident, &s[end..]))
1414 }
1415 }
1416}
1417
1418fn parse_type_token(s: &str) -> Result<(String, &str), String> {
1421 let s = s.trim_start();
1422 if s.is_empty() {
1423 return Err("expected column type".to_string());
1424 }
1425
1426 let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
1427 if end == 0 {
1428 return Err(format!("invalid type: '{}'", s));
1429 }
1430 let type_name = s[..end].to_string();
1431 let mut rest = &s[end..];
1432
1433 rest = rest.trim_start();
1435 if rest.starts_with('(') {
1436 let close = rest
1437 .find(')')
1438 .ok_or_else(|| "unterminated type parameter list".to_string())?;
1439 rest = &rest[close + 1..];
1440 }
1441
1442 Ok((type_name, rest))
1443}
1444
1445fn split_top_level_commas(s: &str) -> Vec<String> {
1447 let mut parts = Vec::new();
1448 let mut depth: i32 = 0;
1449 let mut current = String::new();
1450
1451 for ch in s.chars() {
1452 match ch {
1453 '(' => {
1454 depth += 1;
1455 current.push(ch);
1456 }
1457 ')' => {
1458 depth -= 1;
1459 current.push(ch);
1460 }
1461 ',' if depth == 0 => {
1462 parts.push(std::mem::take(&mut current));
1463 }
1464 _ => {
1465 current.push(ch);
1466 }
1467 }
1468 }
1469
1470 if !current.trim().is_empty() {
1471 parts.push(current);
1472 }
1473
1474 parts
1475}
1476
1477fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
1482 let upper = sql_type.to_uppercase();
1483 let rust = match upper.as_str() {
1484 "BIGINT" | "INT8" => "i64",
1486 "INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
1488 "SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
1490 "TINYINT" => "i8",
1492 "FLOAT" | "REAL" | "FLOAT4" => "f32",
1494 "DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
1496 "BOOLEAN" | "BOOL" => "bool",
1498 "BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
1500 "VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
1502 | "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
1503 _ => "String",
1504 };
1505
1506 if nullable {
1507 format!("Option<{}>", rust)
1508 } else {
1509 rust.to_string()
1510 }
1511}
1512
1513#[proc_macro_derive(Schema, attributes(table, column))]
1543pub fn derive_schema(input: TokenStream) -> TokenStream {
1544 let input = parse_macro_input!(input as syn::DeriveInput);
1545 derive::derive_schema_impl(input).into()
1546}
1547
1548#[proc_macro_derive(Builder, attributes(builder))]
1582pub fn derive_builder(input: TokenStream) -> TokenStream {
1583 let input = parse_macro_input!(input as syn::DeriveInput);
1584 derive::derive_builder_impl(input).into()
1585}
1586
1587#[cfg(test)]
1592mod tests {
1593 use super::*;
1594
1595 #[test]
1598 fn test_strip_plain_double_quoted() {
1599 assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
1600 }
1601
1602 #[test]
1603 fn test_strip_raw_double_hash() {
1604 assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
1605 }
1606
1607 #[test]
1608 fn test_strip_raw_double_no_hash() {
1609 assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
1610 }
1611
1612 #[test]
1613 fn test_strip_byte_string() {
1614 assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
1615 assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
1616 }
1617
1618 #[test]
1619 fn test_strip_non_string_returns_none() {
1620 assert_eq!(strip_string_literal("123"), None);
1621 assert_eq!(strip_string_literal("foo"), None);
1622 }
1623
1624 #[test]
1627 fn test_validate_select_with_from_ok() {
1628 assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
1629 }
1630
1631 #[test]
1632 fn test_validate_select_missing_from_fails() {
1633 assert!(validate_sql_content("SELECT * users", None).is_err());
1634 }
1635
1636 #[test]
1637 fn test_validate_insert_missing_into_fails() {
1638 assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
1639 assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
1640 }
1641
1642 #[test]
1643 fn test_validate_update_missing_set_fails() {
1644 assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
1645 assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
1646 }
1647
1648 #[test]
1649 fn test_validate_delete_missing_from_fails() {
1650 assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
1651 assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
1652 }
1653
1654 #[test]
1655 fn test_validate_empty_sql_fails() {
1656 assert!(validate_sql_content("", None).is_err());
1657 assert!(validate_sql_content(" ", None).is_err());
1658 }
1659
1660 #[test]
1663 fn test_validate_balanced_parens_ok() {
1664 assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
1665 }
1666
1667 #[test]
1668 fn test_validate_balanced_parens_unbalanced() {
1669 assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
1670 assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
1671 }
1672
1673 #[test]
1676 fn test_validate_no_injection_clean() {
1677 assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
1678 }
1679
1680 #[test]
1681 fn test_validate_no_injection_drop_table() {
1682 assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
1683 }
1684
1685 #[test]
1686 fn test_validate_no_injection_or_1_1() {
1687 assert!(validate_no_injection("' OR 1=1").is_err());
1691 assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
1692 }
1693
1694 #[test]
1695 fn test_validate_no_injection_drop_database() {
1696 assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
1697 }
1698
1699 #[test]
1700 fn test_validate_no_injection_information_schema() {
1701 assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
1702 }
1703
1704 #[test]
1705 fn test_validate_no_injection_xp_cmdshell() {
1706 assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
1707 }
1708
1709 #[test]
1710 fn test_validate_no_injection_union_select() {
1711 assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
1712 }
1713
1714 #[test]
1715 fn test_validate_no_injection_comment_dashes() {
1716 assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
1717 }
1718
1719 #[test]
1720 fn test_validate_no_injection_block_comment() {
1721 assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
1722 }
1723
1724 #[test]
1727 fn test_validate_string_literals_closed_ok() {
1728 assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
1729 assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
1730 }
1731
1732 #[test]
1733 fn test_validate_string_literals_closed_unclosed_single() {
1734 assert!(validate_string_literals_closed("'hello").is_err());
1735 }
1736
1737 #[test]
1738 fn test_validate_string_literals_closed_unclosed_double() {
1739 assert!(validate_string_literals_closed(r#""hello"#).is_err());
1740 }
1741
1742 #[test]
1745 fn test_validate_param_count_match() {
1746 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
1747 assert!(
1748 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
1749 );
1750 }
1751
1752 #[test]
1753 fn test_validate_param_count_mismatch() {
1754 assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
1755 assert!(
1756 validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
1757 );
1758 }
1759
1760 #[cfg(feature = "db-verify")]
1763 #[test]
1764 fn test_detect_db_kind_mysql() {
1765 assert_eq!(
1766 detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
1767 DbKind::MySql
1768 );
1769 }
1770
1771 #[cfg(feature = "db-verify")]
1772 #[test]
1773 fn test_detect_db_kind_postgres() {
1774 assert_eq!(
1775 detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
1776 DbKind::Postgres
1777 );
1778 assert_eq!(
1779 detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
1780 DbKind::Postgres
1781 );
1782 }
1783
1784 #[cfg(feature = "db-verify")]
1785 #[test]
1786 fn test_detect_db_kind_sqlite() {
1787 assert_eq!(
1788 detect_db_kind("sqlite://path/to/db.db").unwrap(),
1789 DbKind::Sqlite
1790 );
1791 assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
1792 }
1793
1794 #[cfg(feature = "db-verify")]
1795 #[test]
1796 fn test_detect_db_kind_oracle() {
1797 assert_eq!(
1798 detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
1799 DbKind::Oracle
1800 );
1801 assert_eq!(
1802 detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
1803 DbKind::Oracle
1804 );
1805 }
1806
1807 #[cfg(feature = "db-verify")]
1808 #[test]
1809 fn test_detect_db_kind_sqlserver() {
1810 assert_eq!(
1811 detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
1812 DbKind::SqlServer
1813 );
1814 assert_eq!(
1815 detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
1816 DbKind::SqlServer
1817 );
1818 assert_eq!(
1819 detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
1820 DbKind::SqlServer
1821 );
1822 }
1823
1824 #[cfg(feature = "db-verify")]
1825 #[test]
1826 fn test_detect_db_kind_unsupported() {
1827 assert!(detect_db_kind("redis://user:pass@host/db").is_err());
1828 assert!(detect_db_kind("not-a-url").is_err());
1829 }
1830
1831 #[cfg(feature = "db-verify")]
1832 #[test]
1833 fn test_parse_oracle_dsn_basic() {
1834 let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
1835 let p = parse_oracle_dsn(dsn).unwrap();
1836 assert_eq!(p.user, "sys");
1837 assert_eq!(p.password, "test123");
1838 assert_eq!(p.host, "127.0.0.1");
1839 assert_eq!(p.port, 1521);
1840 assert_eq!(p.service, "freepdb1.FALSE");
1841 assert!(p.sysdba);
1842 }
1843
1844 #[cfg(feature = "db-verify")]
1845 #[test]
1846 fn test_parse_oracle_dsn_default_port() {
1847 let dsn = "oracle://sys:test123@127.0.0.1/FREE";
1849 let p = parse_oracle_dsn(dsn).unwrap();
1850 assert_eq!(p.port, 1521);
1851 assert_eq!(p.service, "FREE");
1852 assert!(!p.sysdba);
1853 }
1854
1855 #[cfg(feature = "db-verify")]
1856 #[test]
1857 fn test_parse_sqlserver_dsn_basic() {
1858 let dsn =
1859 "sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
1860 let p = parse_sqlserver_dsn(dsn).unwrap();
1861 assert_eq!(p.user, "test");
1862 assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
1863 assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
1864 assert_eq!(p.port, 22527);
1865 assert_eq!(p.database, "test");
1866 }
1867
1868 #[cfg(feature = "db-verify")]
1869 #[test]
1870 fn test_parse_sqlserver_dsn_default_port() {
1871 let dsn = "mssql://user:pass@host/db";
1872 let p = parse_sqlserver_dsn(dsn).unwrap();
1873 assert_eq!(p.port, 1433);
1874 assert_eq!(p.database, "db");
1875 }
1876
1877 #[test]
1880 fn test_parse_create_table_basic() {
1881 let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
1882 let (table, cols) = parse_create_table(sql).unwrap();
1883 assert_eq!(table, "users");
1884 assert_eq!(
1885 cols,
1886 vec![
1887 ("id".to_string(), "i32".to_string()),
1888 ("name".to_string(), "String".to_string())
1889 ]
1890 );
1891 }
1892
1893 #[test]
1894 fn test_parse_create_table_with_if_not_exists() {
1895 let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
1896 let (table, cols) = parse_create_table(sql).unwrap();
1897 assert_eq!(table, "orders");
1898 assert_eq!(
1899 cols,
1900 vec![
1901 ("id".to_string(), "i64".to_string()),
1902 ("total".to_string(), "f64".to_string())
1903 ]
1904 );
1905 }
1906
1907 #[test]
1908 fn test_parse_create_table_nullable() {
1909 let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
1910 let (_, cols) = parse_create_table(sql).unwrap();
1911 assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
1912 assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
1913 }
1914
1915 #[test]
1916 fn test_parse_create_table_skip_constraints() {
1917 let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
1918 let (_, cols) = parse_create_table(sql).unwrap();
1919 assert_eq!(cols.len(), 2);
1920 assert_eq!(cols[0].0, "id");
1921 assert_eq!(cols[1].0, "name");
1922 }
1923
1924 #[test]
1925 fn test_parse_create_table_varchar_with_len() {
1926 let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
1927 let (_, cols) = parse_create_table(sql).unwrap();
1928 assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
1929 assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
1930 }
1931
1932 #[test]
1933 fn test_sql_type_to_rust_mappings() {
1934 assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
1936 assert_eq!(sql_type_to_rust("INT8", false), "i64");
1937 assert_eq!(sql_type_to_rust("INT", false), "i32");
1938 assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
1939 assert_eq!(sql_type_to_rust("INT4", false), "i32");
1940 assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
1941 assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
1942 assert_eq!(sql_type_to_rust("INT2", false), "i16");
1943 assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
1944 assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
1945 assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
1947 assert_eq!(sql_type_to_rust("REAL", false), "f32");
1948 assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
1949 assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
1950 assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
1951 assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
1952 assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
1953 assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
1954 assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
1956 assert_eq!(sql_type_to_rust("BOOL", false), "bool");
1957 assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
1959 assert_eq!(sql_type_to_rust("TEXT", false), "String");
1960 assert_eq!(sql_type_to_rust("CHAR", false), "String");
1961 assert_eq!(sql_type_to_rust("UUID", false), "String");
1962 assert_eq!(sql_type_to_rust("DATE", false), "String");
1963 assert_eq!(sql_type_to_rust("DATETIME", false), "String");
1964 assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
1965 assert_eq!(sql_type_to_rust("JSON", false), "String");
1966 assert_eq!(sql_type_to_rust("JSONB", false), "String");
1967 assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
1969 assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
1970 assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
1971 assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
1972 assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
1974 assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
1975 assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
1976 assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
1977 assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
1979 }
1980
1981 #[test]
1982 fn test_parse_create_table_error_no_create() {
1983 assert!(parse_create_table("SELECT * FROM users").is_err());
1984 }
1985
1986 #[test]
1987 fn test_parse_create_table_error_no_parens() {
1988 assert!(parse_create_table("CREATE TABLE foo").is_err());
1989 }
1990}