#![forbid(unsafe_code)]
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use std::time::SystemTime;
use ariadne::{Color, Label, Report, ReportKind, Source};
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::{format_ident, quote, quote_spanned};
use qusql_type::schema::{parse_schemas, Schemas};
use qusql_type::{
type_statement, ByteToChar, Issue, SQLArguments, SQLDialect, SelectTypeColumn, TypeOptions,
};
use syn::spanned::Spanned;
use syn::{parse::Parse, punctuated::Punctuated, Expr, Ident, LitStr, Token};
use yoke::{Yoke, Yokeable};
static RESOLVED_SCHEMA_PATHS: Mutex<Option<HashMap<PathBuf, PathBuf>>> = Mutex::new(None);
fn resolve_schema_path() -> PathBuf {
let manifest_dir: PathBuf = std::env::var("CARGO_MANIFEST_DIR")
.expect("`CARGO_MANIFEST_DIR` must be set")
.into();
let mut cache_guard = RESOLVED_SCHEMA_PATHS
.lock()
.expect("resolved schema paths lock poisoned");
let cache = cache_guard.get_or_insert_with(HashMap::new);
if let Some(cached) = cache.get(&manifest_dir) {
return cached.clone();
}
let mut schema_path = manifest_dir.join("sqlx-type-schema.sql");
if !schema_path.exists() {
use serde::Deserialize;
use std::process::Command;
let cargo = std::env::var("CARGO").expect("`CARGO` must be set");
let output = Command::new(cargo)
.args(["metadata", "--format-version=1"])
.current_dir(&manifest_dir)
.env_remove("__CARGO_FIX_PLZ")
.output()
.expect("Could not fetch metadata");
#[derive(Deserialize)]
struct CargoMetadata {
workspace_root: PathBuf,
}
let metadata: CargoMetadata =
serde_json::from_slice(&output.stdout).expect("Invalid `cargo metadata` output");
schema_path = metadata.workspace_root.join("sqlx-type-schema.sql");
}
if !schema_path.exists() {
panic!("Unable to locate sqlx-type-schema.sql");
}
cache.insert(manifest_dir, schema_path.clone());
schema_path
}
fn issue_to_report(issue: Issue, b2c: &ByteToChar) -> Report<'static, std::ops::Range<usize>> {
let span = b2c.map_span(issue.span);
let kind = match issue.level {
qusql_type::Level::Warning => ReportKind::Warning,
qusql_type::Level::Error => ReportKind::Error,
};
let mut builder = Report::build(kind, span.clone())
.with_config(ariadne::Config::default().with_color(false))
.with_message(&issue.message)
.with_label(
Label::new(span)
.with_order(-1)
.with_priority(-1)
.with_message(issue.message),
);
for frag in issue.fragments {
builder =
builder.with_label(Label::new(b2c.map_span(frag.span)).with_message(frag.message));
}
if let Some(help) = issue.help {
builder = builder.with_help(help);
}
builder.finish()
}
fn issue_to_report_color(
issue: Issue,
b2c: &ByteToChar,
) -> Report<'static, std::ops::Range<usize>> {
let span = b2c.map_span(issue.span);
let err_color = match issue.level {
qusql_type::Level::Warning => Color::Yellow,
qusql_type::Level::Error => Color::Red,
};
let mut builder = Report::build(
match issue.level {
qusql_type::Level::Warning => ReportKind::Warning,
qusql_type::Level::Error => ReportKind::Error,
},
span.clone(),
)
.with_config(ariadne::Config::default().with_compact(true))
.with_message(&issue.message)
.with_label(
Label::new(span)
.with_color(err_color)
.with_order(-1)
.with_priority(-1)
.with_message(issue.message),
);
for frag in issue.fragments {
builder = builder.with_label(
Label::new(b2c.map_span(frag.span))
.with_color(Color::Blue)
.with_message(frag.message),
);
}
if let Some(help) = issue.help {
builder = builder.with_help(help);
}
builder.finish()
}
struct NamedSource<'a>(&'a str, Source<&'a str>);
impl<'a> ariadne::Cache<()> for &NamedSource<'a> {
type Storage = &'a str;
fn display<'b>(&self, _: &'b ()) -> Option<impl std::fmt::Display + 'b> {
Some(self.0.to_string())
}
fn fetch(&mut self, _: &()) -> Result<&Source<Self::Storage>, impl std::fmt::Debug> {
Ok::<_, ()>(&self.1)
}
}
#[derive(Yokeable)]
struct SchemasYoke<'a>(Schemas<'a>);
struct SchemaCacheEntry {
schemas: Yoke<SchemasYoke<'static>, String>,
dialect: SQLDialect,
path: PathBuf,
file_len: u64,
modified: SystemTime,
hash: u64,
}
static SCHEMA_CACHE: Mutex<Option<HashMap<PathBuf, Arc<SchemaCacheEntry>>>> = Mutex::new(None);
fn get_schemas() -> Arc<SchemaCacheEntry> {
let path = resolve_schema_path();
let meta = std::fs::metadata(&path).unwrap_or_else(|e| panic!("Cannot stat {path:?}: {e}"));
let file_len = meta.len();
let modified = meta.modified().unwrap_or(SystemTime::UNIX_EPOCH);
let mut cache_guard = SCHEMA_CACHE.lock().expect("schema cache lock poisoned");
let cache = cache_guard.get_or_insert_with(HashMap::new);
if let Some(entry) = cache.get(&path) {
if entry.file_len == file_len && entry.modified == modified {
return Arc::clone(entry);
}
}
let src_string = std::fs::read_to_string(&path)
.unwrap_or_else(|e| panic!("Unable to read schema from {path:?}: {e}"));
let dialect = {
let header = src_string.lines().take(2).collect::<Vec<_>>().join(" ");
if header.contains("qusql-type-variant: postgis") || header.contains("sql-product: postgis")
{
SQLDialect::PostGIS
} else if header.contains("sql-product: postgres") {
SQLDialect::PostgreSQL
} else if header.contains("sql-product: sqlite") {
SQLDialect::Sqlite
} else {
SQLDialect::MariaDB
}
};
let dialect_for_closure = dialect.clone();
let schema_hash = {
use std::hash::{DefaultHasher, Hash, Hasher};
let mut hasher = DefaultHasher::new();
src_string.hash(&mut hasher);
hasher.finish()
};
let schemas = Yoke::attach_to_cart(src_string, move |src| {
let options = TypeOptions::new().dialect(dialect_for_closure);
let mut issues = qusql_type::Issues::new(src);
let parsed = parse_schemas(src, &mut issues, &options);
if !issues.is_ok() {
let b2c = ByteToChar::new(src.as_bytes());
let source = NamedSource("sqlx-type-schema.sql", Source::from(src));
let mut err = false;
for issue in issues.into_vec() {
if issue.level == qusql_type::Level::Error {
err = true;
}
issue_to_report_color(issue, &b2c).eprint(&source).unwrap();
}
if err {
panic!("Errors processing sqlx-type-schema.sql");
}
}
SchemasYoke(parsed)
});
let entry = Arc::new(SchemaCacheEntry {
schemas,
dialect,
path: path.clone(),
file_len,
modified,
hash: schema_hash,
});
cache.insert(path, Arc::clone(&entry));
entry
}
fn quote_args(
errors: &mut Vec<proc_macro2::TokenStream>,
query: &str,
last_span: Span,
args: &[Expr],
arguments: &[(qusql_type::ArgumentKey<'_>, qusql_type::FullType)],
dialect: &SQLDialect,
) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
let cls = match dialect {
SQLDialect::MariaDB => quote!(sqlx::mysql::MySql),
SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite),
SQLDialect::PostgreSQL | SQLDialect::PostGIS => quote!(sqlx::postgres::Postgres),
};
let mut at = Vec::new();
let inv = qusql_type::FullType::invalid();
for (k, v) in arguments {
match k {
qusql_type::ArgumentKey::Index(i) => {
while at.len() <= *i {
at.push(&inv);
}
at[*i] = v;
}
qusql_type::ArgumentKey::Identifier(_) => {
errors.push(
syn::Error::new(last_span.span(), "Named arguments not supported")
.to_compile_error(),
);
}
}
}
if at.len() > args.len() {
errors.push(
syn::Error::new(
last_span,
format!("Expected {} additional arguments", at.len() - args.len()),
)
.to_compile_error(),
);
}
if let Some(args) = args.get(at.len()..) {
for arg in args {
errors.push(syn::Error::new(arg.span(), "unexpected argument").to_compile_error());
}
}
let arg_names = (0..args.len())
.map(|i| format_ident!("arg{}", i))
.collect::<Vec<_>>();
let mut arg_bindings = Vec::new();
let mut arg_add = Vec::new();
let mut list_lengths = Vec::new();
for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) {
let mut t = match ta.t {
qusql_type::Type::U8 => quote! {u8},
qusql_type::Type::I8 => quote! {i8},
qusql_type::Type::U16 => quote! {u16},
qusql_type::Type::I16 => quote! {i16},
qusql_type::Type::U24 => quote! {u32},
qusql_type::Type::I24 => quote! {i32},
qusql_type::Type::U32 => quote! {u32},
qusql_type::Type::I32 => quote! {i32},
qusql_type::Type::U64 => quote! {u64},
qusql_type::Type::I64 => quote! {i64},
qusql_type::Type::Base(qusql_type::BaseType::Any) => quote! {qusql_sqlx_type::Any},
qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {&[u8]},
qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {qusql_sqlx_type::Date},
qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
quote! {qusql_sqlx_type::DateTime}
}
qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {qusql_sqlx_type::Float},
qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
quote! {qusql_sqlx_type::Integer}
}
qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {&str},
qusql_type::Type::Base(qusql_type::BaseType::Time) => quote! {qusql_sqlx_type::Time},
qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => todo!("time_interval"),
qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
quote! {qusql_sqlx_type::Timestamp}
}
qusql_type::Type::Base(qusql_type::BaseType::Uuid) => quote! {qusql_sqlx_type::Uuid},
qusql_type::Type::Null => todo!("null"),
qusql_type::Type::Invalid => quote! {std::convert::Infallible},
qusql_type::Type::Enum(_) => quote! {&str},
qusql_type::Type::Set(_) => quote! {&str},
qusql_type::Type::Args(_, _) => todo!("args"),
qusql_type::Type::F32 => quote! {f32},
qusql_type::Type::F64 => quote! {f64},
qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any},
qusql_type::Type::Geometry => quote! {qusql_sqlx_type::Any},
qusql_type::Type::Range(_) => quote! {qusql_sqlx_type::Any},
qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
};
if !ta.not_null {
t = quote! {Option<#t>}
}
let span = qa.span();
if ta.list_hack {
list_lengths.push(quote!(#name.len()));
arg_bindings.push(quote_spanned! {span=>
let #name = &(#qa);
args_count += #name.len();
for v in #name.iter() {
size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(v);
}
if false {
qusql_sqlx_type::check_arg_list_hack::<#t, _>(#name);
::std::panic!();
}
});
arg_add.push(quote!(
for v in #name.iter() {
e = e.and_then(|()| query_args.add(v));
}
));
} else {
arg_bindings.push(quote_spanned! {span=>
let #name = &(#qa);
args_count += 1;
size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(#name);
if false {
qusql_sqlx_type::check_arg::<#t, _>(#name);
::std::panic!();
}
});
arg_add.push(quote!(e = e.and_then(|()| query_args.add(#name));));
}
}
let query = if list_lengths.is_empty() {
quote!(#query)
} else {
quote!(
&qusql_sqlx_type::convert_list_query(#query, &[#(#list_lengths),*])
)
};
(
quote! {
let mut size_hints = 0;
let mut args_count = 0;
#(#arg_bindings)*
let mut query_args = <#cls as ::sqlx::database::Database>::Arguments::default();
query_args.reserve(args_count, size_hints);
let mut e = Ok(());
#(#arg_add)*
let query_args = e.and_then(|()| Ok(query_args));
},
query,
)
}
fn issues_to_errors(issues: Vec<Issue>, source: &str, span: Span) -> Vec<proc_macro2::TokenStream> {
if !issues.is_empty() {
let b2c = ByteToChar::new(source.as_bytes());
let source = NamedSource("query", Source::from(source));
let mut err = false;
let mut out = Vec::new();
for issue in issues {
if issue.level == qusql_type::Level::Error {
err = true;
}
let r = issue_to_report(issue, &b2c);
r.write(&source, &mut out).unwrap();
}
if err {
let raw = String::from_utf8(out).unwrap();
let body = raw
.find('\n')
.map(|i| raw[i + 1..].trim_start_matches('\n').trim_end())
.unwrap_or(raw.trim_end());
return vec![syn::Error::new(span, body).to_compile_error()];
}
}
Vec::new()
}
fn construct_row(
columns: &[SelectTypeColumn],
) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
let mut row_members = Vec::new();
let mut row_construct = Vec::new();
for (i, c) in columns.iter().enumerate() {
let mut t = match c.type_.t {
qusql_type::Type::U8 => quote! {u8},
qusql_type::Type::I8 => quote! {i8},
qusql_type::Type::U16 => quote! {u16},
qusql_type::Type::I16 => quote! {i16},
qusql_type::Type::U24 => quote! {u32},
qusql_type::Type::I24 => quote! {i32},
qusql_type::Type::U32 => quote! {u32},
qusql_type::Type::I32 => quote! {i32},
qusql_type::Type::U64 => quote! {u64},
qusql_type::Type::I64 => quote! {i64},
qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
quote! {chrono::NaiveDateTime}
}
qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
qusql_type::Type::Base(qusql_type::BaseType::Integer) => quote! {i64},
qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
todo!("from_time_interval")
}
qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
}
qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
quote! {qusql_sqlx_type::UuidValue}
}
qusql_type::Type::Null => todo!("from_null"),
qusql_type::Type::Invalid => quote! {i64},
qusql_type::Type::Enum(_) => quote! {String},
qusql_type::Type::Set(_) => quote! {String},
qusql_type::Type::Args(_, _) => todo!("from_args"),
qusql_type::Type::F32 => quote! {f32},
qusql_type::Type::F64 => quote! {f64},
qusql_type::Type::JSON => quote! {String},
qusql_type::Type::Geometry => quote! {Vec<u8>},
qusql_type::Type::Range(_) => quote! {Vec<u8>},
qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
};
let name = match &c.name {
Some(v) => v,
None => continue,
};
let name_str = name.value;
let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
(stripped, true)
} else {
(name_str, false)
};
let ident = String::from("r#") + name_str;
let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
ident
} else {
continue;
};
if !c.type_.not_null && !force_not_null {
t = quote! {Option<#t>};
}
row_members.push(quote! {
#ident : #t
});
row_construct.push(quote! {
#ident: sqlx::Row::get(&row, #i)
});
}
(row_members, row_construct)
}
struct Query {
query: String,
query_span: Span,
args: Vec<Expr>,
last_span: Span,
}
impl Parse for Query {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
let query: String = query_.iter().map(LitStr::value).collect();
let query_span = query_.span();
let mut last_span = query_span;
let mut args = Vec::new();
while !input.is_empty() {
let _ = input.parse::<syn::token::Comma>()?;
if input.is_empty() {
break;
}
let arg = input.parse::<Expr>()?;
last_span = arg.span();
args.push(arg);
}
Ok(Self {
query,
query_span,
args,
last_span,
})
}
}
#[proc_macro]
pub fn query(input: TokenStream) -> TokenStream {
let query = syn::parse_macro_input!(input as Query);
let cache = get_schemas();
let (schemas, dialect, schema_hash) = (cache.schemas.get(), &cache.dialect, cache.hash);
let options = TypeOptions::new()
.dialect(dialect.clone())
.arguments(match &dialect {
SQLDialect::MariaDB => SQLArguments::QuestionMark,
SQLDialect::Sqlite => SQLArguments::QuestionMark,
SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
})
.list_hack(true);
let mut issues = qusql_type::Issues::new(&query.query);
let stmt = type_statement(&schemas.0, &query.query, &mut issues, &options);
let sp = cache.path.to_str().unwrap();
let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
match &stmt {
qusql_type::StatementType::Select { columns, arguments } => {
let (args_tokens, q) = quote_args(
&mut errors,
&query.query,
query.last_span,
&query.args,
arguments,
dialect,
);
let (row_members, row_construct) = construct_row(columns);
let s = quote! { {
use ::sqlx::Arguments as _;
let _ = std::include_bytes!(#sp);
#(#errors; )*
#args_tokens;
struct Row {
#(#row_members),*
};
sqlx::__query_with_result(#q, query_args).map(|row|
Row{
#(#row_construct),*
}
)
}};
s.into()
}
qusql_type::StatementType::Delete {
arguments,
returning,
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query.query,
query.last_span,
&query.args,
arguments,
dialect,
);
let s = match returning.as_ref() {
Some(returning) => {
let (row_members, row_construct) = construct_row(returning);
quote! { {
use ::sqlx::Arguments as _;
let _ = std::include_bytes!(#sp);
#(#errors; )*
#args_tokens
struct Row {
#(#row_members),*
};
sqlx::__query_with_result(#q, query_args).map(|row|
Row{
#(#row_construct),*
}
)
}}
}
None => quote! { {
use ::sqlx::Arguments as _;
const _SCHEMA_HASH: u64 = #schema_hash;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args)
}
},
};
s.into()
}
qusql_type::StatementType::Insert {
arguments,
returning,
..
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query.query,
query.last_span,
&query.args,
arguments,
dialect,
);
let s = match returning.as_ref() {
Some(returning) => {
let (row_members, row_construct) = construct_row(returning);
quote! { {
use ::sqlx::Arguments as _;
let _ = std::include_bytes!(#sp);
#(#errors; )*
#args_tokens
struct Row {
#(#row_members),*
};
sqlx::__query_with_result(#q, query_args).map(|row|
Row{
#(#row_construct),*
}
)
}}
}
None => quote! { {
use ::sqlx::Arguments as _;
const _SCHEMA_HASH: u64 = #schema_hash;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args)
}
},
};
s.into()
}
qusql_type::StatementType::Update {
arguments,
returning,
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query.query,
query.last_span,
&query.args,
arguments,
dialect,
);
let s = match returning.as_ref() {
Some(returning) => {
let (row_members, row_construct) = construct_row(returning);
quote! { {
use ::sqlx::Arguments as _;
let _ = std::include_bytes!(#sp);
#(#errors; )*
#args_tokens
struct Row {
#(#row_members),*
};
sqlx::__query_with_result(#q, query_args).map(|row|
Row{
#(#row_construct),*
}
)
}}
}
None => quote! { {
use ::sqlx::Arguments as _;
const _SCHEMA_HASH: u64 = #schema_hash;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args)
}
},
};
s.into()
}
qusql_type::StatementType::Replace {
arguments,
returning,
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query.query,
query.last_span,
&query.args,
arguments,
dialect,
);
let s = match returning.as_ref() {
Some(returning) => {
let (row_members, row_construct) = construct_row(returning);
quote! { {
use ::sqlx::Arguments as _;
const _SCHEMA_HASH: u64 = #schema_hash;
let _ = std::include_bytes!(#sp);
#(#errors; )*
#args_tokens
struct Row {
#(#row_members),*
};
sqlx::__query_with_result(#q, query_args).map(|row|
Row{
#(#row_construct),*
}
)
}}
}
None => quote! { {
use ::sqlx::Arguments as _;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args)
}
},
};
s.into()
}
qusql_type::StatementType::Truncate => {
errors.push(
syn::Error::new(query.query_span, "TRUNCATE not supported in query!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("truncate")
}}
.into()
}
qusql_type::StatementType::Call { .. } => {
errors.push(
syn::Error::new(query.query_span, "CALL not supported in query!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("call")
}}
.into()
}
qusql_type::StatementType::Transaction => {
errors.push(
syn::Error::new(
query.query_span,
"Transaction control not supported in query!",
)
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("transaction")
}}
.into()
}
qusql_type::StatementType::Set => {
errors.push(
syn::Error::new(query.query_span, "SET not supported in query!").to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("set")
}}
.into()
}
qusql_type::StatementType::Lock => {
errors.push(
syn::Error::new(query.query_span, "LOCK not supported in query!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("lock")
}}
.into()
}
qusql_type::StatementType::Invalid => {
let s = quote! { {
#(#errors; )*;
todo!("Invalid")
}};
s.into()
}
}
}
fn construct_row2(columns: &[SelectTypeColumn]) -> Vec<proc_macro2::TokenStream> {
let mut row_construct = Vec::new();
for (i, c) in columns.iter().enumerate() {
let mut t = match c.type_.t {
qusql_type::Type::U8 => quote! {u8},
qusql_type::Type::I8 => quote! {i8},
qusql_type::Type::U16 => quote! {u16},
qusql_type::Type::I16 => quote! {i16},
qusql_type::Type::U24 => quote! {u32},
qusql_type::Type::I24 => quote! {i32},
qusql_type::Type::U32 => quote! {u32},
qusql_type::Type::I32 => quote! {i32},
qusql_type::Type::U64 => quote! {u64},
qusql_type::Type::I64 => quote! {i64},
qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
quote! {chrono::NaiveDateTime}
}
qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
qusql_type::Type::Base(qusql_type::BaseType::Integer) => quote! {i64},
qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
todo!("from_time_interval")
}
qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
}
qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
quote! {qusql_sqlx_type::UuidValue}
}
qusql_type::Type::Null => todo!("from_null"),
qusql_type::Type::Invalid => quote! {i64},
qusql_type::Type::Enum(_) => quote! {String},
qusql_type::Type::Set(_) => quote! {String},
qusql_type::Type::Args(_, _) => todo!("from_args"),
qusql_type::Type::F32 => quote! {f32},
qusql_type::Type::F64 => quote! {f64},
qusql_type::Type::JSON => quote! {String},
qusql_type::Type::Geometry => quote! {Vec<u8>},
qusql_type::Type::Range(_) => quote! {Vec<u8>},
qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
};
let name = match &c.name {
Some(v) => v,
None => continue,
};
let name_str = name.value;
let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
(stripped, true)
} else {
(name_str, false)
};
let ident = String::from("r#") + name_str;
let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
ident
} else {
continue;
};
if !c.type_.not_null && !force_not_null {
t = quote! {Option<#t>};
}
row_construct.push(quote! {
#ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
});
}
row_construct
}
struct QueryAs {
as_: Ident,
query: String,
query_span: Span,
args: Vec<Expr>,
last_span: Span,
}
impl Parse for QueryAs {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let as_ = input.parse::<Ident>()?;
let _ = input.parse::<syn::token::Comma>()?;
let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
let query: String = query_.iter().map(LitStr::value).collect();
let query_span = query_.span();
let mut last_span = query_span;
let mut args = Vec::new();
while !input.is_empty() {
let _ = input.parse::<syn::token::Comma>()?;
if input.is_empty() {
break;
}
let arg = input.parse::<Expr>()?;
last_span = arg.span();
args.push(arg);
}
Ok(Self {
as_,
query,
query_span,
args,
last_span,
})
}
}
#[proc_macro]
pub fn query_as(input: TokenStream) -> TokenStream {
let query_as = syn::parse_macro_input!(input as QueryAs);
let cache = get_schemas();
let (schemas, dialect) = (cache.schemas.get(), &cache.dialect);
let options = TypeOptions::new()
.dialect(dialect.clone())
.arguments(match &dialect {
SQLDialect::MariaDB => SQLArguments::QuestionMark,
SQLDialect::Sqlite => SQLArguments::QuestionMark,
SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
})
.list_hack(true);
let mut issues = qusql_type::Issues::new(&query_as.query);
let stmt = type_statement(&schemas.0, &query_as.query, &mut issues, &options);
let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
match &stmt {
qusql_type::StatementType::Select { columns, arguments } => {
let (args_tokens, q) = quote_args(
&mut errors,
&query_as.query,
query_as.last_span,
&query_as.args,
arguments,
dialect,
);
let row_construct = construct_row2(columns);
let row = query_as.as_;
let s = quote! { {
use ::sqlx::Arguments as _;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args).map(|row|
#row{
#(#row_construct),*
}
)
}};
s.into()
}
qusql_type::StatementType::Delete { .. } => {
errors.push(
syn::Error::new(query_as.query_span, "DELETE not support in query_as")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("delete")
}}
.into()
}
qusql_type::StatementType::Insert {
returning: None, ..
} => {
errors.push(
syn::Error::new(
query_as.query_span,
"INSERT without RETURNING not support in query_as",
)
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("insert")
}}
.into()
}
qusql_type::StatementType::Insert {
arguments,
returning: Some(returning),
..
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query_as.query,
query_as.last_span,
&query_as.args,
arguments,
dialect,
);
let row_construct = construct_row2(returning);
let row = query_as.as_;
let s = quote! { {
use ::sqlx::Arguments as _;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args).map(|row|
#row{
#(#row_construct),*
}
)
}};
s.into()
}
qusql_type::StatementType::Update {
returning: None, ..
} => {
errors.push(
syn::Error::new(
query_as.query_span,
"UPDATE without RETURNING not support in query_as",
)
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("update")
}}
.into()
}
qusql_type::StatementType::Update {
arguments,
returning: Some(returning),
..
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query_as.query,
query_as.last_span,
&query_as.args,
arguments,
dialect,
);
let row_construct = construct_row2(returning);
let row = query_as.as_;
let s = quote! { {
use ::sqlx::Arguments as _;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args).map(|row|
#row{
#(#row_construct),*
}
)
}};
s.into()
}
qusql_type::StatementType::Replace {
returning: None, ..
} => {
errors.push(
syn::Error::new(
query_as.query_span,
"REPLACE without RETURNING not support in query_as",
)
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("replace")
}}
.into()
}
qusql_type::StatementType::Replace {
arguments,
returning: Some(returning),
..
} => {
let (args_tokens, q) = quote_args(
&mut errors,
&query_as.query,
query_as.last_span,
&query_as.args,
arguments,
dialect,
);
let row_construct = construct_row2(returning);
let row = query_as.as_;
let s = quote! { {
use ::sqlx::Arguments as _;
#(#errors; )*
#args_tokens
sqlx::__query_with_result(#q, query_args).map(|row|
#row{
#(#row_construct),*
}
)
}};
s.into()
}
qusql_type::StatementType::Truncate => {
errors.push(
syn::Error::new(query_as.query_span, "TRUNCATE not supported in query_as!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("truncate")
}}
.into()
}
qusql_type::StatementType::Call { .. } => {
errors.push(
syn::Error::new(query_as.query_span, "CALL not supported in query_as!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("call")
}}
.into()
}
qusql_type::StatementType::Transaction => {
errors.push(
syn::Error::new(
query_as.query_span,
"Transaction control not supported in query_as!",
)
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("transaction")
}}
.into()
}
qusql_type::StatementType::Set => {
errors.push(
syn::Error::new(query_as.query_span, "SET not supported in query_as!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("set")
}}
.into()
}
qusql_type::StatementType::Lock => {
errors.push(
syn::Error::new(query_as.query_span, "LOCK not supported in query_as!")
.to_compile_error(),
);
quote! { {
#(#errors; )*
todo!("lock")
}}
.into()
}
qusql_type::StatementType::Invalid => quote! { {
#(#errors; )*;
todo!("invalid")
}}
.into(),
}
}