1#![forbid(unsafe_code)]
6
7use std::collections::HashMap;
8use std::path::PathBuf;
9use std::sync::{Arc, Mutex};
10use std::time::SystemTime;
11
12use ariadne::{Color, Label, Report, ReportKind, Source};
13use proc_macro::TokenStream;
14use proc_macro2::Span;
15use quote::{format_ident, quote, quote_spanned};
16use qusql_type::schema::{parse_schemas, Schemas};
17use qusql_type::{
18 type_statement, ByteToChar, Issue, SQLArguments, SQLDialect, SelectTypeColumn, TypeOptions,
19};
20use syn::spanned::Spanned;
21use syn::{parse::Parse, punctuated::Punctuated, Expr, Ident, LitStr, Token};
22use yoke::{Yoke, Yokeable};
23
24static RESOLVED_SCHEMA_PATHS: Mutex<Option<HashMap<PathBuf, PathBuf>>> = Mutex::new(None);
33
34fn resolve_schema_path() -> PathBuf {
39 let manifest_dir: PathBuf = std::env::var("CARGO_MANIFEST_DIR")
40 .expect("`CARGO_MANIFEST_DIR` must be set")
41 .into();
42
43 let mut cache_guard = RESOLVED_SCHEMA_PATHS
44 .lock()
45 .expect("resolved schema paths lock poisoned");
46 let cache = cache_guard.get_or_insert_with(HashMap::new);
47
48 if let Some(cached) = cache.get(&manifest_dir) {
49 return cached.clone();
50 }
51
52 let mut schema_path = manifest_dir.join("sqlx-type-schema.sql");
53
54 if !schema_path.exists() {
55 use serde::Deserialize;
56 use std::process::Command;
57
58 let cargo = std::env::var("CARGO").expect("`CARGO` must be set");
59
60 let output = Command::new(cargo)
61 .args(["metadata", "--format-version=1"])
62 .current_dir(&manifest_dir)
63 .env_remove("__CARGO_FIX_PLZ")
64 .output()
65 .expect("Could not fetch metadata");
66
67 #[derive(Deserialize)]
69 struct CargoMetadata {
70 workspace_root: PathBuf,
72 }
73
74 let metadata: CargoMetadata =
75 serde_json::from_slice(&output.stdout).expect("Invalid `cargo metadata` output");
76
77 schema_path = metadata.workspace_root.join("sqlx-type-schema.sql");
78 }
79 if !schema_path.exists() {
80 panic!("Unable to locate sqlx-type-schema.sql");
81 }
82
83 cache.insert(manifest_dir, schema_path.clone());
84 schema_path
85}
86
87fn issue_to_report(issue: Issue, b2c: &ByteToChar) -> Report<'static, std::ops::Range<usize>> {
89 let span = b2c.map_span(issue.span);
90 let kind = match issue.level {
91 qusql_type::Level::Warning => ReportKind::Warning,
92 qusql_type::Level::Error => ReportKind::Error,
93 };
94 let mut builder = Report::build(kind, span.clone())
95 .with_config(ariadne::Config::default().with_color(false))
96 .with_message(&issue.message)
97 .with_label(
98 Label::new(span)
99 .with_order(-1)
100 .with_priority(-1)
101 .with_message(issue.message),
102 );
103 for frag in issue.fragments {
104 builder =
105 builder.with_label(Label::new(b2c.map_span(frag.span)).with_message(frag.message));
106 }
107 if let Some(help) = issue.help {
108 builder = builder.with_help(help);
109 }
110 builder.finish()
111}
112
113fn issue_to_report_color(
115 issue: Issue,
116 b2c: &ByteToChar,
117) -> Report<'static, std::ops::Range<usize>> {
118 let span = b2c.map_span(issue.span);
119 let err_color = match issue.level {
120 qusql_type::Level::Warning => Color::Yellow,
121 qusql_type::Level::Error => Color::Red,
122 };
123 let mut builder = Report::build(
124 match issue.level {
125 qusql_type::Level::Warning => ReportKind::Warning,
126 qusql_type::Level::Error => ReportKind::Error,
127 },
128 span.clone(),
129 )
130 .with_config(ariadne::Config::default().with_compact(true))
131 .with_message(&issue.message)
132 .with_label(
133 Label::new(span)
134 .with_color(err_color)
135 .with_order(-1)
136 .with_priority(-1)
137 .with_message(issue.message),
138 );
139 for frag in issue.fragments {
140 builder = builder.with_label(
141 Label::new(b2c.map_span(frag.span))
142 .with_color(Color::Blue)
143 .with_message(frag.message),
144 );
145 }
146 if let Some(help) = issue.help {
147 builder = builder.with_help(help);
148 }
149 builder.finish()
150}
151
152struct NamedSource<'a>(&'a str, Source<&'a str>);
154
155impl<'a> ariadne::Cache<()> for &NamedSource<'a> {
156 type Storage = &'a str;
157
158 fn display<'b>(&self, _: &'b ()) -> Option<impl std::fmt::Display + 'b> {
159 Some(self.0.to_string())
160 }
161
162 fn fetch(&mut self, _: &()) -> Result<&Source<Self::Storage>, impl std::fmt::Debug> {
163 Ok::<_, ()>(&self.1)
164 }
165}
166
167#[derive(Yokeable)]
168struct SchemasYoke<'a>(Schemas<'a>);
170
171struct SchemaCacheEntry {
175 schemas: Yoke<SchemasYoke<'static>, String>,
177 dialect: SQLDialect,
179 path: PathBuf,
181 file_len: u64,
183 modified: SystemTime,
185 hash: u64,
187}
188
189static SCHEMA_CACHE: Mutex<Option<HashMap<PathBuf, Arc<SchemaCacheEntry>>>> = Mutex::new(None);
194
195fn get_schemas() -> Arc<SchemaCacheEntry> {
198 let path = resolve_schema_path();
199 let meta = std::fs::metadata(&path).unwrap_or_else(|e| panic!("Cannot stat {path:?}: {e}"));
200 let file_len = meta.len();
201 let modified = meta.modified().unwrap_or(SystemTime::UNIX_EPOCH);
202
203 let mut cache_guard = SCHEMA_CACHE.lock().expect("schema cache lock poisoned");
204 let cache = cache_guard.get_or_insert_with(HashMap::new);
205 if let Some(entry) = cache.get(&path) {
206 if entry.file_len == file_len && entry.modified == modified {
207 return Arc::clone(entry);
208 }
209 }
210
211 let src_string = std::fs::read_to_string(&path)
213 .unwrap_or_else(|e| panic!("Unable to read schema from {path:?}: {e}"));
214
215 let dialect = {
218 let header = src_string.lines().take(2).collect::<Vec<_>>().join(" ");
219 if header.contains("qusql-type-variant: postgis") || header.contains("sql-product: postgis")
220 {
221 SQLDialect::PostGIS
222 } else if header.contains("sql-product: postgres") {
223 SQLDialect::PostgreSQL
224 } else if header.contains("sql-product: sqlite") {
225 SQLDialect::Sqlite
226 } else {
227 SQLDialect::MariaDB
228 }
229 };
230
231 let dialect_for_closure = dialect.clone();
232
233 let schema_hash = {
235 use std::hash::{DefaultHasher, Hash, Hasher};
236 let mut hasher = DefaultHasher::new();
237 src_string.hash(&mut hasher);
238 hasher.finish()
239 };
240
241 let schemas = Yoke::attach_to_cart(src_string, move |src| {
242 let options = TypeOptions::new().dialect(dialect_for_closure);
243 let mut issues = qusql_type::Issues::new(src);
244 let parsed = parse_schemas(src, &mut issues, &options);
245 if !issues.is_ok() {
246 let b2c = ByteToChar::new(src.as_bytes());
247 let source = NamedSource("sqlx-type-schema.sql", Source::from(src));
248 let mut err = false;
249 for issue in issues.into_vec() {
250 if issue.level == qusql_type::Level::Error {
251 err = true;
252 }
253 issue_to_report_color(issue, &b2c).eprint(&source).unwrap();
254 }
255 if err {
256 panic!("Errors processing sqlx-type-schema.sql");
257 }
258 }
259 SchemasYoke(parsed)
260 });
261
262 let entry = Arc::new(SchemaCacheEntry {
263 schemas,
264 dialect,
265 path: path.clone(),
266 file_len,
267 modified,
268 hash: schema_hash,
269 });
270 cache.insert(path, Arc::clone(&entry));
271 entry
272}
273
274fn quote_args(
276 errors: &mut Vec<proc_macro2::TokenStream>,
277 query: &str,
278 last_span: Span,
279 args: &[Expr],
280 arguments: &[(qusql_type::ArgumentKey<'_>, qusql_type::FullType)],
281 dialect: &SQLDialect,
282) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
283 let cls = match dialect {
284 SQLDialect::MariaDB => quote!(sqlx::mysql::MySql),
285 SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite),
286 SQLDialect::PostgreSQL | SQLDialect::PostGIS => quote!(sqlx::postgres::Postgres),
287 };
288
289 let mut at = Vec::new();
290 let inv = qusql_type::FullType::invalid();
291 for (k, v) in arguments {
292 match k {
293 qusql_type::ArgumentKey::Index(i) => {
294 while at.len() <= *i {
295 at.push(&inv);
296 }
297 at[*i] = v;
298 }
299 qusql_type::ArgumentKey::Identifier(_) => {
300 errors.push(
301 syn::Error::new(last_span.span(), "Named arguments not supported")
302 .to_compile_error(),
303 );
304 }
305 }
306 }
307
308 if at.len() > args.len() {
309 errors.push(
310 syn::Error::new(
311 last_span,
312 format!("Expected {} additional arguments", at.len() - args.len()),
313 )
314 .to_compile_error(),
315 );
316 }
317
318 if let Some(args) = args.get(at.len()..) {
319 for arg in args {
320 errors.push(syn::Error::new(arg.span(), "unexpected argument").to_compile_error());
321 }
322 }
323
324 let arg_names = (0..args.len())
325 .map(|i| format_ident!("arg{}", i))
326 .collect::<Vec<_>>();
327
328 let mut arg_bindings = Vec::new();
329 let mut arg_add = Vec::new();
330
331 let mut list_lengths = Vec::new();
332
333 for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) {
334 let mut t = match ta.t {
335 qusql_type::Type::U8 => quote! {u8},
336 qusql_type::Type::I8 => quote! {i8},
337 qusql_type::Type::U16 => quote! {u16},
338 qusql_type::Type::I16 => quote! {i16},
339 qusql_type::Type::U24 => quote! {u32},
340 qusql_type::Type::I24 => quote! {i32},
341 qusql_type::Type::U32 => quote! {u32},
342 qusql_type::Type::I32 => quote! {i32},
343 qusql_type::Type::U64 => quote! {u64},
344 qusql_type::Type::I64 => quote! {i64},
345 qusql_type::Type::Base(qusql_type::BaseType::Any) => quote! {qusql_sqlx_type::Any},
346 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
347 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {&[u8]},
348 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {qusql_sqlx_type::Date},
349 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
350 quote! {qusql_sqlx_type::DateTime}
351 }
352 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {qusql_sqlx_type::Float},
353 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
354 quote! {qusql_sqlx_type::Integer}
355 }
356 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {&str},
357 qusql_type::Type::Base(qusql_type::BaseType::Time) => quote! {qusql_sqlx_type::Time},
358 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => todo!("time_interval"),
359 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
360 quote! {qusql_sqlx_type::Timestamp}
361 }
362 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => quote! {qusql_sqlx_type::Uuid},
363 qusql_type::Type::Null => todo!("null"),
364 qusql_type::Type::Invalid => quote! {std::convert::Infallible},
365 qusql_type::Type::Enum(_) => quote! {&str},
366 qusql_type::Type::Set(_) => quote! {&str},
367 qusql_type::Type::Args(_, _) => todo!("args"),
368 qusql_type::Type::F32 => quote! {f32},
369 qusql_type::Type::F64 => quote! {f64},
370 qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any},
371 qusql_type::Type::Geometry => quote! {qusql_sqlx_type::Any},
372 qusql_type::Type::Range(_) => quote! {qusql_sqlx_type::Any},
373 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
374 };
375 if !ta.not_null {
376 t = quote! {Option<#t>}
377 }
378 let span = qa.span();
379 if ta.list_hack {
380 list_lengths.push(quote!(#name.len()));
381 arg_bindings.push(quote_spanned! {span=>
382 let #name = &(#qa);
383 args_count += #name.len();
384 for v in #name.iter() {
385 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(v);
386 }
387 if false {
388 qusql_sqlx_type::check_arg_list_hack::<#t, _>(#name);
389 ::std::panic!();
390 }
391 });
392 arg_add.push(quote!(
393 for v in #name.iter() {
394 e = e.and_then(|()| query_args.add(v));
395 }
396 ));
397 } else {
398 arg_bindings.push(quote_spanned! {span=>
399 let #name = &(#qa);
400 args_count += 1;
401 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(#name);
402 if false {
403 qusql_sqlx_type::check_arg::<#t, _>(#name);
404 ::std::panic!();
405 }
406 });
407 arg_add.push(quote!(e = e.and_then(|()| query_args.add(#name));));
408 }
409 }
410
411 let query = if list_lengths.is_empty() {
412 quote!(#query)
413 } else {
414 quote!(
415 qusql_sqlx_type::convert_list_query(#query, &[#(#list_lengths),*])
416 )
417 };
418
419 (
420 quote! {
421 let mut size_hints = 0;
422 let mut args_count = 0;
423 #(#arg_bindings)*
424
425 let mut query_args = <#cls as ::sqlx::database::Database>::Arguments::default();
426 query_args.reserve(args_count, size_hints);
427 let mut e = Ok(());
428 #(#arg_add)*
429 let query_args = e.and_then(|()| Ok(query_args));
430 },
431 query,
432 )
433}
434
435fn issues_to_errors(issues: Vec<Issue>, source: &str, span: Span) -> Vec<proc_macro2::TokenStream> {
437 if !issues.is_empty() {
438 let b2c = ByteToChar::new(source.as_bytes());
439 let source = NamedSource("query", Source::from(source));
440 let mut err = false;
441 let mut out = Vec::new();
442 for issue in issues {
443 if issue.level == qusql_type::Level::Error {
444 err = true;
445 }
446 let r = issue_to_report(issue, &b2c);
447 r.write(&source, &mut out).unwrap();
448 }
449 if err {
450 let raw = String::from_utf8(out).unwrap();
451 let body = raw
454 .find('\n')
455 .map(|i| raw[i + 1..].trim_start_matches('\n').trim_end())
456 .unwrap_or(raw.trim_end());
457 return vec![syn::Error::new(span, body).to_compile_error()];
458 }
459 }
460 Vec::new()
461}
462
463fn construct_row(
465 columns: &[SelectTypeColumn],
466 is_postgres: bool,
467) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
468 let mut row_members = Vec::new();
469 let mut row_construct = Vec::new();
470 for (i, c) in columns.iter().enumerate() {
471 let mut t = match c.type_.t {
472 qusql_type::Type::U8 => quote! {u8},
473 qusql_type::Type::I8 => quote! {i8},
474 qusql_type::Type::U16 => quote! {u16},
475 qusql_type::Type::I16 => quote! {i16},
476 qusql_type::Type::U24 => quote! {u32},
477 qusql_type::Type::I24 => quote! {i32},
478 qusql_type::Type::U32 => quote! {u32},
479 qusql_type::Type::I32 => quote! {i32},
480 qusql_type::Type::U64 => quote! {u64},
481 qusql_type::Type::I64 => quote! {i64},
482 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
483 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
484 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
485 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
486 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
487 quote! {chrono::NaiveDateTime}
488 }
489 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
490 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
491 if is_postgres {
492 quote! {i32}
493 } else {
494 quote! {i64}
495 }
496 }
497 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
498 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
499 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
500 todo!("from_time_interval")
501 }
502 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
503 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
504 }
505 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
506 quote! {qusql_sqlx_type::UuidValue}
507 }
508 qusql_type::Type::Null => todo!("from_null"),
509 qusql_type::Type::Invalid => quote! {i64},
510 qusql_type::Type::Enum(_) => quote! {String},
511 qusql_type::Type::Set(_) => quote! {String},
512 qusql_type::Type::Args(_, _) => todo!("from_args"),
513 qusql_type::Type::F32 => quote! {f32},
514 qusql_type::Type::F64 => quote! {f64},
515 qusql_type::Type::JSON => quote! {String},
516 qusql_type::Type::Geometry => quote! {Vec<u8>},
517 qusql_type::Type::Range(_) => quote! {Vec<u8>},
518 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
519 };
520 let name = match &c.name {
521 Some(v) => v,
522 None => continue,
523 };
524
525 let name_str = name.value;
527 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
528 (stripped, true)
529 } else {
530 (name_str, false)
531 };
532
533 let ident = String::from("r#") + name_str;
534 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
535 ident
536 } else {
537 continue;
540 };
541
542 if !c.type_.not_null && !force_not_null {
543 t = quote! {Option<#t>};
544 }
545 row_members.push(quote! {
546 #ident : #t
547 });
548 row_construct.push(quote! {
549 #ident: sqlx::Row::get(&row, #i)
550 });
551 }
552 (row_members, row_construct)
553}
554
555struct Query {
557 query: String,
559 query_span: Span,
561 args: Vec<Expr>,
563 last_span: Span,
565}
566
567impl Parse for Query {
568 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
569 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
570 let query: String = query_.iter().map(LitStr::value).collect();
571 let query_span = query_.span();
572 let mut last_span = query_span;
573 let mut args = Vec::new();
574 while !input.is_empty() {
575 let _ = input.parse::<syn::token::Comma>()?;
576 if input.is_empty() {
577 break;
578 }
579 let arg = input.parse::<Expr>()?;
580 last_span = arg.span();
581 args.push(arg);
582 }
583 Ok(Self {
584 query,
585 query_span,
586 args,
587 last_span,
588 })
589 }
590}
591
592#[proc_macro]
596pub fn query(input: TokenStream) -> TokenStream {
597 let query = syn::parse_macro_input!(input as Query);
598 let cache = get_schemas();
599 let (schemas, dialect, schema_hash) = (cache.schemas.get(), &cache.dialect, cache.hash);
600 let options = TypeOptions::new()
601 .dialect(dialect.clone())
602 .arguments(match &dialect {
603 SQLDialect::MariaDB => SQLArguments::QuestionMark,
604 SQLDialect::Sqlite => SQLArguments::QuestionMark,
605 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
606 })
607 .list_hack(true);
608 let mut issues = qusql_type::Issues::new(&query.query);
609 let stmt = type_statement(&schemas.0, &query.query, &mut issues, &options);
610 let sp = cache.path.to_str().unwrap();
611 let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
612 match &stmt {
613 qusql_type::StatementType::Select { columns, arguments } => {
614 let (args_tokens, q) = quote_args(
615 &mut errors,
616 &query.query,
617 query.last_span,
618 &query.args,
619 arguments,
620 dialect,
621 );
622 let (row_members, row_construct) = construct_row(columns, dialect.is_postgresql());
623 let s = quote! { {
624 use ::sqlx::Arguments as _;
625 use ::sqlx::SqlSafeStr;
626 let _ = std::include_bytes!(#sp);
627 #(#errors; )*
628 #args_tokens;
629
630 struct Row {
631 #(#row_members),*
632 };
633 sqlx::__query_with_result(#q, query_args).map(|row|
634 Row{
635 #(#row_construct),*
636 }
637 )
638 }};
639 s.into()
640 }
641 qusql_type::StatementType::Delete {
642 arguments,
643 returning,
644 } => {
645 let (args_tokens, q) = quote_args(
646 &mut errors,
647 &query.query,
648 query.last_span,
649 &query.args,
650 arguments,
651 dialect,
652 );
653 let s = match returning.as_ref() {
654 Some(returning) => {
655 let (row_members, row_construct) =
656 construct_row(returning, dialect.is_postgresql());
657 quote! { {
658 use ::sqlx::Arguments as _;
659 let _ = std::include_bytes!(#sp);
660 #(#errors; )*
661 #args_tokens
662
663 struct Row {
664 #(#row_members),*
665 };
666 sqlx::__query_with_result(#q, query_args).map(|row|
667 Row{
668 #(#row_construct),*
669 }
670 )
671 }}
672 }
673 None => quote! { {
674 use ::sqlx::Arguments as _;
675 const _SCHEMA_HASH: u64 = #schema_hash;
676 #(#errors; )*
677 #args_tokens
678 sqlx::__query_with_result(#q, query_args)
679 }
680 },
681 };
682 s.into()
683 }
684 qusql_type::StatementType::Insert {
685 arguments,
686 returning,
687 ..
688 } => {
689 let (args_tokens, q) = quote_args(
690 &mut errors,
691 &query.query,
692 query.last_span,
693 &query.args,
694 arguments,
695 dialect,
696 );
697 let s = match returning.as_ref() {
698 Some(returning) => {
699 let (row_members, row_construct) =
700 construct_row(returning, dialect.is_postgresql());
701 quote! { {
702 use ::sqlx::Arguments as _;
703 let _ = std::include_bytes!(#sp);
704 #(#errors; )*
705 #args_tokens
706
707 struct Row {
708 #(#row_members),*
709 };
710 sqlx::__query_with_result(#q, query_args).map(|row|
711 Row{
712 #(#row_construct),*
713 }
714 )
715 }}
716 }
717 None => quote! { {
718 use ::sqlx::Arguments as _;
719 const _SCHEMA_HASH: u64 = #schema_hash;
720 #(#errors; )*
721 #args_tokens
722 sqlx::__query_with_result(#q, query_args)
723 }
724 },
725 };
726 s.into()
727 }
728 qusql_type::StatementType::Update {
729 arguments,
730 returning,
731 } => {
732 let (args_tokens, q) = quote_args(
733 &mut errors,
734 &query.query,
735 query.last_span,
736 &query.args,
737 arguments,
738 dialect,
739 );
740
741 let s = match returning.as_ref() {
742 Some(returning) => {
743 let (row_members, row_construct) =
744 construct_row(returning, dialect.is_postgresql());
745 quote! { {
746 use ::sqlx::Arguments as _;
747 let _ = std::include_bytes!(#sp);
748 #(#errors; )*
749 #args_tokens
750
751 struct Row {
752 #(#row_members),*
753 };
754 sqlx::__query_with_result(#q, query_args).map(|row|
755 Row{
756 #(#row_construct),*
757 }
758 )
759 }}
760 }
761 None => quote! { {
762 use ::sqlx::Arguments as _;
763 const _SCHEMA_HASH: u64 = #schema_hash;
764 #(#errors; )*
765 #args_tokens
766 sqlx::__query_with_result(#q, query_args)
767 }
768 },
769 };
770 s.into()
771 }
772 qusql_type::StatementType::Replace {
773 arguments,
774 returning,
775 } => {
776 let (args_tokens, q) = quote_args(
777 &mut errors,
778 &query.query,
779 query.last_span,
780 &query.args,
781 arguments,
782 dialect,
783 );
784 let s = match returning.as_ref() {
785 Some(returning) => {
786 let (row_members, row_construct) =
787 construct_row(returning, dialect.is_postgresql());
788 quote! { {
789 use ::sqlx::Arguments as _;
790 const _SCHEMA_HASH: u64 = #schema_hash;
791 let _ = std::include_bytes!(#sp);
792 #(#errors; )*
793 #args_tokens
794
795 struct Row {
796 #(#row_members),*
797 };
798 sqlx::__query_with_result(#q, query_args).map(|row|
799 Row{
800 #(#row_construct),*
801 }
802 )
803 }}
804 }
805 None => quote! { {
806 use ::sqlx::Arguments as _;
807 #(#errors; )*
808 #args_tokens
809 sqlx::__query_with_result(#q, query_args)
810 }
811 },
812 };
813 s.into()
814 }
815 qusql_type::StatementType::Truncate => {
816 errors.push(
817 syn::Error::new(query.query_span, "TRUNCATE not supported in query!")
818 .to_compile_error(),
819 );
820 quote! { {
821 #(#errors; )*
822 todo!("truncate")
823 }}
824 .into()
825 }
826 qusql_type::StatementType::Call { .. } => {
827 errors.push(
828 syn::Error::new(query.query_span, "CALL not supported in query!")
829 .to_compile_error(),
830 );
831 quote! { {
832 #(#errors; )*
833 todo!("call")
834 }}
835 .into()
836 }
837 qusql_type::StatementType::Transaction => {
838 errors.push(
839 syn::Error::new(
840 query.query_span,
841 "Transaction control not supported in query!",
842 )
843 .to_compile_error(),
844 );
845 quote! { {
846 #(#errors; )*
847 todo!("transaction")
848 }}
849 .into()
850 }
851 qusql_type::StatementType::Set => {
852 errors.push(
853 syn::Error::new(query.query_span, "SET not supported in query!").to_compile_error(),
854 );
855 quote! { {
856 #(#errors; )*
857 todo!("set")
858 }}
859 .into()
860 }
861 qusql_type::StatementType::Lock => {
862 errors.push(
863 syn::Error::new(query.query_span, "LOCK not supported in query!")
864 .to_compile_error(),
865 );
866 quote! { {
867 #(#errors; )*
868 todo!("lock")
869 }}
870 .into()
871 }
872 qusql_type::StatementType::Invalid => {
873 let s = quote! { {
874 #(#errors; )*;
875 todo!("Invalid")
876 }};
877 s.into()
878 }
879 }
880}
881
882fn construct_row2(
884 columns: &[SelectTypeColumn],
885 is_postgres: bool,
886) -> Vec<proc_macro2::TokenStream> {
887 let mut row_construct = Vec::new();
888 for (i, c) in columns.iter().enumerate() {
889 let mut t = match c.type_.t {
890 qusql_type::Type::U8 => quote! {u8},
891 qusql_type::Type::I8 => quote! {i8},
892 qusql_type::Type::U16 => quote! {u16},
893 qusql_type::Type::I16 => quote! {i16},
894 qusql_type::Type::U24 => quote! {u32},
895 qusql_type::Type::I24 => quote! {i32},
896 qusql_type::Type::U32 => quote! {u32},
897 qusql_type::Type::I32 => quote! {i32},
898 qusql_type::Type::U64 => quote! {u64},
899 qusql_type::Type::I64 => quote! {i64},
900 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
901 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
902 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
903 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
904 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
905 quote! {chrono::NaiveDateTime}
906 }
907 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
908 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
909 if is_postgres {
910 quote! {i32}
911 } else {
912 quote! {i64}
913 }
914 }
915 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
916 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
917 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
918 todo!("from_time_interval")
919 }
920 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
921 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
922 }
923 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
924 quote! {qusql_sqlx_type::UuidValue}
925 }
926 qusql_type::Type::Null => todo!("from_null"),
927 qusql_type::Type::Invalid => quote! {i64},
928 qusql_type::Type::Enum(_) => quote! {String},
929 qusql_type::Type::Set(_) => quote! {String},
930 qusql_type::Type::Args(_, _) => todo!("from_args"),
931 qusql_type::Type::F32 => quote! {f32},
932 qusql_type::Type::F64 => quote! {f64},
933 qusql_type::Type::JSON => quote! {String},
934 qusql_type::Type::Geometry => quote! {Vec<u8>},
935 qusql_type::Type::Range(_) => quote! {Vec<u8>},
936 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
937 };
938 let name = match &c.name {
939 Some(v) => v,
940 None => continue,
941 };
942
943 let name_str = name.value;
945 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
946 (stripped, true)
947 } else {
948 (name_str, false)
949 };
950
951 let ident = String::from("r#") + name_str;
952 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
953 ident
954 } else {
955 continue;
958 };
959
960 if !c.type_.not_null && !force_not_null {
961 t = quote! {Option<#t>};
962 }
963 row_construct.push(quote! {
964 #ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
965 });
966 }
967 row_construct
968}
969
970struct QueryAs {
972 as_: Ident,
974 query: String,
976 query_span: Span,
978 args: Vec<Expr>,
980 last_span: Span,
982}
983
984impl Parse for QueryAs {
985 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
986 let as_ = input.parse::<Ident>()?;
987 let _ = input.parse::<syn::token::Comma>()?;
988
989 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
990 let query: String = query_.iter().map(LitStr::value).collect();
991 let query_span = query_.span();
992
993 let mut last_span = query_span;
994 let mut args = Vec::new();
995 while !input.is_empty() {
996 let _ = input.parse::<syn::token::Comma>()?;
997 if input.is_empty() {
998 break;
999 }
1000 let arg = input.parse::<Expr>()?;
1001 last_span = arg.span();
1002 args.push(arg);
1003 }
1004 Ok(Self {
1005 as_,
1006 query,
1007 query_span,
1008 args,
1009 last_span,
1010 })
1011 }
1012}
1013
1014#[proc_macro]
1018pub fn query_as(input: TokenStream) -> TokenStream {
1019 let query_as = syn::parse_macro_input!(input as QueryAs);
1020 let cache = get_schemas();
1021 let (schemas, dialect) = (cache.schemas.get(), &cache.dialect);
1022 let options = TypeOptions::new()
1023 .dialect(dialect.clone())
1024 .arguments(match &dialect {
1025 SQLDialect::MariaDB => SQLArguments::QuestionMark,
1026 SQLDialect::Sqlite => SQLArguments::QuestionMark,
1027 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
1028 })
1029 .list_hack(true);
1030 let mut issues = qusql_type::Issues::new(&query_as.query);
1031 let stmt = type_statement(&schemas.0, &query_as.query, &mut issues, &options);
1032
1033 let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
1034 match &stmt {
1035 qusql_type::StatementType::Select { columns, arguments } => {
1036 let (args_tokens, q) = quote_args(
1037 &mut errors,
1038 &query_as.query,
1039 query_as.last_span,
1040 &query_as.args,
1041 arguments,
1042 dialect,
1043 );
1044
1045 let row_construct = construct_row2(columns, dialect.is_postgresql());
1046 let row = query_as.as_;
1047 let s = quote! { {
1048 use ::sqlx::Arguments as _;
1049 #(#errors; )*
1050 #args_tokens
1051 sqlx::__query_with_result(#q, query_args).map(|row|
1052 #row{
1053 #(#row_construct),*
1054 }
1055 )
1056 }};
1057 s.into()
1059 }
1060 qusql_type::StatementType::Delete { .. } => {
1061 errors.push(
1062 syn::Error::new(query_as.query_span, "DELETE not support in query_as")
1063 .to_compile_error(),
1064 );
1065 quote! { {
1066 #(#errors; )*
1067 todo!("delete")
1068 }}
1069 .into()
1070 }
1071 qusql_type::StatementType::Insert {
1072 returning: None, ..
1073 } => {
1074 errors.push(
1075 syn::Error::new(
1076 query_as.query_span,
1077 "INSERT without RETURNING not support in query_as",
1078 )
1079 .to_compile_error(),
1080 );
1081 quote! { {
1082 #(#errors; )*
1083 todo!("insert")
1084 }}
1085 .into()
1086 }
1087 qusql_type::StatementType::Insert {
1088 arguments,
1089 returning: Some(returning),
1090 ..
1091 } => {
1092 let (args_tokens, q) = quote_args(
1093 &mut errors,
1094 &query_as.query,
1095 query_as.last_span,
1096 &query_as.args,
1097 arguments,
1098 dialect,
1099 );
1100
1101 let row_construct = construct_row2(returning, dialect.is_postgresql());
1102 let row = query_as.as_;
1103 let s = quote! { {
1104 use ::sqlx::Arguments as _;
1105 #(#errors; )*
1106 #args_tokens
1107 sqlx::__query_with_result(#q, query_args).map(|row|
1108 #row{
1109 #(#row_construct),*
1110 }
1111 )
1112 }};
1113 s.into()
1114 }
1115 qusql_type::StatementType::Update {
1116 returning: None, ..
1117 } => {
1118 errors.push(
1119 syn::Error::new(
1120 query_as.query_span,
1121 "UPDATE without RETURNING not support in query_as",
1122 )
1123 .to_compile_error(),
1124 );
1125 quote! { {
1126 #(#errors; )*
1127 todo!("update")
1128 }}
1129 .into()
1130 }
1131 qusql_type::StatementType::Update {
1132 arguments,
1133 returning: Some(returning),
1134 ..
1135 } => {
1136 let (args_tokens, q) = quote_args(
1137 &mut errors,
1138 &query_as.query,
1139 query_as.last_span,
1140 &query_as.args,
1141 arguments,
1142 dialect,
1143 );
1144
1145 let row_construct = construct_row2(returning, dialect.is_postgresql());
1146 let row = query_as.as_;
1147 let s = quote! { {
1148 use ::sqlx::Arguments as _;
1149 #(#errors; )*
1150 #args_tokens
1151 sqlx::__query_with_result(#q, query_args).map(|row|
1152 #row{
1153 #(#row_construct),*
1154 }
1155 )
1156 }};
1157 s.into()
1158 }
1159 qusql_type::StatementType::Replace {
1160 returning: None, ..
1161 } => {
1162 errors.push(
1163 syn::Error::new(
1164 query_as.query_span,
1165 "REPLACE without RETURNING not support in query_as",
1166 )
1167 .to_compile_error(),
1168 );
1169 quote! { {
1170 #(#errors; )*
1171 todo!("replace")
1172 }}
1173 .into()
1174 }
1175 qusql_type::StatementType::Replace {
1176 arguments,
1177 returning: Some(returning),
1178 ..
1179 } => {
1180 let (args_tokens, q) = quote_args(
1181 &mut errors,
1182 &query_as.query,
1183 query_as.last_span,
1184 &query_as.args,
1185 arguments,
1186 dialect,
1187 );
1188
1189 let row_construct = construct_row2(returning, dialect.is_postgresql());
1190 let row = query_as.as_;
1191 let s = quote! { {
1192 use ::sqlx::Arguments as _;
1193 #(#errors; )*
1194 #args_tokens
1195 sqlx::__query_with_result(#q, query_args).map(|row|
1196 #row{
1197 #(#row_construct),*
1198 }
1199 )
1200 }};
1201 s.into()
1202 }
1203 qusql_type::StatementType::Truncate => {
1204 errors.push(
1205 syn::Error::new(query_as.query_span, "TRUNCATE not supported in query_as!")
1206 .to_compile_error(),
1207 );
1208 quote! { {
1209 #(#errors; )*
1210 todo!("truncate")
1211 }}
1212 .into()
1213 }
1214 qusql_type::StatementType::Call { .. } => {
1215 errors.push(
1216 syn::Error::new(query_as.query_span, "CALL not supported in query_as!")
1217 .to_compile_error(),
1218 );
1219 quote! { {
1220 #(#errors; )*
1221 todo!("call")
1222 }}
1223 .into()
1224 }
1225 qusql_type::StatementType::Transaction => {
1226 errors.push(
1227 syn::Error::new(
1228 query_as.query_span,
1229 "Transaction control not supported in query_as!",
1230 )
1231 .to_compile_error(),
1232 );
1233 quote! { {
1234 #(#errors; )*
1235 todo!("transaction")
1236 }}
1237 .into()
1238 }
1239 qusql_type::StatementType::Set => {
1240 errors.push(
1241 syn::Error::new(query_as.query_span, "SET not supported in query_as!")
1242 .to_compile_error(),
1243 );
1244 quote! { {
1245 #(#errors; )*
1246 todo!("set")
1247 }}
1248 .into()
1249 }
1250 qusql_type::StatementType::Lock => {
1251 errors.push(
1252 syn::Error::new(query_as.query_span, "LOCK not supported in query_as!")
1253 .to_compile_error(),
1254 );
1255 quote! { {
1256 #(#errors; )*
1257 todo!("lock")
1258 }}
1259 .into()
1260 }
1261 qusql_type::StatementType::Invalid => quote! { {
1262 #(#errors; )*;
1263 todo!("invalid")
1264 }}
1265 .into(),
1266 }
1267}