diff --git a/qusql-mysql-type-macro/Cargo.toml b/qusql-mysql-type-macro/Cargo.toml index d8e43667..a3e5776e 100644 --- a/qusql-mysql-type-macro/Cargo.toml +++ b/qusql-mysql-type-macro/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-mysql-type-macro" -version = "0.1.9" +version = "0.1.10" authors = ["Jakob Truelsen "] edition = "2021" license = "Apache-2.0" @@ -28,9 +28,9 @@ list_hack = [] [dependencies] quote = "1" -syn = { version = "2", features = ["full", "parsing"] } +syn = { version = "3", features = ["full", "parsing"] } proc-macro2 = "1" -qusql-type = { path="../qusql-type", version="0.7.1" } +qusql-type = { path="../qusql-type", version="0.8" } ariadne = "0.6" serde = { version = "1", features = ["derive"] } serde_json = "1" diff --git a/qusql-mysql-type-macro/src/lib.rs b/qusql-mysql-type-macro/src/lib.rs index 448233b2..9a60bc80 100644 --- a/qusql-mysql-type-macro/src/lib.rs +++ b/qusql-mysql-type-macro/src/lib.rs @@ -301,6 +301,7 @@ fn map_type(ta: &FullType<'_>) -> proc_macro2::TokenStream { qusql_type::Type::Geometry => quote! {qusql_mysql_type::Any}, qusql_type::Type::Array(_) => quote! {qusql_mysql_type::Any}, qusql_type::Type::Range(_) => todo!(), + qusql_type::Type::MultiRange(_) => todo!(), }; if !ta.not_null { quote! {Option<#t>} @@ -489,6 +490,7 @@ fn construct_row( qusql_type::Type::Geometry => quote! {Vec}, qusql_type::Type::Array(_) => quote! {qusql_mysql_type::Any}, qusql_type::Type::Range(_) => todo!(), + qusql_type::Type::MultiRange(_) => todo!(), }; let name = match &c.name { Some(v) => v, diff --git a/qusql-mysql-type/Cargo.toml b/qusql-mysql-type/Cargo.toml index 471dd597..a81c4fde 100644 --- a/qusql-mysql-type/Cargo.toml +++ b/qusql-mysql-type/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-mysql-type" -version = "0.1.9" +version = "0.1.10" edition = "2024" license = "Apache-2.0" keywords = [ "mysql" ] @@ -26,6 +26,6 @@ chrono = ["dep:chrono", "qusql-mysql/chrono"] list_hack = ["qusql-mysql/list_hack", "qusql-mysql-type-macro/list_hack"] [dependencies] -qusql-mysql-type-macro={path="../qusql-mysql-type-macro", version="0.1.9"} -qusql-mysql={path="../qusql-mysql", version="0.1.0"} +qusql-mysql-type-macro={path="../qusql-mysql-type-macro", version="0.1.10"} +qusql-mysql={path="../qusql-mysql", version="0.1.1"} chrono = {version = "0.4", optional=true} \ No newline at end of file diff --git a/qusql-parse/Cargo.toml b/qusql-parse/Cargo.toml index 8a66f224..bff70778 100644 --- a/qusql-parse/Cargo.toml +++ b/qusql-parse/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-parse" -version = "0.10.0" +version = "0.11.0" edition = "2024" authors = ["Jakob Truelsen "] keywords = [ "mysql", "postgresql", "sql", "lexer", "parser" ] diff --git a/qusql-parse/src/function_expression.rs b/qusql-parse/src/function_expression.rs index 6333f740..8ff82e12 100644 --- a/qusql-parse/src/function_expression.rs +++ b/qusql-parse/src/function_expression.rs @@ -832,6 +832,26 @@ pub enum Function<'a> { ArrayUpper, Cardinality, TrimArray, + // PostgreSQL range/multirange functions + IsEmpty, + LowerInc, + LowerInf, + UpperInc, + UpperInf, + RangeMerge, + Multirange, + Int4Range, + Int8Range, + NumRange, + TsRange, + TstzRange, + DateRange, + Int4Multirange, + Int8Multirange, + NumMultirange, + TsMultirange, + TstzMultirange, + DateMultirange, // PostgreSQL text search functions ArrayToTsvector, GetCurrentTsConfig, @@ -2605,6 +2625,65 @@ pub(crate) fn parse_function<'a>( Function::TrimArray } + // PostgreSQL range/multirange functions + Token::Ident(_, Keyword::ISEMPTY) if parser.options.dialect.is_postgresql() => { + Function::IsEmpty + } + Token::Ident(_, Keyword::LOWER_INC) if parser.options.dialect.is_postgresql() => { + Function::LowerInc + } + Token::Ident(_, Keyword::LOWER_INF) if parser.options.dialect.is_postgresql() => { + Function::LowerInf + } + Token::Ident(_, Keyword::UPPER_INC) if parser.options.dialect.is_postgresql() => { + Function::UpperInc + } + Token::Ident(_, Keyword::UPPER_INF) if parser.options.dialect.is_postgresql() => { + Function::UpperInf + } + Token::Ident(_, Keyword::RANGE_MERGE) if parser.options.dialect.is_postgresql() => { + Function::RangeMerge + } + Token::Ident(_, Keyword::MULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::Multirange + } + Token::Ident(_, Keyword::INT4RANGE) if parser.options.dialect.is_postgresql() => { + Function::Int4Range + } + Token::Ident(_, Keyword::INT8RANGE) if parser.options.dialect.is_postgresql() => { + Function::Int8Range + } + Token::Ident(_, Keyword::NUMRANGE) if parser.options.dialect.is_postgresql() => { + Function::NumRange + } + Token::Ident(_, Keyword::TSRANGE) if parser.options.dialect.is_postgresql() => { + Function::TsRange + } + Token::Ident(_, Keyword::TSTZRANGE) if parser.options.dialect.is_postgresql() => { + Function::TstzRange + } + Token::Ident(_, Keyword::DATERANGE) if parser.options.dialect.is_postgresql() => { + Function::DateRange + } + Token::Ident(_, Keyword::INT4MULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::Int4Multirange + } + Token::Ident(_, Keyword::INT8MULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::Int8Multirange + } + Token::Ident(_, Keyword::NUMMULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::NumMultirange + } + Token::Ident(_, Keyword::TSMULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::TsMultirange + } + Token::Ident(_, Keyword::TSTZMULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::TstzMultirange + } + Token::Ident(_, Keyword::DATEMULTIRANGE) if parser.options.dialect.is_postgresql() => { + Function::DateMultirange + } + // PostgreSQL text search functions Token::Ident(_, Keyword::ARRAY_TO_TSVECTOR) if parser.options.dialect.is_postgresql() => { Function::ArrayToTsvector diff --git a/qusql-parse/src/keywords.rs b/qusql-parse/src/keywords.rs index 2aaa2a4e..610478bf 100644 --- a/qusql-parse/src/keywords.rs +++ b/qusql-parse/src/keywords.rs @@ -741,6 +741,7 @@ pub(crate) enum Keyword { IS_IPV6, IS_USED_LOCK, IS_UUID, + ISEMPTY, ISOLATION, ISCLOSED, ISOPEN, @@ -879,6 +880,8 @@ pub(crate) enum Keyword { LOGIN, LOGS, LOWER, + LOWER_INC, + LOWER_INF, LPAD, LSEG, LTRIM, @@ -945,6 +948,7 @@ pub(crate) enum Keyword { MONITOR, MONTH, MONTHNAME, + MULTIRANGE, MUTEX, MXID_AGE, MYSQL_ERRNO, @@ -1159,6 +1163,7 @@ pub(crate) enum Keyword { RAND, RANDOM_BYTES, RANDOM_NORMAL, + RANGE_MERGE, RANK, RAW, READ_ONLY, @@ -1655,6 +1660,8 @@ pub(crate) enum Keyword { UPDATEXML, UPGRADE, UPPER, + UPPER_INC, + UPPER_INF, USE_FRM, USER_RESOURCES, USER, diff --git a/qusql-py-mysql-type-plugin/Cargo.toml b/qusql-py-mysql-type-plugin/Cargo.toml index 6d8414d2..dc40f0e1 100644 --- a/qusql-py-mysql-type-plugin/Cargo.toml +++ b/qusql-py-mysql-type-plugin/Cargo.toml @@ -14,7 +14,7 @@ name = "qusql_mysql_type_plugin" crate-type = ["cdylib"] [dependencies] -pyo3 = { version = "0.28", features = ["extension-module"] } -qusql-type = {version="0.7.0", path="../qusql-type" } +pyo3 = { version = "0.29", features = ["extension-module"] } +qusql-type = {version="0.8.0", path="../qusql-type" } ariadne = "0.6" yoke = { version ="0.8", features = ["derive", "alloc"] } diff --git a/qusql-py-mysql-type-plugin/src/lib.rs b/qusql-py-mysql-type-plugin/src/lib.rs index 65190e0f..9c21063e 100644 --- a/qusql-py-mysql-type-plugin/src/lib.rs +++ b/qusql-py-mysql-type-plugin/src/lib.rs @@ -271,6 +271,7 @@ fn map_type(t: &qusql_type::FullType<'_>) -> Type { qusql_type::Type::Geometry => Type::Any, qusql_type::Type::Array(_) => Type::Any, qusql_type::Type::Range(_) => Type::Any, + qusql_type::Type::MultiRange(_) => Type::Any, qusql_type::Type::Set(_) => Type::String, qusql_type::Type::U16 => Type::Integer, qusql_type::Type::U24 => Type::Integer, diff --git a/qusql-sqlx-type-macro/Cargo.toml b/qusql-sqlx-type-macro/Cargo.toml index 25cec2ef..46ab7055 100644 --- a/qusql-sqlx-type-macro/Cargo.toml +++ b/qusql-sqlx-type-macro/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-sqlx-type-macro" -version = "0.4.2" +version = "0.4.3" authors = ["Jakob Truelsen "] edition = "2021" license = "Apache-2.0" @@ -25,9 +25,9 @@ proc-macro = true [dependencies] quote = "1" -syn = { version = "2", features = ["full", "parsing"] } +syn = { version = "3", features = ["full", "parsing"] } proc-macro2 = "1" -qusql-type = { path = "../qusql-type", version="0.7.2"} +qusql-type = { path = "../qusql-type", version="0.8.0"} ariadne = "0.6" serde = { version = "1", features = ["derive"] } serde_json = "1" diff --git a/qusql-sqlx-type-macro/src/lib.rs b/qusql-sqlx-type-macro/src/lib.rs index 3e7daecc..33843454 100644 --- a/qusql-sqlx-type-macro/src/lib.rs +++ b/qusql-sqlx-type-macro/src/lib.rs @@ -271,6 +271,41 @@ fn get_schemas() -> Arc { entry } +/// Rust type used to bind/read a PostgreSQL range column, keyed by the (boxed) element +/// `qusql_type::Type` of a `qusql_type::Type::Range`/`Type::MultiRange`. +/// +/// `sqlx::postgres::types::PgRange` (re-exported as `qusql_sqlx_type::PgRange`) only has +/// built-in `sqlx` support for a handful of element types without pulling in extra optional +/// dependencies (in particular `NUMRANGE`/numeric would need the `bigdecimal` or +/// `rust_decimal` crate, which this crate does not integrate). Dialects other than +/// PostgreSQL never produce range types, and unsupported element types fall back to +/// `fallback`. +fn quote_range_elem_type( + elem: &qusql_type::Type<'_>, + is_postgres: bool, + fallback: proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + if !is_postgres { + return fallback; + } + match elem { + // `int4range`/`int4multirange` + qusql_type::Type::I32 => quote! { qusql_sqlx_type::PgRange }, + // `int8range`/`int8multirange` + qusql_type::Type::I64 => quote! { qusql_sqlx_type::PgRange }, + qusql_type::Type::Base(qusql_type::BaseType::Date) => { + quote! { qusql_sqlx_type::PgRange } + } + qusql_type::Type::Base(qusql_type::BaseType::DateTime) => { + quote! { qusql_sqlx_type::PgRange } + } + qusql_type::Type::Base(qusql_type::BaseType::TimeStamp) => { + quote! { qusql_sqlx_type::PgRange> } + } + _ => fallback, + } +} + /// Produce quoted arguments for a query fn quote_args( errors: &mut Vec, @@ -285,6 +320,7 @@ fn quote_args( SQLDialect::Sqlite => quote!(sqlx::sqlite::Sqlite), SQLDialect::PostgreSQL | SQLDialect::PostGIS => quote!(sqlx::postgres::Postgres), }; + let is_postgres = dialect.is_postgresql(); let mut at = Vec::new(); let inv = qusql_type::FullType::invalid(); @@ -331,7 +367,7 @@ fn quote_args( let mut list_lengths = Vec::new(); for ((qa, ta), name) in args.iter().zip(at).zip(&arg_names) { - let mut t = match ta.t { + let mut t = match &ta.t { qusql_type::Type::U8 => quote! {u8}, qusql_type::Type::I8 => quote! {i8}, qusql_type::Type::U16 => quote! {u16}, @@ -369,7 +405,10 @@ fn quote_args( qusql_type::Type::F64 => quote! {f64}, qusql_type::Type::JSON => quote! {qusql_sqlx_type::Any}, qusql_type::Type::Geometry => quote! {qusql_sqlx_type::Any}, - qusql_type::Type::Range(_) => quote! {qusql_sqlx_type::Any}, + qusql_type::Type::Range(elem) => { + quote_range_elem_type(elem, is_postgres, quote! {qusql_sqlx_type::Any}) + } + qusql_type::Type::MultiRange(_) => quote! {qusql_sqlx_type::Any}, qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any}, }; if !ta.not_null { @@ -468,7 +507,7 @@ fn construct_row( let mut row_members = Vec::new(); let mut row_construct = Vec::new(); for (i, c) in columns.iter().enumerate() { - let mut t = match c.type_.t { + let mut t = match &c.type_.t { qusql_type::Type::U8 => quote! {u8}, qusql_type::Type::I8 => quote! {i8}, qusql_type::Type::U16 => quote! {u16}, @@ -521,7 +560,10 @@ fn construct_row( } } qusql_type::Type::Geometry => quote! {Vec}, - qusql_type::Type::Range(_) => quote! {Vec}, + qusql_type::Type::MultiRange(_) => quote! {Vec}, + qusql_type::Type::Range(elem) => { + quote_range_elem_type(elem, is_postgres, quote! {Vec}) + } qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any}, }; let name = match &c.name { @@ -893,7 +935,7 @@ fn construct_row2( ) -> Vec { let mut row_construct = Vec::new(); for (i, c) in columns.iter().enumerate() { - let mut t = match c.type_.t { + let mut t = match &c.type_.t { qusql_type::Type::U8 => quote! {u8}, qusql_type::Type::I8 => quote! {i8}, qusql_type::Type::U16 => quote! {u16}, @@ -946,7 +988,10 @@ fn construct_row2( } } qusql_type::Type::Geometry => quote! {Vec}, - qusql_type::Type::Range(_) => quote! {Vec}, + qusql_type::Type::MultiRange(_) => quote! {Vec}, + qusql_type::Type::Range(elem) => { + quote_range_elem_type(elem, is_postgres, quote! {Vec}) + } qusql_type::Type::Array(_) => quote! {qusql_sqlx_type::Any}, }; let name = match &c.name { diff --git a/qusql-sqlx-type-test/Cargo.toml b/qusql-sqlx-type-test/Cargo.toml index 63d0db64..b05ea2ba 100644 --- a/qusql-sqlx-type-test/Cargo.toml +++ b/qusql-sqlx-type-test/Cargo.toml @@ -5,7 +5,7 @@ edition = "2024" publish = false [dev-dependencies] -qusql-sqlx-type = { path = "../qusql-sqlx-type", features = ["json"] } +qusql-sqlx-type = { path = "../qusql-sqlx-type", features = ["json", "postgres", "uuid"] } sqlx = { version = "0.9", default-features = false, features = [ "postgres", "runtime-tokio", @@ -13,6 +13,9 @@ sqlx = { version = "0.9", default-features = false, features = [ "json", "migrate", "macros", + "uuid", ] } tokio = { version = "1", features = ["full"] } serde_json = "1" +chrono = "0.4" +uuid = "1" diff --git a/qusql-sqlx-type-test/sqlx-type-schema.sql b/qusql-sqlx-type-test/sqlx-type-schema.sql index f3143361..00363062 100644 --- a/qusql-sqlx-type-test/sqlx-type-schema.sql +++ b/qusql-sqlx-type-test/sqlx-type-schema.sql @@ -8,7 +8,15 @@ CREATE TABLE IF NOT EXISTS type_test_items ( score integer NOT NULL DEFAULT 0, active boolean NOT NULL DEFAULT true, ratio float8 NOT NULL DEFAULT 0.0, - props jsonb NOT NULL DEFAULT '{}' + props jsonb NOT NULL DEFAULT '{}', + validity daterange ); CREATE SEQUENCE IF NOT EXISTS type_test_seq; + +CREATE TABLE IF NOT EXISTS partial_files ( + id uuid PRIMARY KEY DEFAULT gen_random_uuid(), + uploaded_bytes int8multirange NOT NULL DEFAULT '{}', + last_modified timestamptz NOT NULL DEFAULT now() +); + diff --git a/qusql-sqlx-type-test/src/test.rs b/qusql-sqlx-type-test/src/test.rs index 1d8b43bd..b8a0e980 100644 --- a/qusql-sqlx-type-test/src/test.rs +++ b/qusql-sqlx-type-test/src/test.rs @@ -227,3 +227,92 @@ async fn test_json_extract_text_operator_is_string(pool: PgPool) { let k: Option = row.k; assert_eq!(k.as_deref(), Some("hello")); } + +/// A `daterange` column must decode as `qusql_sqlx_type::PgRange`. +#[sqlx::test] +async fn test_daterange_column_is_pgrange(pool: PgPool) { + setup(&pool).await; + let range = qusql_sqlx_type::PgRange::from( + chrono::NaiveDate::from_ymd_opt(2024, 1, 1).unwrap() + ..chrono::NaiveDate::from_ymd_opt(2024, 2, 1).unwrap(), + ); + sqlx::query("UPDATE type_test_items SET validity = $1 WHERE name = 'alpha'") + .bind(range.clone()) + .execute(&pool) + .await + .unwrap(); + let row = query!("SELECT validity FROM type_test_items WHERE name = 'alpha'") + .fetch_one(&pool) + .await + .unwrap(); + let validity: Option> = row.validity; + assert_eq!(validity, Some(range)); +} + +/// `int4range(1, 10)` must decode as `qusql_sqlx_type::PgRange`. +#[sqlx::test] +async fn test_int4range_function_is_pgrange(pool: PgPool) { + let row = query!("SELECT int4range(1, 10) AS r") + .fetch_one(&pool) + .await + .unwrap(); + let r: qusql_sqlx_type::PgRange = row.r; + assert_eq!(r, qusql_sqlx_type::PgRange::from(1..10)); +} + +/// `isempty()`/`lower()`/`upper()` on a range must type as `bool`/`i32`. +#[sqlx::test] +async fn test_range_functions(pool: PgPool) { + let row = query!( + "SELECT isempty(int4range(1, 10)) AS is_empty, \ + lower(int4range(1, 10)) AS lo, \ + upper(int4range(1, 10)) AS hi" + ) + .fetch_one(&pool) + .await + .unwrap(); + let is_empty: bool = row.is_empty; + let lo: Option = row.lo; + let hi: Option = row.hi; + assert!(!is_empty); + assert_eq!(lo, Some(1)); + assert_eq!(hi, Some(10)); +} + +/// Reported real-world usage: a writable CTE that unions a range into an +/// `int8multirange` column via `+ multirange(int8range(...))`, then `unnest()`s the +/// updated value. Also verifies that `int8multirange`/`int8range` decode with `i64` +/// bounds (not `i32`, which is what `int4range`/`int4multirange` use). +#[sqlx::test] +async fn test_int8multirange_update_returning_unnest(pool: PgPool) { + setup(&pool).await; + sqlx::query("DELETE FROM partial_files") + .execute(&pool) + .await + .unwrap(); + let id: uuid::Uuid = + sqlx::query_scalar("INSERT INTO partial_files DEFAULT VALUES RETURNING id") + .fetch_one(&pool) + .await + .unwrap(); + + let row = query!( + "WITH the_update AS ( + UPDATE partial_files + SET last_modified = now(), + uploaded_bytes = uploaded_bytes + multirange(int8range($2, $3, '[)')) + WHERE id = $1 + RETURNING uploaded_bytes + ) + SELECT unnest(uploaded_bytes) AS \"range\" + FROM the_update", + id, + 0i64, + 4_294_967_296i64 // larger than i32::MAX, to prove i64 bounds are actually used + ) + .fetch_one(&pool) + .await + .unwrap(); + let range: qusql_sqlx_type::PgRange = row.range.expect("unnest() row must be present"); + assert_eq!(range, qusql_sqlx_type::PgRange::from(0..4_294_967_296i64)); +} diff --git a/qusql-sqlx-type/Cargo.toml b/qusql-sqlx-type/Cargo.toml index 00764b5e..cdaf689f 100644 --- a/qusql-sqlx-type/Cargo.toml +++ b/qusql-sqlx-type/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-sqlx-type" -version = "0.4.2" +version = "0.4.3" authors = ["Jakob Truelsen "] edition = "2021" license = "Apache-2.0" @@ -22,10 +22,10 @@ multiple_unsafe_ops_per_block = "deny" missing_docs_in_private_items = "deny" [dev-dependencies] -sqlx = { version = "0.9", default-features = false, features = ["chrono", "runtime-tokio", "mysql"] } +sqlx = { version = "0.9", default-features = false, features = ["chrono", "runtime-tokio", "mysql", "postgres"] } [dependencies] -qusql-sqlx-type-macro = { version = "0.4.2", path = "../qusql-sqlx-type-macro"} +qusql-sqlx-type-macro = { version = "0.4.3", path = "../qusql-sqlx-type-macro"} chrono = "0.4" uuid = { version = "1", optional = true } serde_json = { version = "1", optional = true } @@ -34,3 +34,4 @@ sqlx = { version = "0.9", default-features = false} [features] uuid = ["dep:uuid"] json = ["dep:serde_json"] +postgres = ["sqlx/postgres"] diff --git a/qusql-sqlx-type/src/lib.rs b/qusql-sqlx-type/src/lib.rs index 79973398..ca1c3777 100644 --- a/qusql-sqlx-type/src/lib.rs +++ b/qusql-sqlx-type/src/lib.rs @@ -252,6 +252,33 @@ pub type JsonValue = serde_json::Value; #[doc(hidden)] pub type JsonValue = String; +#[cfg(feature = "postgres")] +mod postgres_support { + //! PostgreSQL `RANGE` bindings when the `postgres` feature is enabled. + use super::*; + + // `PgRange` is used both as the argument tag and as the concrete Rust + // representation of a range column, so it is registered as its own tag (as is + // done for e.g. `chrono::DateTime` above). + arg_io!(PgRange, PgRange); + arg_io!(PgRange, PgRange); + arg_io!(PgRange, PgRange); + arg_io!( + PgRange, + PgRange + ); + arg_io!( + PgRange>, + PgRange> + ); +} + +/// Re-export of [`sqlx::postgres::types::PgRange`], available when the `postgres` +/// feature is enabled. Used to bind/read PostgreSQL `RANGE` columns (e.g. `int4range`, +/// `int8range`, `daterange`, `tsrange`, `tstzrange`). +#[cfg(feature = "postgres")] +pub use sqlx::postgres::types::PgRange; + #[doc(hidden)] pub fn check_arg>(_: &T2) {} diff --git a/qusql-type/Cargo.toml b/qusql-type/Cargo.toml index 36973335..08df0376 100644 --- a/qusql-type/Cargo.toml +++ b/qusql-type/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "qusql-type" -version = "0.7.2" +version = "0.8.0" edition = "2024" authors = ["Jakob Truelsen "] keywords = [ "mysql", "postgesql", "sql", "typer" ] @@ -14,4 +14,4 @@ description = "Typer for sql" codespan-reporting = "0.13" [dependencies] -qusql-parse = {path="../qusql-parse", version="0.10.0"} +qusql-parse = {path="../qusql-parse", version="0.11.0"} diff --git a/qusql-type/src/lib.rs b/qusql-type/src/lib.rs index 13006fbf..bfa0f007 100644 --- a/qusql-type/src/lib.rs +++ b/qusql-type/src/lib.rs @@ -1975,4 +1975,262 @@ mod tests { assert!(!issues.is_ok(), "schema.table in FROM should fail in MySQL"); } } + + /// PostgreSQL range and multirange types: parsing, operators and functions. + /// See https://www.postgresql.org/docs/15/rangetypes.html and + /// https://www.postgresql.org/docs/15/functions-range.html + #[test] + fn postgresql_range_types() { + let schema_src = " + CREATE TABLE reservation ( + id bigint NOT NULL PRIMARY KEY GENERATED ALWAYS AS IDENTITY, + during tsrange NOT NULL, + ids int4multirange NOT NULL + ); + "; + + let opts = TypeOptions::new() + .dialect(SQLDialect::PostgreSQL) + .arguments(SQLArguments::Dollar); + let mut issues = Issues::new(schema_src); + let schema = parse_schemas(schema_src, &mut issues, &opts); + let mut errors = 0; + check_no_errors("schema", schema_src, issues.get(), &mut errors); + + // Checks that `src` types as a single-column SELECT with the given type, or, + // when `expected` is `None`, that typing `src` produces a type error. + let mut check_expr = |src: &'static str, expected: Option>| { + let mut issues = Issues::new(src); + let q = type_statement(&schema, src, &mut issues, &opts); + let Some(expected) = expected else { + if issues.is_ok() { + println!("{src}: expected a type error"); + errors += 1; + } + return; + }; + check_no_errors(src, src, issues.get(), &mut errors); + if let StatementType::Select { columns, .. } = q { + match columns.first() { + Some(c) if c.type_.t == expected => {} + Some(c) => { + println!("{src}: expected type {expected} got {}", c.type_.t); + errors += 1; + } + None => { + println!("{src}: expected a column, got none"); + errors += 1; + } + } + } else { + println!("{src}: expected Select"); + errors += 1; + } + }; + // Range constructors (8.17.6) + check_expr( + "SELECT int4range(1, 10)", + Some(Type::Range(Box::new(Type::I32))), + ); + check_expr( + "SELECT int4range(1, 10, '[]')", + Some(Type::Range(Box::new(Type::I32))), + ); + check_expr( + "SELECT int8range(1, 10)", + Some(Type::Range(Box::new(Type::I64))), + ); + check_expr( + "SELECT numrange(1.1, 2.2)", + Some(Type::Range(Box::new(BaseType::Float.into()))), + ); + check_expr( + "SELECT daterange('2010-01-01'::date, '2010-01-02'::date)", + Some(Type::Range(Box::new(BaseType::Date.into()))), + ); + check_expr( + "SELECT tstzrange(now(), now())", + Some(Type::Range(Box::new(BaseType::TimeStamp.into()))), + ); + + // Multirange constructors + check_expr( + "SELECT int4multirange(int4range(1, 2), int4range(3, 4))", + Some(Type::MultiRange(Box::new(Type::I32))), + ); + check_expr( + "SELECT int4multirange()", + Some(Type::MultiRange(Box::new(Type::I32))), + ); + check_expr( + "SELECT multirange(int4range(1, 2))", + Some(Type::MultiRange(Box::new(Type::I32))), + ); + + // Range/multirange operators (Table 9.54, 9.55) + check_expr( + "SELECT int4range(10, 20) @> 3", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT numrange(11.1, 22.2) && numrange(20.0, 30.0)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT int8range(1, 10) << int8range(100, 110)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT int8range(50, 60) >> int8range(20, 30)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT int8range(1, 20) &< int8range(18, 20)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT int8range(7, 20) &> int8range(5, 10)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT numrange(1.1, 2.2) -|- numrange(2.2, 3.3)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT numrange(5.0, 15.0) + numrange(10.0, 20.0)", + Some(Type::Range(Box::new(BaseType::Float.into()))), + ); + check_expr( + "SELECT int8range(5, 15) * int8range(10, 20)", + Some(Type::Range(Box::new(Type::I64))), + ); + check_expr( + "SELECT int8range(5, 15) - int8range(10, 20)", + Some(Type::Range(Box::new(Type::I64))), + ); + // Plain integer bit-shift/bitwise ops must still work as before. + check_expr("SELECT 1 << 2", Some(Type::Base(BaseType::Integer))); + check_expr("SELECT 1 & 2", Some(Type::Base(BaseType::Integer))); + + // Range/multirange functions (Table 9.56, 9.57) + check_expr("SELECT upper(int8range(15, 25))", Some(Type::I64)); + check_expr( + "SELECT lower(numrange(1.1, 2.2))", + Some(Type::Base(BaseType::Float)), + ); + check_expr( + "SELECT isempty(numrange(1.0, 5.0))", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT lower_inc(numrange(1.1, 2.2))", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT upper_inc(numrange(1.1, 2.2))", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT lower_inf('(,)'::daterange)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT upper_inf('(,)'::daterange)", + Some(Type::Base(BaseType::Bool)), + ); + check_expr( + "SELECT range_merge(int4range(1, 2), int4range(3, 4))", + Some(Type::Range(Box::new(Type::I32))), + ); + check_expr( + "SELECT range_merge('{[1,2), [3,4)}'::int4multirange)", + Some(Type::Range(Box::new(Type::I32))), + ); + check_expr( + "SELECT unnest('{[1,2), [3,4)}'::int4multirange)", + Some(Type::Range(Box::new(Type::I32))), + ); + + // Columns typed from the schema, including lower()/upper() on a `tsrange` column + // and the `int4multirange` -> `int4range` unnest() table-function overload. + check_expr( + "SELECT lower(during) FROM reservation", + Some(Type::Base(BaseType::DateTime)), + ); + check_expr( + "SELECT r FROM unnest('{[1,2), [3,4)}'::int4multirange) AS u(r)", + Some(Type::Range(Box::new(Type::I32))), + ); + + // Mismatched element types / mismatched range vs. multirange must be type errors. + check_expr("SELECT int4range(1, 2) + numrange(1.0, 2.0)", None); + check_expr("SELECT int4range(1, 2) + '{[1,2)}'::int4multirange", None); + + if errors != 0 { + panic!("{errors} errors in postgresql_range_types test"); + } + } + + /// A writable CTE that unions a range into an `int8multirange` column and then + /// `unnest()`s the result. Reported as a real-world usage that should type-check. + #[test] + fn postgresql_multirange_update_returning_unnest() { + let schema_src = " + CREATE TABLE partial_files ( + id uuid NOT NULL PRIMARY KEY DEFAULT gen_random_uuid(), + uploaded_bytes int8multirange NOT NULL DEFAULT '{}', + last_modified timestamptz NOT NULL DEFAULT now() + ); + "; + + let opts = TypeOptions::new() + .dialect(SQLDialect::PostgreSQL) + .arguments(SQLArguments::Dollar); + let mut issues = Issues::new(schema_src); + let schema = parse_schemas(schema_src, &mut issues, &opts); + let mut errors = 0; + check_no_errors("schema", schema_src, issues.get(), &mut errors); + + let src = " + WITH the_update AS ( + UPDATE partial_files + SET last_modified = now(), + uploaded_bytes = uploaded_bytes + multirange(int8range($2, $3, '[)')) + WHERE id = $1 + RETURNING uploaded_bytes + ) + SELECT unnest(uploaded_bytes) AS \"range\" + FROM the_update + "; + let mut issues = Issues::new(src); + let q = type_statement(&schema, src, &mut issues, &opts); + check_no_errors("query", src, issues.get(), &mut errors); + if let StatementType::Select { columns, arguments } = q { + match columns.first() { + Some(c) if c.type_.t == Type::Range(Box::new(Type::I64)) => {} + Some(c) => { + println!( + "query: expected column 0 of type range(i64) got {}", + c.type_.t + ); + errors += 1; + } + None => { + println!("query: expected a column, got none"); + errors += 1; + } + } + if arguments.len() != 3 { + println!("query: expected 3 arguments, got {}", arguments.len()); + errors += 1; + } + } else { + println!("query: expected Select"); + errors += 1; + } + + if errors != 0 { + panic!("{errors} errors in postgresql_multirange_update_returning_unnest test"); + } + } } diff --git a/qusql-type/src/schema.rs b/qusql-type/src/schema.rs index b60af3fb..a6a93590 100644 --- a/qusql-type/src/schema.rs +++ b/qusql-type/src/schema.rs @@ -365,6 +365,21 @@ fn try_parse_body<'a>( }) } +/// Map a PostgreSQL range/multirange subtype to its canonical element `Type`, preserving +/// the concrete integer width (`I32` for `int4range`, `I64` for `int8range`, ...) so that +/// e.g. `int4range` and `int8range` are not conflated. +pub(crate) fn range_subtype_to_type<'a>(sub: &qusql_parse::RangeSubtype) -> Type<'a> { + use qusql_parse::RangeSubtype; + match sub { + RangeSubtype::Int4 => Type::I32, + RangeSubtype::Int8 => Type::I64, + RangeSubtype::Num => BaseType::Float.into(), + RangeSubtype::Ts => BaseType::DateTime.into(), + RangeSubtype::Tstz => BaseType::TimeStamp.into(), + RangeSubtype::Date => BaseType::Date.into(), + } +} + fn type_kind_from_parse<'a, S: SearchPath<'a>>( type_: qusql_parse::Type<'a>, unsigned: bool, @@ -513,17 +528,9 @@ fn type_kind_from_parse<'a, S: SearchPath<'a>>( qusql_parse::Type::TsVector => BaseType::String.into(), qusql_parse::Type::Uuid => BaseType::Uuid.into(), qusql_parse::Type::Xml => BaseType::String.into(), - qusql_parse::Type::Range(sub) | qusql_parse::Type::MultiRange(sub) => { - use qusql_parse::RangeSubtype; - let elem = match sub { - RangeSubtype::Int4 => BaseType::Integer, - RangeSubtype::Int8 => BaseType::Integer, - RangeSubtype::Num => BaseType::Float, - RangeSubtype::Ts => BaseType::DateTime, - RangeSubtype::Tstz => BaseType::TimeStamp, - RangeSubtype::Date => BaseType::Date, - }; - Type::Range(elem) + qusql_parse::Type::Range(sub) => Type::Range(Box::new(range_subtype_to_type(&sub))), + qusql_parse::Type::MultiRange(sub) => { + Type::MultiRange(Box::new(range_subtype_to_type(&sub))) } qusql_parse::Type::Point | qusql_parse::Type::Line diff --git a/qusql-type/src/type_.rs b/qusql-type/src/type_.rs index 6e0f29be..4756fdb8 100644 --- a/qusql-type/src/type_.rs +++ b/qusql-type/src/type_.rs @@ -84,8 +84,11 @@ pub enum Type<'a> { Invalid, JSON, Geometry, - /// A PostgreSQL range type. The inner BaseType is the element type. - Range(BaseType), + /// A PostgreSQL range type. The inner type is the element type (e.g. `I32` for + /// `int4range`, `I64` for `int8range`, `Base(Date)` for `daterange`, ...). + Range(Box>), + /// A PostgreSQL multirange type. The inner type is the element type. + MultiRange(Box>), Array(Box>), Set(Arc>>), U16, @@ -120,6 +123,7 @@ impl<'a> Display for Type<'a> { Type::JSON => f.write_str("json"), Type::Geometry => f.write_str("geometry"), Type::Range(inner) => write!(f, "range({inner})"), + Type::MultiRange(inner) => write!(f, "multirange({inner})"), Type::Array(inner) => { inner.fmt(f)?; f.write_str("[]") @@ -172,6 +176,7 @@ impl<'a> Type<'a> { Type::JSON => BaseType::Any, Type::Geometry => BaseType::Any, Type::Range(_) => BaseType::Any, + Type::MultiRange(_) => BaseType::Any, Type::Array(_) => BaseType::Any, Type::Null => BaseType::Any, Type::Set(_) => BaseType::String, diff --git a/qusql-type/src/type_binary_expression.rs b/qusql-type/src/type_binary_expression.rs index 86a299bf..099feef2 100644 --- a/qusql-type/src/type_binary_expression.rs +++ b/qusql-type/src/type_binary_expression.rs @@ -119,11 +119,14 @@ pub(crate) fn type_binary_expression<'a>( (flags, BaseType::String) } } - BinaryOperator::ShiftLeft(_) - | BinaryOperator::ShiftRight(_) - | BinaryOperator::BitAnd(_) - | BinaryOperator::BitOr(_) - | BinaryOperator::BitXor(_) => { + BinaryOperator::ShiftLeft(_) | BinaryOperator::ShiftRight(_) => { + // Overloaded: integer bit-shift (`<<`/`>>` between integers) and the + // PostgreSQL range "strictly left of"/"strictly right of" operators + // (`<<`/`>>` between ranges/multiranges). Use `Any` here and decide the + // concrete meaning after both operand types are known. + (flags.without_values(), BaseType::Any) + } + BinaryOperator::BitAnd(_) | BinaryOperator::BitOr(_) | BinaryOperator::BitXor(_) => { if flags.true_ { ( flags.with_not_null(true).with_true(false), @@ -207,11 +210,28 @@ pub(crate) fn type_binary_expression<'a>( } FullType::new(BaseType::Bool, true) } - BinaryOperator::ShiftLeft(_) - | BinaryOperator::ShiftRight(_) - | BinaryOperator::BitAnd(_) - | BinaryOperator::BitOr(_) - | BinaryOperator::BitXor(_) => { + BinaryOperator::ShiftLeft(_) | BinaryOperator::ShiftRight(_) => { + // Range/multirange "strictly left of" / "strictly right of" when both sides are + // compatible ranges/multiranges, otherwise the usual integer bit-shift. + if matches!(lhs_type.t, Type::Range(_) | Type::MultiRange(_)) + || matches!(rhs_type.t, Type::Range(_) | Type::MultiRange(_)) + { + if typer.matched_type(&lhs_type, &rhs_type).is_none() { + typer + .err("Type error in range operator", &op_span) + .frag(format!("Of type {}", lhs_type.t), lhs) + .frag(format!("Of type {}", rhs_type.t), rhs); + FullType::invalid() + } else { + FullType::new(BaseType::Bool, lhs_type.not_null && rhs_type.not_null) + } + } else { + typer.ensure_base(lhs, &lhs_type, BaseType::Integer); + typer.ensure_base(rhs, &rhs_type, BaseType::Integer); + FullType::new(BaseType::Integer, lhs_type.not_null && rhs_type.not_null) + } + } + BinaryOperator::BitAnd(_) | BinaryOperator::BitOr(_) | BinaryOperator::BitXor(_) => { typer.ensure_base(lhs, &lhs_type, BaseType::Integer); typer.ensure_base(rhs, &rhs_type, BaseType::Integer); FullType::new(BaseType::Integer, lhs_type.not_null && rhs_type.not_null) @@ -340,8 +360,22 @@ pub(crate) fn type_binary_expression<'a>( // Returns the type of the value being assigned (rhs) rhs_type } - BinaryOperator::User(_, _) => { - FullType::new(BaseType::Any, lhs_type.not_null && rhs_type.not_null) + BinaryOperator::User(name, _) => { + // PostgreSQL range/multirange operators that don't have a dedicated + // `BinaryOperator` variant and fall back to `User`. + if matches!(*name, "&<" | "&>" | "-|-") { + if typer.matched_type(&lhs_type, &rhs_type).is_none() { + typer + .err("Type error in range operator", &op_span) + .frag(format!("Of type {}", lhs_type.t), lhs) + .frag(format!("Of type {}", rhs_type.t), rhs); + FullType::invalid() + } else { + FullType::new(BaseType::Bool, lhs_type.not_null && rhs_type.not_null) + } + } else { + FullType::new(BaseType::Any, lhs_type.not_null && rhs_type.not_null) + } } o @ BinaryOperator::Operator(_, _) => { typer.err("Not supported", o); diff --git a/qusql-type/src/type_expression.rs b/qusql-type/src/type_expression.rs index 0658a96a..f12c17c3 100644 --- a/qusql-type/src/type_expression.rs +++ b/qusql-type/src/type_expression.rs @@ -82,6 +82,7 @@ fn type_unary_expression<'a>( | Type::JSON | Type::Geometry | Type::Range(..) + | Type::MultiRange(..) | Type::Array(..) | Type::Set(..) => { typer.err(format!("Expected numeric type got {}", op_type.t), &op_span); diff --git a/qusql-type/src/type_function.rs b/qusql-type/src/type_function.rs index a29f2a24..fa84f027 100644 --- a/qusql-type/src/type_function.rs +++ b/qusql-type/src/type_function.rs @@ -72,6 +72,53 @@ fn typed_args<'a, 'b, 'c>( typed } +/// Type a PostgreSQL range constructor function, e.g. `int4range(lower, upper [, bounds])`. +/// Takes 2 or 3 arguments: the lower and upper bounds (of the range's element type), and +/// an optional bounds-inclusivity string (`"()"`, `"(]"`, `"[)"`, or `"[]"`). +fn range_constructor<'a, 'b>( + typer: &mut Typer<'a, 'b>, + args: &[Expression<'a>], + span: &Span, + flags: ExpressionFlags, + elem: Type<'a>, +) -> FullType<'a> { + arg_cnt(typer, 2..3, args, span); + let expected = FullType::new(elem.clone(), false); + let mut arg_iter = args.iter(); + for _ in 0..2 { + if let Some(arg) = arg_iter.next() { + let t = type_expression(typer, arg, flags.without_values(), elem.base()); + typer.ensure_type(arg, &t, &expected); + } + } + if let Some(arg) = arg_iter.next() { + let t = type_expression(typer, arg, flags.without_values(), BaseType::String); + typer.ensure_base(arg, &t, BaseType::String); + } + for arg in arg_iter { + type_expression(typer, arg, flags.without_values(), BaseType::Any); + } + // The constructed range is never SQL NULL, even if a bound argument is NULL + // (a NULL bound means an unbounded/infinite side of the range). + FullType::new(Type::Range(Box::new(elem)), true) +} + +/// Type a PostgreSQL multirange constructor function, e.g. `int4multirange(r1, r2, ...)`. +/// Takes zero or more arguments, each of which must be a range of the given element type. +fn multirange_constructor<'a, 'b>( + typer: &mut Typer<'a, 'b>, + args: &[Expression<'a>], + flags: ExpressionFlags, + elem: Type<'a>, +) -> FullType<'a> { + let expected = FullType::new(Type::Range(Box::new(elem.clone())), false); + for arg in args { + let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); + typer.ensure_type(arg, &t, &expected); + } + FullType::new(Type::MultiRange(Box::new(elem)), true) +} + pub(crate) fn type_function<'a, 'b>( typer: &mut Typer<'a, 'b>, func: &Function<'a>, @@ -1326,15 +1373,15 @@ pub(crate) fn type_function<'a, 'b>( FullType::new(BaseType::String, false) } Function::LCase | Function::Lower => { - // PostgreSQL overloads lower(): string lowercase AND range lower-bound. - // For a range arg, return the element type; for string/any, return String. + // PostgreSQL overloads lower(): string lowercase AND range/multirange lower-bound. + // For a range/multirange arg, return the element type (nullable, since the + // bound may be NULL for an empty or unbounded range); for string/any, return String. arg_cnt(typer, 1..1, args, span); if let Some(arg) = args.first() { let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); - if let Type::Range(elem) = t.t { - FullType::new(elem, false) - } else { - FullType::new(BaseType::String, t.not_null) + match t.t { + Type::Range(elem) | Type::MultiRange(elem) => FullType::new(*elem, false), + _ => FullType::new(BaseType::String, t.not_null), } } else { FullType::invalid() @@ -1457,13 +1504,13 @@ pub(crate) fn type_function<'a, 'b>( &[], ), Function::UCase | Function::Upper => { - // PostgreSQL overloads upper(): string uppercase AND range upper-bound. + // PostgreSQL overloads upper(): string uppercase AND range/multirange upper-bound. arg_cnt(typer, 1..1, args, span); if let Some(arg) = args.first() { let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); - match t.base() { - BaseType::Any | BaseType::String => FullType::new(BaseType::String, t.not_null), - _ => FullType::new(BaseType::Any, t.not_null), + match t.t { + Type::Range(elem) | Type::MultiRange(elem) => FullType::new(*elem, false), + _ => FullType::new(BaseType::String, t.not_null), } } else { FullType::invalid() @@ -1760,7 +1807,20 @@ pub(crate) fn type_function<'a, 'b>( tf(BaseType::Float.into(), &[BaseType::Any], &[BaseType::Any]) } Function::TsvectorToArray => tf(BaseType::Any.into(), &[BaseType::Any], &[]), - Function::Unnest => tf(BaseType::Any.into(), &[BaseType::Any], &[]), + Function::Unnest => { + // PostgreSQL overloads unnest(): array element AND multirange -> range. + arg_cnt(typer, 1..1, args, span); + if let Some(arg) = args.first() { + let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); + match t.t { + Type::Array(inner) => FullType::new(*inner, false), + Type::MultiRange(elem) => FullType::new(Type::Range(elem), false), + _ => FullType::new(BaseType::Any, false), + } + } else { + FullType::invalid() + } + } // Text search debug functions Function::TsDebug | Function::TsLexize | Function::TsParse | Function::TsTokenType | Function::TsStat => tf(BaseType::Any.into(), &[BaseType::Any], &[BaseType::Any]), @@ -1808,6 +1868,94 @@ pub(crate) fn type_function<'a, 'b>( tf(BaseType::Integer.into(), &[BaseType::Any], &[BaseType::Any]) } Function::ArrayPositions => tf(BaseType::Any.into(), &[BaseType::Any, BaseType::Any], &[]), + // PostgreSQL range/multirange functions + Function::IsEmpty + | Function::LowerInc + | Function::UpperInc + | Function::LowerInf + | Function::UpperInf => { + arg_cnt(typer, 1..1, args, span); + if let Some(arg) = args.first() { + let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); + if matches!(t.t, Type::Range(_) | Type::MultiRange(_)) || t.base() == BaseType::Any { + FullType::new(BaseType::Bool, t.not_null) + } else { + typer.err(format!("Expected range type got {}", t.t), arg); + FullType::invalid() + } + } else { + FullType::invalid() + } + } + Function::RangeMerge => match args { + [a] => { + let t = type_expression(typer, a, flags.without_values(), BaseType::Any); + match t.t { + Type::MultiRange(elem) => FullType::new(Type::Range(elem), t.not_null), + _ => { + typer.err(format!("Expected multirange type got {}", t.t), a); + FullType::invalid() + } + } + } + [a, b] => { + let ta = type_expression(typer, a, flags.without_values(), BaseType::Any); + let tb = type_expression(typer, b, flags.without_values(), BaseType::Any); + if let Some(t) = typer.matched_type(&ta, &tb) { + FullType::new(t, ta.not_null && tb.not_null) + } else { + typer + .err("Type error in range_merge", span) + .frag(format!("Of type {}", ta.t), a) + .frag(format!("Of type {}", tb.t), b); + FullType::invalid() + } + } + _ => { + arg_cnt(typer, 1..2, args, span); + typed_args(typer, args, flags); + FullType::invalid() + } + }, + Function::Multirange => { + arg_cnt(typer, 1..1, args, span); + if let Some(arg) = args.first() { + let t = type_expression(typer, arg, flags.without_values(), BaseType::Any); + match t.t { + Type::Range(elem) => FullType::new(Type::MultiRange(elem), t.not_null), + _ => { + typer.err(format!("Expected range type got {}", t.t), arg); + FullType::invalid() + } + } + } else { + FullType::invalid() + } + } + Function::Int4Range => range_constructor(typer, args, span, flags, Type::I32), + Function::Int8Range => range_constructor(typer, args, span, flags, Type::I64), + Function::NumRange => range_constructor(typer, args, span, flags, BaseType::Float.into()), + Function::TsRange => { + range_constructor(typer, args, span, flags, BaseType::DateTime.into()) + } + Function::TstzRange => { + range_constructor(typer, args, span, flags, BaseType::TimeStamp.into()) + } + Function::DateRange => { + range_constructor(typer, args, span, flags, BaseType::Date.into()) + } + Function::Int4Multirange => multirange_constructor(typer, args, flags, Type::I32), + Function::Int8Multirange => multirange_constructor(typer, args, flags, Type::I64), + Function::NumMultirange => multirange_constructor(typer, args, flags, BaseType::Float.into()), + Function::TsMultirange => { + multirange_constructor(typer, args, flags, BaseType::DateTime.into()) + } + Function::TstzMultirange => { + multirange_constructor(typer, args, flags, BaseType::TimeStamp.into()) + } + Function::DateMultirange => { + multirange_constructor(typer, args, flags, BaseType::Date.into()) + } // PostgreSQL system information functions (9.27) Function::CurrentDatabase => { arg_cnt(typer, 0..0, args, span); diff --git a/qusql-type/src/type_insert_replace.rs b/qusql-type/src/type_insert_replace.rs index 3da91459..8918b977 100644 --- a/qusql-type/src/type_insert_replace.rs +++ b/qusql-type/src/type_insert_replace.rs @@ -70,7 +70,7 @@ pub(crate) fn type_insert_replace<'a>( typer.err( format!( "No value for column {} provided, but it has no default value", - &col.identifier + col.identifier ), set, ); @@ -89,7 +89,7 @@ pub(crate) fn type_insert_replace<'a>( typer.err( format!( "No value for column {} provided, but it has no default value", - &col.identifier + col.identifier ), &columns.opt_span().unwrap_or(table.span()), ); diff --git a/qusql-type/src/type_reference.rs b/qusql-type/src/type_reference.rs index 9481fdae..d82db190 100644 --- a/qusql-type/src/type_reference.rs +++ b/qusql-type/src/type_reference.rs @@ -137,16 +137,21 @@ pub(crate) fn type_reference<'a>( match name { qusql_parse::TableFunctionName::Unnest(unnest_span) => { // Each argument to UNNEST expands to one column. - // The column type is the element type of the array argument. + // The column type is the element type of an array argument, or + // the range type of a multirange argument. let mut columns: Vec<(Identifier<'a>, FullType<'a>)> = Vec::new(); for (idx, arg) in args.iter().enumerate() { let arr_type = type_expression(typer, arg, ExpressionFlags::default(), BaseType::Any); - let elem_type = if let crate::type_::Type::Array(inner) = arr_type.t { - FullType::new(*inner, false) - } else { - // If we can't determine it's an array, use Any/nullable - FullType::new(BaseType::Any, false) + let elem_type = match arr_type.t { + crate::type_::Type::Array(inner) => FullType::new(*inner, false), + crate::type_::Type::MultiRange(elem) => { + FullType::new(crate::type_::Type::Range(elem), false) + } + _ => { + // If we can't determine it's an array or multirange, use Any/nullable + FullType::new(BaseType::Any, false) + } }; // Use col_list alias if provided, otherwise generate "unnest1", "unnest2", ... let col_name = if let Some(alias) = col_list.get(idx) { diff --git a/qusql-type/src/typer.rs b/qusql-type/src/typer.rs index b62e48f7..3ff3c3b2 100644 --- a/qusql-type/src/typer.rs +++ b/qusql-type/src/typer.rs @@ -113,6 +113,25 @@ impl<'a, 'b> Typer<'a, 'b> { } (Type::Array(_), other) if other.base() != BaseType::Any => return None, (other, Type::Array(_)) if other.base() != BaseType::Any => return None, + // Ranges/multiranges match recursively on their element types; a range + // never matches a multirange, nor a concrete non-range/multirange type. + (Type::Range(i1), Type::Range(i2)) => { + return self + .matched_type(i1, i2) + .map(|inner| Type::Range(Box::new(inner))); + } + (Type::MultiRange(i1), Type::MultiRange(i2)) => { + return self + .matched_type(i1, i2) + .map(|inner| Type::MultiRange(Box::new(inner))); + } + (Type::Range(_), Type::MultiRange(_)) | (Type::MultiRange(_), Type::Range(_)) => { + return None; + } + (Type::Range(_), other) if other.base() != BaseType::Any => return None, + (other, Type::Range(_)) if other.base() != BaseType::Any => return None, + (Type::MultiRange(_), other) if other.base() != BaseType::Any => return None, + (other, Type::MultiRange(_)) if other.base() != BaseType::Any => return None, _ => {} }