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::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
24/// Cache of resolved schema paths, keyed by `CARGO_MANIFEST_DIR`.
25///
26/// Path resolution is cheap when the schema sits next to `Cargo.toml`, but
27/// falls back to `cargo metadata` (a subprocess) when it doesn't.  Either
28/// way the result is stable for the lifetime of the proc-macro server process,
29/// so we cache it per manifest-dir.  A plain `Lazy<PathBuf>` would be wrong
30/// here because rust-analyzer reuses the same process for multiple crates and
31/// each crate has a different `CARGO_MANIFEST_DIR`.
32static RESOLVED_SCHEMA_PATHS: Mutex<Option<HashMap<PathBuf, PathBuf>>> = Mutex::new(None);
33
34/// Resolve the path to `sqlx-type-schema.sql` for the crate currently being
35/// compiled.  The result is cached per `CARGO_MANIFEST_DIR` so the
36/// (potentially expensive) `cargo metadata` fallback runs at most once per
37/// crate per proc-macro server process.
38fn 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        /// Struct used to deserialize the cargo metadata
68        #[derive(Deserialize)]
69        struct CargoMetadata {
70            /// The root of the workspace
71            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
87/// Construct a none color report for an issue (used as a compile_error! message)
88fn 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
113/// Construct a color report for an issue (printed to stderr during schema errors)
114fn 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
152/// Source name adaptor for ariadne
153struct 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)]
168/// `Yokeable` wrapper that lets `Schemas<'a>` be stored in a `Yoke`.
169struct SchemasYoke<'a>(Schemas<'a>);
170
171/// Parsed schema and its source string, kept alive together so neither
172/// outlives the other.  Stored behind an `Arc` so callers can cheaply
173/// clone a handle without holding the mutex.
174struct SchemaCacheEntry {
175    /// The parsed schema yoked to the source it borrows from.
176    schemas: Yoke<SchemasYoke<'static>, String>,
177    /// SQL dialect detected from the schema file header.
178    dialect: SQLDialect,
179    /// Resolved path to the schema file, used in generated include_bytes! calls.
180    path: PathBuf,
181    /// File byte-length at the time of parsing, used for cache invalidation.
182    file_len: u64,
183    /// File modification time at the time of parsing, used for cache invalidation.
184    modified: SystemTime,
185    /// Hash of the schema content
186    hash: u64,
187}
188
189/// Per-path cache of parsed schemas, invalidated by file mtime/size.
190///
191/// Keyed by resolved schema path so that multiple crates compiled in the same
192/// proc-macro server process (e.g. rust-analyzer) each get their own entry.
193static SCHEMA_CACHE: Mutex<Option<HashMap<PathBuf, Arc<SchemaCacheEntry>>>> = Mutex::new(None);
194
195/// Return the parsed schema for the crate currently being compiled, reparsing
196/// from disk only when the file's mtime or size has changed.
197fn 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    // Cache miss or stale - read and reparse.
212    let src_string = std::fs::read_to_string(&path)
213        .unwrap_or_else(|e| panic!("Unable to read schema from {path:?}: {e}"));
214
215    // Detect dialect from the first two lines before moving src_box into the
216    // Yoke cart.  The result is moved into the closure by value.
217    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    // Compute the schema hash from the fresh source before we move it into the cart.
234    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
274/// Produce quoted arguments for a query
275fn quote_args(
276    errors: &mut Vec<proc_macro2::TokenStream>,
277    query: &str,
278    last_span: Span,
279    args: &[Expr],
280    arguments: &[(qusql_type::ArgumentKey<'_>, qusql_type::FullType)],
281    dialect: &SQLDialect,
282) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
283    let cls = match dialect {
284        SQLDialect::MariaDB => quote!(sqlx::mysql::MySql),
285        SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite),
286        SQLDialect::PostgreSQL | SQLDialect::PostGIS => quote!(sqlx::postgres::Postgres),
287    };
288
289    let mut at = Vec::new();
290    let inv = qusql_type::FullType::invalid();
291    for (k, v) in arguments {
292        match k {
293            qusql_type::ArgumentKey::Index(i) => {
294                while at.len() <= *i {
295                    at.push(&inv);
296                }
297                at[*i] = v;
298            }
299            qusql_type::ArgumentKey::Identifier(_) => {
300                errors.push(
301                    syn::Error::new(last_span.span(), "Named arguments not supported")
302                        .to_compile_error(),
303                );
304            }
305        }
306    }
307
308    if at.len() > args.len() {
309        errors.push(
310            syn::Error::new(
311                last_span,
312                format!("Expected {} additional arguments", at.len() - args.len()),
313            )
314            .to_compile_error(),
315        );
316    }
317
318    if let Some(args) = args.get(at.len()..) {
319        for arg in args {
320            errors.push(syn::Error::new(arg.span(), "unexpected argument").to_compile_error());
321        }
322    }
323
324    let arg_names = (0..args.len())
325        .map(|i| format_ident!("arg{}", i))
326        .collect::<Vec<_>>();
327
328    let mut arg_bindings = Vec::new();
329    let mut arg_add = Vec::new();
330
331    let mut list_lengths = Vec::new();
332
333    for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) {
334        let mut t = match ta.t {
335            qusql_type::Type::U8 => quote! {u8},
336            qusql_type::Type::I8 => quote! {i8},
337            qusql_type::Type::U16 => quote! {u16},
338            qusql_type::Type::I16 => quote! {i16},
339            qusql_type::Type::U24 => quote! {u32},
340            qusql_type::Type::I24 => quote! {i32},
341            qusql_type::Type::U32 => quote! {u32},
342            qusql_type::Type::I32 => quote! {i32},
343            qusql_type::Type::U64 => quote! {u64},
344            qusql_type::Type::I64 => quote! {i64},
345            qusql_type::Type::Base(qusql_type::BaseType::Any) => quote! {qusql_sqlx_type::Any},
346            qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
347            qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {&[u8]},
348            qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {qusql_sqlx_type::Date},
349            qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
350                quote! {qusql_sqlx_type::DateTime}
351            }
352            qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {qusql_sqlx_type::Float},
353            qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
354                quote! {qusql_sqlx_type::Integer}
355            }
356            qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {&str},
357            qusql_type::Type::Base(qusql_type::BaseType::Time) => quote! {qusql_sqlx_type::Time},
358            qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => todo!("time_interval"),
359            qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
360                quote! {qusql_sqlx_type::Timestamp}
361            }
362            qusql_type::Type::Base(qusql_type::BaseType::Uuid) => quote! {qusql_sqlx_type::Uuid},
363            qusql_type::Type::Null => todo!("null"),
364            qusql_type::Type::Invalid => quote! {std::convert::Infallible},
365            qusql_type::Type::Enum(_) => quote! {&str},
366            qusql_type::Type::Set(_) => quote! {&str},
367            qusql_type::Type::Args(_, _) => todo!("args"),
368            qusql_type::Type::F32 => quote! {f32},
369            qusql_type::Type::F64 => quote! {f64},
370            qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any},
371            qusql_type::Type::Geometry => quote! {qusql_sqlx_type::Any},
372            qusql_type::Type::Range(_) => quote! {qusql_sqlx_type::Any},
373            qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
374        };
375        if !ta.not_null {
376            t = quote! {Option<#t>}
377        }
378        let span = qa.span();
379        if ta.list_hack {
380            list_lengths.push(quote!(#name.len()));
381            arg_bindings.push(quote_spanned! {span=>
382                let #name = &(#qa);
383                args_count += #name.len();
384                for v in #name.iter() {
385                    size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(v);
386                }
387                if false {
388                    qusql_sqlx_type::check_arg_list_hack::<#t, _>(#name);
389                    ::std::panic!();
390                }
391            });
392            arg_add.push(quote!(
393                for v in #name.iter() {
394                    e = e.and_then(|()| query_args.add(v));
395                }
396            ));
397        } else {
398            arg_bindings.push(quote_spanned! {span=>
399                let #name = &(#qa);
400                args_count += 1;
401                size_hints += ::sqlx::encode::Encode::<#cls>::size_hint(#name);
402                if false {
403                    qusql_sqlx_type::check_arg::<#t, _>(#name);
404                    ::std::panic!();
405                }
406            });
407            arg_add.push(quote!(e = e.and_then(|()| query_args.add(#name));));
408        }
409    }
410
411    let query = if list_lengths.is_empty() {
412        quote!(#query)
413    } else {
414        quote!(
415            qusql_sqlx_type::convert_list_query(#query, &[#(#list_lengths),*])
416        )
417    };
418
419    (
420        quote! {
421            let mut size_hints = 0;
422            let mut args_count = 0;
423            #(#arg_bindings)*
424
425            let mut query_args = <#cls as ::sqlx::database::Database>::Arguments::default();
426            query_args.reserve(args_count, size_hints);
427            let mut e = Ok(());
428            #(#arg_add)*
429            let query_args = e.and_then(|()| Ok(query_args));
430        },
431        query,
432    )
433}
434
435/// Output an [Issue] as a compile error
436fn issues_to_errors(issues: Vec<Issue>, source: &str, span: Span) -> Vec<proc_macro2::TokenStream> {
437    if !issues.is_empty() {
438        let b2c = ByteToChar::new(source.as_bytes());
439        let source = NamedSource("query", Source::from(source));
440        let mut err = false;
441        let mut out = Vec::new();
442        for issue in issues {
443            if issue.level == qusql_type::Level::Error {
444                err = true;
445            }
446            let r = issue_to_report(issue, &b2c);
447            r.write(&source, &mut out).unwrap();
448        }
449        if err {
450            let raw = String::from_utf8(out).unwrap();
451            // Strip ariadne's first "Error: <message>" line — rustc provides
452            // its own "error:" heading, so keeping ariadne's makes it double.
453            let body = raw
454                .find('\n')
455                .map(|i| raw[i + 1..].trim_start_matches('\n').trim_end())
456                .unwrap_or(raw.trim_end());
457            return vec![syn::Error::new(span, body).to_compile_error()];
458        }
459    }
460    Vec::new()
461}
462
463/// Construct row struct members, and fill in statements
464fn construct_row(
465    columns: &[SelectTypeColumn],
466    is_postgres: bool,
467) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
468    let mut row_members = Vec::new();
469    let mut row_construct = Vec::new();
470    for (i, c) in columns.iter().enumerate() {
471        let mut t = match c.type_.t {
472            qusql_type::Type::U8 => quote! {u8},
473            qusql_type::Type::I8 => quote! {i8},
474            qusql_type::Type::U16 => quote! {u16},
475            qusql_type::Type::I16 => quote! {i16},
476            qusql_type::Type::U24 => quote! {u32},
477            qusql_type::Type::I24 => quote! {i32},
478            qusql_type::Type::U32 => quote! {u32},
479            qusql_type::Type::I32 => quote! {i32},
480            qusql_type::Type::U64 => quote! {u64},
481            qusql_type::Type::I64 => quote! {i64},
482            qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
483            qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
484            qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
485            qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
486            qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
487                quote! {chrono::NaiveDateTime}
488            }
489            qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
490            qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
491                if is_postgres {
492                    quote! {i32}
493                } else {
494                    quote! {i64}
495                }
496            }
497            qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
498            qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
499            qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
500                todo!("from_time_interval")
501            }
502            qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
503                quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
504            }
505            qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
506                quote! {qusql_sqlx_type::UuidValue}
507            }
508            qusql_type::Type::Null => todo!("from_null"),
509            qusql_type::Type::Invalid => quote! {i64},
510            qusql_type::Type::Enum(_) => quote! {String},
511            qusql_type::Type::Set(_) => quote! {String},
512            qusql_type::Type::Args(_, _) => todo!("from_args"),
513            qusql_type::Type::F32 => quote! {f32},
514            qusql_type::Type::F64 => quote! {f64},
515            qusql_type::Type::JSON => quote! {String},
516            qusql_type::Type::Geometry => quote! {Vec<u8>},
517            qusql_type::Type::Range(_) => quote! {Vec<u8>},
518            qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
519        };
520        let name = match &c.name {
521            Some(v) => v,
522            None => continue,
523        };
524
525        // Handle sqlx "column!" convention: strip trailing ! and force not_null
526        let name_str = name.value;
527        let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
528            (stripped, true)
529        } else {
530            (name_str, false)
531        };
532
533        let ident = String::from("r#") + name_str;
534        let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
535            ident
536        } else {
537            // TODO error
538            //errors.push(syn::Error::new(span, String::from_utf8(out).unwrap()).to_compile_error().into());
539            continue;
540        };
541
542        if !c.type_.not_null && !force_not_null {
543            t = quote! {Option<#t>};
544        }
545        row_members.push(quote! {
546            #ident : #t
547        });
548        row_construct.push(quote! {
549            #ident: sqlx::Row::get(&row, #i)
550        });
551    }
552    (row_members, row_construct)
553}
554
555/// Parsed query! macro
556struct Query {
557    /// The query expression
558    query: String,
559    /// The span of the query expression
560    query_span: Span,
561    /// The arguments to supply to the query
562    args: Vec<Expr>,
563    /// The last span parsed
564    last_span: Span,
565}
566
567impl Parse for Query {
568    fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
569        let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
570        let query: String = query_.iter().map(LitStr::value).collect();
571        let query_span = query_.span();
572        let mut last_span = query_span;
573        let mut args = Vec::new();
574        while !input.is_empty() {
575            let _ = input.parse::<syn::token::Comma>()?;
576            if input.is_empty() {
577                break;
578            }
579            let arg = input.parse::<Expr>()?;
580            last_span = arg.span();
581            args.push(arg);
582        }
583        Ok(Self {
584            query,
585            query_span,
586            args,
587            last_span,
588        })
589    }
590}
591
592/// Statically checked SQL query, similarly to sqlx::query!.
593///
594/// This expands to an instance of query::Map that outputs an ad-hoc anonymous struct type.
595#[proc_macro]
596pub fn query(input: TokenStream) -> TokenStream {
597    let query = syn::parse_macro_input!(input as Query);
598    let cache = get_schemas();
599    let (schemas, dialect, schema_hash) = (cache.schemas.get(), &cache.dialect, cache.hash);
600    let options = TypeOptions::new()
601        .dialect(dialect.clone())
602        .arguments(match &dialect {
603            SQLDialect::MariaDB => SQLArguments::QuestionMark,
604            SQLDialect::Sqlite => SQLArguments::QuestionMark,
605            SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
606        })
607        .list_hack(true);
608    let mut issues = qusql_type::Issues::new(&query.query);
609    let stmt = type_statement(&schemas.0, &query.query, &mut issues, &options);
610    let sp = cache.path.to_str().unwrap();
611    let mut errors = issues_to_errors(issues.into_vec(), &query.query, query.query_span);
612    match &stmt {
613        qusql_type::StatementType::Select { columns, arguments } => {
614            let (args_tokens, q) = quote_args(
615                &mut errors,
616                &query.query,
617                query.last_span,
618                &query.args,
619                arguments,
620                dialect,
621            );
622            let (row_members, row_construct) = construct_row(columns, dialect.is_postgresql());
623            let s = quote! { {
624                use ::sqlx::Arguments as _;
625                use ::sqlx::SqlSafeStr;
626                let _ = std::include_bytes!(#sp);
627                #(#errors; )*
628                #args_tokens;
629
630                struct Row {
631                    #(#row_members),*
632                };
633                sqlx::__query_with_result(#q, query_args).map(|row|
634                    Row{
635                        #(#row_construct),*
636                    }
637                )
638            }};
639            s.into()
640        }
641        qusql_type::StatementType::Delete {
642            arguments,
643            returning,
644        } => {
645            let (args_tokens, q) = quote_args(
646                &mut errors,
647                &query.query,
648                query.last_span,
649                &query.args,
650                arguments,
651                dialect,
652            );
653            let s = match returning.as_ref() {
654                Some(returning) => {
655                    let (row_members, row_construct) =
656                        construct_row(returning, dialect.is_postgresql());
657                    quote! { {
658                        use ::sqlx::Arguments as _;
659                        let _ = std::include_bytes!(#sp);
660                        #(#errors; )*
661                        #args_tokens
662
663                        struct Row {
664                            #(#row_members),*
665                        };
666                        sqlx::__query_with_result(#q, query_args).map(|row|
667                            Row{
668                                #(#row_construct),*
669                            }
670                        )
671                    }}
672                }
673                None => quote! { {
674                    use ::sqlx::Arguments as _;
675                    const _SCHEMA_HASH: u64 = #schema_hash;
676                    #(#errors; )*
677                    #args_tokens
678                    sqlx::__query_with_result(#q, query_args)
679                }
680                },
681            };
682            s.into()
683        }
684        qusql_type::StatementType::Insert {
685            arguments,
686            returning,
687            ..
688        } => {
689            let (args_tokens, q) = quote_args(
690                &mut errors,
691                &query.query,
692                query.last_span,
693                &query.args,
694                arguments,
695                dialect,
696            );
697            let s = match returning.as_ref() {
698                Some(returning) => {
699                    let (row_members, row_construct) =
700                        construct_row(returning, dialect.is_postgresql());
701                    quote! { {
702                        use ::sqlx::Arguments as _;
703                        let _ = std::include_bytes!(#sp);
704                        #(#errors; )*
705                        #args_tokens
706
707                        struct Row {
708                            #(#row_members),*
709                        };
710                        sqlx::__query_with_result(#q, query_args).map(|row|
711                            Row{
712                                #(#row_construct),*
713                            }
714                        )
715                    }}
716                }
717                None => quote! { {
718                    use ::sqlx::Arguments as _;
719                    const _SCHEMA_HASH: u64 = #schema_hash;
720                    #(#errors; )*
721                    #args_tokens
722                    sqlx::__query_with_result(#q, query_args)
723                }
724                },
725            };
726            s.into()
727        }
728        qusql_type::StatementType::Update {
729            arguments,
730            returning,
731        } => {
732            let (args_tokens, q) = quote_args(
733                &mut errors,
734                &query.query,
735                query.last_span,
736                &query.args,
737                arguments,
738                dialect,
739            );
740
741            let s = match returning.as_ref() {
742                Some(returning) => {
743                    let (row_members, row_construct) =
744                        construct_row(returning, dialect.is_postgresql());
745                    quote! { {
746                        use ::sqlx::Arguments as _;
747                        let _ = std::include_bytes!(#sp);
748                        #(#errors; )*
749                        #args_tokens
750
751                        struct Row {
752                            #(#row_members),*
753                        };
754                        sqlx::__query_with_result(#q, query_args).map(|row|
755                            Row{
756                                #(#row_construct),*
757                            }
758                        )
759                    }}
760                }
761                None => quote! { {
762                    use ::sqlx::Arguments as _;
763                    const _SCHEMA_HASH: u64 = #schema_hash;
764                    #(#errors; )*
765                    #args_tokens
766                    sqlx::__query_with_result(#q, query_args)
767                }
768                },
769            };
770            s.into()
771        }
772        qusql_type::StatementType::Replace {
773            arguments,
774            returning,
775        } => {
776            let (args_tokens, q) = quote_args(
777                &mut errors,
778                &query.query,
779                query.last_span,
780                &query.args,
781                arguments,
782                dialect,
783            );
784            let s = match returning.as_ref() {
785                Some(returning) => {
786                    let (row_members, row_construct) =
787                        construct_row(returning, dialect.is_postgresql());
788                    quote! { {
789                        use ::sqlx::Arguments as _;
790                        const _SCHEMA_HASH: u64 = #schema_hash;
791                        let _ = std::include_bytes!(#sp);
792                        #(#errors; )*
793                        #args_tokens
794
795                        struct Row {
796                            #(#row_members),*
797                        };
798                        sqlx::__query_with_result(#q, query_args).map(|row|
799                            Row{
800                                #(#row_construct),*
801                            }
802                        )
803                    }}
804                }
805                None => quote! { {
806                    use ::sqlx::Arguments as _;
807                    #(#errors; )*
808                    #args_tokens
809                    sqlx::__query_with_result(#q, query_args)
810                }
811                },
812            };
813            s.into()
814        }
815        qusql_type::StatementType::Truncate => {
816            errors.push(
817                syn::Error::new(query.query_span, "TRUNCATE not supported in query!")
818                    .to_compile_error(),
819            );
820            quote! { {
821                #(#errors; )*
822                todo!("truncate")
823            }}
824            .into()
825        }
826        qusql_type::StatementType::Call { .. } => {
827            errors.push(
828                syn::Error::new(query.query_span, "CALL not supported in query!")
829                    .to_compile_error(),
830            );
831            quote! { {
832                #(#errors; )*
833                todo!("call")
834            }}
835            .into()
836        }
837        qusql_type::StatementType::Transaction => {
838            errors.push(
839                syn::Error::new(
840                    query.query_span,
841                    "Transaction control not supported in query!",
842                )
843                .to_compile_error(),
844            );
845            quote! { {
846                #(#errors; )*
847                todo!("transaction")
848            }}
849            .into()
850        }
851        qusql_type::StatementType::Set => {
852            errors.push(
853                syn::Error::new(query.query_span, "SET not supported in query!").to_compile_error(),
854            );
855            quote! { {
856                #(#errors; )*
857                todo!("set")
858            }}
859            .into()
860        }
861        qusql_type::StatementType::Lock => {
862            errors.push(
863                syn::Error::new(query.query_span, "LOCK not supported in query!")
864                    .to_compile_error(),
865            );
866            quote! { {
867                #(#errors; )*
868                todo!("lock")
869            }}
870            .into()
871        }
872        qusql_type::StatementType::Invalid => {
873            let s = quote! { {
874                #(#errors; )*;
875                todo!("Invalid")
876            }};
877            s.into()
878        }
879    }
880}
881
882/// Fill in row values in a query_as struct
883fn construct_row2(
884    columns: &[SelectTypeColumn],
885    is_postgres: bool,
886) -> Vec<proc_macro2::TokenStream> {
887    let mut row_construct = Vec::new();
888    for (i, c) in columns.iter().enumerate() {
889        let mut t = match c.type_.t {
890            qusql_type::Type::U8 => quote! {u8},
891            qusql_type::Type::I8 => quote! {i8},
892            qusql_type::Type::U16 => quote! {u16},
893            qusql_type::Type::I16 => quote! {i16},
894            qusql_type::Type::U24 => quote! {u32},
895            qusql_type::Type::I24 => quote! {i32},
896            qusql_type::Type::U32 => quote! {u32},
897            qusql_type::Type::I32 => quote! {i32},
898            qusql_type::Type::U64 => quote! {u64},
899            qusql_type::Type::I64 => quote! {i64},
900            qusql_type::Type::Base(qusql_type::BaseType::Any) => todo!("from_any"),
901            qusql_type::Type::Base(qusql_type::BaseType::Bool) => quote! {bool},
902            qusql_type::Type::Base(qusql_type::BaseType::Bytes) => quote! {Vec<u8>},
903            qusql_type::Type::Base(qusql_type::BaseType::Date) => quote! {chrono::NaiveDate},
904            qusql_type::Type::Base(qusql_type::BaseType::DateTime) => {
905                quote! {chrono::NaiveDateTime}
906            }
907            qusql_type::Type::Base(qusql_type::BaseType::Float) => quote! {f64},
908            qusql_type::Type::Base(qusql_type::BaseType::Integer) => {
909                if is_postgres {
910                    quote! {i32}
911                } else {
912                    quote! {i64}
913                }
914            }
915            qusql_type::Type::Base(qusql_type::BaseType::String) => quote! {String},
916            qusql_type::Type::Base(qusql_type::BaseType::Time) => todo!("from_time"),
917            qusql_type::Type::Base(qusql_type::BaseType::TimeInterval) => {
918                todo!("from_time_interval")
919            }
920            qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => {
921                quote! {sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>}
922            }
923            qusql_type::Type::Base(qusql_type::BaseType::Uuid) => {
924                quote! {qusql_sqlx_type::UuidValue}
925            }
926            qusql_type::Type::Null => todo!("from_null"),
927            qusql_type::Type::Invalid => quote! {i64},
928            qusql_type::Type::Enum(_) => quote! {String},
929            qusql_type::Type::Set(_) => quote! {String},
930            qusql_type::Type::Args(_, _) => todo!("from_args"),
931            qusql_type::Type::F32 => quote! {f32},
932            qusql_type::Type::F64 => quote! {f64},
933            qusql_type::Type::JSON => quote! {String},
934            qusql_type::Type::Geometry => quote! {Vec<u8>},
935            qusql_type::Type::Range(_) => quote! {Vec<u8>},
936            qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any},
937        };
938        let name = match &c.name {
939            Some(v) => v,
940            None => continue,
941        };
942
943        // Handle sqlx "column!" convention: strip trailing ! and force not_null
944        let name_str = name.value;
945        let (name_str, force_not_null) = if let Some(stripped) = name_str.strip_suffix('!') {
946            (stripped, true)
947        } else {
948            (name_str, false)
949        };
950
951        let ident = String::from("r#") + name_str;
952        let ident: Ident = if let Ok(ident) = syn::parse_str(&ident) {
953            ident
954        } else {
955            // TODO error
956            //errors.push(syn::Error::new(span, String::from_utf8(out).unwrap()).to_compile_error().into());
957            continue;
958        };
959
960        if !c.type_.not_null && !force_not_null {
961            t = quote! {Option<#t>};
962        }
963        row_construct.push(quote! {
964            #ident: qusql_sqlx_type::arg_out::<#t, _, #i>(sqlx::Row::get(&row, #i))
965        });
966    }
967    row_construct
968}
969
970/// Parse result of a query_as! macro
971struct QueryAs {
972    /// Name of output type
973    as_: Ident,
974    /// The query to execute
975    query: String,
976    /// The span of the query to execute
977    query_span: Span,
978    /// The arguments to supply
979    args: Vec<Expr>,
980    /// The last span parsed
981    last_span: Span,
982}
983
984impl Parse for QueryAs {
985    fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
986        let as_ = input.parse::<Ident>()?;
987        let _ = input.parse::<syn::token::Comma>()?;
988
989        let query_ = Punctuated::<LitStr, Token![+]>::parse_separated_nonempty(input)?;
990        let query: String = query_.iter().map(LitStr::value).collect();
991        let query_span = query_.span();
992
993        let mut last_span = query_span;
994        let mut args = Vec::new();
995        while !input.is_empty() {
996            let _ = input.parse::<syn::token::Comma>()?;
997            if input.is_empty() {
998                break;
999            }
1000            let arg = input.parse::<Expr>()?;
1001            last_span = arg.span();
1002            args.push(arg);
1003        }
1004        Ok(Self {
1005            as_,
1006            query,
1007            query_span,
1008            args,
1009            last_span,
1010        })
1011    }
1012}
1013
1014/// A variant of query! which takes a path to an explicitly defined struct as the output type.
1015///
1016/// This lets you return the struct from a function or add your own trait implementations.
1017#[proc_macro]
1018pub fn query_as(input: TokenStream) -> TokenStream {
1019    let query_as = syn::parse_macro_input!(input as QueryAs);
1020    let cache = get_schemas();
1021    let (schemas, dialect) = (cache.schemas.get(), &cache.dialect);
1022    let options = TypeOptions::new()
1023        .dialect(dialect.clone())
1024        .arguments(match &dialect {
1025            SQLDialect::MariaDB => SQLArguments::QuestionMark,
1026            SQLDialect::Sqlite => SQLArguments::QuestionMark,
1027            SQLDialect::PostgreSQL | SQLDialect::PostGIS => SQLArguments::Dollar,
1028        })
1029        .list_hack(true);
1030    let mut issues = qusql_type::Issues::new(&query_as.query);
1031    let stmt = type_statement(&schemas.0, &query_as.query, &mut issues, &options);
1032
1033    let mut errors = issues_to_errors(issues.into_vec(), &query_as.query, query_as.query_span);
1034    match &stmt {
1035        qusql_type::StatementType::Select { columns, arguments } => {
1036            let (args_tokens, q) = quote_args(
1037                &mut errors,
1038                &query_as.query,
1039                query_as.last_span,
1040                &query_as.args,
1041                arguments,
1042                dialect,
1043            );
1044
1045            let row_construct = construct_row2(columns, dialect.is_postgresql());
1046            let row = query_as.as_;
1047            let s = quote! { {
1048                use ::sqlx::Arguments as _;
1049                #(#errors; )*
1050                #args_tokens
1051                sqlx::__query_with_result(#q, query_args).map(|row|
1052                    #row{
1053                        #(#row_construct),*
1054                    }
1055                )
1056            }};
1057            //println!("TOKENS: {}", s);
1058            s.into()
1059        }
1060        qusql_type::StatementType::Delete { .. } => {
1061            errors.push(
1062                syn::Error::new(query_as.query_span, "DELETE not support in query_as")
1063                    .to_compile_error(),
1064            );
1065            quote! { {
1066                #(#errors; )*
1067                todo!("delete")
1068            }}
1069            .into()
1070        }
1071        qusql_type::StatementType::Insert {
1072            returning: None, ..
1073        } => {
1074            errors.push(
1075                syn::Error::new(
1076                    query_as.query_span,
1077                    "INSERT without RETURNING not support in query_as",
1078                )
1079                .to_compile_error(),
1080            );
1081            quote! { {
1082                #(#errors; )*
1083                todo!("insert")
1084            }}
1085            .into()
1086        }
1087        qusql_type::StatementType::Insert {
1088            arguments,
1089            returning: Some(returning),
1090            ..
1091        } => {
1092            let (args_tokens, q) = quote_args(
1093                &mut errors,
1094                &query_as.query,
1095                query_as.last_span,
1096                &query_as.args,
1097                arguments,
1098                dialect,
1099            );
1100
1101            let row_construct = construct_row2(returning, dialect.is_postgresql());
1102            let row = query_as.as_;
1103            let s = quote! { {
1104                use ::sqlx::Arguments as _;
1105                #(#errors; )*
1106                #args_tokens
1107                sqlx::__query_with_result(#q, query_args).map(|row|
1108                    #row{
1109                        #(#row_construct),*
1110                    }
1111                )
1112            }};
1113            s.into()
1114        }
1115        qusql_type::StatementType::Update {
1116            returning: None, ..
1117        } => {
1118            errors.push(
1119                syn::Error::new(
1120                    query_as.query_span,
1121                    "UPDATE without RETURNING not support in query_as",
1122                )
1123                .to_compile_error(),
1124            );
1125            quote! { {
1126                #(#errors; )*
1127                todo!("update")
1128            }}
1129            .into()
1130        }
1131        qusql_type::StatementType::Update {
1132            arguments,
1133            returning: Some(returning),
1134            ..
1135        } => {
1136            let (args_tokens, q) = quote_args(
1137                &mut errors,
1138                &query_as.query,
1139                query_as.last_span,
1140                &query_as.args,
1141                arguments,
1142                dialect,
1143            );
1144
1145            let row_construct = construct_row2(returning, dialect.is_postgresql());
1146            let row = query_as.as_;
1147            let s = quote! { {
1148                use ::sqlx::Arguments as _;
1149                #(#errors; )*
1150                #args_tokens
1151                sqlx::__query_with_result(#q, query_args).map(|row|
1152                    #row{
1153                        #(#row_construct),*
1154                    }
1155                )
1156            }};
1157            s.into()
1158        }
1159        qusql_type::StatementType::Replace {
1160            returning: None, ..
1161        } => {
1162            errors.push(
1163                syn::Error::new(
1164                    query_as.query_span,
1165                    "REPLACE without RETURNING not support in query_as",
1166                )
1167                .to_compile_error(),
1168            );
1169            quote! { {
1170                #(#errors; )*
1171                todo!("replace")
1172            }}
1173            .into()
1174        }
1175        qusql_type::StatementType::Replace {
1176            arguments,
1177            returning: Some(returning),
1178            ..
1179        } => {
1180            let (args_tokens, q) = quote_args(
1181                &mut errors,
1182                &query_as.query,
1183                query_as.last_span,
1184                &query_as.args,
1185                arguments,
1186                dialect,
1187            );
1188
1189            let row_construct = construct_row2(returning, dialect.is_postgresql());
1190            let row = query_as.as_;
1191            let s = quote! { {
1192                use ::sqlx::Arguments as _;
1193                #(#errors; )*
1194                #args_tokens
1195                sqlx::__query_with_result(#q, query_args).map(|row|
1196                    #row{
1197                        #(#row_construct),*
1198                    }
1199                )
1200            }};
1201            s.into()
1202        }
1203        qusql_type::StatementType::Truncate => {
1204            errors.push(
1205                syn::Error::new(query_as.query_span, "TRUNCATE not supported in query_as!")
1206                    .to_compile_error(),
1207            );
1208            quote! { {
1209                #(#errors; )*
1210                todo!("truncate")
1211            }}
1212            .into()
1213        }
1214        qusql_type::StatementType::Call { .. } => {
1215            errors.push(
1216                syn::Error::new(query_as.query_span, "CALL not supported in query_as!")
1217                    .to_compile_error(),
1218            );
1219            quote! { {
1220                #(#errors; )*
1221                todo!("call")
1222            }}
1223            .into()
1224        }
1225        qusql_type::StatementType::Transaction => {
1226            errors.push(
1227                syn::Error::new(
1228                    query_as.query_span,
1229                    "Transaction control not supported in query_as!",
1230                )
1231                .to_compile_error(),
1232            );
1233            quote! { {
1234                #(#errors; )*
1235                todo!("transaction")
1236            }}
1237            .into()
1238        }
1239        qusql_type::StatementType::Set => {
1240            errors.push(
1241                syn::Error::new(query_as.query_span, "SET not supported in query_as!")
1242                    .to_compile_error(),
1243            );
1244            quote! { {
1245                #(#errors; )*
1246                todo!("set")
1247            }}
1248            .into()
1249        }
1250        qusql_type::StatementType::Lock => {
1251            errors.push(
1252                syn::Error::new(query_as.query_span, "LOCK not supported in query_as!")
1253                    .to_compile_error(),
1254            );
1255            quote! { {
1256                #(#errors; )*
1257                todo!("lock")
1258            }}
1259            .into()
1260        }
1261        qusql_type::StatementType::Invalid => quote! { {
1262            #(#errors; )*;
1263            todo!("invalid")
1264        }}
1265        .into(),
1266    }
1267}