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