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_range_elem_type(
284 elem: &qusql_type::Type<'_>,
285 is_postgres: bool,
286 fallback: proc_macro2::TokenStream,
287) -> proc_macro2::TokenStream {
288 if !is_postgres {
289 return fallback;
290 }
291 match elem {
292 qusql_type::Type::I32 => quote! { qusql_sqlx_type::PgRange<i32> },
294 qusql_type::Type::I64 => quote! { qusql_sqlx_type::PgRange<i64> },
296 qusql_type::Type::Base(qusql_type::BaseType::Date) => {
297 quote! { qusql_sqlx_type::PgRange<chrono::NaiveDate> }
298 }
299 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
300 quote! { qusql_sqlx_type::PgRange<chrono::NaiveDateTime> }
301 }
302 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
303 quote! { qusql_sqlx_type::PgRange<chrono::DateTime<chrono::Utc>> }
304 }
305 _ => fallback,
306 }
307}
308
309fn quote_args(
311 errors: &mut Vec<proc_macro2::TokenStream>,
312 query: &str,
313 last_span: Span,
314 args: &[Expr],
315 arguments: &[(qusql_type::ArgumentKey<'_>, qusql_type::FullType)],
316 dialect: &SQLDialect,
317) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
318 let cls = match dialect {
319 SQLDialect::MariaDB => quote!(sqlx::mysql::MySql),
320 SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite),
321 SQLDialect::PostgreSQL | SQLDialect::PostGIS => quote!(sqlx::postgres::Postgres),
322 };
323 let is_postgres = dialect.is_postgresql();
324
325 let mut at = Vec::new();
326 let inv = qusql_type::FullType::invalid();
327 for (k, v) in arguments {
328 match k {
329 qusql_type::ArgumentKey::Index(i) => {
330 while at.len() <= *i {
331 at.push(&inv);
332 }
333 at[*i] = v;
334 }
335 qusql_type::ArgumentKey::Identifier(_) => {
336 errors.push(
337 syn::Error::new(last_span.span(), "Named arguments not supported")
338 .to_compile_error(),
339 );
340 }
341 }
342 }
343
344 if at.len() > args.len() {
345 errors.push(
346 syn::Error::new(
347 last_span,
348 format!("Expected {} additional arguments", at.len() - args.len()),
349 )
350 .to_compile_error(),
351 );
352 }
353
354 if let Some(args) = args.get(at.len()..) {
355 for arg in args {
356 errors.push(syn::Error::new(arg.span(), "unexpected argument").to_compile_error());
357 }
358 }
359
360 let arg_names = (0..args.len())
361 .map(|i| format_ident!("arg{}", i))
362 .collect::<Vec<_>>();
363
364 let mut arg_bindings = Vec::new();
365 let mut arg_add = Vec::new();
366
367 let mut list_lengths = Vec::new();
368
369 for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) {
370 let mut t = match &ta.t {
371 qusql_type::Type::U8 => quote! {u8},
372 qusql_type::Type::I8 => quote! {i8},
373 qusql_type::Type::U16 => quote! {u16},
374 qusql_type::Type::I16 => quote! {i16},
375 qusql_type::Type::U24 => quote! {u32},
376 qusql_type::Type::I24 => quote! {i32},
377 qusql_type::Type::U32 => quote! {u32},
378 qusql_type::Type::I32 => quote! {i32},
379 qusql_type::Type::U64 => quote! {u64},
380 qusql_type::Type::I64 => quote! {i64},
381 qusql_type::Type::Base(qusql_type::BaseType::Any) => quote! {qusql_sqlx_type::Any},
382 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
383 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {&[u8]},
384 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {qusql_sqlx_type::Date},
385 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
386 quote! {qusql_sqlx_type::DateTime}
387 }
388 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {qusql_sqlx_type::Float},
389 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
390 quote! {qusql_sqlx_type::Integer}
391 }
392 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {&str},
393 qusql_type::Type::Base(qusql_type::BaseType::Time) => quote! {qusql_sqlx_type::Time},
394 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => todo!("time_interval"),
395 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
396 quote! {qusql_sqlx_type::Timestamp}
397 }
398 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => quote! {qusql_sqlx_type::Uuid},
399 qusql_type::Type::Null => todo!("null"),
400 qusql_type::Type::Invalid => quote! {std::convert::Infallible},
401 qusql_type::Type::Enum(_) => quote! {&str},
402 qusql_type::Type::Set(_) => quote! {&str},
403 qusql_type::Type::Args(_, _) => todo!("args"),
404 qusql_type::Type::F32 => quote! {f32},
405 qusql_type::Type::F64 => quote! {f64},
406 qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any},
407 qusql_type::Type::Geometry => quote! {qusql_sqlx_type::Any},
408 qusql_type::Type::Range(elem) => {
409 quote_range_elem_type(elem, is_postgres, quote! {qusql_sqlx_type::Any})
410 }
411 qusql_type::Type::MultiRange(_) => quote! {qusql_sqlx_type::Any},
412 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
413 };
414 if !ta.not_null {
415 t = quote! {Option<#t>}
416 }
417 let span = qa.span();
418 if ta.list_hack {
419 list_lengths.push(quote!(#name.len()));
420 arg_bindings.push(quote_spanned! {span=>
421 let #name = &(#qa);
422 args_count += #name.len();
423 for v in #name.iter() {
424 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(v);
425 }
426 if false {
427 qusql_sqlx_type::check_arg_list_hack::<#t, _>(#name);
428 ::std::panic!();
429 }
430 });
431 arg_add.push(quote!(
432 for v in #name.iter() {
433 e = e.and_then(|()| query_args.add(v));
434 }
435 ));
436 } else {
437 arg_bindings.push(quote_spanned! {span=>
438 let #name = &(#qa);
439 args_count += 1;
440 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(#name);
441 if false {
442 qusql_sqlx_type::check_arg::<#t, _>(#name);
443 ::std::panic!();
444 }
445 });
446 arg_add.push(quote!(e = e.and_then(|()| query_args.add(#name));));
447 }
448 }
449
450 let query = if list_lengths.is_empty() {
451 quote!(#query)
452 } else {
453 quote!(
454 qusql_sqlx_type::convert_list_query(#query, &[#(#list_lengths),*])
455 )
456 };
457
458 (
459 quote! {
460 let mut size_hints = 0;
461 let mut args_count = 0;
462 #(#arg_bindings)*
463
464 let mut query_args = <#cls as ::sqlx::database::Database>::Arguments::default();
465 query_args.reserve(args_count, size_hints);
466 let mut e = Ok(());
467 #(#arg_add)*
468 let query_args = e.and_then(|()| Ok(query_args));
469 },
470 query,
471 )
472}
473
474fn issues_to_errors(issues: Vec<Issue>, source: &str, span: Span) -> Vec<proc_macro2::TokenStream> {
476 if !issues.is_empty() {
477 let b2c = ByteToChar::new(source.as_bytes());
478 let source = NamedSource("query", Source::from(source));
479 let mut err = false;
480 let mut out = Vec::new();
481 for issue in issues {
482 if issue.level == qusql_type::Level::Error {
483 err = true;
484 }
485 let r = issue_to_report(issue, &b2c);
486 r.write(&source, &mut out).unwrap();
487 }
488 if err {
489 let raw = String::from_utf8(out).unwrap();
490 let body = raw
493 .find('\n')
494 .map(|i| raw[i + 1..].trim_start_matches('\n').trim_end())
495 .unwrap_or(raw.trim_end());
496 return vec![syn::Error::new(span, body).to_compile_error()];
497 }
498 }
499 Vec::new()
500}
501
502fn construct_row(
504 columns: &[SelectTypeColumn],
505 is_postgres: bool,
506) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
507 let mut row_members = Vec::new();
508 let mut row_construct = Vec::new();
509 for (i, c) in columns.iter().enumerate() {
510 let mut t = match &c.type_.t {
511 qusql_type::Type::U8 => quote! {u8},
512 qusql_type::Type::I8 => quote! {i8},
513 qusql_type::Type::U16 => quote! {u16},
514 qusql_type::Type::I16 => quote! {i16},
515 qusql_type::Type::U24 => quote! {u32},
516 qusql_type::Type::I24 => quote! {i32},
517 qusql_type::Type::U32 => quote! {u32},
518 qusql_type::Type::I32 => quote! {i32},
519 qusql_type::Type::U64 => quote! {u64},
520 qusql_type::Type::I64 => quote! {i64},
521 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
522 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
523 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
524 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
525 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
526 quote! {chrono::NaiveDateTime}
527 }
528 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
529 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
530 if is_postgres {
531 quote! {i32}
532 } else {
533 quote! {i64}
534 }
535 }
536 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
537 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
538 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
539 todo!("from_time_interval")
540 }
541 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
542 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
543 }
544 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
545 quote! {qusql_sqlx_type::UuidValue}
546 }
547 qusql_type::Type::Null => todo!("from_null"),
548 qusql_type::Type::Invalid => quote! {i64},
549 qusql_type::Type::Enum(_) => quote! {String},
550 qusql_type::Type::Set(_) => quote! {String},
551 qusql_type::Type::Args(_, _) => todo!("from_args"),
552 qusql_type::Type::F32 => quote! {f32},
553 qusql_type::Type::F64 => quote! {f64},
554 qusql_type::Type::JSON => {
555 if is_postgres {
556 quote! {qusql_sqlx_type::JsonValue}
558 } else {
559 quote! {String}
560 }
561 }
562 qusql_type::Type::Geometry => quote! {Vec<u8>},
563 qusql_type::Type::MultiRange(_) => quote! {Vec<u8>},
564 qusql_type::Type::Range(elem) => {
565 quote_range_elem_type(elem, is_postgres, quote! {Vec<u8>})
566 }
567 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
568 };
569 let name = match &c.name {
570 Some(v) => v,
571 None => continue,
572 };
573
574 let name_str = name.value;
576 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
577 (stripped, true)
578 } else {
579 (name_str, false)
580 };
581
582 let ident = String::from("r#") + name_str;
583 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
584 ident
585 } else {
586 continue;
589 };
590
591 if !c.type_.not_null && !force_not_null {
592 t = quote! {Option<#t>};
593 }
594 row_members.push(quote! {
595 #ident : #t
596 });
597 row_construct.push(quote! {
598 #ident: sqlx::Row::get(&row, #i)
599 });
600 }
601 (row_members, row_construct)
602}
603
604struct Query {
606 query: String,
608 query_span: Span,
610 args: Vec<Expr>,
612 last_span: Span,
614}
615
616impl Parse for Query {
617 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
618 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
619 let query: String = query_.iter().map(LitStr::value).collect();
620 let query_span = query_.span();
621 let mut last_span = query_span;
622 let mut args = Vec::new();
623 while !input.is_empty() {
624 let _ = input.parse::<syn::token::Comma>()?;
625 if input.is_empty() {
626 break;
627 }
628 let arg = input.parse::<Expr>()?;
629 last_span = arg.span();
630 args.push(arg);
631 }
632 Ok(Self {
633 query,
634 query_span,
635 args,
636 last_span,
637 })
638 }
639}
640
641#[proc_macro]
645pub fn query(input: TokenStream) -> TokenStream {
646 let query = syn::parse_macro_input!(input as Query);
647 let cache = get_schemas();
648 let (schemas, dialect, schema_hash) = (cache.schemas.get(), &cache.dialect, cache.hash);
649 let options = TypeOptions::new()
650 .dialect(dialect.clone())
651 .arguments(match &dialect {
652 SQLDialect::MariaDB => SQLArguments::QuestionMark,
653 SQLDialect::Sqlite => SQLArguments::QuestionMark,
654 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
655 })
656 .list_hack(true);
657 let mut issues = qusql_type::Issues::new(&query.query);
658 let stmt = type_statement(&schemas.0, &query.query, &mut issues, &options);
659 let sp = cache.path.to_str().unwrap();
660 let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
661 match &stmt {
662 qusql_type::StatementType::Select { columns, arguments } => {
663 let (args_tokens, q) = quote_args(
664 &mut errors,
665 &query.query,
666 query.last_span,
667 &query.args,
668 arguments,
669 dialect,
670 );
671 let (row_members, row_construct) = construct_row(columns, dialect.is_postgresql());
672 let s = quote! { {
673 use ::sqlx::Arguments as _;
674 use ::sqlx::SqlSafeStr;
675 let _ = std::include_bytes!(#sp);
676 #(#errors; )*
677 #args_tokens;
678
679 struct Row {
680 #(#row_members),*
681 };
682 sqlx::__query_with_result(#q, query_args).map(|row|
683 Row{
684 #(#row_construct),*
685 }
686 )
687 }};
688 s.into()
689 }
690 qusql_type::StatementType::Delete {
691 arguments,
692 returning,
693 } => {
694 let (args_tokens, q) = quote_args(
695 &mut errors,
696 &query.query,
697 query.last_span,
698 &query.args,
699 arguments,
700 dialect,
701 );
702 let s = match returning.as_ref() {
703 Some(returning) => {
704 let (row_members, row_construct) =
705 construct_row(returning, dialect.is_postgresql());
706 quote! { {
707 use ::sqlx::Arguments as _;
708 let _ = std::include_bytes!(#sp);
709 #(#errors; )*
710 #args_tokens
711
712 struct Row {
713 #(#row_members),*
714 };
715 sqlx::__query_with_result(#q, query_args).map(|row|
716 Row{
717 #(#row_construct),*
718 }
719 )
720 }}
721 }
722 None => quote! { {
723 use ::sqlx::Arguments as _;
724 const _SCHEMA_HASH: u64 = #schema_hash;
725 #(#errors; )*
726 #args_tokens
727 sqlx::__query_with_result(#q, query_args)
728 }
729 },
730 };
731 s.into()
732 }
733 qusql_type::StatementType::Insert {
734 arguments,
735 returning,
736 ..
737 } => {
738 let (args_tokens, q) = quote_args(
739 &mut errors,
740 &query.query,
741 query.last_span,
742 &query.args,
743 arguments,
744 dialect,
745 );
746 let s = match returning.as_ref() {
747 Some(returning) => {
748 let (row_members, row_construct) =
749 construct_row(returning, dialect.is_postgresql());
750 quote! { {
751 use ::sqlx::Arguments as _;
752 let _ = std::include_bytes!(#sp);
753 #(#errors; )*
754 #args_tokens
755
756 struct Row {
757 #(#row_members),*
758 };
759 sqlx::__query_with_result(#q, query_args).map(|row|
760 Row{
761 #(#row_construct),*
762 }
763 )
764 }}
765 }
766 None => quote! { {
767 use ::sqlx::Arguments as _;
768 const _SCHEMA_HASH: u64 = #schema_hash;
769 #(#errors; )*
770 #args_tokens
771 sqlx::__query_with_result(#q, query_args)
772 }
773 },
774 };
775 s.into()
776 }
777 qusql_type::StatementType::Update {
778 arguments,
779 returning,
780 } => {
781 let (args_tokens, q) = quote_args(
782 &mut errors,
783 &query.query,
784 query.last_span,
785 &query.args,
786 arguments,
787 dialect,
788 );
789
790 let s = match returning.as_ref() {
791 Some(returning) => {
792 let (row_members, row_construct) =
793 construct_row(returning, dialect.is_postgresql());
794 quote! { {
795 use ::sqlx::Arguments as _;
796 let _ = std::include_bytes!(#sp);
797 #(#errors; )*
798 #args_tokens
799
800 struct Row {
801 #(#row_members),*
802 };
803 sqlx::__query_with_result(#q, query_args).map(|row|
804 Row{
805 #(#row_construct),*
806 }
807 )
808 }}
809 }
810 None => quote! { {
811 use ::sqlx::Arguments as _;
812 const _SCHEMA_HASH: u64 = #schema_hash;
813 #(#errors; )*
814 #args_tokens
815 sqlx::__query_with_result(#q, query_args)
816 }
817 },
818 };
819 s.into()
820 }
821 qusql_type::StatementType::Replace {
822 arguments,
823 returning,
824 } => {
825 let (args_tokens, q) = quote_args(
826 &mut errors,
827 &query.query,
828 query.last_span,
829 &query.args,
830 arguments,
831 dialect,
832 );
833 let s = match returning.as_ref() {
834 Some(returning) => {
835 let (row_members, row_construct) =
836 construct_row(returning, dialect.is_postgresql());
837 quote! { {
838 use ::sqlx::Arguments as _;
839 const _SCHEMA_HASH: u64 = #schema_hash;
840 let _ = std::include_bytes!(#sp);
841 #(#errors; )*
842 #args_tokens
843
844 struct Row {
845 #(#row_members),*
846 };
847 sqlx::__query_with_result(#q, query_args).map(|row|
848 Row{
849 #(#row_construct),*
850 }
851 )
852 }}
853 }
854 None => quote! { {
855 use ::sqlx::Arguments as _;
856 #(#errors; )*
857 #args_tokens
858 sqlx::__query_with_result(#q, query_args)
859 }
860 },
861 };
862 s.into()
863 }
864 qusql_type::StatementType::Truncate => {
865 errors.push(
866 syn::Error::new(query.query_span, "TRUNCATE not supported in query!")
867 .to_compile_error(),
868 );
869 quote! { {
870 #(#errors; )*
871 todo!("truncate")
872 }}
873 .into()
874 }
875 qusql_type::StatementType::Call { .. } => {
876 errors.push(
877 syn::Error::new(query.query_span, "CALL not supported in query!")
878 .to_compile_error(),
879 );
880 quote! { {
881 #(#errors; )*
882 todo!("call")
883 }}
884 .into()
885 }
886 qusql_type::StatementType::Transaction => {
887 errors.push(
888 syn::Error::new(
889 query.query_span,
890 "Transaction control not supported in query!",
891 )
892 .to_compile_error(),
893 );
894 quote! { {
895 #(#errors; )*
896 todo!("transaction")
897 }}
898 .into()
899 }
900 qusql_type::StatementType::Set => {
901 errors.push(
902 syn::Error::new(query.query_span, "SET not supported in query!").to_compile_error(),
903 );
904 quote! { {
905 #(#errors; )*
906 todo!("set")
907 }}
908 .into()
909 }
910 qusql_type::StatementType::Lock => {
911 errors.push(
912 syn::Error::new(query.query_span, "LOCK not supported in query!")
913 .to_compile_error(),
914 );
915 quote! { {
916 #(#errors; )*
917 todo!("lock")
918 }}
919 .into()
920 }
921 qusql_type::StatementType::Invalid => {
922 let s = quote! { {
923 #(#errors; )*;
924 todo!("Invalid")
925 }};
926 s.into()
927 }
928 }
929}
930
931fn construct_row2(
933 columns: &[SelectTypeColumn],
934 is_postgres: bool,
935) -> Vec<proc_macro2::TokenStream> {
936 let mut row_construct = Vec::new();
937 for (i, c) in columns.iter().enumerate() {
938 let mut t = match &c.type_.t {
939 qusql_type::Type::U8 => quote! {u8},
940 qusql_type::Type::I8 => quote! {i8},
941 qusql_type::Type::U16 => quote! {u16},
942 qusql_type::Type::I16 => quote! {i16},
943 qusql_type::Type::U24 => quote! {u32},
944 qusql_type::Type::I24 => quote! {i32},
945 qusql_type::Type::U32 => quote! {u32},
946 qusql_type::Type::I32 => quote! {i32},
947 qusql_type::Type::U64 => quote! {u64},
948 qusql_type::Type::I64 => quote! {i64},
949 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
950 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
951 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
952 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
953 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
954 quote! {chrono::NaiveDateTime}
955 }
956 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
957 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
958 if is_postgres {
959 quote! {i32}
960 } else {
961 quote! {i64}
962 }
963 }
964 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
965 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
966 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
967 todo!("from_time_interval")
968 }
969 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
970 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
971 }
972 qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
973 quote! {qusql_sqlx_type::UuidValue}
974 }
975 qusql_type::Type::Null => todo!("from_null"),
976 qusql_type::Type::Invalid => quote! {i64},
977 qusql_type::Type::Enum(_) => quote! {String},
978 qusql_type::Type::Set(_) => quote! {String},
979 qusql_type::Type::Args(_, _) => todo!("from_args"),
980 qusql_type::Type::F32 => quote! {f32},
981 qusql_type::Type::F64 => quote! {f64},
982 qusql_type::Type::JSON => {
983 if is_postgres {
984 quote! {qusql_sqlx_type::JsonValue}
986 } else {
987 quote! {String}
988 }
989 }
990 qusql_type::Type::Geometry => quote! {Vec<u8>},
991 qusql_type::Type::MultiRange(_) => quote! {Vec<u8>},
992 qusql_type::Type::Range(elem) => {
993 quote_range_elem_type(elem, is_postgres, quote! {Vec<u8>})
994 }
995 qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
996 };
997 let name = match &c.name {
998 Some(v) => v,
999 None => continue,
1000 };
1001
1002 let name_str = name.value;
1004 let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
1005 (stripped, true)
1006 } else {
1007 (name_str, false)
1008 };
1009
1010 let ident = String::from("r#") + name_str;
1011 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
1012 ident
1013 } else {
1014 continue;
1017 };
1018
1019 if !c.type_.not_null && !force_not_null {
1020 t = quote! {Option<#t>};
1021 }
1022 row_construct.push(quote! {
1023 #ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
1024 });
1025 }
1026 row_construct
1027}
1028
1029struct QueryAs {
1031 as_: Ident,
1033 query: String,
1035 query_span: Span,
1037 args: Vec<Expr>,
1039 last_span: Span,
1041}
1042
1043impl Parse for QueryAs {
1044 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
1045 let as_ = input.parse::<Ident>()?;
1046 let _ = input.parse::<syn::token::Comma>()?;
1047
1048 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
1049 let query: String = query_.iter().map(LitStr::value).collect();
1050 let query_span = query_.span();
1051
1052 let mut last_span = query_span;
1053 let mut args = Vec::new();
1054 while !input.is_empty() {
1055 let _ = input.parse::<syn::token::Comma>()?;
1056 if input.is_empty() {
1057 break;
1058 }
1059 let arg = input.parse::<Expr>()?;
1060 last_span = arg.span();
1061 args.push(arg);
1062 }
1063 Ok(Self {
1064 as_,
1065 query,
1066 query_span,
1067 args,
1068 last_span,
1069 })
1070 }
1071}
1072
1073#[proc_macro]
1077pub fn query_as(input: TokenStream) -> TokenStream {
1078 let query_as = syn::parse_macro_input!(input as QueryAs);
1079 let cache = get_schemas();
1080 let (schemas, dialect) = (cache.schemas.get(), &cache.dialect);
1081 let options = TypeOptions::new()
1082 .dialect(dialect.clone())
1083 .arguments(match &dialect {
1084 SQLDialect::MariaDB => SQLArguments::QuestionMark,
1085 SQLDialect::Sqlite => SQLArguments::QuestionMark,
1086 SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
1087 })
1088 .list_hack(true);
1089 let mut issues = qusql_type::Issues::new(&query_as.query);
1090 let stmt = type_statement(&schemas.0, &query_as.query, &mut issues, &options);
1091
1092 let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
1093 match &stmt {
1094 qusql_type::StatementType::Select { columns, arguments } => {
1095 let (args_tokens, q) = quote_args(
1096 &mut errors,
1097 &query_as.query,
1098 query_as.last_span,
1099 &query_as.args,
1100 arguments,
1101 dialect,
1102 );
1103
1104 let row_construct = construct_row2(columns, dialect.is_postgresql());
1105 let row = query_as.as_;
1106 let s = quote! { {
1107 use ::sqlx::Arguments as _;
1108 #(#errors; )*
1109 #args_tokens
1110 sqlx::__query_with_result(#q, query_args).map(|row|
1111 #row{
1112 #(#row_construct),*
1113 }
1114 )
1115 }};
1116 s.into()
1118 }
1119 qusql_type::StatementType::Delete { .. } => {
1120 errors.push(
1121 syn::Error::new(query_as.query_span, "DELETE not support in query_as")
1122 .to_compile_error(),
1123 );
1124 quote! { {
1125 #(#errors; )*
1126 todo!("delete")
1127 }}
1128 .into()
1129 }
1130 qusql_type::StatementType::Insert {
1131 returning: None, ..
1132 } => {
1133 errors.push(
1134 syn::Error::new(
1135 query_as.query_span,
1136 "INSERT without RETURNING not support in query_as",
1137 )
1138 .to_compile_error(),
1139 );
1140 quote! { {
1141 #(#errors; )*
1142 todo!("insert")
1143 }}
1144 .into()
1145 }
1146 qusql_type::StatementType::Insert {
1147 arguments,
1148 returning: Some(returning),
1149 ..
1150 } => {
1151 let (args_tokens, q) = quote_args(
1152 &mut errors,
1153 &query_as.query,
1154 query_as.last_span,
1155 &query_as.args,
1156 arguments,
1157 dialect,
1158 );
1159
1160 let row_construct = construct_row2(returning, dialect.is_postgresql());
1161 let row = query_as.as_;
1162 let s = quote! { {
1163 use ::sqlx::Arguments as _;
1164 #(#errors; )*
1165 #args_tokens
1166 sqlx::__query_with_result(#q, query_args).map(|row|
1167 #row{
1168 #(#row_construct),*
1169 }
1170 )
1171 }};
1172 s.into()
1173 }
1174 qusql_type::StatementType::Update {
1175 returning: None, ..
1176 } => {
1177 errors.push(
1178 syn::Error::new(
1179 query_as.query_span,
1180 "UPDATE without RETURNING not support in query_as",
1181 )
1182 .to_compile_error(),
1183 );
1184 quote! { {
1185 #(#errors; )*
1186 todo!("update")
1187 }}
1188 .into()
1189 }
1190 qusql_type::StatementType::Update {
1191 arguments,
1192 returning: Some(returning),
1193 ..
1194 } => {
1195 let (args_tokens, q) = quote_args(
1196 &mut errors,
1197 &query_as.query,
1198 query_as.last_span,
1199 &query_as.args,
1200 arguments,
1201 dialect,
1202 );
1203
1204 let row_construct = construct_row2(returning, dialect.is_postgresql());
1205 let row = query_as.as_;
1206 let s = quote! { {
1207 use ::sqlx::Arguments as _;
1208 #(#errors; )*
1209 #args_tokens
1210 sqlx::__query_with_result(#q, query_args).map(|row|
1211 #row{
1212 #(#row_construct),*
1213 }
1214 )
1215 }};
1216 s.into()
1217 }
1218 qusql_type::StatementType::Replace {
1219 returning: None, ..
1220 } => {
1221 errors.push(
1222 syn::Error::new(
1223 query_as.query_span,
1224 "REPLACE without RETURNING not support in query_as",
1225 )
1226 .to_compile_error(),
1227 );
1228 quote! { {
1229 #(#errors; )*
1230 todo!("replace")
1231 }}
1232 .into()
1233 }
1234 qusql_type::StatementType::Replace {
1235 arguments,
1236 returning: Some(returning),
1237 ..
1238 } => {
1239 let (args_tokens, q) = quote_args(
1240 &mut errors,
1241 &query_as.query,
1242 query_as.last_span,
1243 &query_as.args,
1244 arguments,
1245 dialect,
1246 );
1247
1248 let row_construct = construct_row2(returning, dialect.is_postgresql());
1249 let row = query_as.as_;
1250 let s = quote! { {
1251 use ::sqlx::Arguments as _;
1252 #(#errors; )*
1253 #args_tokens
1254 sqlx::__query_with_result(#q, query_args).map(|row|
1255 #row{
1256 #(#row_construct),*
1257 }
1258 )
1259 }};
1260 s.into()
1261 }
1262 qusql_type::StatementType::Truncate => {
1263 errors.push(
1264 syn::Error::new(query_as.query_span, "TRUNCATE not supported in query_as!")
1265 .to_compile_error(),
1266 );
1267 quote! { {
1268 #(#errors; )*
1269 todo!("truncate")
1270 }}
1271 .into()
1272 }
1273 qusql_type::StatementType::Call { .. } => {
1274 errors.push(
1275 syn::Error::new(query_as.query_span, "CALL not supported in query_as!")
1276 .to_compile_error(),
1277 );
1278 quote! { {
1279 #(#errors; )*
1280 todo!("call")
1281 }}
1282 .into()
1283 }
1284 qusql_type::StatementType::Transaction => {
1285 errors.push(
1286 syn::Error::new(
1287 query_as.query_span,
1288 "Transaction control not supported in query_as!",
1289 )
1290 .to_compile_error(),
1291 );
1292 quote! { {
1293 #(#errors; )*
1294 todo!("transaction")
1295 }}
1296 .into()
1297 }
1298 qusql_type::StatementType::Set => {
1299 errors.push(
1300 syn::Error::new(query_as.query_span, "SET not supported in query_as!")
1301 .to_compile_error(),
1302 );
1303 quote! { {
1304 #(#errors; )*
1305 todo!("set")
1306 }}
1307 .into()
1308 }
1309 qusql_type::StatementType::Lock => {
1310 errors.push(
1311 syn::Error::new(query_as.query_span, "LOCK not supported in query_as!")
1312 .to_compile_error(),
1313 );
1314 quote! { {
1315 #(#errors; )*
1316 todo!("lock")
1317 }}
1318 .into()
1319 }
1320 qusql_type::StatementType::Invalid => quote! { {
1321 #(#errors; )*;
1322 todo!("invalid")
1323 }}
1324 .into(),
1325 }
1326}