#![allow(linker_messages)]
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;
use syn::parse_macro_input;
mod derive;
#[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 sql_no_placeholders = replace_placeholders_with_null(sql);
let explain_sql = match db_kind {
DbKind::MySql | DbKind::Postgres => format!("EXPLAIN {}", sql_no_placeholders),
DbKind::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql_no_placeholders),
DbKind::Oracle => format!("EXPLAIN PLAN FOR {}", sql_no_placeholders),
DbKind::SqlServer => sql_no_placeholders,
};
if matches!(db_kind, DbKind::MySql | DbKind::Postgres | DbKind::Sqlite) {
let rt = tokio::runtime::Runtime::new()
.map_err(|e| format!("Failed to create tokio runtime: {}", e))?;
return 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,
_ => unreachable!(),
}
});
}
match db_kind {
DbKind::Oracle => verify_oracle(&dsn, &explain_sql),
DbKind::SqlServer => verify_sqlserver(&dsn, &explain_sql),
_ => unreachable!(),
}
}
#[cfg(feature = "db-verify")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DbKind {
MySql,
Postgres,
Sqlite,
Oracle,
SqlServer,
}
#[cfg(feature = "db-verify")]
fn replace_placeholders_with_null(sql: &str) -> String {
let mut result = String::with_capacity(sql.len() + 16);
let mut in_single_quote = false;
let mut in_double_quote = false;
let mut prev = '\0';
for ch in sql.chars() {
if prev == '\\' {
result.push(ch);
prev = ch;
continue;
}
match ch {
'\'' if !in_double_quote => in_single_quote = !in_single_quote,
'"' if !in_single_quote => in_double_quote = !in_double_quote,
'?' if !in_single_quote && !in_double_quote => {
result.push_str("NULL");
prev = ch;
continue;
}
_ => {}
}
result.push(ch);
prev = ch;
}
result
}
#[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 if lower.starts_with("oracle://") || lower.starts_with("oracle:") {
Ok(DbKind::Oracle)
} else if lower.starts_with("sqlserver://")
|| lower.starts_with("mssql://")
|| lower.starts_with("tds://")
{
Ok(DbKind::SqlServer)
} 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(())
}
#[cfg(feature = "db-verify")]
fn verify_oracle(dsn: &str, explain_sql: &str) -> Result<(), String> {
let parsed = parse_oracle_dsn(dsn)?;
let mut conn_str = format!(
"{}/{}@{}:{}/{}",
parsed.user, parsed.password, parsed.host, parsed.port, parsed.service
);
if parsed.sysdba {
conn_str.push_str(" AS SYSDBA");
}
let full_script = format!(
"SET HEADING OFF FEEDBACK OFF ECHO OFF;\n\
EXPLAIN PLAN FOR {};\n\
SELECT COUNT(*) FROM plan_table WHERE statement_id = (SELECT MAX(statement_id) FROM plan_table);\n\
EXIT;\n",
explain_sql
);
let output = std::process::Command::new("sqlplus")
.args(["-S", "-L", &conn_str])
.stdin(std::process::Stdio::piped())
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.spawn()
.map_err(|e| format!("sqlplus not found (Oracle client required): {}", e))?;
use std::io::Write;
let mut child = output;
if let Some(mut stdin) = child.stdin.take() {
stdin
.write_all(full_script.as_bytes())
.map_err(|e| format!("sqlplus stdin write failed: {}", e))?;
}
let out = child
.wait_with_output()
.map_err(|e| format!("sqlplus wait failed: {}", e))?;
let stdout = String::from_utf8_lossy(&out.stdout);
let stderr = String::from_utf8_lossy(&out.stderr);
if !out.status.success() || stdout.contains("ORA-") || stdout.contains("SP2-") {
return Err(format!(
"Oracle EXPLAIN failed: stdout={} stderr={}",
stdout.trim(),
stderr.trim()
));
}
Ok(())
}
#[cfg(feature = "db-verify")]
fn verify_sqlserver(dsn: &str, explain_sql: &str) -> Result<(), String> {
let parsed = parse_sqlserver_dsn(dsn)?;
let query = format!("SET SHOWPLAN_TEXT ON;\n{}", explain_sql);
let out = std::process::Command::new("sqlcmd")
.args([
"-S",
&format!("{},{}", parsed.host, parsed.port),
"-U",
&parsed.user,
"-P",
&parsed.password,
"-d",
&parsed.database,
"-Q",
&query,
"-h",
"-1",
"-W",
])
.output()
.map_err(|e| format!("sqlcmd not found (SQL Server client required): {}", e))?;
let stdout = String::from_utf8_lossy(&out.stdout);
let stderr = String::from_utf8_lossy(&out.stderr);
if !out.status.success() || stdout.contains("Msg ") || stdout.contains("Level ") {
return Err(format!(
"SQL Server SHOWPLAN failed: stdout={} stderr={}",
stdout.trim(),
stderr.trim()
));
}
Ok(())
}
#[cfg(feature = "db-verify")]
struct OracleDsn {
user: String,
password: String,
host: String,
port: u16,
service: String,
sysdba: bool,
}
#[cfg(feature = "db-verify")]
fn parse_oracle_dsn(dsn: &str) -> Result<OracleDsn, String> {
let raw = dsn
.strip_prefix("oracle://")
.or_else(|| dsn.strip_prefix("oracle:"))
.ok_or_else(|| format!("Invalid Oracle DSN: {}", dsn))?;
let (auth_host_service, query) = match raw.find('?') {
Some(idx) => (&raw[..idx], &raw[idx + 1..]),
None => (raw, ""),
};
let sysdba = query
.split('&')
.any(|p| p == "sysdba=1" || p == "sysdba=true");
let at = auth_host_service
.find('@')
.ok_or_else(|| format!("Oracle DSN missing '@': {}", dsn))?;
let (user_pass, host_port_service) = (&auth_host_service[..at], &auth_host_service[at + 1..]);
let colon = user_pass
.find(':')
.ok_or_else(|| format!("Oracle DSN missing password separator: {}", dsn))?;
let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
let (host_port, service) = match host_port_service.rfind('/') {
Some(idx) => (&host_port_service[..idx], &host_port_service[idx + 1..]),
None => return Err(format!("Oracle DSN missing service name: {}", dsn)),
};
let (host, port) = match host_port.find(':') {
Some(idx) => (
&host_port[..idx],
host_port[idx + 1..]
.parse::<u16>()
.map_err(|_| format!("Oracle DSN invalid port: {}", dsn))?,
),
None => (host_port, 1521u16),
};
Ok(OracleDsn {
user: user.to_string(),
password: password.to_string(),
host: host.to_string(),
port,
service: service.to_string(),
sysdba,
})
}
#[cfg(feature = "db-verify")]
struct SqlServerDsn {
user: String,
password: String,
host: String,
port: u16,
database: String,
}
#[cfg(feature = "db-verify")]
fn parse_sqlserver_dsn(dsn: &str) -> Result<SqlServerDsn, String> {
let raw = dsn
.strip_prefix("sqlserver://")
.or_else(|| dsn.strip_prefix("mssql://"))
.or_else(|| dsn.strip_prefix("tds://"))
.ok_or_else(|| format!("Invalid SQL Server DSN: {}", dsn))?;
let at = raw
.find('@')
.ok_or_else(|| format!("SQL Server DSN missing '@': {}", dsn))?;
let (user_pass, host_port_db) = (&raw[..at], &raw[at + 1..]);
let colon = user_pass
.find(':')
.ok_or_else(|| format!("SQL Server DSN missing password separator: {}", dsn))?;
let (user, password) = (&user_pass[..colon], &user_pass[colon + 1..]);
let (host_port, database) = match host_port_db.rfind('/') {
Some(idx) => (&host_port_db[..idx], &host_port_db[idx + 1..]),
None => return Err(format!("SQL Server DSN missing database: {}", dsn)),
};
let (host, port) = match host_port.find(':') {
Some(idx) => (
&host_port[..idx],
host_port[idx + 1..]
.parse::<u16>()
.map_err(|_| format!("SQL Server DSN invalid port: {}", dsn))?,
),
None => (host_port, 1433u16),
};
Ok(SqlServerDsn {
user: user.to_string(),
password: password.to_string(),
host: host.to_string(),
port,
database: database.to_string(),
})
}
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 = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
}
}
})
.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
}
#[proc_macro]
pub fn schema(input: TokenStream) -> TokenStream {
let mut tokens = input.into_iter().peekable();
let sql_raw = match tokens.next() {
Some(TokenTree::Literal(lit)) => lit.to_string(),
Some(other) => {
return compile_error(
other.span(),
"Expected a string literal as the argument to schema!",
);
}
None => {
return compile_error(
Span::call_site(),
"Expected a string literal argument to schema!",
);
}
};
let sql = match strip_string_literal(&sql_raw) {
Some(s) => s,
None => {
return compile_error(
Span::call_site(),
"schema! requires a string literal argument",
);
}
};
let (table_name, columns) = match parse_create_table(sql) {
Ok(v) => v,
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 = <#rust_type as ::sz_orm_core::typed_ast::InferSqlType>::SqlType;
}
}
})
.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_create_table(sql: &str) -> Result<(String, Vec<(String, String)>), String> {
let trimmed = sql.trim();
let upper = trimmed.to_uppercase();
if !upper.starts_with("CREATE TABLE") {
return Err("schema! expects a CREATE TABLE statement".to_string());
}
let mut rest = &trimmed["CREATE TABLE".len()..];
let rest_upper = rest.trim_start().to_uppercase();
if rest_upper.starts_with("IF NOT EXISTS") {
rest = &rest.trim_start()["IF NOT EXISTS".len()..];
}
rest = rest.trim_start();
let (table_name, after_name) = parse_identifier(rest)?;
let rest = after_name.trim_start();
let paren_start = rest
.find('(')
.ok_or_else(|| "CREATE TABLE missing '(' for column definitions".to_string())?;
let paren_end = rest
.rfind(')')
.ok_or_else(|| "CREATE TABLE missing ')' for column definitions".to_string())?;
if paren_end <= paren_start {
return Err("CREATE TABLE has malformed parentheses".to_string());
}
let cols_str = &rest[paren_start + 1..paren_end];
let col_defs = split_top_level_commas(cols_str);
let mut columns = Vec::new();
for def in col_defs {
let def = def.trim();
if def.is_empty() {
continue;
}
let def_upper = def.to_uppercase();
if def_upper.starts_with("PRIMARY KEY")
|| def_upper.starts_with("FOREIGN KEY")
|| def_upper.starts_with("CONSTRAINT")
|| def_upper.starts_with("UNIQUE")
|| def_upper.starts_with("INDEX")
|| def_upper.starts_with("KEY")
{
continue;
}
let (col_name, after_col) = parse_identifier(def)?;
let rest = after_col.trim_start();
let (sql_type, after_type) = parse_type_token(rest)?;
let rest = after_type.trim();
let rest_upper = rest.to_uppercase();
let not_null = rest_upper.contains("NOT NULL") || rest_upper.contains("PRIMARY KEY");
let rust_type = sql_type_to_rust(&sql_type, !not_null);
columns.push((col_name, rust_type));
}
Ok((table_name, columns))
}
fn parse_identifier(s: &str) -> Result<(String, &str), String> {
let s = s.trim_start();
if s.is_empty() {
return Err("expected identifier".to_string());
}
let bytes = s.as_bytes();
match bytes[0] {
b'`' => {
let end = s[1..]
.find('`')
.ok_or_else(|| "unterminated backtick-quoted identifier".to_string())?;
let ident = s[1..1 + end].to_string();
Ok((ident, &s[1 + end + 1..]))
}
b'"' => {
let end = s[1..]
.find('"')
.ok_or_else(|| "unterminated double-quoted identifier".to_string())?;
let ident = s[1..1 + end].to_string();
Ok((ident, &s[1 + end + 1..]))
}
_ => {
let end = s
.find(|c: char| !c.is_alphanumeric() && c != '_')
.unwrap_or(s.len());
if end == 0 {
return Err(format!("invalid identifier: '{}'", s));
}
let ident = s[..end].to_string();
Ok((ident, &s[end..]))
}
}
}
fn parse_type_token(s: &str) -> Result<(String, &str), String> {
let s = s.trim_start();
if s.is_empty() {
return Err("expected column type".to_string());
}
let end = s.find(|c: char| !c.is_alphabetic()).unwrap_or(s.len());
if end == 0 {
return Err(format!("invalid type: '{}'", s));
}
let type_name = s[..end].to_string();
let mut rest = &s[end..];
rest = rest.trim_start();
if rest.starts_with('(') {
let close = rest
.find(')')
.ok_or_else(|| "unterminated type parameter list".to_string())?;
rest = &rest[close + 1..];
}
Ok((type_name, rest))
}
fn split_top_level_commas(s: &str) -> Vec<String> {
let mut parts = Vec::new();
let mut depth: i32 = 0;
let mut current = String::new();
for ch in s.chars() {
match ch {
'(' => {
depth += 1;
current.push(ch);
}
')' => {
depth -= 1;
current.push(ch);
}
',' if depth == 0 => {
parts.push(std::mem::take(&mut current));
}
_ => {
current.push(ch);
}
}
}
if !current.trim().is_empty() {
parts.push(current);
}
parts
}
fn sql_type_to_rust(sql_type: &str, nullable: bool) -> String {
let upper = sql_type.to_uppercase();
let rust = match upper.as_str() {
"BIGINT" | "INT8" => "i64",
"INT" | "INTEGER" | "INT4" | "SERIAL" => "i32",
"SMALLINT" | "INT2" | "SMALLSERIAL" => "i16",
"TINYINT" => "i8",
"FLOAT" | "REAL" | "FLOAT4" => "f32",
"DOUBLE" | "DOUBLE PRECISION" | "FLOAT8" | "DECIMAL" | "NUMERIC" => "f64",
"BOOLEAN" | "BOOL" => "bool",
"BLOB" | "BYTEA" | "BINARY" | "VARBINARY" => "Vec<u8>",
"VARCHAR" | "TEXT" | "CHAR" | "CHARACTER" | "CLOB" | "UUID" | "DATE" | "TIME"
| "DATETIME" | "TIMESTAMP" | "JSON" | "JSONB" => "String",
_ => "String",
};
if nullable {
format!("Option<{}>", rust)
} else {
rust.to_string()
}
}
#[proc_macro_derive(Schema, attributes(table, column))]
pub fn derive_schema(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as syn::DeriveInput);
derive::derive_schema_impl(input).into()
}
#[proc_macro_derive(Builder, attributes(builder))]
pub fn derive_builder(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as syn::DeriveInput);
derive::derive_builder_impl(input).into()
}
#[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_oracle() {
assert_eq!(
detect_db_kind("oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1").unwrap(),
DbKind::Oracle
);
assert_eq!(
detect_db_kind("oracle:sys:test123@127.0.0.1:1521/FREE").unwrap(),
DbKind::Oracle
);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_sqlserver() {
assert_eq!(
detect_db_kind("sqlserver://test:pass@host:1433/db").unwrap(),
DbKind::SqlServer
);
assert_eq!(
detect_db_kind("mssql://test:pass@host:1433/db").unwrap(),
DbKind::SqlServer
);
assert_eq!(
detect_db_kind("tds://test:pass@host:1433/db").unwrap(),
DbKind::SqlServer
);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_detect_db_kind_unsupported() {
assert!(detect_db_kind("redis://user:pass@host/db").is_err());
assert!(detect_db_kind("not-a-url").is_err());
}
#[cfg(feature = "db-verify")]
#[test]
fn test_parse_oracle_dsn_basic() {
let dsn = "oracle://sys:test123@127.0.0.1:1521/freepdb1.FALSE?sysdba=1";
let p = parse_oracle_dsn(dsn).unwrap();
assert_eq!(p.user, "sys");
assert_eq!(p.password, "test123");
assert_eq!(p.host, "127.0.0.1");
assert_eq!(p.port, 1521);
assert_eq!(p.service, "freepdb1.FALSE");
assert!(p.sysdba);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_parse_oracle_dsn_default_port() {
let dsn = "oracle://sys:test123@127.0.0.1/FREE";
let p = parse_oracle_dsn(dsn).unwrap();
assert_eq!(p.port, 1521);
assert_eq!(p.service, "FREE");
assert!(!p.sysdba);
}
#[cfg(feature = "db-verify")]
#[test]
fn test_parse_sqlserver_dsn_basic() {
let dsn =
"sqlserver://test:JkbC2jsaWAYDe2Gz@sh-mssql-adrul9nm.sql.tencentcdb.com:22527/test";
let p = parse_sqlserver_dsn(dsn).unwrap();
assert_eq!(p.user, "test");
assert_eq!(p.password, "JkbC2jsaWAYDe2Gz");
assert_eq!(p.host, "sh-mssql-adrul9nm.sql.tencentcdb.com");
assert_eq!(p.port, 22527);
assert_eq!(p.database, "test");
}
#[cfg(feature = "db-verify")]
#[test]
fn test_parse_sqlserver_dsn_default_port() {
let dsn = "mssql://user:pass@host/db";
let p = parse_sqlserver_dsn(dsn).unwrap();
assert_eq!(p.port, 1433);
assert_eq!(p.database, "db");
}
#[test]
fn test_parse_create_table_basic() {
let sql = "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)";
let (table, cols) = parse_create_table(sql).unwrap();
assert_eq!(table, "users");
assert_eq!(
cols,
vec![
("id".to_string(), "i32".to_string()),
("name".to_string(), "String".to_string())
]
);
}
#[test]
fn test_parse_create_table_with_if_not_exists() {
let sql = "CREATE TABLE IF NOT EXISTS `orders` (`id` BIGINT PRIMARY KEY, `total` DECIMAL(10,2) NOT NULL)";
let (table, cols) = parse_create_table(sql).unwrap();
assert_eq!(table, "orders");
assert_eq!(
cols,
vec![
("id".to_string(), "i64".to_string()),
("total".to_string(), "f64".to_string())
]
);
}
#[test]
fn test_parse_create_table_nullable() {
let sql = "CREATE TABLE t (a INT NOT NULL, b INT)";
let (_, cols) = parse_create_table(sql).unwrap();
assert_eq!(cols[0], ("a".to_string(), "i32".to_string()));
assert_eq!(cols[1], ("b".to_string(), "Option<i32>".to_string()));
}
#[test]
fn test_parse_create_table_skip_constraints() {
let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT, PRIMARY KEY (id), CONSTRAINT fk1 FOREIGN KEY (x) REFERENCES y(id))";
let (_, cols) = parse_create_table(sql).unwrap();
assert_eq!(cols.len(), 2);
assert_eq!(cols[0].0, "id");
assert_eq!(cols[1].0, "name");
}
#[test]
fn test_parse_create_table_varchar_with_len() {
let sql = "CREATE TABLE t (name VARCHAR(255) NOT NULL, code CHAR(10))";
let (_, cols) = parse_create_table(sql).unwrap();
assert_eq!(cols[0], ("name".to_string(), "String".to_string()));
assert_eq!(cols[1], ("code".to_string(), "Option<String>".to_string()));
}
#[test]
fn test_sql_type_to_rust_mappings() {
assert_eq!(sql_type_to_rust("BIGINT", false), "i64");
assert_eq!(sql_type_to_rust("INT8", false), "i64");
assert_eq!(sql_type_to_rust("INT", false), "i32");
assert_eq!(sql_type_to_rust("INTEGER", false), "i32");
assert_eq!(sql_type_to_rust("INT4", false), "i32");
assert_eq!(sql_type_to_rust("SERIAL", false), "i32");
assert_eq!(sql_type_to_rust("SMALLINT", false), "i16");
assert_eq!(sql_type_to_rust("INT2", false), "i16");
assert_eq!(sql_type_to_rust("SMALLSERIAL", false), "i16");
assert_eq!(sql_type_to_rust("TINYINT", false), "i8");
assert_eq!(sql_type_to_rust("FLOAT", false), "f32");
assert_eq!(sql_type_to_rust("REAL", false), "f32");
assert_eq!(sql_type_to_rust("FLOAT4", false), "f32");
assert_eq!(sql_type_to_rust("DOUBLE", false), "f64");
assert_eq!(sql_type_to_rust("DOUBLE PRECISION", false), "f64");
assert_eq!(sql_type_to_rust("FLOAT8", false), "f64");
assert_eq!(sql_type_to_rust("DECIMAL", false), "f64");
assert_eq!(sql_type_to_rust("NUMERIC", false), "f64");
assert_eq!(sql_type_to_rust("BOOLEAN", false), "bool");
assert_eq!(sql_type_to_rust("BOOL", false), "bool");
assert_eq!(sql_type_to_rust("VARCHAR", false), "String");
assert_eq!(sql_type_to_rust("TEXT", false), "String");
assert_eq!(sql_type_to_rust("CHAR", false), "String");
assert_eq!(sql_type_to_rust("UUID", false), "String");
assert_eq!(sql_type_to_rust("DATE", false), "String");
assert_eq!(sql_type_to_rust("DATETIME", false), "String");
assert_eq!(sql_type_to_rust("TIMESTAMP", false), "String");
assert_eq!(sql_type_to_rust("JSON", false), "String");
assert_eq!(sql_type_to_rust("JSONB", false), "String");
assert_eq!(sql_type_to_rust("BLOB", false), "Vec<u8>");
assert_eq!(sql_type_to_rust("BYTEA", false), "Vec<u8>");
assert_eq!(sql_type_to_rust("BINARY", false), "Vec<u8>");
assert_eq!(sql_type_to_rust("VARBINARY", false), "Vec<u8>");
assert_eq!(sql_type_to_rust("INT", true), "Option<i32>");
assert_eq!(sql_type_to_rust("BIGINT", true), "Option<i64>");
assert_eq!(sql_type_to_rust("VARCHAR", true), "Option<String>");
assert_eq!(sql_type_to_rust("BLOB", true), "Option<Vec<u8>>");
assert_eq!(sql_type_to_rust("UNKNOWNTYPE", false), "String");
}
#[test]
fn test_parse_create_table_error_no_create() {
assert!(parse_create_table("SELECT * FROM users").is_err());
}
#[test]
fn test_parse_create_table_error_no_parens() {
assert!(parse_create_table("CREATE TABLE foo").is_err());
}
}