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