extern crate proc_macro;
use proc_macro::{Delimiter, Group, Ident, Literal, Punct, Spacing, Span, TokenStream, TokenTree};
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
#[proc_macro]
pub fn sql_string(input: TokenStream) -> TokenStream {
let mut tokens = input.into_iter().peekable();
let sql = match tokens.next() {
Some(TokenTree::Literal(lit)) => lit.to_string(),
Some(other) => {
return compile_error(
other.span(),
"Expected a string literal as the first argument to sql_string!",
);
}
None => {
return compile_error(
Span::call_site(),
"Expected a string literal argument to sql_string!",
);
}
};
let sql_content = if sql.starts_with("r#\"") {
&sql[3..sql.len() - 2]
} else if sql.starts_with("r\"") {
&sql[2..sql.len() - 1]
} else if sql.starts_with('"') {
&sql[1..sql.len() - 1]
} else if sql.starts_with("b\"") || sql.starts_with("b\'") {
&sql[2..sql.len() - 1]
} else {
return compile_error(
Span::call_site(),
"sql_string! requires a string literal argument",
);
};
let mut expected_params = None;
if tokens.peek().is_some() {
match tokens.next() {
Some(TokenTree::Punct(p)) if p.as_char() == ';' => {}
Some(other) => {
return compile_error(
other.span(),
"Expected `;` before param count, e.g. sql_string!(\"...\"; params: 2)",
);
}
None => {}
}
match tokens.next() {
Some(TokenTree::Ident(id)) if id.to_string() == "params" => {}
Some(other) => {
return compile_error(
other.span(),
"Expected `params:` keyword, e.g. sql_string!(\"...\"; params: 2)",
);
}
None => {
return compile_error(Span::call_site(), "Expected param count after `;`");
}
}
match tokens.next() {
Some(TokenTree::Punct(p)) if p.as_char() == ':' => {}
Some(other) => {
return compile_error(
other.span(),
"Expected `:` after `params`, e.g. sql_string!(\"...\"; params: 2)",
);
}
None => {
return compile_error(Span::call_site(), "Expected param count after `params`");
}
}
match tokens.next() {
Some(TokenTree::Literal(lit)) => {
let num_str = lit.to_string();
if let Ok(n) = num_str.parse::<usize>() {
expected_params = Some(n);
} else {
return compile_error(
lit.span(),
"Expected a positive integer for param count",
);
}
}
Some(other) => {
return compile_error(
other.span(),
"Expected a number after `params:`, e.g. sql_string!(\"...\"; params: 2)",
);
}
None => {
return compile_error(Span::call_site(), "Expected a number after `params:`");
}
}
}
if let Err(err_msg) = validate_sql_content(sql_content, expected_params) {
return compile_error(Span::call_site(), &err_msg);
}
let output = format!("\"{}\"", sql_content.escape_default());
output
.parse()
.unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
}
fn validate_sql_content(sql: &str, expected_params: Option<usize>) -> Result<(), String> {
let trimmed = sql.trim();
if trimmed.is_empty() {
return Err("SQL statement is empty".to_string());
}
validate_balanced_parens(trimmed)?;
validate_string_literals_closed(trimmed)?;
validate_no_injection(trimmed)?;
let sql_upper = trimmed.to_uppercase();
if sql_upper.starts_with("SELECT") {
if !sql_upper.contains("FROM") {
return Err("SELECT statement missing FROM clause".to_string());
}
} else if sql_upper.starts_with("INSERT") {
if !sql_upper.contains("INTO") {
return Err("INSERT statement missing INTO clause".to_string());
}
if !sql_upper.contains("VALUES") {
return Err("INSERT statement missing VALUES clause".to_string());
}
} else if sql_upper.starts_with("UPDATE") {
if !sql_upper.contains("SET") {
return Err("UPDATE statement missing SET clause".to_string());
}
} else if sql_upper.starts_with("DELETE") && !sql_upper.contains("FROM") {
return Err("DELETE statement missing FROM clause".to_string());
}
if let Some(expected) = expected_params {
let actual = sql.chars().filter(|&c| c == '?').count();
if actual != expected {
return Err(format!(
"Parameter count mismatch: expected {} parameters, found {}",
expected, actual
));
}
}
Ok(())
}
fn validate_balanced_parens(sql: &str) -> Result<(), String> {
let mut depth: i32 = 0;
for (i, ch) in sql.char_indices() {
match ch {
'(' => depth += 1,
')' => {
depth -= 1;
if depth < 0 {
return Err(format!(
"Unbalanced parentheses: unexpected ')' at position {}",
i
));
}
}
_ => {}
}
}
if depth != 0 {
return Err(format!("Unbalanced parentheses: {} unclosed '('", depth));
}
Ok(())
}
fn validate_string_literals_closed(sql: &str) -> Result<(), String> {
let mut in_single = false;
let mut in_double = false;
let mut prev = '\0';
for ch in sql.chars() {
if prev == '\\' {
prev = ch;
continue;
}
match ch {
'\'' if !in_double => in_single = !in_single,
'"' if !in_single => in_double = !in_double,
_ => {}
}
prev = ch;
}
if in_single {
return Err("Unclosed single-quoted string literal".to_string());
}
if in_double {
return Err("Unclosed double-quoted string literal".to_string());
}
Ok(())
}
fn validate_no_injection(sql: &str) -> Result<(), String> {
let sql_lower = sql.to_lowercase();
let injection_patterns: &[&str] = &[
"drop table",
"drop database",
"; drop",
"or 1=1",
"or 1 = 1",
"union select",
"union all select",
"--",
"/*",
"*/",
"xp_cmdshell",
"sp_executesql",
"exec(",
"execute(",
"information_schema",
"sys.tables",
"sys.columns",
];
for pattern in injection_patterns {
if sql_lower.contains(pattern) {
return Err(format!("潜在的 SQL 注入模式被检测到: '{}'", pattern));
}
}
Ok(())
}
#[proc_macro]
pub fn query(input: TokenStream) -> TokenStream {
let mut tokens = input.into_iter().peekable();
let sql = match tokens.next() {
Some(TokenTree::Literal(lit)) => lit.to_string(),
Some(other) => {
return compile_error(
other.span(),
"Expected a string literal as the first argument to query!",
);
}
None => {
return compile_error(
Span::call_site(),
"Expected a string literal argument to query!",
);
}
};
let sql_content = match strip_string_literal(&sql) {
Some(s) => s,
None => {
return compile_error(
Span::call_site(),
"query! requires a string literal argument",
);
}
};
if let Err(err_msg) = validate_sql_content(sql_content, None) {
return compile_error(Span::call_site(), &err_msg);
}
#[cfg(feature = "db-verify")]
{
if std::env::var("SZ_ORM_QUERY_VERIFY").ok().as_deref() == Some("1") {
if let Err(err) = verify_with_real_db(sql_content) {
return compile_error(
Span::call_site(),
&format!("query! real DB verification failed: {}", err),
);
}
}
}
let output = format!("\"{}\"", sql_content.escape_default());
output
.parse()
.unwrap_or_else(|_| compile_error(Span::call_site(), "Failed to generate output token"))
}
fn strip_string_literal(raw: &str) -> Option<&str> {
if raw.starts_with("r#\"") {
Some(&raw[3..raw.len() - 2])
} else if raw.starts_with("r\"") {
Some(&raw[2..raw.len() - 1])
} else if raw.starts_with('"') {
Some(&raw[1..raw.len() - 1])
} else if raw.starts_with("b\"") || raw.starts_with("b\'") {
Some(&raw[2..raw.len() - 1])
} else {
None
}
}
#[cfg(feature = "db-verify")]
fn verify_with_real_db(sql: &str) -> Result<(), String> {
let dsn = std::env::var("DATABASE_URL")
.map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
let db_kind =
detect_db_kind(&dsn).map_err(|e| format!("Failed to detect DB kind from DSN: {}", e))?;
let explain_sql = match db_kind {
DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql),
DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
};
let rt = tokio::runtime::Runtime::new()
.map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
rt.block_on(async {
match db_kind {
DbKind::MySql => verify_mysql(&dsn, &explain_sql).await,
DbKind::Postgres => verify_postgres(&dsn, &explain_sql).await,
DbKind::Sqlite => verify_sqlite(&dsn, &explain_sql).await,
}
})
}
#[cfg(feature = "db-verify")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DbKind {
MySql,
Postgres,
Sqlite,
}
#[cfg(feature = "db-verify")]
fn detect_db_kind(dsn: &str) -> Result<DbKind, String> {
let lower = dsn.to_lowercase();
if lower.starts_with("mysql://") {
Ok(DbKind::MySql)
} else if lower.starts_with("postgres://") || lower.starts_with("postgresql://") {
Ok(DbKind::Postgres)
} else if lower.starts_with("sqlite://") || lower.starts_with("sqlite:") {
Ok(DbKind::Sqlite)
} else {
Err(format!("Unsupported DSN scheme: {}", dsn))
}
}
#[cfg(feature = "db-verify")]
async fn verify_mysql(dsn: &str, explain_sql: &str) -> Result<(), String> {
let pool = sqlx::MySqlPool::connect(dsn)
.await
.map_err(|e| format!("MySQL connect failed: {}", e))?;
sqlx::query(sqlx::AssertSqlSafe(explain_sql))
.execute(&pool)
.await
.map_err(|e| format!("MySQL EXPLAIN failed: {}", e))?;
Ok(())
}
#[cfg(feature = "db-verify")]
async fn verify_postgres(dsn: &str, explain_sql: &str) -> Result<(), String> {
let pool = sqlx::PgPool::connect(dsn)
.await
.map_err(|e| format!("PostgreSQL connect failed: {}", e))?;
sqlx::query(sqlx::AssertSqlSafe(explain_sql))
.execute(&pool)
.await
.map_err(|e| format!("PostgreSQL EXPLAIN failed: {}", e))?;
Ok(())
}
#[cfg(feature = "db-verify")]
async fn verify_sqlite(dsn: &str, explain_sql: &str) -> Result<(), String> {
let pool = sqlx::SqlitePool::connect(dsn)
.await
.map_err(|e| format!("SQLite connect failed: {}", e))?;
sqlx::query(sqlx::AssertSqlSafe(explain_sql))
.execute(&pool)
.await
.map_err(|e| format!("SQLite EXPLAIN failed: {}", e))?;
Ok(())
}
fn compile_error(span: Span, msg: &str) -> TokenStream {
let mut ts = TokenStream::new();
ts.extend([
TokenTree::Ident(Ident::new("compile_error", span)),
TokenTree::Punct(Punct::new('!', Spacing::Alone)),
TokenTree::Group(Group::new(
Delimiter::Parenthesis,
TokenStream::from(TokenTree::Literal(Literal::string(msg))),
)),
]);
ts
}
#[proc_macro]
pub fn typed_query(input: TokenStream) -> TokenStream {
let tokens: Vec<TokenTree> = input.into_iter().collect();
if tokens.iter().any(|t| {
if let TokenTree::Ident(id) = t {
id.to_string() == "table"
} else {
false
}
}) {
return parse_table_decl(&tokens);
}
if tokens.iter().any(|t| {
if let TokenTree::Ident(id) = t {
id.to_string().eq_ignore_ascii_case("SELECT")
} else {
false
}
}) {
return parse_typed_select(&tokens);
}
compile_error(
Span::call_site(),
"typed_query! expects either `table name { ... }` declaration or `SELECT ... FROM ...` expression",
)
}
fn parse_table_decl(tokens: &[TokenTree]) -> TokenStream {
let mut idx = 0;
if idx >= tokens.len() {
return compile_error(Span::call_site(), "expected table name after 'table'");
}
if let TokenTree::Ident(id) = &tokens[idx] {
if id.to_string() != "table" {
return compile_error(id.span(), "expected 'table' keyword");
}
}
idx += 1;
let table_name = if idx < tokens.len() {
if let TokenTree::Ident(id) = &tokens[idx] {
id.to_string()
} else {
return compile_error(tokens[idx].span(), "expected table name identifier");
}
} else {
return compile_error(Span::call_site(), "expected table name");
};
idx += 1;
let body_group = if idx < tokens.len() {
if let TokenTree::Group(g) = &tokens[idx] {
if g.delimiter() != Delimiter::Brace {
return compile_error(g.span(), "expected '{' after table name");
}
g.clone()
} else {
return compile_error(tokens[idx].span(), "expected '{' after table name");
}
} else {
return compile_error(Span::call_site(), "expected table body in '{ }'");
};
let body_tokens: Vec<TokenTree> = body_group.stream().into_iter().collect();
let columns = match parse_column_list(&body_tokens) {
Ok(c) => c,
Err(e) => return compile_error(Span::call_site(), &e),
};
let table_ident = proc_macro2::Ident::new(&table_name, Span::call_site().into());
let table_name_lit = table_name.as_str();
let col_impls: Vec<TokenStream2> = columns
.iter()
.map(|(col_name, col_type)| {
let col_ident =
proc_macro2::Ident::new(&format!("col_{}", col_name), Span::call_site().into());
let col_name_lit = col_name.as_str();
let rust_type: TokenStream2 = col_type.parse().unwrap_or_else(|_| quote! { () });
quote! {
#[derive(Debug, Clone, Copy)]
pub struct #col_ident;
impl ::sz_orm_core::typed::TypedColumn for #col_ident {
const NAME: &'static str = #col_name_lit;
type Table = table;
type RustType = #rust_type;
type SqlType = ::sz_orm_core::typed_ast::Untyped;
}
}
})
.collect();
let schema_entries: Vec<TokenStream2> = columns
.iter()
.map(|(n, t)| {
let n_lit = n.as_str();
let t_lit = t.as_str();
quote! { (#n_lit, #t_lit) }
})
.collect();
let schema_const_ident = proc_macro2::Ident::new(
&format!("__SZ_ORM_TYPED_SCHEMA_{}", table_name.to_uppercase()),
Span::call_site().into(),
);
let expanded = quote! {
pub mod #table_ident {
use super::*;
pub struct table;
impl ::sz_orm_core::typed::TypedTable for table {
const NAME: &'static str = #table_name_lit;
}
#(#col_impls)*
}
const #schema_const_ident: &[(&str, &str)] = &[#(#schema_entries),*];
};
expanded.into()
}
fn parse_column_list(tokens: &[TokenTree]) -> Result<Vec<(String, String)>, String> {
let mut cols = Vec::new();
let mut i = 0;
while i < tokens.len() {
let col_name = if let TokenTree::Ident(id) = &tokens[i] {
id.to_string()
} else {
return Err(format!("expected column name at position {}", i));
};
i += 1;
if i >= tokens.len() {
return Err(format!("expected ':' after column '{}'", col_name));
}
if let TokenTree::Punct(p) = &tokens[i] {
if p.as_char() != ':' {
return Err(format!("expected ':' after column '{}'", col_name));
}
} else {
return Err(format!("expected ':' after column '{}'", col_name));
}
i += 1;
let mut type_str = String::new();
let mut depth = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Punct(p) => {
if p.as_char() == ',' && depth == 0 {
i += 1;
break;
} else if p.as_char() == '<' || p.as_char() == '(' {
depth += 1;
type_str.push(p.as_char());
} else if p.as_char() == '>' || p.as_char() == ')' {
depth -= 1;
type_str.push(p.as_char());
} else {
type_str.push(p.as_char());
}
}
TokenTree::Ident(id) => {
if !type_str.is_empty() && !type_str.ends_with('<') && !type_str.ends_with('(')
{
type_str.push(' ');
}
type_str.push_str(&id.to_string());
}
_ => {}
}
i += 1;
}
cols.push((col_name, type_str.trim().to_string()));
}
Ok(cols)
}
fn parse_typed_select(tokens: &[TokenTree]) -> TokenStream {
let mut sql_parts: Vec<String> = Vec::new();
let mut table_name: Option<String> = None;
let mut in_from = false;
for (i, t) in tokens.iter().enumerate() {
match t {
TokenTree::Ident(id) => {
let s = id.to_string();
if s.eq_ignore_ascii_case("SELECT") {
sql_parts.push("SELECT".to_string());
} else if s.eq_ignore_ascii_case("FROM") {
in_from = true;
sql_parts.push("FROM".to_string());
} else if s.eq_ignore_ascii_case("WHERE")
|| s.eq_ignore_ascii_case("AND")
|| s.eq_ignore_ascii_case("OR")
|| s.eq_ignore_ascii_case("LIMIT")
|| s.eq_ignore_ascii_case("OFFSET")
|| s.eq_ignore_ascii_case("ORDER")
|| s.eq_ignore_ascii_case("BY")
|| s.eq_ignore_ascii_case("GROUP")
|| s.eq_ignore_ascii_case("HAVING")
|| s.eq_ignore_ascii_case("JOIN")
|| s.eq_ignore_ascii_case("INNER")
|| s.eq_ignore_ascii_case("LEFT")
|| s.eq_ignore_ascii_case("RIGHT")
|| s.eq_ignore_ascii_case("ON")
|| s.eq_ignore_ascii_case("AS")
|| s.eq_ignore_ascii_case("ASC")
|| s.eq_ignore_ascii_case("DESC")
|| s.eq_ignore_ascii_case("DISTINCT")
|| s.eq_ignore_ascii_case("NOT")
|| s.eq_ignore_ascii_case("NULL")
|| s.eq_ignore_ascii_case("IN")
|| s.eq_ignore_ascii_case("BETWEEN")
|| s.eq_ignore_ascii_case("LIKE")
|| s.eq_ignore_ascii_case("IS")
{
sql_parts.push(s.to_uppercase());
} else if in_from && table_name.is_none() {
table_name = Some(s.clone());
sql_parts.push(s.clone());
} else {
sql_parts.push(s.clone());
}
}
TokenTree::Literal(lit) => {
sql_parts.push(lit.to_string());
}
TokenTree::Punct(p) => {
let c = p.as_char();
let part = if c == ',' {
",".to_string()
} else if c == '?' {
"?".to_string()
} else if c == '*' {
"*".to_string()
} else if c == '=' {
"=".to_string()
} else if c == '>' {
">".to_string()
} else if c == '<' {
"<".to_string()
} else if c == '.' {
".".to_string()
} else if c == ';' {
";".to_string()
} else {
c.to_string()
};
sql_parts.push(part);
}
TokenTree::Group(g) => {
let inner: String = g.stream().to_string();
let delim = match g.delimiter() {
Delimiter::Parenthesis => "(",
Delimiter::Brace => "{",
Delimiter::Bracket => "[",
Delimiter::None => "",
};
let close = match g.delimiter() {
Delimiter::Parenthesis => ")",
Delimiter::Brace => "}",
Delimiter::Bracket => "]",
Delimiter::None => "",
};
sql_parts.push(format!("{}{}{}", delim, inner, close));
}
}
let _ = i;
}
let sql = sql_parts
.join(" ")
.replace(", ", ",")
.replace(" ,", ",")
.replace("= ", "=")
.replace(" =", "=")
.replace("> ", ">")
.replace(" >", ">")
.replace("< ", "<")
.replace(" <", "<")
.replace(" ", " ");
if let Err(e) = validate_sql_content(&sql, None) {
return compile_error(
Span::call_site(),
&format!("typed_query! SQL validation failed: {}", e),
);
}
let mut ts = TokenStream::new();
let lit = Literal::string(&sql);
ts.extend([TokenTree::Literal(lit)]);
ts
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_strip_plain_double_quoted() {
assert_eq!(strip_string_literal(r#""hello""#), Some("hello"));
}
#[test]
fn test_strip_raw_double_hash() {
assert_eq!(strip_string_literal(r###"r#"hello"#"###), Some("hello"));
}
#[test]
fn test_strip_raw_double_no_hash() {
assert_eq!(strip_string_literal(r#"r"hello""#), Some("hello"));
}
#[test]
fn test_strip_byte_string() {
assert_eq!(strip_string_literal(r#"b"hello""#), Some("hello"));
assert_eq!(strip_string_literal(r#"b'hello'"#), Some("hello"));
}
#[test]
fn test_strip_non_string_returns_none() {
assert_eq!(strip_string_literal("123"), None);
assert_eq!(strip_string_literal("foo"), None);
}
#[test]
fn test_validate_select_with_from_ok() {
assert!(validate_sql_content("SELECT * FROM users", None).is_ok());
}
#[test]
fn test_validate_select_missing_from_fails() {
assert!(validate_sql_content("SELECT * users", None).is_err());
}
#[test]
fn test_validate_insert_missing_into_fails() {
assert!(validate_sql_content("INSERT INTO users VALUES (1)", None).is_ok());
assert!(validate_sql_content("INSERT users VALUES (1)", None).is_err());
}
#[test]
fn test_validate_update_missing_set_fails() {
assert!(validate_sql_content("UPDATE users SET name='a'", None).is_ok());
assert!(validate_sql_content("UPDATE users name='a'", None).is_err());
}
#[test]
fn test_validate_delete_missing_from_fails() {
assert!(validate_sql_content("DELETE FROM users WHERE id=1", None).is_ok());
assert!(validate_sql_content("DELETE users WHERE id=1", None).is_err());
}
#[test]
fn test_validate_empty_sql_fails() {
assert!(validate_sql_content("", None).is_err());
assert!(validate_sql_content(" ", None).is_err());
}
#[test]
fn test_validate_balanced_parens_ok() {
assert!(validate_balanced_parens("SELECT * FROM (SELECT * FROM t)").is_ok());
}
#[test]
fn test_validate_balanced_parens_unbalanced() {
assert!(validate_balanced_parens("SELECT * FROM (t").is_err());
assert!(validate_balanced_parens("SELECT * FROM t)").is_err());
}
#[test]
fn test_validate_no_injection_clean() {
assert!(validate_no_injection("SELECT * FROM users WHERE id = 1").is_ok());
}
#[test]
fn test_validate_no_injection_drop_table() {
assert!(validate_no_injection("'; DROP TABLE users; --").is_err());
}
#[test]
fn test_validate_no_injection_or_1_1() {
assert!(validate_no_injection("' OR 1=1").is_err());
assert!(validate_no_injection("WHERE id = 1 OR 1=1").is_err());
}
#[test]
fn test_validate_no_injection_drop_database() {
assert!(validate_no_injection("SELECT x; DROP DATABASE db").is_err());
}
#[test]
fn test_validate_no_injection_information_schema() {
assert!(validate_no_injection("SELECT * FROM information_schema.tables").is_err());
}
#[test]
fn test_validate_no_injection_xp_cmdshell() {
assert!(validate_no_injection("EXEC xp_cmdshell 'dir'").is_err());
}
#[test]
fn test_validate_no_injection_union_select() {
assert!(validate_no_injection("1 UNION SELECT * FROM users").is_err());
}
#[test]
fn test_validate_no_injection_comment_dashes() {
assert!(validate_no_injection("SELECT * FROM users -- comment").is_err());
}
#[test]
fn test_validate_no_injection_block_comment() {
assert!(validate_no_injection("SELECT /* x */ * FROM users").is_err());
}
#[test]
fn test_validate_string_literals_closed_ok() {
assert!(validate_string_literals_closed("'hello' = 'world'").is_ok());
assert!(validate_string_literals_closed(r#""foo" = "bar""#).is_ok());
}
#[test]
fn test_validate_string_literals_closed_unclosed_single() {
assert!(validate_string_literals_closed("'hello").is_err());
}
#[test]
fn test_validate_string_literals_closed_unclosed_double() {
assert!(validate_string_literals_closed(r#""hello"#).is_err());
}
#[test]
fn test_validate_param_count_match() {
assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(1)).is_ok());
assert!(
validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(2)).is_ok()
);
}
#[test]
fn test_validate_param_count_mismatch() {
assert!(validate_sql_content("SELECT * FROM users WHERE id = ?", Some(2)).is_err());
assert!(
validate_sql_content("SELECT * FROM users WHERE id = ? AND name = ?", Some(1)).is_err()
);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_mysql() {
assert_eq!(
detect_db_kind("mysql://user:pass@host:3306/db").unwrap(),
DbKind::MySql
);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_postgres() {
assert_eq!(
detect_db_kind("postgres://user:pass@host:5432/db").unwrap(),
DbKind::Postgres
);
assert_eq!(
detect_db_kind("postgresql://user:pass@host:5432/db").unwrap(),
DbKind::Postgres
);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_sqlite() {
assert_eq!(
detect_db_kind("sqlite://path/to/db.db").unwrap(),
DbKind::Sqlite
);
assert_eq!(detect_db_kind("sqlite::memory:").unwrap(), DbKind::Sqlite);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_unsupported() {
assert!(detect_db_kind("oracle://user:pass@host/db").is_err());
assert!(detect_db_kind("not-a-url").is_err());
}
}