Skip to main content

qusql_sqlx_type_macro/
lib.rs

1//! Implementation of typed db query macros
2//!
3//! Used the exposed macros from the qusql-sqlx-type crate
4//!
5#![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
20/// Path where the sqlx-type-schema.sql file can be found
21static 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        /// Struct used to deserialize the cargo betadata
43        #[derive(Deserialize)]
44        struct CargoMetadata {
45            /// The root of the workspace
46            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
61/// Source of the sqlx-type-schema.sql
62// If we are in a workspace, lookup `workspace_root` since `CARGO_MANIFEST_DIR` won't
63// reflect the workspace dir: https://github.com/rust-lang/cargo/issues/3946
64static 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
74/// Construct a none color report for an issue
75fn 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
96/// Construct a color report for an issue
97fn 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
125/// Source name adaptor for ariadne
126struct 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
140/// Parsed version of the sqlx-type-schema.sql
141static 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
175/// Produce quoted arguments for a query
176fn 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
332/// Output an [Issue] as a compile error
333fn 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
352/// Construct row struct members, and fill in statements
353fn 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            // TODO error
406            //errors.push(syn::Error::new(span, String::from_utf8(out).unwrap()).to_compile_error().into());
407            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
423/// Parsed query! macro
424struct Query {
425    /// The query expression
426    query: String,
427    /// The span of the query expression
428    query_span: Span,
429    /// The arguments to supply to the query
430    args: Vec<Expr>,
431    /// The last span parsed
432    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/// Statically checked SQL query, similarly to sqlx::query!.
461///
462/// This expands to an instance of query::Map that outputs an ad-hoc anonymous struct type.
463#[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
684/// Fill in row values in a query_as struct
685fn 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            // TODO error
735            //errors.push(syn::Error::new(span, String::from_utf8(out).unwrap()).to_compile_error().into());
736            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
749/// Parse result of a query_as! macro
750struct QueryAs {
751    /// Name of output type
752    as_: Ident,
753    /// The query to execute
754    query: String,
755    /// The span of the query to execute
756    query_span: Span,
757    /// The arguments to supply
758    args: Vec<Expr>,
759    /// The last span parsed
760    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/// A variant of query! which takes a path to an explicitly defined struct as the output type.
794///
795/// This lets you return the struct from a function or add your own trait implementations.
796#[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            //println!("TOKENS: {}", s);
836            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}