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