1#![forbid(unsafe_code)]
6
7use std::ops::Deref;
8use std::path::PathBuf;
9
10use ariadne::{Color, Label, Report, ReportKind, Source};
11use once_cell::sync::Lazy;
12use proc_macro::TokenStream;
13use proc_macro2::Span;
14use quote::{format_ident, quote, quote_spanned};
15use qusql_type::schema::{parse_schemas, Schemas};
16use qusql_type::{type_statement, Issue, SQLArguments, SQLDialect, SelectTypeColumn, TypeOptions};
17use syn::spanned::Spanned;
18use syn::{parse::Parse, punctuated::Punctuated, Expr, Ident, LitStr, Token};
19
20static SCHEMA_PATH: Lazy<PathBuf> = Lazy::new(|| {
22 let mut schema_path: PathBuf = std::env::var("CARGO_MANIFEST_DIR")
23 .expect("`CARGO_schema_path` must be set")
24 .into();
25
26 schema_path.push("sqlx-type-schema.sql");
27
28 if !schema_path.exists() {
29 use serde::Deserialize;
30 use std::process::Command;
31
32 let cargo = std::env::var("CARGO").expect("`CARGO` must be set");
33 schema_path.pop();
34
35 let output = Command::new(cargo)
36 .args(["metadata", "--format-version=1"])
37 .current_dir(&schema_path)
38 .env_remove("__CARGO_FIX_PLZ")
39 .output()
40 .expect("Could not fetch metadata");
41
42 #[derive(Deserialize)]
44 struct CargoMetadata {
45 workspace_root: PathBuf,
47 }
48
49 let metadata: CargoMetadata =
50 serde_json::from_slice(&output.stdout).expect("Invalid `cargo metadata` output");
51
52 schema_path = metadata.workspace_root;
53 schema_path.push("sqlx-type-schema.sql");
54 }
55 if !schema_path.exists() {
56 panic!("Unable to locate sqlx-type-schema.sql");
57 }
58 schema_path
59});
60
61static SCHEMA_SRC: Lazy<String> =
65 Lazy::new(|| match std::fs::read_to_string(SCHEMA_PATH.as_path()) {
66 Ok(v) => v,
67 Err(e) => panic!(
68 "Unable to read schema from {:?}: {}",
69 SCHEMA_PATH.as_path(),
70 e
71 ),
72 });
73
74fn issue_to_report(issue: Issue) -> Report<'static, std::ops::Range<usize>> {
76 let mut builder = Report::build(
77 match issue.level {
78 qusql_type::Level::Warning => ReportKind::Warning,
79 qusql_type::Level::Error => ReportKind::Error,
80 },
81 issue.span.clone(),
82 )
83 .with_config(ariadne::Config::default().with_color(false))
84 .with_label(
85 Label::new(issue.span)
86 .with_order(-1)
87 .with_priority(-1)
88 .with_message(issue.message),
89 );
90 for frag in issue.fragments {
91 builder = builder.with_label(Label::new(frag.span).with_message(frag.message));
92 }
93 builder.finish()
94}
95
96fn issue_to_report_color(issue: Issue) -> Report<'static, std::ops::Range<usize>> {
98 let mut builder = Report::build(
99 match issue.level {
100 qusql_type::Level::Warning => ReportKind::Warning,
101 qusql_type::Level::Error => ReportKind::Error,
102 },
103 issue.span.clone(),
104 )
105 .with_label(
106 Label::new(issue.span)
107 .with_color(match issue.level {
108 qusql_type::Level::Warning => Color::Yellow,
109 qusql_type::Level::Error => Color::Red,
110 })
111 .with_order(-1)
112 .with_priority(-1)
113 .with_message(issue.message),
114 );
115 for frag in issue.fragments {
116 builder = builder.with_label(
117 Label::new(frag.span)
118 .with_color(Color::Blue)
119 .with_message(frag.message),
120 );
121 }
122 builder.finish()
123}
124
125struct NamedSource<'a>(&'a str, Source<&'a str>);
127
128impl<'a> ariadne::Cache<()> for &NamedSource<'a> {
129 type Storage = &'a str;
130
131 fn display<'b>(&self, _: &'b ()) -> Option<impl std::fmt::Display + 'b> {
132 Some(self.0.to_string())
133 }
134
135 fn fetch(&mut self, _: &()) -> Result<&Source<Self::Storage>, impl std::fmt::Debug> {
136 Ok::<_, ()>(&self.1)
137 }
138}
139
140static SCHEMAS: Lazy<(Schemas, SQLDialect)> = Lazy::new(|| {
142 let schema_src = SCHEMA_SRC.as_str();
143 let dialect = if let Some(first_line) = schema_src.lines().next() {
144 if first_line.contains("sql-product: postgres") {
145 SQLDialect::PostgreSQL
146 } else if first_line.contains("sql-product: sqlite") {
147 SQLDialect::Sqlite
148 } else {
149 SQLDialect::MariaDB
150 }
151 } else {
152 SQLDialect::MariaDB
153 };
154
155 let options = TypeOptions::new().dialect(dialect.clone());
156 let mut issues = qusql_type::Issues::new(schema_src);
157 let schemas = parse_schemas(schema_src, &mut issues, &options);
158 if !issues.is_ok() {
159 let source = NamedSource("sqlx-type-schema.sql", Source::from(schema_src));
160 let mut err = false;
161 for issue in issues.into_vec() {
162 if issue.level == qusql_type::Level::Error {
163 err = true;
164 }
165 let r = issue_to_report_color(issue);
166 r.eprint(&source).unwrap();
167 }
168 if err {
169 panic!("Errors processing sqlx-type-schema.sql");
170 }
171 }
172 (schemas, dialect)
173});
174
175fn quote_args(
177 errors: &mut Vec<proc_macro2::TokenStream>,
178 query: &str,
179 last_span: Span,
180 args: &[Expr],
181 arguments: &[(qusql_type::ArgumentKey<'_>, qusql_type::FullType)],
182 dialect: &SQLDialect,
183) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
184 let cls = match dialect {
185 SQLDialect::MariaDB => quote!(sqlx::mysql::MySql),
186 SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite),
187 SQLDialect::PostgreSQL => quote!(sqlx::postgres::Postgres),
188 };
189
190 let mut at = Vec::new();
191 let inv = qusql_type::FullType::invalid();
192 for (k, v) in arguments {
193 match k {
194 qusql_type::ArgumentKey::Index(i) => {
195 while at.len() <= *i {
196 at.push(&inv);
197 }
198 at[*i] = v;
199 }
200 qusql_type::ArgumentKey::Identifier(_) => {
201 errors.push(
202 syn::Error::new(last_span.span(), "Named arguments not supported")
203 .to_compile_error(),
204 );
205 }
206 }
207 }
208
209 if at.len() > args.len() {
210 errors.push(
211 syn::Error::new(
212 last_span,
213 format!("Expected {} additional arguments", at.len() - args.len()),
214 )
215 .to_compile_error(),
216 );
217 }
218
219 if let Some(args) = args.get(at.len()..) {
220 for arg in args {
221 errors.push(syn::Error::new(arg.span(), "unexpected argument").to_compile_error());
222 }
223 }
224
225 let arg_names = (0..args.len())
226 .map(|i| format_ident!("arg{}", i))
227 .collect::<Vec<_>>();
228
229 let mut arg_bindings = Vec::new();
230 let mut arg_add = Vec::new();
231
232 let mut list_lengths = Vec::new();
233
234 for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) {
235 let mut t = match ta.t {
236 qusql_type::Type::U8 => quote! {u8},
237 qusql_type::Type::I8 => quote! {i8},
238 qusql_type::Type::U16 => quote! {u16},
239 qusql_type::Type::I16 => quote! {i16},
240 qusql_type::Type::U24 => quote! {u32},
241 qusql_type::Type::I24 => quote! {i32},
242 qusql_type::Type::U32 => quote! {u32},
243 qusql_type::Type::I32 => quote! {i32},
244 qusql_type::Type::U64 => quote! {u64},
245 qusql_type::Type::I64 => quote! {i64},
246 qusql_type::Type::Base(qusql_type::BaseType::Any) => quote! {qusql_sqlx_type::Any},
247 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
248 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {&[u8]},
249 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {qusql_sqlx_type::Date},
250 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
251 quote! {qusql_sqlx_type::DateTime}
252 }
253 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {qusql_sqlx_type::Float},
254 qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
255 quote! {qusql_sqlx_type::Integer}
256 }
257 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {&str},
258 qusql_type::Type::Base(qusql_type::BaseType::Time) => quote! {qusql_sqlx_type::Time},
259 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => todo!("time_interval"),
260 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
261 quote! {qusql_sqlx_type::Timestamp}
262 }
263 qusql_type::Type::Null => todo!("null"),
264 qusql_type::Type::Invalid => quote! {std::convert::Infallible},
265 qusql_type::Type::Enum(_) => quote! {&str},
266 qusql_type::Type::Set(_) => quote! {&str},
267 qusql_type::Type::Args(_, _) => todo!("args"),
268 qusql_type::Type::F32 => quote! {f32},
269 qusql_type::Type::F64 => quote! {f64},
270 qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any},
271 };
272 if !ta.not_null {
273 t = quote! {Option<#t>}
274 }
275 let span = qa.span();
276 if ta.list_hack {
277 list_lengths.push(quote!(#name.len()));
278 arg_bindings.push(quote_spanned! {span=>
279 let #name = &(#qa);
280 args_count += #name.len();
281 for v in #name.iter() {
282 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(v);
283 }
284 if false {
285 qusql_sqlx_type::check_arg_list_hack::<#t, _>(#name);
286 ::std::panic!();
287 }
288 });
289 arg_add.push(quote!(
290 for v in #name.iter() {
291 e = e.and_then(|()| query_args.add(v));
292 }
293 ));
294 } else {
295 arg_bindings.push(quote_spanned! {span=>
296 let #name = &(#qa);
297 args_count += 1;
298 size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(#name);
299 if false {
300 qusql_sqlx_type::check_arg::<#t, _>(#name);
301 ::std::panic!();
302 }
303 });
304 arg_add.push(quote!(e = e.and_then(|()| query_args.add(#name));));
305 }
306 }
307
308 let query = if list_lengths.is_empty() {
309 quote!(#query)
310 } else {
311 quote!(
312 &qusql_sqlx_type::convert_list_query(#query, &[#(#list_lengths),*])
313 )
314 };
315
316 (
317 quote! {
318 let mut size_hints = 0;
319 let mut args_count = 0;
320 #(#arg_bindings)*
321
322 let mut query_args = <#cls as ::sqlx::database::Database>::Arguments::default();
323 query_args.reserve(args_count, size_hints);
324 let mut e = Ok(());
325 #(#arg_add)*
326 let query_args = e.and_then(|()| Ok(query_args));
327 },
328 query,
329 )
330}
331
332fn issues_to_errors(issues: Vec<Issue>, source: &str, span: Span) -> Vec<proc_macro2::TokenStream> {
334 if !issues.is_empty() {
335 let source = NamedSource("", Source::from(source));
336 let mut err = false;
337 let mut out = Vec::new();
338 for issue in issues {
339 if issue.level == qusql_type::Level::Error {
340 err = true;
341 }
342 let r = issue_to_report(issue);
343 r.write(&source, &mut out).unwrap();
344 }
345 if err {
346 return vec![syn::Error::new(span, String::from_utf8(out).unwrap()).to_compile_error()];
347 }
348 }
349 Vec::new()
350}
351
352fn construct_row(
354 columns: &[SelectTypeColumn],
355) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
356 let mut row_members = Vec::new();
357 let mut row_construct = Vec::new();
358 for (i, c) in columns.iter().enumerate() {
359 let mut t = match c.type_.t {
360 qusql_type::Type::U8 => quote! {u8},
361 qusql_type::Type::I8 => quote! {i8},
362 qusql_type::Type::U16 => quote! {u16},
363 qusql_type::Type::I16 => quote! {i16},
364 qusql_type::Type::U24 => quote! {u32},
365 qusql_type::Type::I24 => quote! {i32},
366 qusql_type::Type::U32 => quote! {u32},
367 qusql_type::Type::I32 => quote! {i32},
368 qusql_type::Type::U64 => quote! {u64},
369 qusql_type::Type::I64 => quote! {i64},
370 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
371 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
372 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
373 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
374 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
375 quote! {chrono::NaiveDateTime}
376 }
377 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
378 qusql_type::Type::Base(qusql_type::BaseType::Integer) => quote! {i64},
379 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
380 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
381 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
382 todo!("from_time_interval")
383 }
384 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
385 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
386 }
387 qusql_type::Type::Null => todo!("from_null"),
388 qusql_type::Type::Invalid => quote! {i64},
389 qusql_type::Type::Enum(_) => quote! {String},
390 qusql_type::Type::Set(_) => quote! {String},
391 qusql_type::Type::Args(_, _) => todo!("from_args"),
392 qusql_type::Type::F32 => quote! {f32},
393 qusql_type::Type::F64 => quote! {f64},
394 qusql_type::Type::JSON => quote! {String},
395 };
396 let name = match &c.name {
397 Some(v) => v,
398 None => continue,
399 };
400
401 let ident = String::from("r#") + name.value;
402 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
403 ident
404 } else {
405 continue;
408 };
409
410 if !c.type_.not_null {
411 t = quote! {Option<#t>};
412 }
413 row_members.push(quote! {
414 #ident : #t
415 });
416 row_construct.push(quote! {
417 #ident: sqlx::Row::get(&row, #i)
418 });
419 }
420 (row_members, row_construct)
421}
422
423struct Query {
425 query: String,
427 query_span: Span,
429 args: Vec<Expr>,
431 last_span: Span,
433}
434
435impl Parse for Query {
436 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
437 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
438 let query: String = query_.iter().map(LitStr::value).collect();
439 let query_span = query_.span();
440 let mut last_span = query_span;
441 let mut args = Vec::new();
442 while !input.is_empty() {
443 let _ = input.parse::<syn::token::Comma>()?;
444 if input.is_empty() {
445 break;
446 }
447 let arg = input.parse::<Expr>()?;
448 last_span = arg.span();
449 args.push(arg);
450 }
451 Ok(Self {
452 query,
453 query_span,
454 args,
455 last_span,
456 })
457 }
458}
459
460#[proc_macro]
464pub fn query(input: TokenStream) -> TokenStream {
465 let query = syn::parse_macro_input!(input as Query);
466 let (schemas, dialect) = SCHEMAS.deref();
467 let options = TypeOptions::new()
468 .dialect(dialect.clone())
469 .arguments(match &dialect {
470 SQLDialect::MariaDB => SQLArguments::QuestionMark,
471 SQLDialect::Sqlite => SQLArguments::QuestionMark,
472 SQLDialect::PostgreSQL => SQLArguments::Dollar,
473 })
474 .list_hack(true);
475 let mut issues = qusql_type::Issues::new(&query.query);
476 let stmt = type_statement(schemas, &query.query, &mut issues, &options);
477 let sp = SCHEMA_PATH.as_path().to_str().unwrap();
478
479 let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
480 match &stmt {
481 qusql_type::StatementType::Select { columns, arguments } => {
482 let (args_tokens, q) = quote_args(
483 &mut errors,
484 &query.query,
485 query.last_span,
486 &query.args,
487 arguments,
488 dialect,
489 );
490 let (row_members, row_construct) = construct_row(columns);
491 let s = quote! { {
492 use ::sqlx::Arguments as _;
493 let _ = std::include_bytes!(#sp);
494 #(#errors; )*
495 #args_tokens;
496
497 struct Row {
498 #(#row_members),*
499 };
500 sqlx::__query_with_result(#q, query_args).map(|row|
501 Row{
502 #(#row_construct),*
503 }
504 )
505 }};
506 s.into()
507 }
508 qusql_type::StatementType::Delete {
509 arguments,
510 returning,
511 } => {
512 let (args_tokens, q) = quote_args(
513 &mut errors,
514 &query.query,
515 query.last_span,
516 &query.args,
517 arguments,
518 dialect,
519 );
520 let s = match returning.as_ref() {
521 Some(returning) => {
522 let (row_members, row_construct) = construct_row(returning);
523 quote! { {
524 use ::sqlx::Arguments as _;
525 let _ = std::include_bytes!(#sp);
526 #(#errors; )*
527 #args_tokens
528
529 struct Row {
530 #(#row_members),*
531 };
532 sqlx::__query_with_result(#q, query_args).map(|row|
533 Row{
534 #(#row_construct),*
535 }
536 )
537 }}
538 }
539 None => quote! { {
540 use ::sqlx::Arguments as _;
541 #(#errors; )*
542 #args_tokens
543 sqlx::__query_with_result(#q, query_args)
544 }
545 },
546 };
547 s.into()
548 }
549 qusql_type::StatementType::Insert {
550 arguments,
551 returning,
552 ..
553 } => {
554 let (args_tokens, q) = quote_args(
555 &mut errors,
556 &query.query,
557 query.last_span,
558 &query.args,
559 arguments,
560 dialect,
561 );
562 let s = match returning.as_ref() {
563 Some(returning) => {
564 let (row_members, row_construct) = construct_row(returning);
565 quote! { {
566 use ::sqlx::Arguments as _;
567 let _ = std::include_bytes!(#sp);
568 #(#errors; )*
569 #args_tokens
570
571 struct Row {
572 #(#row_members),*
573 };
574 sqlx::__query_with_result(#q, query_args).map(|row|
575 Row{
576 #(#row_construct),*
577 }
578 )
579 }}
580 }
581 None => quote! { {
582 use ::sqlx::Arguments as _;
583 #(#errors; )*
584 #args_tokens
585 sqlx::__query_with_result(#q, query_args)
586 }
587 },
588 };
589 s.into()
590 }
591 qusql_type::StatementType::Update {
592 arguments,
593 returning,
594 } => {
595 let (args_tokens, q) = quote_args(
596 &mut errors,
597 &query.query,
598 query.last_span,
599 &query.args,
600 arguments,
601 dialect,
602 );
603
604 let s = match returning.as_ref() {
605 Some(returning) => {
606 let (row_members, row_construct) = construct_row(returning);
607 quote! { {
608 use ::sqlx::Arguments as _;
609 let _ = std::include_bytes!(#sp);
610 #(#errors; )*
611 #args_tokens
612
613 struct Row {
614 #(#row_members),*
615 };
616 sqlx::__query_with_result(#q, query_args).map(|row|
617 Row{
618 #(#row_construct),*
619 }
620 )
621 }}
622 }
623 None => quote! { {
624 use ::sqlx::Arguments as _;
625 #(#errors; )*
626 #args_tokens
627 sqlx::__query_with_result(#q, query_args)
628 }
629 },
630 };
631 s.into()
632 }
633 qusql_type::StatementType::Replace {
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 #(#errors; )*
667 #args_tokens
668 sqlx::__query_with_result(#q, query_args)
669 }
670 },
671 };
672 s.into()
673 }
674 qusql_type::StatementType::Invalid => {
675 let s = quote! { {
676 #(#errors; )*;
677 todo!("Invalid")
678 }};
679 s.into()
680 }
681 }
682}
683
684fn construct_row2(columns: &[SelectTypeColumn]) -> Vec<proc_macro2::TokenStream> {
686 let mut row_construct = Vec::new();
687 for (i, c) in columns.iter().enumerate() {
688 let mut t = match c.type_.t {
689 qusql_type::Type::U8 => quote! {u8},
690 qusql_type::Type::I8 => quote! {i8},
691 qusql_type::Type::U16 => quote! {u16},
692 qusql_type::Type::I16 => quote! {i16},
693 qusql_type::Type::U24 => quote! {u32},
694 qusql_type::Type::I24 => quote! {i32},
695 qusql_type::Type::U32 => quote! {u32},
696 qusql_type::Type::I32 => quote! {i32},
697 qusql_type::Type::U64 => quote! {u64},
698 qusql_type::Type::I64 => quote! {i64},
699 qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
700 qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
701 qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
702 qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
703 qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
704 quote! {chrono::NaiveDateTime}
705 }
706 qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
707 qusql_type::Type::Base(qusql_type::BaseType::Integer) => quote! {i64},
708 qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
709 qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
710 qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
711 todo!("from_time_interval")
712 }
713 qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
714 quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
715 }
716 qusql_type::Type::Null => todo!("from_null"),
717 qusql_type::Type::Invalid => quote! {i64},
718 qusql_type::Type::Enum(_) => quote! {String},
719 qusql_type::Type::Set(_) => quote! {String},
720 qusql_type::Type::Args(_, _) => todo!("from_args"),
721 qusql_type::Type::F32 => quote! {f32},
722 qusql_type::Type::F64 => quote! {f64},
723 qusql_type::Type::JSON => quote! {String},
724 };
725 let name = match &c.name {
726 Some(v) => v,
727 None => continue,
728 };
729
730 let ident = String::from("r#") + name.value;
731 let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
732 ident
733 } else {
734 continue;
737 };
738
739 if !c.type_.not_null {
740 t = quote! {Option<#t>};
741 }
742 row_construct.push(quote! {
743 #ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
744 });
745 }
746 row_construct
747}
748
749struct QueryAs {
751 as_: Ident,
753 query: String,
755 query_span: Span,
757 args: Vec<Expr>,
759 last_span: Span,
761}
762
763impl Parse for QueryAs {
764 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
765 let as_ = input.parse::<Ident>()?;
766 let _ = input.parse::<syn::token::Comma>()?;
767
768 let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
769 let query: String = query_.iter().map(LitStr::value).collect();
770 let query_span = query_.span();
771
772 let mut last_span = query_span;
773 let mut args = Vec::new();
774 while !input.is_empty() {
775 let _ = input.parse::<syn::token::Comma>()?;
776 if input.is_empty() {
777 break;
778 }
779 let arg = input.parse::<Expr>()?;
780 last_span = arg.span();
781 args.push(arg);
782 }
783 Ok(Self {
784 as_,
785 query,
786 query_span,
787 args,
788 last_span,
789 })
790 }
791}
792
793#[proc_macro]
797pub fn query_as(input: TokenStream) -> TokenStream {
798 let query_as = syn::parse_macro_input!(input as QueryAs);
799 let (schemas, dialect) = SCHEMAS.deref();
800 let options = TypeOptions::new()
801 .dialect(dialect.clone())
802 .arguments(match &dialect {
803 SQLDialect::MariaDB => SQLArguments::QuestionMark,
804 SQLDialect::Sqlite => SQLArguments::QuestionMark,
805 SQLDialect::PostgreSQL => SQLArguments::Dollar,
806 })
807 .list_hack(true);
808 let mut issues = qusql_type::Issues::new(&query_as.query);
809 let stmt = type_statement(schemas, &query_as.query, &mut issues, &options);
810
811 let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
812 match &stmt {
813 qusql_type::StatementType::Select { columns, arguments } => {
814 let (args_tokens, q) = quote_args(
815 &mut errors,
816 &query_as.query,
817 query_as.last_span,
818 &query_as.args,
819 arguments,
820 dialect,
821 );
822
823 let row_construct = construct_row2(columns);
824 let row = query_as.as_;
825 let s = quote! { {
826 use ::sqlx::Arguments as _;
827 #(#errors; )*
828 #args_tokens
829 sqlx::__query_with_result(#q, query_args).map(|row|
830 #row{
831 #(#row_construct),*
832 }
833 )
834 }};
835 s.into()
837 }
838 qusql_type::StatementType::Delete { .. } => {
839 errors.push(
840 syn::Error::new(query_as.query_span, "DELETE not support in query_as")
841 .to_compile_error(),
842 );
843 quote! { {
844 #(#errors; )*
845 todo!("delete")
846 }}
847 .into()
848 }
849 qusql_type::StatementType::Insert {
850 returning: None, ..
851 } => {
852 errors.push(
853 syn::Error::new(
854 query_as.query_span,
855 "INSERT without RETURNING not support in query_as",
856 )
857 .to_compile_error(),
858 );
859 quote! { {
860 #(#errors; )*
861 todo!("insert")
862 }}
863 .into()
864 }
865 qusql_type::StatementType::Insert {
866 arguments,
867 returning: Some(returning),
868 ..
869 } => {
870 let (args_tokens, q) = quote_args(
871 &mut errors,
872 &query_as.query,
873 query_as.last_span,
874 &query_as.args,
875 arguments,
876 dialect,
877 );
878
879 let row_construct = construct_row2(returning);
880 let row = query_as.as_;
881 let s = quote! { {
882 use ::sqlx::Arguments as _;
883 #(#errors; )*
884 #args_tokens
885 sqlx::__query_with_result(#q, query_args).map(|row|
886 #row{
887 #(#row_construct),*
888 }
889 )
890 }};
891 s.into()
892 }
893 qusql_type::StatementType::Update { .. } => {
894 errors.push(
895 syn::Error::new(query_as.query_span, "UPDATE not support in query_as")
896 .to_compile_error(),
897 );
898 quote! { {
899 #(#errors; )*
900 todo!("update")
901 }}
902 .into()
903 }
904 qusql_type::StatementType::Replace {
905 returning: None, ..
906 } => {
907 errors.push(
908 syn::Error::new(
909 query_as.query_span,
910 "REPLACE without RETURNING not support in query_as",
911 )
912 .to_compile_error(),
913 );
914 quote! { {
915 #(#errors; )*
916 todo!("replace")
917 }}
918 .into()
919 }
920 qusql_type::StatementType::Replace {
921 arguments,
922 returning: Some(returning),
923 ..
924 } => {
925 let (args_tokens, q) = quote_args(
926 &mut errors,
927 &query_as.query,
928 query_as.last_span,
929 &query_as.args,
930 arguments,
931 dialect,
932 );
933
934 let row_construct = construct_row2(returning);
935 let row = query_as.as_;
936 let s = quote! { {
937 use ::sqlx::Arguments as _;
938 #(#errors; )*
939 #args_tokens
940 sqlx::__query_with_result(#q, query_args).map(|row|
941 #row{
942 #(#row_construct),*
943 }
944 )
945 }};
946 s.into()
947 }
948 qusql_type::StatementType::Invalid => quote! { {
949 #(#errors; )*;
950 todo!("invalid")
951 }}
952 .into(),
953 }
954}