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 => {
516 if is_postgres {
517 quote! {qusql_sqlx_type::JsonValue}
519 } else {
520 quote! {String}
521 }
522 }
523 qusql_type::Type::Geometry => quote! {Vec<u8>},
524 qusql_type::Type::Range(_) => quote! {Vec<u8>},
525 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
526 };
527 let name = match &c.name {
528 Some(v) => v,
529 None => continue,
530 };
531
532 let name_str = name.value;
534 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
535 (stripped, true)
536 } else {
537 (name_str, false)
538 };
539
540 let ident = String::from("r#") + name_str;
541 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
542 ident
543 } else {
544 continue;
547 };
548
549 if !c.type_.not_null && !force_not_null {
550 t = quote! {Option<#t>};
551 }
552 row_members.push(quote! {
553 #ident : #t
554 });
555 row_construct.push(quote! {
556 #ident: sqlx::Row::get(&row, #i)
557 });
558 }
559 (row_members, row_construct)
560}
561
562struct Query {
564 query: String,
566 query_span: Span,
568 args: Vec<Expr>,
570 last_span: Span,
572}
573
574impl Parse for Query {
575 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
576 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
577 let query: String = query_.iter().map(LitStr::value).collect();
578 let query_span = query_.span();
579 let mut last_span = query_span;
580 let mut args = Vec::new();
581 while !input.is_empty() {
582 let _ = input.parse::<syn::token::Comma>()?;
583 if input.is_empty() {
584 break;
585 }
586 let arg = input.parse::<Expr>()?;
587 last_span = arg.span();
588 args.push(arg);
589 }
590 Ok(Self {
591 query,
592 query_span,
593 args,
594 last_span,
595 })
596 }
597}
598
599#[proc_macro]
603pub fn query(input: TokenStream) -> TokenStream {
604 let query = syn::parse_macro_input!(input as Query);
605 let cache = get_schemas();
606 let (schemas, dialect, schema_hash) = (cache.schemas.get(), &cache.dialect, cache.hash);
607 let options = TypeOptions::new()
608 .dialect(dialect.clone())
609 .arguments(match &dialect {
610 SQLDialect::MariaDB => SQLArguments::QuestionMark,
611 SQLDialect::Sqlite => SQLArguments::QuestionMark,
612 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
613 })
614 .list_hack(true);
615 let mut issues = qusql_type::Issues::new(&query.query);
616 let stmt = type_statement(&schemas.0, &query.query, &mut issues, &options);
617 let sp = cache.path.to_str().unwrap();
618 let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
619 match &stmt {
620 qusql_type::StatementType::Select { columns, arguments } => {
621 let (args_tokens, q) = quote_args(
622 &mut errors,
623 &query.query,
624 query.last_span,
625 &query.args,
626 arguments,
627 dialect,
628 );
629 let (row_members, row_construct) = construct_row(columns, dialect.is_postgresql());
630 let s = quote! { {
631 use ::sqlx::Arguments as _;
632 use ::sqlx::SqlSafeStr;
633 let _ = std::include_bytes!(#sp);
634 #(#errors; )*
635 #args_tokens;
636
637 struct Row {
638 #(#row_members),*
639 };
640 sqlx::__query_with_result(#q, query_args).map(|row|
641 Row{
642 #(#row_construct),*
643 }
644 )
645 }};
646 s.into()
647 }
648 qusql_type::StatementType::Delete {
649 arguments,
650 returning,
651 } => {
652 let (args_tokens, q) = quote_args(
653 &mut errors,
654 &query.query,
655 query.last_span,
656 &query.args,
657 arguments,
658 dialect,
659 );
660 let s = match returning.as_ref() {
661 Some(returning) => {
662 let (row_members, row_construct) =
663 construct_row(returning, dialect.is_postgresql());
664 quote! { {
665 use ::sqlx::Arguments as _;
666 let _ = std::include_bytes!(#sp);
667 #(#errors; )*
668 #args_tokens
669
670 struct Row {
671 #(#row_members),*
672 };
673 sqlx::__query_with_result(#q, query_args).map(|row|
674 Row{
675 #(#row_construct),*
676 }
677 )
678 }}
679 }
680 None => quote! { {
681 use ::sqlx::Arguments as _;
682 const _SCHEMA_HASH: u64 = #schema_hash;
683 #(#errors; )*
684 #args_tokens
685 sqlx::__query_with_result(#q, query_args)
686 }
687 },
688 };
689 s.into()
690 }
691 qusql_type::StatementType::Insert {
692 arguments,
693 returning,
694 ..
695 } => {
696 let (args_tokens, q) = quote_args(
697 &mut errors,
698 &query.query,
699 query.last_span,
700 &query.args,
701 arguments,
702 dialect,
703 );
704 let s = match returning.as_ref() {
705 Some(returning) => {
706 let (row_members, row_construct) =
707 construct_row(returning, dialect.is_postgresql());
708 quote! { {
709 use ::sqlx::Arguments as _;
710 let _ = std::include_bytes!(#sp);
711 #(#errors; )*
712 #args_tokens
713
714 struct Row {
715 #(#row_members),*
716 };
717 sqlx::__query_with_result(#q, query_args).map(|row|
718 Row{
719 #(#row_construct),*
720 }
721 )
722 }}
723 }
724 None => quote! { {
725 use ::sqlx::Arguments as _;
726 const _SCHEMA_HASH: u64 = #schema_hash;
727 #(#errors; )*
728 #args_tokens
729 sqlx::__query_with_result(#q, query_args)
730 }
731 },
732 };
733 s.into()
734 }
735 qusql_type::StatementType::Update {
736 arguments,
737 returning,
738 } => {
739 let (args_tokens, q) = quote_args(
740 &mut errors,
741 &query.query,
742 query.last_span,
743 &query.args,
744 arguments,
745 dialect,
746 );
747
748 let s = match returning.as_ref() {
749 Some(returning) => {
750 let (row_members, row_construct) =
751 construct_row(returning, dialect.is_postgresql());
752 quote! { {
753 use ::sqlx::Arguments as _;
754 let _ = std::include_bytes!(#sp);
755 #(#errors; )*
756 #args_tokens
757
758 struct Row {
759 #(#row_members),*
760 };
761 sqlx::__query_with_result(#q, query_args).map(|row|
762 Row{
763 #(#row_construct),*
764 }
765 )
766 }}
767 }
768 None => quote! { {
769 use ::sqlx::Arguments as _;
770 const _SCHEMA_HASH: u64 = #schema_hash;
771 #(#errors; )*
772 #args_tokens
773 sqlx::__query_with_result(#q, query_args)
774 }
775 },
776 };
777 s.into()
778 }
779 qusql_type::StatementType::Replace {
780 arguments,
781 returning,
782 } => {
783 let (args_tokens, q) = quote_args(
784 &mut errors,
785 &query.query,
786 query.last_span,
787 &query.args,
788 arguments,
789 dialect,
790 );
791 let s = match returning.as_ref() {
792 Some(returning) => {
793 let (row_members, row_construct) =
794 construct_row(returning, dialect.is_postgresql());
795 quote! { {
796 use ::sqlx::Arguments as _;
797 const _SCHEMA_HASH: u64 = #schema_hash;
798 let _ = std::include_bytes!(#sp);
799 #(#errors; )*
800 #args_tokens
801
802 struct Row {
803 #(#row_members),*
804 };
805 sqlx::__query_with_result(#q, query_args).map(|row|
806 Row{
807 #(#row_construct),*
808 }
809 )
810 }}
811 }
812 None => quote! { {
813 use ::sqlx::Arguments as _;
814 #(#errors; )*
815 #args_tokens
816 sqlx::__query_with_result(#q, query_args)
817 }
818 },
819 };
820 s.into()
821 }
822 qusql_type::StatementType::Truncate => {
823 errors.push(
824 syn::Error::new(query.query_span, "TRUNCATE not supported in query!")
825 .to_compile_error(),
826 );
827 quote! { {
828 #(#errors; )*
829 todo!("truncate")
830 }}
831 .into()
832 }
833 qusql_type::StatementType::Call { .. } => {
834 errors.push(
835 syn::Error::new(query.query_span, "CALL not supported in query!")
836 .to_compile_error(),
837 );
838 quote! { {
839 #(#errors; )*
840 todo!("call")
841 }}
842 .into()
843 }
844 qusql_type::StatementType::Transaction => {
845 errors.push(
846 syn::Error::new(
847 query.query_span,
848 "Transaction control not supported in query!",
849 )
850 .to_compile_error(),
851 );
852 quote! { {
853 #(#errors; )*
854 todo!("transaction")
855 }}
856 .into()
857 }
858 qusql_type::StatementType::Set => {
859 errors.push(
860 syn::Error::new(query.query_span, "SET not supported in query!").to_compile_error(),
861 );
862 quote! { {
863 #(#errors; )*
864 todo!("set")
865 }}
866 .into()
867 }
868 qusql_type::StatementType::Lock => {
869 errors.push(
870 syn::Error::new(query.query_span, "LOCK not supported in query!")
871 .to_compile_error(),
872 );
873 quote! { {
874 #(#errors; )*
875 todo!("lock")
876 }}
877 .into()
878 }
879 qusql_type::StatementType::Invalid => {
880 let s = quote! { {
881 #(#errors; )*;
882 todo!("Invalid")
883 }};
884 s.into()
885 }
886 }
887}
888
889fn construct_row2(
891 columns: &[SelectTypeColumn],
892 is_postgres: bool,
893) -> Vec<proc_macro2::TokenStream> {
894 let mut row_construct = Vec::new();
895 for (i, c) in columns.iter().enumerate() {
896 let mut t = match c.type_.t {
897 qusql_type::Type::U8 => quote! {u8},
898 qusql_type::Type::I8 => quote! {i8},
899 qusql_type::Type::U16 => quote! {u16},
900 qusql_type::Type::I16 => quote! {i16},
901 qusql_type::Type::U24 => quote! {u32},
902 qusql_type::Type::I24 => quote! {i32},
903 qusql_type::Type::U32 => quote! {u32},
904 qusql_type::Type::I32 => quote! {i32},
905 qusql_type::Type::U64 => quote! {u64},
906 qusql_type::Type::I64 => quote! {i64},
907 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
908 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
909 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
910 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
911 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
912 quote! {chrono::NaiveDateTime}
913 }
914 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
915 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
916 if is_postgres {
917 quote! {i32}
918 } else {
919 quote! {i64}
920 }
921 }
922 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
923 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
924 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
925 todo!("from_time_interval")
926 }
927 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
928 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
929 }
930 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
931 quote! {qusql_sqlx_type::UuidValue}
932 }
933 qusql_type::Type::Null => todo!("from_null"),
934 qusql_type::Type::Invalid => quote! {i64},
935 qusql_type::Type::Enum(_) => quote! {String},
936 qusql_type::Type::Set(_) => quote! {String},
937 qusql_type::Type::Args(_, _) => todo!("from_args"),
938 qusql_type::Type::F32 => quote! {f32},
939 qusql_type::Type::F64 => quote! {f64},
940 qusql_type::Type::JSON => {
941 if is_postgres {
942 quote! {qusql_sqlx_type::JsonValue}
944 } else {
945 quote! {String}
946 }
947 }
948 qusql_type::Type::Geometry => quote! {Vec<u8>},
949 qusql_type::Type::Range(_) => quote! {Vec<u8>},
950 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
951 };
952 let name = match &c.name {
953 Some(v) => v,
954 None => continue,
955 };
956
957 let name_str = name.value;
959 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
960 (stripped, true)
961 } else {
962 (name_str, false)
963 };
964
965 let ident = String::from("r#") + name_str;
966 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
967 ident
968 } else {
969 continue;
972 };
973
974 if !c.type_.not_null && !force_not_null {
975 t = quote! {Option<#t>};
976 }
977 row_construct.push(quote! {
978 #ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
979 });
980 }
981 row_construct
982}
983
984struct QueryAs {
986 as_: Ident,
988 query: String,
990 query_span: Span,
992 args: Vec<Expr>,
994 last_span: Span,
996}
997
998impl Parse for QueryAs {
999 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
1000 let as_ = input.parse::<Ident>()?;
1001 let _ = input.parse::<syn::token::Comma>()?;
1002
1003 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
1004 let query: String = query_.iter().map(LitStr::value).collect();
1005 let query_span = query_.span();
1006
1007 let mut last_span = query_span;
1008 let mut args = Vec::new();
1009 while !input.is_empty() {
1010 let _ = input.parse::<syn::token::Comma>()?;
1011 if input.is_empty() {
1012 break;
1013 }
1014 let arg = input.parse::<Expr>()?;
1015 last_span = arg.span();
1016 args.push(arg);
1017 }
1018 Ok(Self {
1019 as_,
1020 query,
1021 query_span,
1022 args,
1023 last_span,
1024 })
1025 }
1026}
1027
1028#[proc_macro]
1032pub fn query_as(input: TokenStream) -> TokenStream {
1033 let query_as = syn::parse_macro_input!(input as QueryAs);
1034 let cache = get_schemas();
1035 let (schemas, dialect) = (cache.schemas.get(), &cache.dialect);
1036 let options = TypeOptions::new()
1037 .dialect(dialect.clone())
1038 .arguments(match &dialect {
1039 SQLDialect::MariaDB => SQLArguments::QuestionMark,
1040 SQLDialect::Sqlite => SQLArguments::QuestionMark,
1041 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
1042 })
1043 .list_hack(true);
1044 let mut issues = qusql_type::Issues::new(&query_as.query);
1045 let stmt = type_statement(&schemas.0, &query_as.query, &mut issues, &options);
1046
1047 let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
1048 match &stmt {
1049 qusql_type::StatementType::Select { columns, arguments } => {
1050 let (args_tokens, q) = quote_args(
1051 &mut errors,
1052 &query_as.query,
1053 query_as.last_span,
1054 &query_as.args,
1055 arguments,
1056 dialect,
1057 );
1058
1059 let row_construct = construct_row2(columns, dialect.is_postgresql());
1060 let row = query_as.as_;
1061 let s = quote! { {
1062 use ::sqlx::Arguments as _;
1063 #(#errors; )*
1064 #args_tokens
1065 sqlx::__query_with_result(#q, query_args).map(|row|
1066 #row{
1067 #(#row_construct),*
1068 }
1069 )
1070 }};
1071 s.into()
1073 }
1074 qusql_type::StatementType::Delete { .. } => {
1075 errors.push(
1076 syn::Error::new(query_as.query_span, "DELETE not support in query_as")
1077 .to_compile_error(),
1078 );
1079 quote! { {
1080 #(#errors; )*
1081 todo!("delete")
1082 }}
1083 .into()
1084 }
1085 qusql_type::StatementType::Insert {
1086 returning: None, ..
1087 } => {
1088 errors.push(
1089 syn::Error::new(
1090 query_as.query_span,
1091 "INSERT without RETURNING not support in query_as",
1092 )
1093 .to_compile_error(),
1094 );
1095 quote! { {
1096 #(#errors; )*
1097 todo!("insert")
1098 }}
1099 .into()
1100 }
1101 qusql_type::StatementType::Insert {
1102 arguments,
1103 returning: Some(returning),
1104 ..
1105 } => {
1106 let (args_tokens, q) = quote_args(
1107 &mut errors,
1108 &query_as.query,
1109 query_as.last_span,
1110 &query_as.args,
1111 arguments,
1112 dialect,
1113 );
1114
1115 let row_construct = construct_row2(returning, dialect.is_postgresql());
1116 let row = query_as.as_;
1117 let s = quote! { {
1118 use ::sqlx::Arguments as _;
1119 #(#errors; )*
1120 #args_tokens
1121 sqlx::__query_with_result(#q, query_args).map(|row|
1122 #row{
1123 #(#row_construct),*
1124 }
1125 )
1126 }};
1127 s.into()
1128 }
1129 qusql_type::StatementType::Update {
1130 returning: None, ..
1131 } => {
1132 errors.push(
1133 syn::Error::new(
1134 query_as.query_span,
1135 "UPDATE without RETURNING not support in query_as",
1136 )
1137 .to_compile_error(),
1138 );
1139 quote! { {
1140 #(#errors; )*
1141 todo!("update")
1142 }}
1143 .into()
1144 }
1145 qusql_type::StatementType::Update {
1146 arguments,
1147 returning: Some(returning),
1148 ..
1149 } => {
1150 let (args_tokens, q) = quote_args(
1151 &mut errors,
1152 &query_as.query,
1153 query_as.last_span,
1154 &query_as.args,
1155 arguments,
1156 dialect,
1157 );
1158
1159 let row_construct = construct_row2(returning, dialect.is_postgresql());
1160 let row = query_as.as_;
1161 let s = quote! { {
1162 use ::sqlx::Arguments as _;
1163 #(#errors; )*
1164 #args_tokens
1165 sqlx::__query_with_result(#q, query_args).map(|row|
1166 #row{
1167 #(#row_construct),*
1168 }
1169 )
1170 }};
1171 s.into()
1172 }
1173 qusql_type::StatementType::Replace {
1174 returning: None, ..
1175 } => {
1176 errors.push(
1177 syn::Error::new(
1178 query_as.query_span,
1179 "REPLACE without RETURNING not support in query_as",
1180 )
1181 .to_compile_error(),
1182 );
1183 quote! { {
1184 #(#errors; )*
1185 todo!("replace")
1186 }}
1187 .into()
1188 }
1189 qusql_type::StatementType::Replace {
1190 arguments,
1191 returning: Some(returning),
1192 ..
1193 } => {
1194 let (args_tokens, q) = quote_args(
1195 &mut errors,
1196 &query_as.query,
1197 query_as.last_span,
1198 &query_as.args,
1199 arguments,
1200 dialect,
1201 );
1202
1203 let row_construct = construct_row2(returning, dialect.is_postgresql());
1204 let row = query_as.as_;
1205 let s = quote! { {
1206 use ::sqlx::Arguments as _;
1207 #(#errors; )*
1208 #args_tokens
1209 sqlx::__query_with_result(#q, query_args).map(|row|
1210 #row{
1211 #(#row_construct),*
1212 }
1213 )
1214 }};
1215 s.into()
1216 }
1217 qusql_type::StatementType::Truncate => {
1218 errors.push(
1219 syn::Error::new(query_as.query_span, "TRUNCATE not supported in query_as!")
1220 .to_compile_error(),
1221 );
1222 quote! { {
1223 #(#errors; )*
1224 todo!("truncate")
1225 }}
1226 .into()
1227 }
1228 qusql_type::StatementType::Call { .. } => {
1229 errors.push(
1230 syn::Error::new(query_as.query_span, "CALL not supported in query_as!")
1231 .to_compile_error(),
1232 );
1233 quote! { {
1234 #(#errors; )*
1235 todo!("call")
1236 }}
1237 .into()
1238 }
1239 qusql_type::StatementType::Transaction => {
1240 errors.push(
1241 syn::Error::new(
1242 query_as.query_span,
1243 "Transaction control not supported in query_as!",
1244 )
1245 .to_compile_error(),
1246 );
1247 quote! { {
1248 #(#errors; )*
1249 todo!("transaction")
1250 }}
1251 .into()
1252 }
1253 qusql_type::StatementType::Set => {
1254 errors.push(
1255 syn::Error::new(query_as.query_span, "SET not supported in query_as!")
1256 .to_compile_error(),
1257 );
1258 quote! { {
1259 #(#errors; )*
1260 todo!("set")
1261 }}
1262 .into()
1263 }
1264 qusql_type::StatementType::Lock => {
1265 errors.push(
1266 syn::Error::new(query_as.query_span, "LOCK not supported in query_as!")
1267 .to_compile_error(),
1268 );
1269 quote! { {
1270 #(#errors; )*
1271 todo!("lock")
1272 }}
1273 .into()
1274 }
1275 qusql_type::StatementType::Invalid => quote! { {
1276 #(#errors; )*;
1277 todo!("invalid")
1278 }}
1279 .into(),
1280 }
1281}