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