diff --git a/.changeset/v2-2-2-posthog-and-speed.md b/.changeset/v2-2-2-posthog-and-speed.md new file mode 100644 index 00000000..a15abee1 --- /dev/null +++ b/.changeset/v2-2-2-posthog-and-speed.md @@ -0,0 +1,21 @@ +### New Features +- Connect PostHog with a personal API key and browse or query your product analytics with HogQL, read-only +- Railway is available in the provider picker +- Right-click a Docker database to restart or stop it, or copy its connection URL +- Connecting shows a full loading screen with the database's logo, its host, elapsed time and Cancel + +### Bug Fixes +- The expanded row JSON and the row panel update right after a cell edit +- MySQL 8 and 9 tables show their columns when empty, and inserted rows appear without a refresh +- Nile no longer drops its connection after every query +- The connect dialog footer names the host being dialled, and Resume only spins when you pressed it + +### Changes +- Connecting to a far database costs one handshake instead of two (Prisma Postgres: 6.2s to 0.6s) +- Far databases keep a warm pool, so opening a table never waits on new connections +- Every query to a far database is one round trip shorter, and a table's first open skips a lookup +- Provider pickers answer sooner: the API connection opens with the dialog, and hovering a database starts building its connection +- PlanetScale lists databases and connects in fewer API calls +- MySQL tables open in one round trip +- Provider sign-in is sturdier across IPv4 and IPv6 callbacks +- A tidier menu bar, and the sidebar shows its Enter hint only while searching diff --git a/README.md b/README.md index e0f7fd0e..2c21c6e6 100644 --- a/README.md +++ b/README.md @@ -23,10 +23,10 @@ Stroke is a Rust + Svelte app for browsing, editing and querying databases. It s | Turso | libSQL | | Cloudflare D1 | SQLite | | Upstash | Redis | +| Railway | Postgres, MySQL, Redis | +| PostHog | HogQL, read-only | -Railway is next. - -Stroke also finds databases already running on your machine (Docker containers, local Postgres and MySQL, the SQLite file your ORM points at), so a local connection is usually one click. Anything can go through an SSH tunnel. +Stroke also finds databases already running on your machine (Docker containers, local Postgres and MySQL, the SQLite file your ORM points at), so a local connection is usually one click, and a right-click restarts or stops a container. Anything can go through an SSH tunnel. ## What's inside diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 75930a1c..55d2cf12 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -5496,8 +5496,6 @@ dependencies = [ [[package]] name = "sqlx-mysql" version = "0.8.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "aa003f0038df784eb8fecbbac13affe3da23b45194bd57dba231c8f48199c526" dependencies = [ "atoi", "base64 0.22.1", @@ -5541,8 +5539,6 @@ dependencies = [ [[package]] name = "sqlx-postgres" version = "0.8.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "db58fcd5a53cf07c184b154801ff91347e4c30d17a3562a635ff028ad5deda46" dependencies = [ "atoi", "base64 0.22.1", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 7a9bbbe6..5b21f5fc 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -163,3 +163,10 @@ opt-level = 3 opt-level = 3 [profile.dev.package.base64] opt-level = 3 + +# sqlx 0.8.6 drivers with one change each: the pool's on-release ping is skipped +# for a connection with nothing pending (see the "Stroke patch" comments). That +# ping cost a round trip after every query to a far host and broke Nile. +[patch.crates-io] +sqlx-postgres = { path = "vendor/sqlx-postgres" } +sqlx-mysql = { path = "vendor/sqlx-mysql" } diff --git a/src-tauri/src/cloudflare.rs b/src-tauri/src/cloudflare.rs index e867119a..e2d278f8 100644 --- a/src-tauri/src/cloudflare.rs +++ b/src-tauri/src/cloudflare.rs @@ -17,8 +17,6 @@ use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use std::sync::OnceLock; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::TcpListener; // ── Cloudflare OAuth constants ──────────────────────────────────────────────── @@ -107,111 +105,7 @@ fn pkce_pair() -> (String, String) { // ── Local callback server ───────────────────────────────────────────────────── -/// Try to bind to one of the pre-registered Cloudflare callback ports. -/// Returns (listener, redirect_uri) on success. -async fn bind_callback_listener() -> Result<(TcpListener, String), String> { - for &port in CF_CALLBACK_PORTS { - if let Ok(listener) = TcpListener::bind(format!("127.0.0.1:{port}")).await { - let redirect_uri = format!("http://localhost:{port}/oauth/callback"); - return Ok((listener, redirect_uri)); - } - } - Err(format!( - "Could not bind to any of the pre-registered callback ports ({}-{}). \ - Close other Wrangler or Stroke processes and try again.", - CF_CALLBACK_PORTS[0], - CF_CALLBACK_PORTS[CF_CALLBACK_PORTS.len() - 1] - )) -} - -/// Wait for one OAuth callback on the listener and return the authorization code. -async fn await_oauth_callback( - listener: TcpListener, - expected_state: &str, -) -> Result { - let success_html = crate::oauth_page::page(true, "Cloudflare"); - let error_html = crate::oauth_page::page(false, "Cloudflare"); - - let send_html = |html: &str| -> String { - format!( - "HTTP/1.1 200 OK\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", - html.len(), - html - ) - }; - - let (mut stream, _) = listener - .accept() - .await - .map_err(|e| format!("Callback accept failed: {e}"))?; - - let mut buf = vec![0u8; 8192]; - let n = stream - .read(&mut buf) - .await - .map_err(|e| format!("Callback read failed: {e}"))?; - let req = String::from_utf8_lossy(&buf[..n]); - - // Parse the first line: GET /oauth/callback?code=...&state=... HTTP/1.1 - let first_line = req.lines().next().unwrap_or(""); - let path = first_line.split_whitespace().nth(1).unwrap_or(""); - let query = path.split('?').nth(1).unwrap_or(""); - - let mut code = None; - let mut state = None; - let mut error: Option = None; - - for pair in query.split('&') { - let mut kv = pair.splitn(2, '='); - let key = kv.next().unwrap_or(""); - let val = kv - .next() - .map(|v| urlencoding::decode(v).unwrap_or_default().into_owned()) - .unwrap_or_default(); - match key { - "code" => code = Some(val), - "state" => state = Some(val), - "error" => error = Some(val), - "error_description" => { - if error.is_none() { - error = Some(val) - } - } - _ => {} - } - } - - if let Some(err) = &error { - let _ = stream - .write_all(send_html(&error_html).as_bytes()) - .await; - return Err(format!("Cloudflare denied authorization: {err}")); - } - - let code = match code { - Some(c) if !c.is_empty() => c, - _ => { - let _ = stream - .write_all(send_html(&error_html).as_bytes()) - .await; - return Err("No authorization code in callback".to_string()); - } - }; - - if state.as_deref() != Some(expected_state) { - let _ = stream - .write_all(send_html(&error_html).as_bytes()) - .await; - return Err("OAuth state mismatch - possible CSRF".to_string()); - } - let _ = stream - .write_all(send_html(&success_html).as_bytes()) - .await; - let _ = stream.flush().await; - - Ok(code) -} // ── Token exchange ──────────────────────────────────────────────────────────── @@ -381,7 +275,21 @@ pub async fn cloudflare_start_oauth(app: tauri::AppHandle) -> Result Result) - .map_err(|e| format!("Failed to open browser: {e}"))?; + crate::providers::open_sign_in_page(&app, &auth_url); let code = tokio::time::timeout( std::time::Duration::from_secs(AUTH_TIMEOUT_SECS), - await_oauth_callback(listener, &state), + crate::providers::await_oauth_callback(listener, &state, "code", "Cloudflare"), ) .await .map_err(|_| "Authorization timed out - please try again.".to_string())??; diff --git a/src-tauri/src/commands.rs b/src-tauri/src/commands.rs index 9a932b5b..498d8e07 100644 --- a/src-tauri/src/commands.rs +++ b/src-tauri/src/commands.rs @@ -433,6 +433,21 @@ pub async fn connect_libsql_db(state: State<'_, DbState>, config: LibSqlConfig) connect_libsql(state, config).await } +// ── PostHog ─────────────────────────────────────────────────────────────────── + +#[tauri::command] +pub async fn test_posthog(config: crate::db::connection::PosthogConfig) -> Result<(), String> { + crate::db::connection::test_posthog_connection(config).await +} + +#[tauri::command] +pub async fn connect_posthog_db( + state: State<'_, DbState>, + config: crate::db::connection::PosthogConfig, +) -> Result<(), String> { + crate::db::connection::connect_posthog(state, config).await +} + // ── ClickHouse ──────────────────────────────────────────────────────────────── #[tauri::command] diff --git a/src-tauri/src/db/backup.rs b/src-tauri/src/db/backup.rs index c8b4a138..58cfc651 100644 --- a/src-tauri/src/db/backup.rs +++ b/src-tauri/src/db/backup.rs @@ -144,6 +144,7 @@ async fn export_one( ActiveConnection::D1(cfg) => export_d1(app, &cfg, tables.as_deref(), opts).await, ActiveConnection::LibSql(_) => Err("Backup export is not supported for LibSQL/Turso connections".to_string()), ActiveConnection::Clickhouse(_) => Err("Backup export is not supported for ClickHouse connections".to_string()), + ActiveConnection::Posthog(_) => Err("Backup is not supported for PostHog".to_string()), ActiveConnection::Redis(_) => Err("Backup is not supported on Redis".to_string()), ActiveConnection::Duckdb(h) => export_duckdb(app, &h).await, ActiveConnection::Mssql(h) => export_mssql(app, &h).await, @@ -164,6 +165,7 @@ pub async fn backup_import( ActiveConnection::D1(cfg) => import_d1(&app, &cfg, &sql).await, ActiveConnection::LibSql(_) => Err("Backup import is not supported for LibSQL/Turso connections".to_string()), ActiveConnection::Clickhouse(_) => Err("Backup import is not supported for ClickHouse connections".to_string()), + ActiveConnection::Posthog(_) => Err("Backup is not supported for PostHog".to_string()), ActiveConnection::Redis(_) => Err("Backup is not supported on Redis".to_string()), ActiveConnection::Duckdb(h) => import_duckdb(&app, &h, &sql).await, ActiveConnection::Mssql(h) => import_mssql(&app, &h, &sql).await, @@ -1069,7 +1071,7 @@ async fn export_mysql( let create_row = sqlx::query(&format!("SHOW CREATE TABLE `{schema}`.`{table}`")) .fetch_one(pool).await .map_err(|e| format!("SHOW CREATE TABLE `{table}` failed: {e}"))?; - let create_sql: String = create_row.try_get(1).unwrap_or_default(); + let create_sql = super::mysql::my_text(&create_row, 1).unwrap_or_default(); out.push_str(&create_sql.replace("CREATE TABLE ", "CREATE TABLE IF NOT EXISTS ")); out.push_str(";\n\n"); @@ -1110,7 +1112,7 @@ async fn export_mysql( out.push_str(&format!("-- Views - {schema}\n")); for view in &view_names { if let Ok(row) = sqlx::query(&format!("SHOW CREATE VIEW `{schema}`.`{view}`")).fetch_one(pool).await { - let create: String = row.try_get(1).unwrap_or_default(); + let create = super::mysql::my_text(&row, 1).unwrap_or_default(); out.push_str(&create.replace("CREATE ", "CREATE OR REPLACE ")); out.push_str(";\n"); } @@ -1135,7 +1137,7 @@ async fn export_mysql( let keyword = if rtype == "FUNCTION" { "FUNCTION" } else { "PROCEDURE" }; if let Ok(row) = sqlx::query(&format!("SHOW CREATE {keyword} `{schema}`.`{name}`")).fetch_one(pool).await { let col_idx: usize = if rtype == "FUNCTION" { 2 } else { 2 }; - let create: String = row.try_get(col_idx).unwrap_or_default(); + let create = super::mysql::my_text(&row, col_idx).unwrap_or_default(); out.push_str(&create); out.push_str("//\n\n"); } @@ -1157,7 +1159,7 @@ async fn export_mysql( out.push_str(&format!("-- Triggers - {schema}\nDELIMITER //\n")); for trig in &trigger_names { if let Ok(row) = sqlx::query(&format!("SHOW CREATE TRIGGER `{schema}`.`{trig}`")).fetch_one(pool).await { - let create: String = row.try_get(2).unwrap_or_default(); + let create = super::mysql::my_text(&row, 2).unwrap_or_default(); out.push_str(&create); out.push_str("//\n\n"); } diff --git a/src-tauri/src/db/clickhouse.rs b/src-tauri/src/db/clickhouse.rs index 2badbb9f..205f0169 100644 --- a/src-tauri/src/db/clickhouse.rs +++ b/src-tauri/src/db/clickhouse.rs @@ -338,7 +338,7 @@ pub async fn get_table_rows( /// Build a `WHERE` clause from the global search box + structured filters. /// Values are escaped into single-quoted literals (ClickHouse HTTP has no bound /// params here); identifiers are validated against the known column list. -fn build_where(cols: &[ColumnStructureRow], search: Option<&str>, filters: Option<&[RowFilter]>) -> String { +pub(crate) fn build_where(cols: &[ColumnStructureRow], search: Option<&str>, filters: Option<&[RowFilter]>) -> String { let known: std::collections::HashSet<&str> = cols.iter().map(|c| c.name.as_str()).collect(); let mut clauses: Vec = Vec::new(); diff --git a/src-tauri/src/db/connection.rs b/src-tauri/src/db/connection.rs index 7fde28eb..316dbdbb 100644 --- a/src-tauri/src/db/connection.rs +++ b/src-tauri/src/db/connection.rs @@ -197,6 +197,33 @@ impl ClickhouseConfig { } } +// ── PostHog ─────────────────────────────────────────────────────────────────── + +/// A PostHog project, queried with HogQL over PostHog's query API. PostHog keeps +/// its data in ClickHouse, but Cloud offers no direct ClickHouse access; the API +/// is the way in, with a personal API key that has the Query Read scope. +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PosthogConfig { + pub name: String, + /// Base URL: `https://us.posthog.com`, `https://eu.posthog.com`, or a + /// self-hosted instance. + pub host: String, + pub project_id: String, + pub api_key: String, +} + +impl PosthogConfig { + pub fn base_url(&self) -> String { + let h = self.host.trim().trim_end_matches('/'); + if h.starts_with("http://") || h.starts_with("https://") { + h.to_string() + } else { + format!("https://{h}") + } + } +} + // ── Redis ───────────────────────────────────────────────────────────────────── #[derive(Debug, Clone, Serialize, Deserialize)] @@ -277,6 +304,8 @@ pub enum AnyConnectionConfig { Duckdb(DuckdbConfig), #[serde(rename = "mssql")] Mssql(MssqlConfig), + #[serde(rename = "posthog")] + Posthog(PosthogConfig), } // ── Active connection ───────────────────────────────────────────────────────── @@ -292,6 +321,7 @@ pub enum ActiveConnection { Redis(RedisConfig), Duckdb(DuckdbHandle), Mssql(MssqlHandle), + Posthog(PosthogConfig), } impl ActiveConnection { @@ -306,6 +336,7 @@ impl ActiveConnection { Self::Redis(_) => "redis", Self::Duckdb(_) => "duckdb", Self::Mssql(_) => "mssql", + Self::Posthog(_) => "posthog", } } } @@ -728,6 +759,100 @@ macro_rules! ping_if_stale { }}; } +/// Measured handshake cost per `host:port`, in ms, remembered across restarts. +/// +/// Measured with psql from here: Neon, Supabase, Nile and Prisma Postgres all +/// take 1.8-3.4s to open a connection (TCP + TLS + startup + SCRAM, six or seven +/// round trips to a far region) but only 265-535ms to answer a query on an open +/// one. sqlx opens a NEW connection for every query that finds no idle one, so +/// the six queries of a table open each paid a full handshake instead of +/// waiting half a second for a busy connection to come back. A host known to be +/// that far gets a small fixed pool instead (`pg_pool_for`). +static HANDSHAKES: std::sync::OnceLock>> = std::sync::OnceLock::new(); +static DATA_DIR: std::sync::OnceLock = std::sync::OnceLock::new(); +const HANDSHAKE_FILE: &str = "handshake-costs.json"; + +/// Called once from setup, so the handshake memory survives a restart: the +/// reconnect on launch is exactly the connect that needs it. +pub fn set_data_dir(dir: std::path::PathBuf) { + let _ = DATA_DIR.set(dir); +} + +fn handshakes() -> &'static Mutex> { + HANDSHAKES.get_or_init(|| { + let map = DATA_DIR + .get() + .and_then(|d| std::fs::read_to_string(d.join(HANDSHAKE_FILE)).ok()) + .and_then(|t| serde_json::from_str(&t).ok()) + .unwrap_or_default(); + Mutex::new(map) + }) +} + +fn known_handshake_ms(host: &str, port: u16) -> Option { + handshakes().lock().unwrap_or_else(|e| e.into_inner()).get(&format!("{host}:{port}")).copied() +} + +fn record_handshake(host: &str, port: u16, ms: u64) { + let snapshot = { + let mut map = handshakes().lock().unwrap_or_else(|e| e.into_inner()); + let key = format!("{host}:{port}"); + // Half old, half new: one unlucky connect doesn't flip the pool shape. + let v = map.get(&key).map_or(ms, |old| old / 2 + ms / 2); + if map.get(&key) == Some(&v) { + return; + } + map.insert(key, v); + serde_json::to_string(&*map).ok() + }; + if let (Some(dir), Some(text)) = (DATA_DIR.get(), snapshot) { + let _ = std::fs::write(dir.join(HANDSHAKE_FILE), text); + } +} + +/// A handshake slower than this means queueing behind an open connection beats +/// opening another. Nearby hosts measured 320-690ms, far providers 1.8-3.4s. +const FAR_HANDSHAKE_MS: u64 = 1200; + +/// Connections a far host keeps open: one per query of a table open (rows, +/// count and four catalog lookups), so browsing once connected runs every query +/// at once, exactly as on a nearby host. +const FAR_POOL: u32 = 6; + +/// Pool shape for this host. A far host's pool is exactly `FAR_POOL` wide and +/// never shrinks: `warm_in_parallel` opens all of it right after connect, and a +/// query that arrives before those land waits ~300-500ms for a free connection +/// instead of opening one more at 2-3s. Unknown and nearby hosts keep the wide +/// pool, where a handshake is cheap. +fn pg_pool_for(host: &str, port: u16) -> PgPoolOptions { + match known_handshake_ms(host, port) { + Some(ms) if ms >= FAR_HANDSHAKE_MS => { + log::info!("{host}:{port} handshake ~{ms}ms: {FAR_POOL}-connection pool, warmed in parallel"); + pg_pool_builder().min_connections(FAR_POOL).max_connections(FAR_POOL) + } + _ => pg_pool_builder(), + } +} + +/// Open the rest of a far host's pool at once, one handshake's wait in total. +/// +/// sqlx's own min-connections fill opens them one after another (six in a row +/// at 3s is 18s of a pool that isn't ready). These are spawned while the +/// caller still holds the first connection, so none of them can grab it and +/// every one opens a fresh connection; each is handed back the moment it +/// lands, so none is held away from a query that needs it. +fn warm_in_parallel(pool: &sqlx::Pool, total: u32) { + let extra = total.saturating_sub(pool.size()); + for _ in 0..extra { + let pool = pool.clone(); + tokio::spawn(async move { + if let Ok(conn) = pool.acquire().await { + drop(conn); + } + }); + } +} + fn pg_pool_builder() -> PgPoolOptions { let health = std::sync::Arc::new(PathHealth::default()); PgPoolOptions::new() @@ -909,6 +1034,33 @@ where attempt().await } +/// Open a pool and return once ONE connection works; the rest fill in behind it. +/// +/// sqlx's `connect_with` does not do that. It opens every `min_connections` +/// connection one after another before returning (sqlx-core 0.8.6, +/// `PoolOptions::connect_with` → `try_min_connections`), so `min_connections(2)` +/// made every connect pay two full handshakes in a row. Against Prisma Postgres, +/// where psql measures a handshake at 3.3-3.5s and a query at ~500ms, that was +/// the 6.2s connect in the log. A lazy pool starts the same min-connections fill +/// as a background task, and the acquire below races it, so the connect costs +/// one handshake and the second connection is ready moments later. +async fn connect_pg_pool(builder: PgPoolOptions, opts: PgConnectOptions) -> Result { + let (host, port) = (opts.get_host().to_string(), opts.get_port()); + let t0 = std::time::Instant::now(); + let pool = builder.connect_lazy_with(opts); + // Proves the address, TLS and credentials, with the same errors + // `connect_with` gave; the connection goes back to the pool for the first query. + let first = pool.acquire().await?; + let ms = t0.elapsed().as_millis() as u64; + record_handshake(&host, port, ms); + // Far host, first time or not: warm the rest now, before `first` is released. + if ms >= FAR_HANDSHAKE_MS { + warm_in_parallel(&pool, FAR_POOL); + } + drop(first); + Ok(pool) +} + pub(crate) async fn open_pg(config: &PgConfig) -> Result { let opts: PgConnectOptions = config .connection_url() @@ -952,13 +1104,14 @@ pub(crate) async fn open_pg(config: &PgConfig) -> Result { // succeeds, so the retry ladder stops restarting a handshake that is fine. let tcp_ok = std::sync::atomic::AtomicBool::new(false); let connect = async { - match retry_fast(&tcp_ok, || pg_pool_builder().connect_with(fast_opts.clone())).await { + match retry_fast(&tcp_ok, || connect_pg_pool(pg_pool_for(&config.host, config.port), fast_opts.clone())).await { Ok(pool) => Ok(pool), // Some poolers (PgBouncer without `ignore_startup_parameters=options`) // reject the `options` startup parameter outright. Fall back to the // slower after_connect SET so those hosts still connect. Err(e) if e.to_string().contains("unsupported startup parameter") => { - pg_pool_builder() +connect_pg_pool( + pg_pool_for(&config.host, config.port) .after_connect(move |conn, _meta| { let tz_set = tz_set.clone(); Box::pin(async move { @@ -970,8 +1123,9 @@ pub(crate) async fn open_pg(config: &PgConfig) -> Result { } Ok(()) }) - }) - .connect_with(opts) + }), + opts, + ) .await .map_err(|e| format!("Connection failed: {e}")) } @@ -1004,8 +1158,8 @@ pub async fn connect( close_existing(&state).await; set_conn(&state, Some(ActiveConnection::Postgres(pool)))?; // Nothing else to do here: `min_connections` fills the pool from the pool's - // own maintenance task, off the critical path, and the connect returns as - // soon as the first connection is usable. + // own maintenance task, off the critical path (see `connect_pg_pool`), and the + // connect returns as soon as the first connection is usable. tunnel_state.set(tunnel); Ok(()) } @@ -1084,9 +1238,18 @@ pub(crate) async fn open_mysql(config: &MysqlConfig) -> Result= FAR_HANDSHAKE_MS); + if far { + log::info!("{}:{} is a far host: 4-connection pool, warmed in parallel", config.host, config.port); + } + let (host, port) = (config.host.clone(), config.port); + let builder = MySqlPoolOptions::new() // Same rationale as PG: 4 is the real-world ceiling for a desktop app. .max_connections(4) + .min_connections(if far { 4 } else { 0 }) // See pg_pool_builder: must clear a cold-pool handshake on a slow link. .acquire_timeout(Duration::from_secs(10)) // Keep connections warm for the session (see open_pg for the full rationale) @@ -1114,8 +1277,21 @@ pub(crate) async fn open_mysql(config: &MysqlConfig) -> Result= FAR_HANDSHAKE_MS { + warm_in_parallel(&pool, 4); + } + drop(first); + Ok::<_, sqlx::Error>(pool) + }; connect_racing_probe( &config.host, @@ -1210,6 +1386,18 @@ pub async fn connect_clickhouse(state: State<'_, DbState>, config: ClickhouseCon set_conn(&state, Some(ActiveConnection::Clickhouse(config))) } +// ── PostHog connect / test ──────────────────────────────────────────────────── + +pub async fn test_posthog_connection(config: PosthogConfig) -> Result<(), String> { + crate::db::posthog::query(&config, "SELECT 1").await.map(|_| ()) +} + +pub async fn connect_posthog(state: State<'_, DbState>, config: PosthogConfig) -> Result<(), String> { + test_posthog_connection(config.clone()).await?; + close_existing(&state).await; + set_conn(&state, Some(ActiveConnection::Posthog(config))) +} + // ── Redis connect / test ────────────────────────────────────────────────────── pub async fn test_redis_connection(config: RedisConfig) -> Result<(), String> { @@ -1663,3 +1851,59 @@ pub async fn prewarm_dns(hosts: Vec) { }); } } + +/// Live checks for the release-ping patch in `vendor/sqlx-postgres`. Run with +/// `STROKE_PG_URL=postgres://… cargo test --lib pg_release_live -- --ignored --nocapture`. +#[cfg(test)] +mod pg_release_live { + use futures::TryStreamExt; + use sqlx::postgres::PgPoolOptions; + use sqlx::{Connection, Row}; + + async fn pool() -> sqlx::PgPool { + let url = std::env::var("STROKE_PG_URL").expect("STROKE_PG_URL"); + // One connection, so every step below reuses the same one. + PgPoolOptions::new().max_connections(1).test_before_acquire(false).connect(&url).await.unwrap() + } + + #[tokio::test] + #[ignore] + async fn a_dropped_transaction_still_rolls_back() { + let pool = pool().await; + { + let mut tx = pool.begin().await.unwrap(); + // Transaction-local: gone once the transaction ends, still set if it didn't. + sqlx::query("SELECT set_config('stroke.probe', 'in_tx', true)").execute(&mut *tx).await.unwrap(); + // Dropped without commit: sqlx queues a ROLLBACK. + } + let mut conn = pool.acquire().await.unwrap(); + let probe: Option = sqlx::query("SELECT current_setting('stroke.probe', true)").fetch_one(&mut *conn).await.unwrap().get(0); + assert_ne!(probe.as_deref(), Some("in_tx"), "the connection came back still inside the transaction"); + conn.ping().await.unwrap(); + } + + #[tokio::test] + #[ignore] + async fn a_half_read_stream_is_cleaned_up() { + let pool = pool().await; + { + let mut conn = pool.acquire().await.unwrap(); + let mut rows = sqlx::query("SELECT g FROM generate_series(1, 100000) g").fetch(&mut *conn); + let _first = rows.try_next().await.unwrap(); + // Dropped mid result set. + } + let n: i64 = sqlx::query("SELECT 41::bigint + 1").fetch_one(&pool).await.unwrap().get(0); + assert_eq!(n, 42); + } + + #[tokio::test] + #[ignore] + async fn sequential_queries_reuse_one_connection() { + let pool = pool().await; + let t = std::time::Instant::now(); + for _ in 0..5 { + sqlx::query("SELECT 1").execute(&pool).await.unwrap(); + } + println!("5 sequential queries on one pooled connection: {}ms", t.elapsed().as_millis()); + } +} diff --git a/src-tauri/src/db/import.rs b/src-tauri/src/db/import.rs index 557db457..11df9862 100644 --- a/src-tauri/src/db/import.rs +++ b/src-tauri/src/db/import.rs @@ -174,6 +174,7 @@ pub async fn import_rows( ActiveConnection::Clickhouse(_) => { Err("Importing rows is not supported for ClickHouse. Use INSERT INTO … in the SQL console.".into()) } + ActiveConnection::Posthog(_) => Err("PostHog is read-only: its data comes from HogQL queries.".into()), ActiveConnection::Redis(_) => Err("Importing rows is not supported on Redis".into()), }?; diff --git a/src-tauri/src/db/insights.rs b/src-tauri/src/db/insights.rs index 11ac2f49..eb93e26e 100644 --- a/src-tauri/src/db/insights.rs +++ b/src-tauri/src/db/insights.rs @@ -468,9 +468,8 @@ async fn mysql_activity(pool: &MySqlPool) -> InstanceActivity { .fetch_one(pool) .await { - Ok(r) => r - .try_get::(1) - .ok() + // my_text: SHOW STATUS values can arrive VARBINARY on MySQL 8+. + Ok(r) => super::mysql::my_text(&r, 1) .and_then(|s| s.trim().parse::().ok()) .unwrap_or(0), Err(_) => 0, @@ -544,10 +543,8 @@ async fn mysql_status_map(pool: &MySqlPool, names: &[&str]) -> HashMap(0).unwrap_or_default(); - let value = r - .try_get::(1) - .ok() + let name = super::mysql::my_text(r, 0).unwrap_or_default(); + let value = super::mysql::my_text(r, 1) .and_then(|s| s.trim().parse::().ok()) .unwrap_or(0); if !name.is_empty() { diff --git a/src-tauri/src/db/mod.rs b/src-tauri/src/db/mod.rs index b2467be1..af67d5e4 100644 --- a/src-tauri/src/db/mod.rs +++ b/src-tauri/src/db/mod.rs @@ -1,5 +1,6 @@ pub mod backup; pub mod clickhouse; +pub mod posthog; pub mod connection; pub mod d1; pub mod duckdb; diff --git a/src-tauri/src/db/mysql.rs b/src-tauri/src/db/mysql.rs index 5dfbc640..94d18d1b 100644 --- a/src-tauri/src/db/mysql.rs +++ b/src-tauri/src/db/mysql.rs @@ -232,7 +232,7 @@ pub async fn fetch_primary_key(pool: &MySqlPool, schema: &str, table: &str) -> R .fetch_all(pool) .await .map_err(|e| format!("Failed to load primary key: {e}"))?; - Ok(rows.iter().filter_map(|r| r.try_get::(0).ok()).collect()) + Ok(rows.iter().filter_map(|r| my_text(r, 0)).collect()) } fn escape_like(input: &str) -> String { @@ -385,67 +385,93 @@ pub async fn get_table_rows( // name/type/nullability projection upfront and reuse it below for the // nullable map and the empty-table column fallback (instead of a second // information_schema.COLUMNS query inside the join). - let meta_rows = sqlx::query( + // Every catalog and data query of a table open, in as few round trips as + // the dependencies allow. Against a remote MySQL each sequential query costs + // a full round trip (~300ms to a US region from South Asia), and this used to + // run four in a row: columns, then count + rows, then primary key, then + // foreign keys. PK and FK depend on nothing, so they always join the page + // fetch; and a plain open (no search, filter or sort) doesn't need the column + // list to build its query, so the column metadata joins it too: one stage. + let meta_q = sqlx::query( "SELECT COLUMN_NAME, DATA_TYPE, IS_NULLABLE, EXTRA, COLUMN_DEFAULT \ FROM information_schema.COLUMNS \ WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? ORDER BY ORDINAL_POSITION", ) .bind(schema) - .bind(table) - .fetch_all(pool) - .await - .map_err(|e| format!("Failed to load columns: {e}"))?; - let table_columns: Vec = meta_rows - .iter() - .filter_map(|r| r.try_get::(0).ok()) - .collect(); - let filters = filters.unwrap_or_default(); - let where_clause = build_where(&table_columns, search.as_deref(), search_is_regex, search_case_sensitive, &filters)?; - - let order_by = if let Some(col) = sort_column.as_deref().map(str::trim).filter(|s| !s.is_empty()) { - // Validate the sort column against the fetched columns so an unknown name - // never reaches the query (mirrors the Postgres ensure_column check). - if !table_columns.iter().any(|c| c == col) { - return Err(format!("Unknown column: {col}")); - } - let dir = match sort_direction.as_deref().unwrap_or("asc") { - "desc" => "DESC", - _ => "ASC", - }; - // Emulate NULLS FIRST/LAST - real MySQL (unlike MariaDB) rejects the - // `NULLS FIRST/LAST` syntax. `ISNULL(col)` yields 0 for non-NULLs and 1 - // for NULLs: ordering it ASC keeps NULLs last, DESC puts NULLs first. - let qc = bt(col); - match nulls_order.as_deref() { - Some("first") => format!(" ORDER BY ISNULL({qc}) DESC, {qc} {dir}"), - _ => format!(" ORDER BY ISNULL({qc}), {qc} {dir}"), + .bind(table); + let pk_fk = async { + if include_meta { + let (pk, fks) = tokio::join!(fetch_primary_key(pool, schema, table), fetch_foreign_keys(pool, schema, table)); + (pk.unwrap_or_default(), fks.unwrap_or_default()) + } else { + // The frontend keeps the values it already loaded for this table. + (Vec::new(), Vec::new()) } - } else { - String::new() }; + let filters = filters.unwrap_or_default(); + let sort_col = sort_column.as_deref().map(str::trim).filter(|s| !s.is_empty()).map(str::to_string); + let needs_columns_first = + sort_col.is_some() || !filters.is_empty() || search.as_deref().is_some_and(|q| !q.trim().is_empty()); let table_ref = format!("{}.{}", bt(schema), bt(table)); // MySQL types COUNT(*) as BIGINT UNSIGNED, but MariaDB types it as signed // BIGINT - decoding the wrong signedness is a hard type-mismatch in sqlx. // CAST(... AS SIGNED) normalizes both to i64, which comfortably holds any // real row count. - let count_sql = format!("SELECT CAST(COUNT(*) AS SIGNED) FROM {table_ref}{}", where_clause.sql); - let mut count_q = sqlx::query_scalar::<_, i64>(&count_sql); - for b in &where_clause.binds { - count_q = count_q.bind(b.as_str()); - } - - let data_sql = format!("SELECT * FROM {table_ref}{}{} LIMIT ? OFFSET ?", where_clause.sql, order_by); - let mut data_q = sqlx::query(&data_sql); - for b in &where_clause.binds { - data_q = data_q.bind(b.as_str()); - } - data_q = data_q.bind(limit).bind(offset); + let build = |table_columns: &[String]| -> Result<(String, String, Vec), String> { + let where_clause = build_where(table_columns, search.as_deref(), search_is_regex, search_case_sensitive, &filters)?; + let order_by = if let Some(col) = sort_col.as_deref() { + // Validate the sort column against the fetched columns so an unknown + // name never reaches the query (mirrors the Postgres ensure_column check). + if !table_columns.iter().any(|c| c == col) { + return Err(format!("Unknown column: {col}")); + } + let dir = match sort_direction.as_deref().unwrap_or("asc") { + "desc" => "DESC", + _ => "ASC", + }; + // Emulate NULLS FIRST/LAST - real MySQL (unlike MariaDB) rejects the + // `NULLS FIRST/LAST` syntax. `ISNULL(col)` yields 0 for non-NULLs and + // 1 for NULLs: ordering it ASC keeps NULLs last, DESC puts NULLs first. + let qc = bt(col); + match nulls_order.as_deref() { + Some("first") => format!(" ORDER BY ISNULL({qc}) DESC, {qc} {dir}"), + _ => format!(" ORDER BY ISNULL({qc}), {qc} {dir}"), + } + } else { + String::new() + }; + Ok(( + format!("SELECT CAST(COUNT(*) AS SIGNED) FROM {table_ref}{}", where_clause.sql), + format!("SELECT * FROM {table_ref}{}{} LIMIT ? OFFSET ?", where_clause.sql, order_by), + where_clause.binds, + )) + }; + let run_page = |count_sql: String, data_sql: String, binds: Vec| async move { + let mut count_q = sqlx::query_scalar::<_, i64>(&count_sql); + let mut data_q = sqlx::query(&data_sql); + for b in &binds { + count_q = count_q.bind(b.clone()); + data_q = data_q.bind(b.clone()); + } + data_q = data_q.bind(limit).bind(offset); + let (t, r) = tokio::join!(count_q.fetch_one(pool), data_q.fetch_all(pool)); + (t, r, count_sql, data_sql) + }; - let (total_res, rows_res) = tokio::join!( - count_q.fetch_one(pool), - data_q.fetch_all(pool), - ); + let (meta_rows, (total_res, rows_res, count_sql, data_sql), (pk, fks)) = if needs_columns_first { + // The WHERE and ORDER BY are built from, and validated against, the + // column list, so it has to land first. + let meta_rows = meta_q.fetch_all(pool).await.map_err(|e| format!("Failed to load columns: {e}"))?; + let names: Vec = meta_rows.iter().filter_map(|r| my_text(r, 0)).collect(); + let (count_sql, data_sql, binds) = build(&names)?; + let (page, meta) = tokio::join!(run_page(count_sql, data_sql, binds), pk_fk); + (meta_rows, page, meta) + } else { + let (count_sql, data_sql, binds) = build(&[])?; + let (meta_rows, page, meta) = tokio::join!(meta_q.fetch_all(pool), run_page(count_sql, data_sql, binds), pk_fk); + (meta_rows.map_err(|e| format!("Failed to load columns: {e}"))?, page, meta) + }; let total: i64 = total_res.map_err(|e| format!("Failed to count rows: {e}"))?; let rows = rows_res.map_err(|e| format!("Failed to fetch rows: {e}"))?; @@ -455,10 +481,10 @@ pub async fn get_table_rows( let flags_map: HashMap = meta_rows .iter() .filter_map(|r| { - let name = r.try_get::(0).ok()?; - let nullable = r.try_get::(2).ok()?; - let extra = r.try_get::(3).unwrap_or_default().to_ascii_lowercase(); - let default = r.try_get::, _>(4).ok().flatten(); + let name = my_text(r, 0)?; + let nullable = my_text(r, 2)?; + let extra = my_text(r, 3).unwrap_or_default().to_ascii_lowercase(); + let default = my_text(r, 4); let auto_generated = extra.contains("auto_increment") || extra.contains("generated"); Some(( name, @@ -482,10 +508,7 @@ pub async fn get_table_rows( meta_rows .iter() .filter_map(|r| { - Some(ColumnInfo::new( - r.try_get::(0).ok()?, - r.try_get::(1).ok()?.to_lowercase(), - )) + Some(ColumnInfo::new(my_text(r, 0)?, my_text(r, 1)?.to_lowercase())) }) .collect() }; @@ -503,17 +526,6 @@ pub async fn get_table_rows( .map(|row| (0..row.len()).map(|i| cell_to_json(row, i)).collect()) .collect(); - // Skip the PK/FK catalog round-trips on metadata-skipping fetches; the - // frontend keeps the values it already loaded for this table. - let (pk, fks) = if include_meta { - ( - fetch_primary_key(pool, schema, table).await.unwrap_or_default(), - fetch_foreign_keys(pool, schema, table).await.unwrap_or_default(), - ) - } else { - (Vec::new(), Vec::new()) - }; - Ok(TableRows { // Preview fetching is a Postgres path (pg_stats + pg_column_size). preview_columns: Vec::new(), @@ -548,11 +560,13 @@ async fn fetch_foreign_keys(pool: &MySqlPool, schema: &str, table: &str) -> Resu let mut out: Vec = Vec::new(); let mut current: Option = None; for row in &rows { - let constraint: String = row.try_get(0).unwrap_or_default(); - let column: String = row.try_get(1).unwrap_or_default(); - let ref_schema: String = row.try_get(2).unwrap_or_default(); - let ref_table: String = row.try_get(3).unwrap_or_default(); - let ref_col: String = row.try_get(4).unwrap_or_default(); + // my_text: a refused VARBINARY decode here blanked every foreign key on + // MySQL 8+, so FK jumps and relation chips silently disappeared. + let constraint = my_text(row, 0).unwrap_or_default(); + let column = my_text(row, 1).unwrap_or_default(); + let ref_schema = my_text(row, 2).unwrap_or_default(); + let ref_table = my_text(row, 3).unwrap_or_default(); + let ref_col = my_text(row, 4).unwrap_or_default(); if current.as_deref() == Some(&constraint) { if let Some(fk) = out.last_mut() { fk.columns.push(column); @@ -722,10 +736,13 @@ pub async fn insert_table_row( let mut auto_increment_col: Option = None; for row in &meta_rows { - let name: String = row.try_get(0).map_err(|e| format!("Invalid column name: {e}"))?; - let is_nullable: String = row.try_get(1).unwrap_or_else(|_| "NO".to_string()); - let default_val: Option = row.try_get::, _>(2).ok().flatten(); - let extra: String = row.try_get::, _>(3).ok().flatten().unwrap_or_default(); + // my_text throughout: on MySQL 8+ these can arrive VARBINARY. A refused + // decode of EXTRA silently lost AUTO_INCREMENT, so the new id was never + // read back and the inserted row came back with a NULL key. + let name = my_text(row, 0).ok_or("Invalid column name in information_schema")?; + let is_nullable = my_text(row, 1).unwrap_or_else(|| "NO".to_string()); + let default_val = my_text(row, 2); + let extra = my_text(row, 3).unwrap_or_default(); let is_auto = extra.to_lowercase().contains("auto_increment"); let opt = is_auto || default_val.is_some() || is_nullable.eq_ignore_ascii_case("YES"); if is_auto { @@ -906,3 +923,59 @@ mod tests { assert!(!is_plain_text("")); } } + +/// Against the dialect-matrix container (`docker compose -f docker/dialects.yml +/// up -d mysql`), like `dialect_matrix`: `cargo test --lib mysql_live -- --ignored`. +#[cfg(test)] +mod mysql_live { + use super::*; + + async fn pool() -> MySqlPool { + MySqlPool::connect("mysql://root:stroke@127.0.0.1:53306/shop").await.expect("stroke-test-mysql is running") + } + + /// A plain open fetches everything in one stage; a sorted one fetches the + /// columns first. Both must return the same rows, columns, key and count. + #[tokio::test] + #[ignore] + async fn plain_and_sorted_opens_agree() { + let pool = pool().await; + let plain = get_table_rows(&pool, "shop", "customers", 50, 0, None, false, false, None, None, None, true, None) + .await + .unwrap(); + assert!(!plain.columns.is_empty() && !plain.rows.is_empty()); + assert_eq!(plain.total, plain.rows.len() as i64); + assert!(!plain.primary_key.is_empty(), "primary key read back"); + let first = plain.columns[0].name.clone(); + let sorted = get_table_rows(&pool, "shop", "customers", 50, 0, None, false, false, Some(first), Some("desc".into()), None, true, None) + .await + .unwrap(); + assert_eq!(sorted.columns.len(), plain.columns.len()); + assert_eq!(sorted.total, plain.total); + assert_eq!(sorted.primary_key, plain.primary_key); + assert!(get_table_rows(&pool, "shop", "customers", 50, 0, None, false, false, Some("nope".into()), None, None, true, None) + .await + .unwrap_err() + .contains("Unknown column")); + } + + /// The bug behind "No columns visible": an empty table must still come back + /// with its columns, from the catalog. + #[tokio::test] + #[ignore] + async fn an_empty_table_still_has_columns() { + let pool = pool().await; + sqlx::query("CREATE TABLE IF NOT EXISTS stroke_empty_probe (id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, name TEXT)") + .execute(&pool) + .await + .unwrap(); + let r = get_table_rows(&pool, "shop", "stroke_empty_probe", 50, 0, None, false, false, None, None, None, true, None) + .await + .unwrap(); + let names: Vec<&str> = r.columns.iter().map(|c| c.name.as_str()).collect(); + assert_eq!(names, ["id", "name"]); + assert_eq!((r.total, r.rows.len()), (0, 0)); + assert_eq!(r.primary_key, ["id"]); + sqlx::query("DROP TABLE stroke_empty_probe").execute(&pool).await.unwrap(); + } +} diff --git a/src-tauri/src/db/posthog.rs b/src-tauri/src/db/posthog.rs new file mode 100644 index 00000000..3d8c5125 --- /dev/null +++ b/src-tauri/src/db/posthog.rs @@ -0,0 +1,446 @@ +//! PostHog, queried with HogQL over PostHog's query API. +//! +//! PostHog stores events, persons and sessions in ClickHouse, but Cloud gives no +//! direct ClickHouse access. The query API (`POST /api/projects/{id}/query/`) +//! takes HogQL, PostHog's SQL dialect over that ClickHouse, and returns +//! `columns` / `types` / `results`. The schema comes from the same endpoint as a +//! `DatabaseSchemaQuery`. Everything here is read-only: the API runs SELECTs. +//! +//! Limits that shape this module (PostHog docs): 10s per query, 50,000 rows at +//! most, 240 requests a minute and 3 concurrent queries per project. Hence the +//! schema cache, one request per page, and a row cap. + +use super::connection::PosthogConfig; +use super::query::{ColumnInfo, ForeignKeyInfo, RowFilter, SqlResult, TableRows}; +use super::schema::{ColumnStructureRow, IndexInfo, TableInfo}; +use serde_json::{json, Value}; +use std::collections::HashMap; +use std::sync::{Mutex, OnceLock}; +use std::time::{Duration, Instant}; + +/// The one schema PostHog exposes to the sidebar. +pub const SCHEMA: &str = "posthog"; +/// HogQL's hard ceiling on returned rows. +const MAX_ROWS: i64 = 50_000; +/// How long a fetched schema is reused before asking PostHog again. +const SCHEMA_TTL: Duration = Duration::from_secs(300); + +fn client() -> &'static reqwest::Client { + static C: OnceLock = OnceLock::new(); + C.get_or_init(|| { + reqwest::Client::builder() + .user_agent("stroke/1.0") + .pool_idle_timeout(Duration::from_secs(20)) + .connect_timeout(Duration::from_secs(10)) + // PostHog stops a query at 10s; leave room for queueing and transfer. + .timeout(Duration::from_secs(45)) + .build() + .expect("failed to build PostHog HTTP client") + }) +} + +/// POST one query object to the project's query endpoint. +async fn post_query(config: &PosthogConfig, query: Value) -> Result { + let url = format!( + "{}/api/projects/{}/query/", + config.base_url(), + urlencoding::encode(config.project_id.trim()) + ); + let resp = client() + .post(&url) + .bearer_auth(config.api_key.trim()) + .json(&json!({ "query": query, "name": "stroke" })) + .send() + .await + .map_err(|e| format!("PostHog request failed: {e}"))?; + let status = resp.status().as_u16(); + let text = resp.text().await.map_err(|e| format!("PostHog read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).unwrap_or(Value::Null); + if !(200..300).contains(&status) { + let detail = body["detail"] + .as_str() + .or_else(|| body["error"].as_str()) + .map(String::from) + .unwrap_or_else(|| text.chars().take(300).collect()); + return Err(match status { + 401 => "PostHog rejected the API key. Check it, and that it has the Query Read scope.".into(), + 403 => format!("PostHog refused access to project {}: {detail}", config.project_id), + 429 => "PostHog's query rate limit was hit (240 a minute, 3 at once). Wait a moment and try again.".into(), + // Seen on internal tables PostHog lists but can't query for a project + // (document_embeddings, some preaggregated and raw tables): even + // count() fails there. Nothing to change in the query. + 500..=599 => "PostHog failed on its side (HTTP 5xx). Some of PostHog's internal tables are listed but can't be queried through the API; the main ones (events, persons, sessions) can.".into(), + _ => format!("PostHog error ({status}): {}", tidy_error(&detail)), + }); + } + if let Some(err) = body["error"].as_str().filter(|e| !e.is_empty()) { + return Err(format!("PostHog error: {err}")); + } + Ok(body) +} + +/// PostHog's 400s can append the whole generated ClickHouse query and its +/// SETTINGS. Keep the sentence that says what's wrong. +fn tidy_error(detail: &str) -> String { + if detail.contains("Unknown table expression identifier") { + return "PostHog lists this table, but it can't be queried on this project.".into(); + } + let cut = [" in scope SELECT", " SETTINGS readonly", "\nSELECT "] + .iter() + .filter_map(|m| detail.find(m)) + .min() + .unwrap_or(detail.len()); + let short = detail[..cut].trim(); + if short.chars().count() > 400 { short.chars().take(400).collect::() + "…" } else { short.to_string() } +} + +/// Only SELECT-shaped HogQL reaches the API: it can't write, so say so up front +/// rather than relaying a parse error. +fn is_read_query(sql: &str) -> bool { + matches!(super::sql_util::statement_head(sql).as_str(), "select" | "with") +} + +/// A `types` entry: `["column", "Nullable(String)"]` or a bare type string. +fn type_of(entry: &Value) -> String { + let raw = entry + .as_array() + .and_then(|pair| pair.get(1)) + .and_then(Value::as_str) + .or_else(|| entry.as_str()) + .unwrap_or(""); + strip_wrappers(raw) +} + +/// A ClickHouse type as a short grid label: wrappers off (`Nullable`, +/// `LowCardinality`, `SimpleAggregateFunction(sum, Int64)` → `Int64`) and +/// parameters dropped (`DateTime64(6, 'UTC')` → `DateTime64`, a long +/// `Enum8('full' = 0, …)` → `Enum8`). +fn strip_wrappers(ty: &str) -> String { + let mut t = ty.trim(); + loop { + let before = t; + for w in ["Nullable(", "LowCardinality("] { + if let Some(inner) = t.strip_prefix(w).and_then(|s| s.strip_suffix(')')) { + t = inner; + } + } + if let Some(inner) = t.strip_prefix("SimpleAggregateFunction(").and_then(|s| s.strip_suffix(')')) { + t = inner.split_once(',').map(|(_, ty)| ty.trim()).unwrap_or(inner); + } + if t == before { + break; + } + } + match t.split_once('(') { + Some((name, _)) if matches!(name, "DateTime64" | "DateTime" | "Enum8" | "Enum16" | "Decimal" | "FixedString") => name.to_string(), + _ => t.to_string(), + } +} + +/// `system.information_schema.columns` → `` `system`.`information_schema`.`columns` ``. +/// Quoted whole, a three-part name made PostHog answer HTTP 500; unquoted it +/// works, and quoting each part works for every shape PostHog lists. +fn quote_table(table: &str) -> String { + table.split('.').map(super::sql_util::quote_backtick).collect::>().join(".") +} + +/// A HogQL response → the app's result shape. +fn to_sql_result(body: &Value, sql: &str, query_ms: u64) -> SqlResult { + let names: Vec = body["columns"] + .as_array() + .map(|c| c.iter().map(|v| v.as_str().unwrap_or("").to_string()).collect()) + .unwrap_or_default(); + let types = body["types"].as_array().cloned().unwrap_or_default(); + let columns: Vec = names + .iter() + .enumerate() + .map(|(i, n)| { + let raw = types.get(i).cloned().unwrap_or(Value::Null); + let mut c = ColumnInfo::new(n.clone(), type_of(&raw)); + // Read-only analytics: NOT NULL is nothing to act on, and marking + // every non-Nullable column required put a `*` on nearly all of them. + c.nullable = true; + c + }) + .collect(); + // Cap oversized cells (event `properties` can be large JSON) before they + // reach the webview, like every other engine. + let rows: Vec> = body["results"] + .as_array() + .map(|rows| { + rows.iter() + .map(|r| { + r.as_array() + .map(|cells| cells.iter().cloned().map(|v| super::sql_util::cap_json_value("text", v)).collect()) + .unwrap_or_default() + }) + .collect() + }) + .unwrap_or_default(); + SqlResult { + row_count: Some(rows.len() as i64), + message: body["hasMore"].as_bool().filter(|m| *m).map(|_| "More rows exist - add a LIMIT to fetch a specific range.".into()), + columns, + rows, + query_ms, + sql: sql.to_string(), + } +} + +/// Run one HogQL statement. +pub async fn query(config: &PosthogConfig, sql: &str) -> Result { + let trimmed = sql.trim().trim_end_matches(';'); + if !is_read_query(trimmed) { + return Err("PostHog is read-only here: HogQL runs SELECT (and WITH) queries.".into()); + } + let t0 = Instant::now(); + let body = post_query(config, json!({ "kind": "HogQLQuery", "query": trimmed })).await?; + Ok(to_sql_result(&body, sql, t0.elapsed().as_millis() as u64)) +} + +// ── Schema ─────────────────────────────────────────────────────────────────── + +fn schema_cache() -> &'static Mutex> { + static C: OnceLock>> = OnceLock::new(); + C.get_or_init(Default::default) +} + +/// The project's `DatabaseSchemaQuery` result (`tables` map), cached for a few +/// minutes: it is the heaviest call here, and the sidebar, structure view and +/// every filtered page all need it. +async fn schema(config: &PosthogConfig) -> Result { + let key = format!("{}|{}", config.base_url(), config.project_id.trim()); + if let Some((at, v)) = schema_cache().lock().ok().and_then(|m| m.get(&key).cloned()) { + if at.elapsed() < SCHEMA_TTL { + return Ok(v); + } + } + let body = post_query(config, json!({ "kind": "DatabaseSchemaQuery" })).await?; + let tables = body["tables"].clone(); + if let Ok(mut m) = schema_cache().lock() { + m.insert(key, (Instant::now(), tables.clone())); + } + Ok(tables) +} + +/// Field types that are real columns. The rest (`lazy_table`, `virtual_table`, +/// `field_traverser`, `expression`, views) are relations and computed paths +/// that `SELECT *` doesn't return. +fn is_column(ty: &str) -> bool { + matches!(ty, "integer" | "float" | "decimal" | "string" | "datetime" | "date" | "boolean" | "array" | "json" | "unknown") +} + +fn table_fields(tables: &Value, table: &str) -> Vec<(String, String)> { + tables[table]["fields"] + .as_object() + .map(|fields| { + fields + .values() + .filter_map(|f| { + let ty = f["type"].as_str().unwrap_or("unknown"); + is_column(ty).then(|| (f["name"].as_str().unwrap_or("").to_string(), ty.to_string())) + }) + .filter(|(n, _)| !n.is_empty()) + .collect() + }) + .unwrap_or_default() +} + +pub async fn list_schemas(_config: &PosthogConfig) -> Result, String> { + Ok(vec![SCHEMA.to_string()]) +} + +pub async fn list_tables(config: &PosthogConfig) -> Result, String> { + let tables = schema(config).await?; + let mut out: Vec = tables + .as_object() + .map(|m| { + m.values() + .filter_map(|t| { + let name = t["name"].as_str()?.to_string(); + let kind = match t["type"].as_str().unwrap_or("") { + "view" | "materialized_view" | "managed_view" => "view", + _ => "table", + }; + let row_count = t["row_count"].as_f64().map(|n| n as i64).unwrap_or(-1); + Some(TableInfo { name, kind: kind.to_string(), row_count, rls_enabled: None }) + }) + .collect() + }) + .unwrap_or_default(); + out.sort_by(|a, b| a.name.cmp(&b.name)); + Ok(out) +} + +pub async fn list_indexes(_config: &PosthogConfig) -> Result, String> { + Ok(vec![]) +} + +pub async fn get_column_structure(config: &PosthogConfig, table: &str) -> Result, String> { + let tables = schema(config).await?; + Ok(table_fields(&tables, table) + .into_iter() + .enumerate() + .map(|(i, (name, ty))| ColumnStructureRow { + ordinal_position: (i + 1) as i32, + name, + data_type: ty, + is_nullable: true, + column_default: None, + foreign_key: None, + fk_constraint_name: None, + comment: None, + }) + .collect()) +} + +// ── Browsing ───────────────────────────────────────────────────────────────── + +/// One page of a table, and its count, as two concurrent HogQL queries. +pub async fn get_table_rows( + config: &PosthogConfig, + table: &str, + limit: i64, + offset: i64, + search: Option, + sort_column: Option, + sort_direction: Option, + filters: Option>, +) -> Result { + let t0 = Instant::now(); + let tq = quote_table(table); + let has_search = search.as_deref().map(str::trim).is_some_and(|s| !s.is_empty()); + let has_filters = filters.as_ref().is_some_and(|f| !f.is_empty()); + let cols = if has_search || has_filters { get_column_structure(config, table).await? } else { Vec::new() }; + let where_clause = super::clickhouse::build_where(&cols, search.as_deref(), filters.as_deref()); + + let order = match (sort_column.as_deref().map(str::trim), sort_direction.as_deref()) { + (Some(c), dir) if !c.is_empty() => { + let d = if dir.is_some_and(|d| d.eq_ignore_ascii_case("desc")) { "DESC" } else { "ASC" }; + format!(" ORDER BY {} {d}", super::sql_util::quote_backtick(c)) + } + // Events are read newest first: the start of a years-long event table is + // rarely what anyone opened it for. + _ if table == "events" => " ORDER BY timestamp DESC".to_string(), + _ => String::new(), + }; + // No OFFSET: PostHog rejects it on personal-API-key queries ("OFFSET is not + // supported ... use keyset pagination on timestamp"), even `OFFSET 0`. + // Keyset needs a timestamp most tables don't have, so a later page asks for + // everything up to its end and drops the rows before it. PostHog's 50,000-row + // ceiling bounds that; past it, say what to do. + let (limit, offset) = (limit.max(1), offset.max(0)); + let window = offset.saturating_add(limit); + if offset >= MAX_ROWS { + return Err(format!( + "PostHog's API returns at most {MAX_ROWS} rows per query, so rows past {MAX_ROWS} can't be paged to. Add a filter or sort to reach them." + )); + } + let fetch = window.min(MAX_ROWS); + let count_sql = format!("SELECT count() FROM {tq}{where_clause}"); + let data_sql = format!("SELECT * FROM {tq}{where_clause}{order} LIMIT {fetch}"); + let (count, page) = tokio::join!(query(config, &count_sql), query(config, &data_sql)); + let total = count + .ok() + .and_then(|r| r.rows.first().and_then(|row| row.first()).and_then(|v| v.as_i64().or_else(|| v.as_str()?.parse().ok()))) + .unwrap_or(-1); + let mut page = page?; + let skip = (offset as usize).min(page.rows.len()); + page.rows.drain(..skip); + Ok(TableRows { + preview_columns: Vec::new(), + columns: page.columns, + rows: page.rows, + total, + query_ms: t0.elapsed().as_millis() as u64, + // Analytics data: browse-only, no row identity to edit by. + primary_key: vec![], + foreign_keys: Vec::::new(), + sql: format!("{data_sql}\n{count_sql}"), + }) +} + +/// The count for the grid's pager (its "All" mode asks for it separately). +pub async fn count_rows( + config: &PosthogConfig, + table: &str, + search: Option, + filters: Option>, +) -> Result { + let has_search = search.as_deref().map(str::trim).is_some_and(|s| !s.is_empty()); + let has_filters = filters.as_ref().is_some_and(|f| !f.is_empty()); + let cols = if has_search || has_filters { get_column_structure(config, table).await? } else { Vec::new() }; + let where_clause = super::clickhouse::build_where(&cols, search.as_deref(), filters.as_deref()); + let r = query(config, &format!("SELECT count() FROM {}{where_clause}", quote_table(table))).await?; + Ok(r.rows + .first() + .and_then(|row| row.first()) + .and_then(|v| v.as_i64().or_else(|| v.as_str()?.parse().ok())) + .unwrap_or(-1)) +} + +/// A readable definition for the DDL view: PostHog tables have no CREATE +/// statement, so describe the columns instead. +pub async fn get_ddl(config: &PosthogConfig, table: &str) -> Result { + let cols = get_column_structure(config, table).await?; + let body: Vec = cols.iter().map(|c| format!(" {} {}", super::sql_util::quote_backtick(&c.name), c.data_type)).collect(); + Ok(format!("-- PostHog table (HogQL), read-only\n{} (\n{}\n)", table, body.join(",\n"))) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn hogql_response_maps_columns_types_and_rows() { + let body = json!({ + "columns": ["event", "timestamp", "distinct_id"], + "types": [["event", "String"], ["timestamp", "DateTime64(6, 'UTC')"], ["distinct_id", "Nullable(String)"]], + "results": [["$pageview", "2026-09-29T10:00:00Z", "u1"]], + "hasMore": false + }); + let r = to_sql_result(&body, "SELECT 1", 12); + let names: Vec<&str> = r.columns.iter().map(|c| c.name.as_str()).collect(); + assert_eq!(names, ["event", "timestamp", "distinct_id"]); + assert_eq!(r.columns[2].data_type, "String"); + // Every PostHog column is nullable for the grid: no `*` markers. + assert!(r.columns.iter().all(|c| c.nullable)); + assert_eq!(r.rows.len(), 1); + assert!(r.message.is_none()); + } + + #[test] + fn schema_fields_keep_real_columns_only() { + let tables = json!({ "events": { "name": "events", "type": "posthog", "fields": { + "event": { "name": "event", "type": "string" }, + "properties": { "name": "properties", "type": "json" }, + "person": { "name": "person", "type": "lazy_table" }, + "pdi": { "name": "pdi", "type": "field_traverser" } + }}}); + let names: Vec = table_fields(&tables, "events").into_iter().map(|(n, _)| n).collect(); + assert_eq!(names, ["event", "properties"]); + } + + #[test] + fn type_labels_and_nested_names_match_what_posthog_sends() { + assert_eq!(strip_wrappers("LowCardinality(String)"), "String"); + assert_eq!(strip_wrappers("SimpleAggregateFunction(sum, Int64)"), "Int64"); + assert_eq!(strip_wrappers("DateTime64(6, 'UTC')"), "DateTime64"); + assert_eq!(strip_wrappers("Nullable(Enum8('full' = 0, 'propertyless' = 1))"), "Enum8"); + assert_eq!(quote_table("system.information_schema.columns"), "`system`.`information_schema`.`columns`"); + assert_eq!(quote_table("events"), "`events`"); + } + + #[test] + fn long_errors_keep_their_first_sentence() { + assert!(tidy_error("Unknown table expression identifier 'x' in scope SELECT a FROM x SETTINGS readonly = 2").contains("can't be queried")); + assert_eq!(tidy_error("Syntax error near FROM in scope SELECT 1 SETTINGS readonly = 2"), "Syntax error near FROM"); + } + + #[test] + fn only_select_shaped_hogql_is_sent() { + assert!(is_read_query("SELECT event FROM events")); + assert!(is_read_query("WITH x AS (SELECT 1) SELECT * FROM x")); + assert!(!is_read_query("DELETE FROM events")); + } +} diff --git a/src-tauri/src/db/query.rs b/src-tauri/src/db/query.rs index e47ce811..08c5af21 100644 --- a/src-tauri/src/db/query.rs +++ b/src-tauri/src/db/query.rs @@ -1677,6 +1677,9 @@ pub async fn get_table_rows( &cfg, &schema, &table, limit, offset, search, sort_column, sort_direction, filters, include_meta, ).await; } + ActiveConnection::Posthog(cfg) => { + return super::posthog::get_table_rows(&cfg, &table, limit, offset, search, sort_column, sort_direction, filters).await; + } ActiveConnection::Redis(cfg) => { return super::redis::get_table_rows( &cfg, &table, limit, offset, search, sort_column, sort_direction, filters, include_meta, @@ -1799,7 +1802,7 @@ pub async fn get_table_rows( // drop a new column, which looks exactly like stale data to whoever is // looking at it. if include_meta { - super::wide_columns::invalidate(&pool, &schema, &table); + super::wide_columns::invalidate_projection(&pool, &schema, &table); } let (mut wide_projection, wide) = if preview_wide { super::wide_columns::page_projection(&pool, &schema, &table).await @@ -2103,6 +2106,9 @@ pub async fn count_table_rows( ) -> Result { match require_conn(&state)? { ActiveConnection::Postgres(_) => {} + // The grid's "All" mode takes its total from here, so PostHog needs its + // own count or the pager showed no total at all. + ActiveConnection::Posthog(cfg) => return super::posthog::count_rows(&cfg, &table, search, filters).await, _ => return Ok(-1), } let pool = require_pool(&state)?; @@ -2194,6 +2200,7 @@ pub async fn update_table_cell( ActiveConnection::Clickhouse(_) => { return Err("Inline row editing is not supported for ClickHouse (OLAP). Use ALTER TABLE … UPDATE in the SQL console.".into()); } + ActiveConnection::Posthog(_) => return Err("PostHog is read-only: its data comes from HogQL queries.".into()), ActiveConnection::Redis(_) => { return Err("Editing is not supported on Redis".into()); } @@ -2412,6 +2419,7 @@ pub async fn insert_table_row( ActiveConnection::Clickhouse(_) => { return Err("Row insertion via the grid is not supported for ClickHouse. Use INSERT INTO … in the SQL console.".into()); } + ActiveConnection::Posthog(_) => return Err("PostHog is read-only: its data comes from HogQL queries.".into()), ActiveConnection::Redis(_) => { return Err("Editing is not supported on Redis".into()); } @@ -2535,6 +2543,7 @@ pub async fn delete_table_rows( ActiveConnection::Clickhouse(_) => { return Err("Row deletion via the grid is not supported for ClickHouse. Use ALTER TABLE … DELETE in the SQL console.".into()); } + ActiveConnection::Posthog(_) => return Err("PostHog is read-only: its data comes from HogQL queries.".into()), ActiveConnection::Redis(_) => { return Err("Editing is not supported on Redis".into()); } @@ -2804,6 +2813,7 @@ pub async fn execute_sql( ActiveConnection::D1(cfg) => super::d1::query(&cfg, &sql_str, vec![]).await, ActiveConnection::LibSql(cfg) => super::libsql::query(&cfg, &sql_str, vec![]).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::query(&cfg, &sql_str).await, + ActiveConnection::Posthog(cfg) => super::posthog::query(&cfg, &sql_str).await, ActiveConnection::Redis(cfg) => super::redis::query(&cfg, &sql_str).await, ActiveConnection::Duckdb(h) => super::duckdb::execute_sql(&h, &sql_str).await, ActiveConnection::Mssql(h) => super::mssql::execute_sql(&h, &sql_str).await, @@ -2854,6 +2864,7 @@ pub async fn execute_sql_on_conn( } AnyConnectionConfig::Libsql(c) => super::libsql::query(&c, sql, vec![]).await, AnyConnectionConfig::Clickhouse(c) => super::clickhouse::query(&c, sql).await, + AnyConnectionConfig::Posthog(c) => super::posthog::query(&c, sql).await, AnyConnectionConfig::Redis(c) => super::redis::query(&c, sql).await, AnyConnectionConfig::Duckdb(c) => { let h = super::connection::open_duckdb(&c).await?; @@ -3355,6 +3366,7 @@ pub async fn execute_sql_multi( ActiveConnection::LibSql(cfg) => super::libsql::query(cfg, stmt, vec![]).await, ActiveConnection::Mysql(pool) => super::mysql::execute_sql(pool, stmt, None).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::query(cfg, stmt).await, + ActiveConnection::Posthog(cfg) => super::posthog::query(cfg, stmt).await, ActiveConnection::Redis(cfg) => super::redis::query(cfg, stmt).await, ActiveConnection::Duckdb(h) => super::duckdb::execute_sql(h, stmt).await, ActiveConnection::Mssql(h) => super::mssql::execute_sql(h, stmt).await, @@ -4083,6 +4095,7 @@ async fn dispatch_stats_sql(conn: &ActiveConnection, sql: &str) -> Result super::d1::query(cfg, sql, vec![]).await, ActiveConnection::LibSql(cfg) => super::libsql::query(cfg, sql, vec![]).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::query(cfg, sql).await, + ActiveConnection::Posthog(cfg) => super::posthog::query(cfg, sql).await, ActiveConnection::Redis(_) => Err("Column statistics are not supported on Redis".into()), ActiveConnection::Duckdb(h) => super::duckdb::execute_sql(h, sql).await, ActiveConnection::Mssql(h) => super::mssql::execute_sql(h, sql).await, @@ -4106,7 +4119,7 @@ pub async fn ping_connection(state: State<'_, DbState>) -> Result<(), String> { sqlx::query("SELECT 1").execute(&pool).await.map(|_| ()).map_err(|e| e.to_string()) } // HTTP-based: stateless, no persistent TCP connection to validate - ActiveConnection::D1(_) | ActiveConnection::LibSql(_) | ActiveConnection::Clickhouse(_) | ActiveConnection::Redis(_) => Ok(()), + ActiveConnection::D1(_) | ActiveConnection::LibSql(_) | ActiveConnection::Clickhouse(_) | ActiveConnection::Redis(_) | ActiveConnection::Posthog(_) => Ok(()), ActiveConnection::Duckdb(h) => super::duckdb::execute_sql(&h, "SELECT 1").await.map(|_| ()), ActiveConnection::Mssql(h) => super::mssql::execute_sql(&h, "SELECT 1").await.map(|_| ()), } diff --git a/src-tauri/src/db/schema.rs b/src-tauri/src/db/schema.rs index d43b0a53..1a784121 100644 --- a/src-tauri/src/db/schema.rs +++ b/src-tauri/src/db/schema.rs @@ -1011,6 +1011,7 @@ pub async fn list_schemas(state: State<'_, DbState>) -> Result, Stri ActiveConnection::Postgres(pool) => list_schemas_pg(&pool).await, ActiveConnection::Mysql(pool) => list_schemas_mysql(&pool).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::list_schemas(&cfg).await, + ActiveConnection::Posthog(cfg) => super::posthog::list_schemas(&cfg).await, ActiveConnection::Redis(_) => Ok(vec![]), ActiveConnection::Mssql(h) => super::mssql::list_schemas(&h).await, ActiveConnection::Sqlite(_) | ActiveConnection::D1(_) | ActiveConnection::LibSql(_) | ActiveConnection::Duckdb(_) => Ok(vec!["main".to_string()]), @@ -1021,13 +1022,20 @@ pub async fn list_tables(state: State<'_, DbState>, schema: String) -> Result { validate_ident(&schema)?; - list_tables_pg(&pool, &schema).await + let tables = list_tables_pg(&pool, &schema).await?; + // Answer the wide-column question for the whole schema now, off the + // critical path, so a table's first open is one round trip, not two. + let names: Vec = tables.iter().map(|t| t.name.clone()).collect(); + let (pool, schema) = (pool.clone(), schema.clone()); + tokio::spawn(async move { super::wide_columns::prefetch_schema(&pool, &schema, &names).await }); + Ok(tables) } ActiveConnection::Mysql(pool) => list_tables_mysql(&pool, &schema).await, ActiveConnection::Sqlite(pool) => list_tables_sqlite(&pool).await, ActiveConnection::D1(cfg) => list_tables_d1(&cfg).await, ActiveConnection::LibSql(cfg) => list_tables_libsql(&cfg).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::list_tables(&cfg, &schema).await, + ActiveConnection::Posthog(cfg) => super::posthog::list_tables(&cfg).await, ActiveConnection::Redis(cfg) => super::redis::list_tables(&cfg).await, ActiveConnection::Duckdb(h) => super::duckdb::list_tables(&h).await, ActiveConnection::Mssql(h) => super::mssql::list_tables(&h, &schema).await, @@ -1045,6 +1053,7 @@ pub async fn list_indexes(state: State<'_, DbState>, schema: String) -> Result list_indexes_d1(&cfg).await, ActiveConnection::LibSql(cfg) => list_indexes_libsql(&cfg).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::list_indexes(&cfg).await, + ActiveConnection::Posthog(cfg) => super::posthog::list_indexes(&cfg).await, ActiveConnection::Redis(cfg) => super::redis::list_indexes(&cfg).await, ActiveConnection::Duckdb(h) => super::duckdb::list_indexes(&h).await, ActiveConnection::Mssql(h) => super::mssql::list_indexes(&h, &schema).await, @@ -1238,6 +1247,7 @@ pub async fn get_table_column_structure( ActiveConnection::Clickhouse(cfg) => { super::clickhouse::get_column_structure(&cfg, &schema, &table).await } + ActiveConnection::Posthog(cfg) => super::posthog::get_column_structure(&cfg, &table).await, ActiveConnection::Redis(cfg) => { super::redis::get_column_structure(&cfg, &table).await } @@ -2220,6 +2230,7 @@ pub async fn get_table_ddl( ActiveConnection::D1(cfg) => get_ddl_d1(&cfg, &table).await, ActiveConnection::LibSql(cfg) => get_ddl_libsql(&cfg, &table).await, ActiveConnection::Clickhouse(cfg) => super::clickhouse::get_ddl(&cfg, &schema, &table).await, + ActiveConnection::Posthog(cfg) => super::posthog::get_ddl(&cfg, &table).await, ActiveConnection::Redis(cfg) => super::redis::get_ddl(&cfg, &table).await, ActiveConnection::Duckdb(h) => super::duckdb::get_ddl(&h, &table).await, ActiveConnection::Mssql(h) => super::mssql::get_ddl(&h, &schema, &table).await, @@ -2279,6 +2290,7 @@ pub async fn list_tables_on_conn( AnyConnectionConfig::D1(c) => list_tables_d1(&c).await.map(to_names), AnyConnectionConfig::Libsql(c) => list_tables_libsql(&c).await.map(to_names), AnyConnectionConfig::Clickhouse(c) => super::clickhouse::list_tables(&c, &schema).await.map(to_names), + AnyConnectionConfig::Posthog(c) => super::posthog::list_tables(&c).await.map(to_names), AnyConnectionConfig::Redis(c) => super::redis::list_tables(&c).await.map(to_names), AnyConnectionConfig::Duckdb(c) => { let h = super::connection::open_duckdb(&c).await?; @@ -2322,6 +2334,7 @@ pub async fn get_table_ddl_on_conn( AnyConnectionConfig::D1(c) => get_ddl_d1(&c, &table).await, AnyConnectionConfig::Libsql(c) => get_ddl_libsql(&c, &table).await, AnyConnectionConfig::Clickhouse(c) => super::clickhouse::get_ddl(&c, &schema, &table).await, + AnyConnectionConfig::Posthog(c) => super::posthog::get_ddl(&c, &table).await, AnyConnectionConfig::Redis(c) => super::redis::get_ddl(&c, &table).await, AnyConnectionConfig::Duckdb(c) => { let h = super::connection::open_duckdb(&c).await?; diff --git a/src-tauri/src/db/tx.rs b/src-tauri/src/db/tx.rs index 50b84237..f780a150 100644 --- a/src-tauri/src/db/tx.rs +++ b/src-tauri/src/db/tx.rs @@ -150,6 +150,7 @@ pub async fn tx_begin( ActiveConnection::Clickhouse(_) => { return Err("ClickHouse does not support transactions".into()) } + ActiveConnection::Posthog(_) => return Err("PostHog does not support transactions".into()), ActiveConnection::Redis(_) => return Err("Redis does not support SQL transactions".into()), }; diff --git a/src-tauri/src/db/wide_columns.rs b/src-tauri/src/db/wide_columns.rs index 96a4caf1..58b65b38 100644 --- a/src-tauri/src/db/wide_columns.rs +++ b/src-tauri/src/db/wide_columns.rs @@ -113,6 +113,16 @@ static WIDE_CACHE: OnceLock>> = Onc /// trip per table per minute and removes the question. const WIDE_CACHE_TTL: Duration = Duration::from_secs(60); +/// How long "nothing wide here" survives. That answer means a plain `SELECT *`, +/// which holds no column list and so can never drop a new column: the only thing +/// it can miss is a column that has grown past 32KB on average since, which +/// takes far longer than this. +const CLEAN_CACHE_TTL: Duration = Duration::from_secs(600); + +fn fresh(at: &Instant, cols: &[WideColumn]) -> bool { + at.elapsed() < if cols.is_empty() { CLEAN_CACHE_TTL } else { WIDE_CACHE_TTL } +} + fn cache_key(pool: &sqlx::PgPool, schema: &str, table: &str) -> String { let opts = pool.connect_options(); format!( @@ -136,7 +146,7 @@ pub async fn wide_columns( { if let Ok(map) = cache.lock() { if let Some((at, cols, _)) = map.get(&key) { - if at.elapsed() < WIDE_CACHE_TTL { + if fresh(at, cols) { return cols.clone(); } } @@ -231,7 +241,7 @@ pub async fn wide_columns( } if let Ok(mut map) = cache.lock() { - if map.len() > 256 { + if map.len() > 4096 { map.clear(); } map.insert(key, (Instant::now(), cols.clone(), None)); @@ -256,7 +266,7 @@ pub async fn page_projection( if let Some((at, cols, sql)) = map.get(&key) { // A cached entry with no SQL yet still has to build one; a // cached entry with no WIDE COLUMNS is already the answer. - if at.elapsed() < WIDE_CACHE_TTL && (sql.is_some() || cols.is_empty()) { + if fresh(at, cols) && (sql.is_some() || cols.is_empty()) { return (sql.clone(), cols.clone()); } } @@ -280,6 +290,79 @@ pub async fn page_projection( (sql, cols) } +/// Decide "nothing wide" for a whole schema in one catalog query, ahead of the +/// first open of any of its tables. +/// +/// The first page of a table otherwise waits on `wide_columns` before its rows +/// query can even be sent: two round trips in a row, which to a far host +/// (Neon, Supabase, Nile at 265-535ms each) is the slow first open. Run in the +/// background when the table list loads. A table with no candidate columns, or +/// whose candidates are all narrow with no TOAST to sample, gets the same empty +/// answer `wide_columns` would give; anything else is left for the per-table +/// path, which samples it properly. +pub async fn prefetch_schema(pool: &sqlx::PgPool, schema: &str, tables: &[String]) { + use sqlx::Row; + let Ok(rows) = sqlx::query( + r#" + SELECT c.relname::text, + MAX(COALESCE(s.avg_width, 0))::bigint, + MAX(COALESCE(pg_total_relation_size(c.reltoastrelid), 0))::bigint + FROM pg_catalog.pg_attribute a + JOIN pg_catalog.pg_class c ON c.oid = a.attrelid + JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace + JOIN pg_catalog.pg_type t ON t.oid = a.atttypid + LEFT JOIN pg_catalog.pg_stats s + ON s.schemaname = n.nspname AND s.tablename = c.relname AND s.attname = a.attname + WHERE n.nspname = $1 + AND a.attnum > 0 AND NOT a.attisdropped + AND t.typname = ANY($2) + GROUP BY c.relname + "#, + ) + .bind(schema) + .bind(WIDE_TYPES) + .fetch_all(pool) + .await + else { + return; + }; + let mut needs_look: std::collections::HashSet = std::collections::HashSet::new(); + for r in &rows { + let (Ok(name), Ok(avg), Ok(toast)) = + (r.try_get::(0), r.try_get::(1), r.try_get::(2)) + else { + continue; + }; + if avg > WIDE_COLUMN_AVG_BYTES || toast > TOAST_SAMPLE_FLOOR { + needs_look.insert(name); + } + } + let cache = WIDE_CACHE.get_or_init(|| std::sync::Mutex::new(HashMap::new())); + let Ok(mut map) = cache.lock() else { return }; + if map.len() + tables.len() > 4096 { + map.clear(); + } + let now = Instant::now(); + for table in tables.iter().filter(|t| !needs_look.contains(*t)) { + map.entry(cache_key(pool, schema, table)).or_insert_with(|| (now, Vec::new(), None)); + } +} + +/// First open of a table: drop a cached PROJECTION, which names columns and can +/// go stale when one is added, but keep a cached "nothing wide", which means +/// `SELECT *` and cannot. Keeping it is what lets `prefetch_schema` save the +/// first open its extra round trip. +pub fn invalidate_projection(pool: &sqlx::PgPool, schema: &str, table: &str) { + if let Some(cache) = WIDE_CACHE.get() { + if let Ok(mut map) = cache.lock() { + let key = cache_key(pool, schema, table); + if map.get(&key).is_some_and(|(_, cols, _)| !cols.is_empty()) { + map.remove(&key); + } + } + } +} + /// Forget the cached stats for a table - after an ANALYZE, or a schema change /// that could have changed what is wide. pub fn invalidate(pool: &sqlx::PgPool, schema: &str, table: &str) { diff --git a/src-tauri/src/docker.rs b/src-tauri/src/docker.rs index a34a7b4a..2b016635 100644 --- a/src-tauri/src/docker.rs +++ b/src-tauri/src/docker.rs @@ -494,6 +494,37 @@ fn database_from_inspect(c: &Value) -> Option { }) } +/// Start, stop or restart one container, from the right-click menu on its card. +/// +/// Only these three: they are reversible, and anything that deletes a container +/// or its volume belongs in Docker's own tools. The name or id is validated +/// against Docker's own character set so it can never be read as a flag +/// (`--rm`, `-f`) by the CLI. +#[tauri::command] +pub async fn docker_container_action(container: String, action: String) -> Result<(), String> { + if !matches!(action.as_str(), "start" | "stop" | "restart") { + return Err(format!("Unsupported Docker action: {action}")); + } + let valid = !container.is_empty() + && !container.starts_with('-') + && container.chars().all(|c| c.is_ascii_alphanumeric() || matches!(c, '_' | '.' | '-')); + if !valid { + return Err("That doesn't look like a container name or id.".into()); + } + let out = docker() + .arg(&action) + .arg(&container) + .output() + .await + .map_err(|e| format!("Could not run docker: {e}"))?; + if out.status.success() { + Ok(()) + } else { + let msg = String::from_utf8_lossy(&out.stderr).trim().to_string(); + Err(if msg.is_empty() { format!("docker {action} failed") } else { msg }) + } +} + /// Every database container running on this machine, with the credentials it was /// started with. Docker missing or not running is not an error - it just means /// there is nothing to offer, and the connection screen stays quiet about it. @@ -631,3 +662,24 @@ mod tests { assert!(database_from_inspect(&caddy).is_none()); } } + +#[cfg(test)] +mod action_tests { + use super::docker_container_action; + + /// Nothing that could be read as a docker flag, and nothing destructive. + #[tokio::test] + async fn rejects_flags_and_unknown_actions_before_running_docker() { + for bad in ["--rm", "-f", "a b", "x;rm -rf /", ""] { + assert!(docker_container_action(bad.into(), "stop".into()).await.is_err(), "{bad:?}"); + } + assert!(docker_container_action("stroke-test-mysql".into(), "rm".into()).await.unwrap_err().contains("Unsupported")); + } + + /// Against the dialect-matrix container: `cargo test --lib docker_live -- --ignored`. + #[tokio::test] + #[ignore] + async fn docker_live_restart_the_test_container() { + docker_container_action("stroke-test-mysql".into(), "restart".into()).await.unwrap(); + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 64835e31..a081bafc 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -330,6 +330,9 @@ pub fn run() { // paths that need them without a State/AppHandle argument: the D1 // driver refreshing an expired Cloudflare token mid-session. cloudflare::set_app_handle(app.handle().clone()); + if let Ok(dir) = app.path().app_data_dir() { + db::connection::set_data_dir(dir); + } db::connection::register_active_conn(std::sync::Arc::clone(&db_conn_for_setup)); let mut window_builder = tauri::WebviewWindowBuilder::new( @@ -506,6 +509,8 @@ pub fn run() { commands::connect_libsql_db, commands::test_clickhouse, commands::connect_clickhouse_db, + commands::test_posthog, + commands::connect_posthog_db, commands::test_redis, commands::connect_redis_db, commands::redis_scan, @@ -566,6 +571,7 @@ pub fn run() { docker::docker_check, docker::docker_run_db, docker::scan_docker_databases, + docker::docker_container_action, db::local_scan::scan_local_studios, db::local_scan::scan_machine_databases, app_lock::app_lock_status, @@ -591,6 +597,7 @@ pub fn run() { providers::provider_store_token, providers::provider_oauth_status, providers::provider_logout, + providers::provider_warm, providers::provider_list_databases, providers::provider_build_connection, db::backup::backup_export, diff --git a/src-tauri/src/mcp/tools.rs b/src-tauri/src/mcp/tools.rs index 64878693..6172bce1 100644 --- a/src-tauri/src/mcp/tools.rs +++ b/src-tauri/src/mcp/tools.rs @@ -446,6 +446,10 @@ async fn execute_sql( let result = crate::db::clickhouse::query(cfg, sql).await?; Ok(truncated_result_json(&result, max_rows)) } + ActiveConnection::Posthog(cfg) => { + let result = crate::db::posthog::query(cfg, sql).await?; + Ok(truncated_result_json(&result, max_rows)) + } ActiveConnection::Redis(cfg) => { let result = crate::db::redis::query(cfg, sql).await?; Ok(truncated_result_json(&result, max_rows)) @@ -575,6 +579,10 @@ async fn list_tables(conn: &ActiveConnection, schema: &str) -> Result>()}).to_string()) } + ActiveConnection::Posthog(cfg) => { + let tables = crate::db::posthog::list_tables(cfg).await?; + Ok(json!({"tables":tables.iter().map(|t|&t.name).collect::>()}).to_string()) + } ActiveConnection::Redis(cfg) => { let tables = crate::db::redis::list_tables(cfg).await?; Ok(json!({"tables":tables.iter().map(|t|&t.name).collect::>()}).to_string()) @@ -691,6 +699,11 @@ async fn describe_table( Ok(json!({"table":table,"pragma_info":r.rows}).to_string()) } ActiveConnection::Mysql(pool) => describe_table_mysql(pool, schema, table).await, + ActiveConnection::Posthog(cfg) => { + let cols = crate::db::posthog::get_column_structure(cfg, table).await?; + let columns: Vec<_> = cols.iter().map(|c| json!({ "name": c.name, "type": c.data_type, "nullable": c.is_nullable })).collect(); + Ok(json!({"table":table,"columns":columns,"note":"PostHog table, query it with HogQL"}).to_string()) + } ActiveConnection::Clickhouse(cfg) => { let cols = crate::db::clickhouse::get_column_structure(cfg, schema, table).await?; let columns: Vec<_> = cols.iter().map(|c| json!({ @@ -977,6 +990,7 @@ async fn check_migrations(conn: &ActiveConnection, schema: &str) -> Result Ok(json!({"migrations":[],"note":"Migration detection not yet supported for LibSQL"}).to_string()), ActiveConnection::Mysql(pool) => check_migrations_mysql(pool, schema).await, ActiveConnection::Clickhouse(_) => Ok(json!({"migrations":[],"note":"Migration detection not supported for ClickHouse"}).to_string()), + ActiveConnection::Posthog(_) => Ok(json!({"migrations":[],"note":"Migration detection not supported for PostHog"}).to_string()), ActiveConnection::Redis(_) => Ok(json!({"migrations":[],"note":"Migration detection not supported for Redis"}).to_string()), ActiveConnection::Duckdb(_) => Ok(json!({"migrations":[],"note":"Migration detection not supported for DuckDB"}).to_string()), ActiveConnection::Mssql(_) => Ok(json!({"migrations":[],"note":"Migration detection not supported for MS SQL Server"}).to_string()), @@ -1186,6 +1200,7 @@ async fn explain_query(conn: &ActiveConnection, sql: &str) -> Result Ok(json!({"plan":[],"database":"redis","note":"EXPLAIN is not supported on Redis"}).to_string()), + ActiveConnection::Posthog(_) => Ok(json!({"plan":[],"database":"posthog","note":"EXPLAIN is not available through PostHog's query API"}).to_string()), ActiveConnection::Duckdb(h) => { let r = crate::db::duckdb::execute_sql(h, &format!("EXPLAIN {sql}")).await?; Ok(json!({"plan":r.rows,"database":"duckdb"}).to_string()) @@ -1269,6 +1284,10 @@ async fn get_database_stats(conn: &ActiveConnection, schema: &str) -> Result { + let tables = crate::db::posthog::list_tables(cfg).await?; + Ok(json!({"database":"posthog","dialect":"HogQL","table_count":tables.len()}).to_string()) + } ActiveConnection::Redis(cfg) => { let tables = crate::db::redis::list_tables(cfg).await?; Ok(json!({"database":"redis","table_count":tables.len()}).to_string()) diff --git a/src-tauri/src/providers/mod.rs b/src-tauri/src/providers/mod.rs index 6c94f8fd..e612e165 100644 --- a/src-tauri/src/providers/mod.rs +++ b/src-tauri/src/providers/mod.rs @@ -15,6 +15,7 @@ mod neon; mod planetscale; +mod posthog; mod prisma; mod supabase; mod nile; @@ -66,6 +67,7 @@ pub enum Provider { Railway, Nile, Upstash, + PostHog, } /// How a provider signs the user in. @@ -93,10 +95,33 @@ impl Provider { "railway" => Ok(Self::Railway), "nile" => Ok(Self::Nile), "upstash" => Ok(Self::Upstash), + "posthog" => Ok(Self::PostHog), other => Err(format!("Unknown provider: {other}")), } } + const ALL: [Provider; 10] = [ + Self::Neon, Self::Supabase, Self::PlanetScale, Self::Prisma, Self::TiDB, + Self::Turso, Self::Railway, Self::Nile, Self::Upstash, Self::PostHog, + ]; + + /// The API host each listing and connect talks to, for `provider_warm`. + /// None for PostHog, whose host is the user's own instance. + fn api_origin(&self) -> Option<&'static str> { + Some(match self { + Self::Neon => "https://console.neon.tech", + Self::Supabase => "https://api.supabase.com", + Self::PlanetScale => "https://api.planetscale.com", + Self::Prisma => "https://api.prisma.io", + Self::TiDB => "https://serverless.tidbapi.com", + Self::Turso => "https://api.turso.tech", + Self::Railway => "https://backboard.railway.com", + Self::Nile => "https://global.thenile.dev", + Self::Upstash => "https://api.upstash.com", + Self::PostHog => return None, + }) + } + /// Stable key used to namespace stored tokens (`__{key}_refresh__`, …). fn key(&self) -> &'static str { match self { @@ -109,6 +134,7 @@ impl Provider { Self::Railway => "railway", Self::Nile => "nile", Self::Upstash => "upstash", + Self::PostHog => "posthog", } } @@ -124,6 +150,7 @@ impl Provider { Self::Railway => "Railway", Self::Nile => "Nile", Self::Upstash => "Upstash", + Self::PostHog => "PostHog", } } @@ -146,13 +173,14 @@ impl Provider { Self::Railway => railway::OAUTH, Self::Nile => nile::OAUTH, Self::Upstash => upstash::OAUTH, + Self::PostHog => posthog::OAUTH, } } /// Whether a provider uses a pasted API credential instead of the browser /// OAuth dance: Upstash, which offers no OAuth to third-party apps. fn is_token_based(&self) -> bool { - matches!(self, Self::Upstash) + matches!(self, Self::Upstash | Self::PostHog) } /// Localhost callback ports to try, in order. PlanetScale accepts only ONE @@ -162,7 +190,7 @@ impl Provider { fn callback_ports(&self) -> &'static [u16] { match self { // Railway, like PlanetScale, matches the redirect URI exactly and - // the app registers one: http://localhost:8989/oauth/callback. + // the app registers one: http://127.0.0.1:8989/oauth/callback. Self::PlanetScale | Self::Railway => &[8989], _ => CALLBACK_PORTS, } @@ -191,6 +219,10 @@ impl Provider { Self::Neon => format!("http://127.0.0.1:{port}/callback"), // nilecli's client allows any localhost port on /callback. Self::Nile => format!("http://localhost:{port}/callback"), + // Stroke's Railway app is registered with the loopback IP, and + // Railway matches the redirect exactly: `localhost` here was + // rejected with invalid_redirect_uri. + Self::Railway => format!("http://127.0.0.1:{port}/oauth/callback"), _ => format!("http://localhost:{port}/oauth/callback"), } } @@ -206,6 +238,7 @@ impl Provider { Self::Railway => railway::list_databases(token).await, Self::Nile => nile::list_databases(token).await, Self::Upstash => upstash::list_databases(token).await, + Self::PostHog => posthog::list_databases(token).await, } } @@ -224,6 +257,7 @@ impl Provider { Self::Railway => railway::build_connection(token, db_ref).await, Self::Nile => nile::build_connection(token, db_ref).await, Self::Upstash => upstash::build_connection(token, db_ref).await, + Self::PostHog => posthog::build_connection(token, db_ref).await, } } } @@ -364,10 +398,38 @@ fn now_secs() -> u64 { // ── Local callback server ──────────────────────────────────────────────────────── -async fn bind_callback_listener(ports: &[u16]) -> Result<(TcpListener, u16), String> { +/// The local end of an OAuth redirect, on both loopback addresses. +/// +/// A redirect to `http://localhost:…` is resolved by the browser, and some +/// setups (IPv6-first resolvers, a proxy, some Linux configs) send it to `::1` +/// before `127.0.0.1`. Listening on IPv4 alone left those sign-ins waiting on a +/// socket nothing connected to until the 5-minute timeout. The IPv6 listener is +/// best effort: a machine without IPv6 loopback just doesn't get one. +pub(crate) struct CallbackListener { + v4: TcpListener, + v6: Option, +} + +impl CallbackListener { + async fn accept(&self) -> std::io::Result<(tokio::net::TcpStream, std::net::SocketAddr)> { + match &self.v6 { + Some(v6) => tokio::select! { + r = self.v4.accept() => r, + r = v6.accept() => r, + }, + None => self.v4.accept().await, + } + } +} + +pub(crate) async fn bind_callback_listener(ports: &[u16]) -> Result<(CallbackListener, u16), String> { for &port in ports { - if let Ok(listener) = TcpListener::bind(format!("127.0.0.1:{port}")).await { - return Ok((listener, port)); + if let Ok(v4) = TcpListener::bind(format!("127.0.0.1:{port}")).await { + // The port IPv4 actually got (differs from `port` only when it is 0), + // so both addresses answer on the same one. + let port = v4.local_addr().map(|a| a.port()).unwrap_or(port); + let v6 = TcpListener::bind(format!("[::1]:{port}")).await.ok(); + return Ok((CallbackListener { v4, v6 }, port)); } } if ports.len() == 1 { @@ -388,8 +450,8 @@ async fn bind_callback_listener(ports: &[u16]) -> Result<(TcpListener, u16), Str /// Wait for one OAuth redirect and return the value of `value_key` from its /// query: the authorization code (`code`), or for a token redirect the token /// itself (`jwt`). -async fn await_oauth_callback( - listener: TcpListener, +pub(crate) async fn await_oauth_callback( + listener: CallbackListener, expected_state: &str, value_key: &str, provider_label: &str, @@ -861,9 +923,7 @@ pub async fn provider_start_oauth( ); eprintln!("[provider oauth] {} authorize URL: {auth_url}", p.key()); - tauri_plugin_opener::OpenerExt::opener(&app) - .open_url(auth_url, None::<&str>) - .map_err(|e| format!("Could not open browser: {e}"))?; + open_sign_in_page(&app, &auth_url); // Register the cancel waiter BEFORE awaiting so a Cancel click can't slip // through between opening the browser and starting to wait. @@ -897,6 +957,20 @@ pub async fn provider_start_oauth( }) } +/// Hand the sign-in page to the UI, then try to open it in the browser. +/// +/// The UI gets the URL first so the "Waiting for…" panel can offer "Open again" +/// and "Copy link". A failure to open the browser is no longer fatal: with no +/// default browser set (common on Linux), or a closed tab, the sign-in used to +/// die with "Could not open browser" or wait five minutes with no way back. Now +/// the flow keeps waiting and the link is on screen. +pub(crate) fn open_sign_in_page(app: &tauri::AppHandle, url: &str) { + let _ = tauri::Emitter::emit(app, "provider-auth-url", serde_json::json!({ "url": url })); + if let Err(e) = tauri_plugin_opener::OpenerExt::opener(app).open_url(url, None::<&str>) { + eprintln!("[provider oauth] could not open the browser: {e}"); + } +} + /// OAuth 2.0 device code grant (RFC 8628). Opens the verification page with the /// code prefilled, then polls the token endpoint until the user approves it, /// declines it, the code expires, or they hit Cancel. @@ -938,9 +1012,11 @@ async fn device_code_sign_in(app: &tauri::AppHandle, cfg: &OAuthConfig) -> Resul "expiresIn": expires, }), ); - tauri_plugin_opener::OpenerExt::opener(app) - .open_url(verify_url, None::<&str>) - .map_err(|e| format!("Could not open browser: {e}"))?; + // The code and a reopen button are already on screen, so a browser that + // won't open is not a reason to abandon the sign-in. + if let Err(e) = tauri_plugin_opener::OpenerExt::opener(app).open_url(verify_url, None::<&str>) { + eprintln!("[provider oauth] could not open the browser: {e}"); + } let started = std::time::Instant::now(); let cancelled = oauth_cancel().notified(); @@ -996,9 +1072,7 @@ async fn token_redirect_sign_in( cfg.auth_url.trim_end_matches('/'), urlencoding::encode(state), ); - tauri_plugin_opener::OpenerExt::opener(app) - .open_url(url, None::<&str>) - .map_err(|e| format!("Could not open browser: {e}"))?; + open_sign_in_page(app, &url); let cancelled = oauth_cancel().notified(); tokio::select! { r = tokio::time::timeout( @@ -1043,6 +1117,28 @@ pub async fn provider_oauth_status( }) } +/// Open the HTTPS connection to every signed-in provider's API ahead of use. +/// +/// The first call to an API pays DNS, TCP and TLS before the request itself: +/// three or four round trips, which to PlanetScale's or Neon's API from a far +/// region is most of a second on top of the request. Called when the connect +/// dialog opens; the shared client keeps the socket for the listing and the +/// connect that follow. Unauthenticated `HEAD /`: no token leaves the app, and +/// whatever it answers is thrown away. +#[tauri::command] +pub async fn provider_warm(app: tauri::AppHandle) { + let map = crate::secrets::read_all_async(&app).await; + for p in Provider::ALL { + let Some(origin) = p.api_origin() else { continue }; + if !map.contains_key(&format!("__{}_access__", p.key())) { + continue; + } + tokio::spawn(async move { + let _ = http().head(origin).timeout(std::time::Duration::from_secs(5)).send().await; + }); + } +} + #[tauri::command] pub async fn provider_logout(app: tauri::AppHandle, provider: String) -> Result<(), String> { clear_tokens(&app, Provider::parse(&provider)?).await @@ -1054,7 +1150,12 @@ pub async fn provider_list_databases( provider: String, ) -> Result, String> { let p = Provider::parse(&provider)?; - with_token(&app, p, |token| async move { p.list_databases(&token).await }).await + let t0 = std::time::Instant::now(); + let r = with_token(&app, p, |token| async move { p.list_databases(&token).await }).await; + // "The provider panel is slow" is otherwise unattributable: this says which + // provider, and whether it was the listing or the connect that took the time. + log::info!("{} list_databases: {} in {}ms", p.key(), if r.is_ok() { "ok" } else { "failed" }, t0.elapsed().as_millis()); + r } #[tauri::command] @@ -1064,11 +1165,14 @@ pub async fn provider_build_connection( db_ref: String, ) -> Result { let p = Provider::parse(&provider)?; - with_token(&app, p, |token| { + let t0 = std::time::Instant::now(); + let r = with_token(&app, p, |token| { let db_ref = db_ref.clone(); async move { p.build_connection(&token, &db_ref).await } }) - .await + .await; + log::info!("{} build_connection: {} in {}ms", p.key(), if r.is_ok() { "ok" } else { "failed" }, t0.elapsed().as_millis()); + r } #[cfg(test)] @@ -1084,12 +1188,26 @@ mod callback_tests { out } + /// A redirect that the browser sends to `::1` is answered too. + #[tokio::test] + async fn the_callback_answers_on_ipv6_loopback() { + let Ok((listener, port)) = super::bind_callback_listener(&[0]).await else { return }; + if listener.v6.is_none() { + return; // no IPv6 loopback on this machine + } + let waiter = tokio::spawn(async move { await_oauth_callback(listener, "s6", "code", "Neon").await }); + let mut s = tokio::net::TcpStream::connect(("::1", port)).await.unwrap(); + s.write_all(b"GET /oauth/callback?code=v6&state=s6 HTTP/1.1\r\nHost: localhost\r\n\r\n").await.unwrap(); + assert_eq!(waiter.await.unwrap().unwrap(), "v6"); + } + /// A stray request (favicon, preconnect) used to spend the only accept and /// fail the sign-in. It must be answered 404 and the wait must go on. #[tokio::test] async fn stray_requests_do_not_consume_the_callback() { let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let port = listener.local_addr().unwrap().port(); + let listener = super::CallbackListener { v4: listener, v6: None }; let waiter = tokio::spawn(async move { await_oauth_callback(listener, "s1", "code", "Neon").await }); // A socket that connects and never sends, then a favicon fetch. let _silent = tokio::net::TcpStream::connect(("127.0.0.1", port)).await.unwrap(); diff --git a/src-tauri/src/providers/planetscale.rs b/src-tauri/src/providers/planetscale.rs index d7cc8248..dccc9e7b 100644 --- a/src-tauri/src/providers/planetscale.rs +++ b/src-tauri/src/providers/planetscale.rs @@ -63,28 +63,85 @@ async fn get(token: &str, path: &str) -> Result { Ok(body) } -/// db_ref encodes "{org}/{database}" so build_connection can act without a -/// second lookup. -pub async fn list_databases(token: &str) -> Result, String> { +/// Organizations the last listing saw, per token, so the next listing can ask +/// for every org's databases at the same moment it re-checks the org list: +/// one round trip in the common case instead of two in a row (3.1s measured). +static ORGS: std::sync::Mutex)>> = std::sync::Mutex::new(None); + +/// Default branch per "{org}/{database}", from the listing, so connecting +/// skips the database lookup and goes straight to minting the password. +static BRANCHES: std::sync::OnceLock>> = + std::sync::OnceLock::new(); + +fn branches() -> &'static std::sync::Mutex> { + BRANCHES.get_or_init(Default::default) +} + +async fn org_names(token: &str) -> Result, String> { let orgs = get(token, "/organizations").await?; - // Every org's database list in flight at once rather than one after another. - let org_names: Vec<&str> = orgs["data"] + Ok(orgs["data"] .as_array() .into_iter() .flatten() - .filter_map(|org| org["name"].as_str()) - .collect(); - let pages = futures::future::join_all(org_names.iter().map(|org_name| { + .filter_map(|org| org["name"].as_str().map(String::from)) + .collect()) +} + +/// Every org's database list in flight at once rather than one after another. +async fn org_pages(token: &str, orgs: &[String]) -> Vec> { + futures::future::join_all(orgs.iter().map(|org_name| { let path = format!("/organizations/{org_name}/databases"); async move { get(token, &path).await } })) - .await; + .await +} + +/// db_ref encodes "{org}/{database}" so build_connection can act without a +/// second lookup. +pub async fn list_databases(token: &str) -> Result, String> { + let known = ORGS + .lock() + .ok() + .and_then(|g| g.as_ref().filter(|(t, _)| t == token).map(|(_, o)| o.clone())); + let (org_list, mut pages) = match known { + // Re-check the org list and fetch the known orgs' databases together. + Some(known) => { + let (fresh, pages) = tokio::join!(org_names(token), org_pages(token, &known)); + let fresh = fresh?; + let mut by_org: std::collections::HashMap> = + known.into_iter().zip(pages).collect(); + // An org that appeared since: fetch it now (rare). + let added: Vec = fresh.iter().filter(|o| !by_org.contains_key(*o)).cloned().collect(); + for (org, page) in added.iter().cloned().zip(org_pages(token, &added).await) { + by_org.insert(org, page); + } + let pages = fresh.iter().map(|o| by_org.remove(o).unwrap_or_else(|| Err("missing".into()))).collect(); + (fresh, pages) + } + None => { + let orgs = org_names(token).await?; + let pages = org_pages(token, &orgs).await; + (orgs, pages) + } + }; + if let Ok(mut g) = ORGS.lock() { + *g = Some((token.to_string(), org_list.clone())); + } // `/organizations` lists every org the user belongs to, but the token only // covers the ones picked on the consent screen: the rest answer 403. let pages = super::merge_partial( - org_names.iter().zip(pages).map(|(org, r)| r.map(|b| (*org, b))).collect(), + org_list.iter().map(String::as_str).zip(pages.drain(..)).map(|(org, r)| r.map(|b| (org, b))).collect(), "Stroke may not have been granted this organization: sign out, sign in again, and select it on PlanetScale's consent screen.", )?; + if let Ok(mut map) = branches().lock() { + for (org, dbs) in &pages { + for db in dbs["data"].as_array().into_iter().flatten() { + if let (Some(name), Some(branch)) = (db["name"].as_str(), db["default_branch"].as_str()) { + map.insert(format!("{org}/{name}"), branch.to_string()); + } + } + } + } Ok(parse_databases(&pages)) } @@ -112,9 +169,16 @@ pub async fn build_connection(token: &str, db_ref: &str) -> Result b, + None => { + let db = get(token, &format!("/organizations/{org}/databases/{database}")).await?; + db["default_branch"].as_str().unwrap_or("main").to_string() + } + }; let resp = http() .post(format!( diff --git a/src-tauri/src/providers/posthog.rs b/src-tauri/src/providers/posthog.rs new file mode 100644 index 00000000..7556d061 --- /dev/null +++ b/src-tauri/src/providers/posthog.rs @@ -0,0 +1,139 @@ +/*! + * PostHog adapter - the provider-picker side of the PostHog engine + * (`db/posthog.rs`). Paste a personal API key and the instance URL, pick a + * project, and connect: no connection string. + * + * Token-based, like Upstash: the stored credential is `{base_url}|{api_key}`. + * A URL never contains `|` and PostHog keys (`phx_…`) don't either, so the first + * one is the split. The key needs the Project Read scope to list projects and + * Query Read to run HogQL. + */ + +use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; +use serde_json::{json, Value}; + +/// Unused for a token provider; present so `Provider::oauth` stays total. +pub const OAUTH: OAuthConfig = OAuthConfig { client_id: "", auth_url: "", token_url: "", scopes: "" }; + +/// Cap on project list pages. +const MAX_PAGES: usize = 20; + +fn credentials(token: &str) -> Result<(String, &str), String> { + let (host, key) = token + .split_once('|') + .filter(|(h, k)| !h.trim().is_empty() && !k.trim().is_empty()) + .ok_or("PostHog needs its URL and a personal API key.")?; + let host = host.trim().trim_end_matches('/'); + let base = if host.starts_with("http://") || host.starts_with("https://") { + host.to_string() + } else { + format!("https://{host}") + }; + Ok((base, key.trim())) +} + +async fn get(token: &str, url_or_path: &str) -> Result { + let (base, key) = credentials(token)?; + let url = if url_or_path.starts_with("http") { url_or_path.to_string() } else { format!("{base}{url_or_path}") }; + let resp = http() + .get(&url) + .bearer_auth(key) + .send() + .await + .map_err(|e| format!("PostHog request failed: {}", super::describe(&e)))?; + let status = resp.status().as_u16(); + // A wrong key is a 401; twice in a row the command layer ends the "session" + // and the key form comes back. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp.text().await.map_err(|e| format!("PostHog read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).unwrap_or(Value::Null); + if !(200..300).contains(&status) { + let detail = body["detail"].as_str().map(String::from).unwrap_or_else(|| text.chars().take(200).collect()); + return Err(if status == 403 { + format!("PostHog refused the key: {detail}. It needs the Project Read and Query Read scopes.") + } else { + format!("PostHog API error ({status}): {detail}") + }); + } + Ok(body) +} + +/// One `/api/projects/` page → picker rows. `db_ref` carries the id and name so +/// connecting needs no second lookup. +fn parse_projects(body: &Value) -> Vec { + body["results"] + .as_array() + .into_iter() + .flatten() + .filter_map(|p| { + let id = p["id"].as_i64().map(|n| n.to_string()).or_else(|| p["id"].as_str().map(String::from))?; + let name = p["name"].as_str().unwrap_or(&id).to_string(); + Some(ProviderDatabase { + db_ref: json!({ "id": id, "n": name }).to_string(), + name, + region: None, + kind: Some("Project".into()), + host: None, + }) + }) + .collect() +} + +pub async fn list_databases(token: &str) -> Result, String> { + let mut out = Vec::new(); + let mut next = "/api/projects/".to_string(); + for _ in 0..MAX_PAGES { + let body = get(token, &next).await?; + out.extend(parse_projects(&body)); + match body["next"].as_str() { + Some(n) if !n.is_empty() => next = n.to_string(), + _ => break, + } + } + Ok(out) +} + +pub async fn build_connection(token: &str, db_ref: &str) -> Result { + let r: Value = serde_json::from_str(db_ref).map_err(|_| "Invalid PostHog project reference")?; + let id = r["id"].as_str().ok_or("Invalid PostHog project reference")?; + let name = r["n"].as_str().unwrap_or(id); + let (base, key) = credentials(token)?; + // The engine's own adapter contract: the base URL in `host`, the project id + // in `database`, the key in `password`. The frontend maps these to a + // `posthog` connection ({ host, projectId, apiKey }). + Ok(ProviderConnection { + db_type: "posthog".into(), + host: base, + port: 443, + username: String::new(), + password: key.to_string(), + database: id.to_string(), + ssl: true, + needs_password: false, + name: format!("PostHog · {name}"), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn credential_splits_url_and_key() { + assert_eq!(credentials("https://eu.posthog.com/|phx_abc").unwrap(), ("https://eu.posthog.com".into(), "phx_abc")); + assert_eq!(credentials("ph.example.com|phx_abc").unwrap().0, "https://ph.example.com"); + assert!(credentials("|phx_abc").is_err()); + assert!(credentials("https://us.posthog.com").is_err()); + } + + #[test] + fn projects_carry_id_and_name() { + let body = json!({ "results": [{ "id": 42, "name": "Stroke" }, { "name": "no id" }], "next": null }); + let rows = parse_projects(&body); + assert_eq!(rows.len(), 1); + let r: Value = serde_json::from_str(&rows[0].db_ref).unwrap(); + assert_eq!((r["id"].as_str(), r["n"].as_str()), (Some("42"), Some("Stroke"))); + } +} diff --git a/src-tauri/src/providers/railway.rs b/src-tauri/src/providers/railway.rs index 802a725b..a60c2fcb 100644 --- a/src-tauri/src/providers/railway.rs +++ b/src-tauri/src/providers/railway.rs @@ -15,10 +15,11 @@ use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; use serde_json::{json, Value}; pub const OAUTH: OAuthConfig = OAuthConfig { - // The native OAuth app registered in Railway (workspace Developer settings), - // redirect URI http://localhost:8989/oauth/callback. Empty until it exists: - // sign-in then says so instead of opening a page Railway would reject. - client_id: "", + // The native (public) OAuth app registered in Railway's workspace Developer + // settings, redirect URI http://127.0.0.1:8989/oauth/callback. Public means + // no secret: Railway rejects a token request that sends one (checked: the + // same code with a secret gets `invalid_client`, without one `invalid_grant`). + client_id: "rlwy_oaci_D7dCmSYWNMbBl3ZDVZc6nd7A", auth_url: "https://backboard.railway.com/oauth/auth", token_url: "https://backboard.railway.com/oauth/token", // workspace:viewer lists workspaces; project:member is what lets the token @@ -149,49 +150,128 @@ pub async fn build_connection(token: &str, db_ref: &str) -> Result = data["variables"].as_object().map(|m| m.keys().map(String::as_str).collect()).unwrap_or_default(); + log::info!( + "railway connect {name}: env={} service={} variables={var_names:?} tcpProxies={}", + field("e"), + field("s"), + data["tcpProxies"] + ); + connection_from_vars(&engine, &name, &data["variables"], &data["tcpProxies"]).map_err(|e| { + // Point at the one screen that fixes it. Railway's public API has no + // mutation to create a TCP proxy, so this can't be done from here. + if e.contains("no public access") { + format!( + "{e} https://railway.com/project/{}/service/{}/settings?environmentId={}", + field("p"), + field("s"), + field("e") + ) + } else { + e + } + }) } -/// A service's variables → a connection, through its PUBLIC url: the plain one -/// points at `*.railway.internal`, which only resolves inside Railway's network. -fn connection_from_vars(engine: &str, name: &str, vars: &Value) -> Result { - let keys: &[&str] = match engine { +/// A service's variables → a connection, through its PUBLIC address: the plain +/// URL points at `*.railway.internal`, which only resolves inside Railway's +/// network. +/// +/// Preferably the template's `*_PUBLIC_URL`. When a service has a TCP proxy but +/// no such variable (custom images, older templates), the proxy's domain and +/// port are combined with the credentials the image reads from its variables. +fn connection_from_vars(engine: &str, name: &str, vars: &Value, proxies: &Value) -> Result { + let default_port = match engine { "postgres" => 5432, "mysql" => 3306, _ => 6379 }; + let named = |parts: (String, u16, String, String, String)| { + let (host, port, username, password, database) = parts; + ProviderConnection { + db_type: engine.to_string(), + host, + port, + username, + password, + database, + // Railway's TCP proxy is plain TCP; the Postgres template serves TLS + // itself (postgres-ssl) but does not require it, and MySQL/Redis there + // are not TLS. + ssl: false, + needs_password: false, + name: format!("Railway · {name}"), + } + }; + + let url_keys: &[&str] = match engine { "postgres" => &["DATABASE_PUBLIC_URL"], "mysql" => &["MYSQL_PUBLIC_URL", "DATABASE_PUBLIC_URL"], _ => &["REDIS_PUBLIC_URL"], }; - let raw = keys + // Right after public access is added, the template's URL can still be + // half-rendered (`mysql://root:…@:/railway`): the proxy domain it references + // only fills in on the next deploy. A URL without a host isn't an error to + // report; fall through to the proxy itself. + let public_url = url_keys .iter() - .find_map(|k| vars[*k].as_str().filter(|v| !v.is_empty())) - .ok_or_else(|| { - format!( - "{name} has no public URL. Turn on its TCP proxy in the service's Railway settings (Networking), then try again." - ) - })?; - let url = reqwest::Url::parse(raw).map_err(|e| format!("Railway returned an unreadable URL for {name}: {e}"))?; - let decode = |s: &str| urlencoding::decode(s).map(|c| c.into_owned()).unwrap_or_else(|_| s.to_string()); + .filter_map(|k| vars[*k].as_str().filter(|v| !v.is_empty())) + .find_map(|raw| reqwest::Url::parse(raw).ok().filter(|u| u.host_str().is_some_and(|h| !h.is_empty()))); + if let Some(url) = public_url { + let decode = |s: &str| urlencoding::decode(s).map(|c| c.into_owned()).unwrap_or_else(|_| s.to_string()); + return Ok(named(( + url.host_str().unwrap_or_default().to_string(), + url.port().unwrap_or(default_port), + decode(url.username()), + decode(url.password().unwrap_or_default()), + decode(url.path().trim_start_matches('/')), + ))); + } - let default_port = match engine { "postgres" => 5432, "mysql" => 3306, _ => 6379 }; - Ok(ProviderConnection { - db_type: engine.to_string(), - host: url.host_str().ok_or_else(|| format!("No host in {name}'s URL"))?.to_string(), - port: url.port().unwrap_or(default_port), - username: decode(url.username()), - password: decode(url.password().unwrap_or_default()), - database: decode(url.path().trim_start_matches('/')), - // Railway's TCP proxy is plain TCP; the Postgres template serves TLS - // itself (postgres-ssl) but does not require it, and MySQL/Redis there - // are not TLS. - ssl: false, - needs_password: false, - name: format!("Railway · {name}"), - }) + let proxy = proxies.as_array().and_then(|list| { + list.iter() + .find(|p| p["applicationPort"].as_u64() == Some(u64::from(default_port))) + .or_else(|| list.first()) + }); + if let Some(p) = proxy { + let host = p["domain"].as_str().filter(|d| !d.is_empty()); + let port = p["proxyPort"].as_u64().and_then(|n| u16::try_from(n).ok()); + if let (Some(host), Some(port)) = (host, port) { + let var = |keys: &[&str]| keys.iter().find_map(|k| vars[*k].as_str().filter(|v| !v.is_empty())).unwrap_or("").to_string(); + let (user, pass, db) = match engine { + "postgres" => ( + var(&["PGUSER", "POSTGRES_USER"]), + var(&["PGPASSWORD", "POSTGRES_PASSWORD"]), + var(&["PGDATABASE", "POSTGRES_DB"]), + ), + "mysql" => ( + var(&["MYSQLUSER", "MYSQL_USER"]), + var(&["MYSQLPASSWORD", "MYSQL_ROOT_PASSWORD", "MYSQL_PASSWORD"]), + var(&["MYSQLDATABASE", "MYSQL_DATABASE"]), + ), + _ => (var(&["REDISUSER"]), var(&["REDISPASSWORD", "REDIS_PASSWORD"]), String::new()), + }; + let user = if !user.is_empty() { user } else { match engine { "postgres" => "postgres", "mysql" => "root", _ => "default" }.into() }; + let db = if !db.is_empty() || engine == "redis" { db } else { match engine { "postgres" => "postgres", _ => "railway" }.into() }; + return Ok(named((host.to_string(), port, user, pass, db))); + } + // The proxy exists but Railway hasn't given it a public domain yet. + return Err(format!( + "{name}'s public access is still being set up: Railway hasn't assigned it an address yet. Give it a minute, then try again." + )); + } + + Err(format!( + "{name} has no public access. Choose Add Public Access under Settings → Networking in Railway, or keep it private with `railway connect --tunnel-only`." + )) } #[cfg(test)] @@ -225,13 +305,31 @@ mod tests { "DATABASE_URL": "postgresql://postgres:x@postgres.railway.internal:5432/railway", "DATABASE_PUBLIC_URL": "postgresql://postgres:p%40ss@shortline.proxy.rlwy.net:41234/railway" }); - let c = connection_from_vars("postgres", "shop / Postgres", &vars).unwrap(); + let c = connection_from_vars("postgres", "shop / Postgres", &vars, &json!([])).unwrap(); assert_eq!((c.host.as_str(), c.port), ("shortline.proxy.rlwy.net", 41234)); assert_eq!((c.username.as_str(), c.password.as_str(), c.database.as_str()), ("postgres", "p@ss", "railway")); let redis = json!({ "REDIS_PUBLIC_URL": "redis://default:pw@x.proxy.rlwy.net:6380" }); - assert_eq!(connection_from_vars("redis", "cache", &redis).unwrap().port, 6380); + assert_eq!(connection_from_vars("redis", "cache", &redis, &json!([])).unwrap().port, 6380); let private_only = json!({ "DATABASE_URL": "postgresql://u:p@postgres.railway.internal:5432/db" }); - assert!(connection_from_vars("postgres", "x", &private_only).unwrap_err().contains("TCP proxy")); + assert!(connection_from_vars("postgres", "x", &private_only, &json!([])).unwrap_err().contains("no public access")); + } + + #[test] + fn a_half_rendered_url_or_a_proxy_without_a_domain_says_not_ready() { + let vars = json!({ "MYSQL_PUBLIC_URL": "mysql://root:pw@:/railway", "MYSQL_ROOT_PASSWORD": "pw" }); + let pending = json!([{ "domain": "", "proxyPort": 0, "applicationPort": 3306 }]); + assert!(connection_from_vars("mysql", "m", &vars, &pending).unwrap_err().contains("still being set up")); + let ready = json!([{ "domain": "x.proxy.rlwy.net", "proxyPort": 4444, "applicationPort": 3306 }]); + assert_eq!(connection_from_vars("mysql", "m", &vars, &ready).unwrap().port, 4444); + } + + #[test] + fn a_tcp_proxy_without_a_public_url_variable_still_connects() { + let vars = json!({ "MYSQLUSER": "root", "MYSQL_ROOT_PASSWORD": "pw", "MYSQL_DATABASE": "railway" }); + let proxies = json!([{ "domain": "nozomi.proxy.rlwy.net", "proxyPort": 21337, "applicationPort": 3306 }]); + let c = connection_from_vars("mysql", "luminous-flexibility / MySQL", &vars, &proxies).unwrap(); + assert_eq!((c.host.as_str(), c.port), ("nozomi.proxy.rlwy.net", 21337)); + assert_eq!((c.username.as_str(), c.password.as_str(), c.database.as_str()), ("root", "pw", "railway")); } #[test] diff --git a/src-tauri/vendor/sqlx-mysql/Cargo.toml b/src-tauri/vendor/sqlx-mysql/Cargo.toml new file mode 100644 index 00000000..e0682a15 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/Cargo.toml @@ -0,0 +1,241 @@ +# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO +# +# When uploading crates to the registry Cargo will automatically +# "normalize" Cargo.toml files for maximal compatibility +# with all versions of Cargo and also rewrite `path` dependencies +# to registry (e.g., crates.io) dependencies. +# +# If you are reading this file be aware that the original Cargo.toml +# will likely look very different (and much more reasonable). +# See Cargo.toml.orig for the original contents. + +[package] +edition = "2021" +name = "sqlx-mysql" +version = "0.8.6" +authors = [ + "Ryan Leckey ", + "Austin Bonander ", + "Chloe Ross ", + "Daniel Akhterov ", +] +description = "MySQL driver implementation for SQLx. Not for direct use; see the `sqlx` crate for details." +documentation = "https://docs.rs/sqlx" +license = "MIT OR Apache-2.0" +repository = "https://github.com/launchbadge/sqlx" + +[dependencies.atoi] +version = "2.0" + +[dependencies.base64] +version = "0.22.0" +features = ["std"] +default-features = false + +[dependencies.bigdecimal] +version = "0.4.0" +optional = true + +[dependencies.bitflags] +version = "2" +features = ["serde"] +default-features = false + +[dependencies.byteorder] +version = "1.4.3" +features = ["std"] +default-features = false + +[dependencies.bytes] +version = "1.1.0" + +[dependencies.chrono] +version = "0.4.34" +features = [ + "std", + "clock", +] +optional = true +default-features = false + +[dependencies.crc] +version = "3.0.0" + +[dependencies.digest] +version = "0.10.0" +features = ["std"] +default-features = false + +[dependencies.dotenvy] +version = "0.15.5" + +[dependencies.either] +version = "1.6.1" + +[dependencies.futures-channel] +version = "0.3.19" +features = [ + "sink", + "alloc", + "std", +] +default-features = false + +[dependencies.futures-core] +version = "0.3.19" +default-features = false + +[dependencies.futures-io] +version = "0.3.24" + +[dependencies.futures-util] +version = "0.3.19" +features = [ + "alloc", + "sink", + "io", +] +default-features = false + +[dependencies.generic-array] +version = "0.14.4" +default-features = false + +[dependencies.hex] +version = "0.4.3" + +[dependencies.hkdf] +version = "0.12.0" + +[dependencies.hmac] +version = "0.12.0" +default-features = false + +[dependencies.itoa] +version = "1.0.1" + +[dependencies.log] +version = "0.4.18" + +[dependencies.md-5] +version = "0.10.0" +default-features = false + +[dependencies.memchr] +version = "2.4.1" +default-features = false + +[dependencies.once_cell] +version = "1.9.0" + +[dependencies.percent-encoding] +version = "2.1.0" + +[dependencies.rand] +version = "0.8.4" +features = [ + "std", + "std_rng", +] +default-features = false + +[dependencies.rsa] +version = "0.9" + +[dependencies.rust_decimal] +version = "1.26.1" +features = ["std"] +optional = true +default-features = false + +[dependencies.serde] +version = "1.0.144" +optional = true + +[dependencies.sha1] +version = "0.10.1" +default-features = false + +[dependencies.sha2] +version = "0.10.0" +default-features = false + +[dependencies.smallvec] +version = "1.7.0" + +[dependencies.sqlx-core] +version = "=0.8.6" + +[dependencies.stringprep] +version = "0.1.2" + +[dependencies.thiserror] +version = "2.0.0" + +[dependencies.time] +version = "0.3.36" +features = [ + "formatting", + "parsing", + "macros", +] +optional = true + +[dependencies.tracing] +version = "0.1.37" +features = ["log"] + +[dependencies.uuid] +version = "1.1.2" +optional = true + +[dependencies.whoami] +version = "1.2.1" +default-features = false + +[dev-dependencies.sqlx] +version = "=0.8.6" +features = ["mysql"] +default-features = false + +[features] +any = ["sqlx-core/any"] +bigdecimal = [ + "dep:bigdecimal", + "sqlx-core/bigdecimal", +] +chrono = [ + "dep:chrono", + "sqlx-core/chrono", +] +json = [ + "sqlx-core/json", + "serde", +] +migrate = ["sqlx-core/migrate"] +offline = [ + "sqlx-core/offline", + "serde/derive", +] +rust_decimal = [ + "dep:rust_decimal", + "rust_decimal/maths", + "sqlx-core/rust_decimal", +] +time = [ + "dep:time", + "sqlx-core/time", +] +uuid = [ + "dep:uuid", + "sqlx-core/uuid", +] + +[lints.clippy] +cast_possible_truncation = "deny" +cast_possible_wrap = "deny" +cast_sign_loss = "deny" +disallowed_methods = "deny" + +[lints.rust] +warnings = "allow" diff --git a/src-tauri/vendor/sqlx-mysql/Cargo.toml.orig b/src-tauri/vendor/sqlx-mysql/Cargo.toml.orig new file mode 100644 index 00000000..3971c2ff --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/Cargo.toml.orig @@ -0,0 +1,79 @@ +[package] +name = "sqlx-mysql" +documentation = "https://docs.rs/sqlx" +description = "MySQL driver implementation for SQLx. Not for direct use; see the `sqlx` crate for details." +version.workspace = true +license.workspace = true +edition.workspace = true +authors.workspace = true +repository.workspace = true +# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html + +[features] +json = ["sqlx-core/json", "serde"] +any = ["sqlx-core/any"] +offline = ["sqlx-core/offline", "serde/derive"] +migrate = ["sqlx-core/migrate"] + +# Type Integration features +bigdecimal = ["dep:bigdecimal", "sqlx-core/bigdecimal"] +chrono = ["dep:chrono", "sqlx-core/chrono"] +rust_decimal = ["dep:rust_decimal", "rust_decimal/maths", "sqlx-core/rust_decimal"] +time = ["dep:time", "sqlx-core/time"] +uuid = ["dep:uuid", "sqlx-core/uuid"] + +[dependencies] +sqlx-core = { workspace = true } + +# Futures crates +futures-channel = { version = "0.3.19", default-features = false, features = ["sink", "alloc", "std"] } +futures-core = { version = "0.3.19", default-features = false } +futures-io = "0.3.24" +futures-util = { version = "0.3.19", default-features = false, features = ["alloc", "sink", "io"] } + +# Cryptographic Primitives +crc = "3.0.0" +digest = { version = "0.10.0", default-features = false, features = ["std"] } +hkdf = "0.12.0" +hmac = { version = "0.12.0", default-features = false } +md-5 = { version = "0.10.0", default-features = false } +rand = { version = "0.8.4", default-features = false, features = ["std", "std_rng"] } +rsa = "0.9" +sha1 = { version = "0.10.1", default-features = false } +sha2 = { version = "0.10.0", default-features = false } + +# Type Integrations (versions inherited from `[workspace.dependencies]`) +bigdecimal = { workspace = true, optional = true } +chrono = { workspace = true, optional = true } +rust_decimal = { workspace = true, optional = true } +time = { workspace = true, optional = true } +uuid = { workspace = true, optional = true } + +# Misc +atoi = "2.0" +base64 = { version = "0.22.0", default-features = false, features = ["std"] } +bitflags = { version = "2", default-features = false, features = ["serde"] } +byteorder = { version = "1.4.3", default-features = false, features = ["std"] } +bytes = "1.1.0" +dotenvy = "0.15.5" +either = "1.6.1" +generic-array = { version = "0.14.4", default-features = false } +hex = "0.4.3" +itoa = "1.0.1" +log = "0.4.18" +memchr = { version = "2.4.1", default-features = false } +once_cell = "1.9.0" +percent-encoding = "2.1.0" +smallvec = "1.7.0" +stringprep = "0.1.2" +thiserror = "2.0.0" +tracing = { version = "0.1.37", features = ["log"] } +whoami = { version = "1.2.1", default-features = false } + +serde = { version = "1.0.144", optional = true } + +[dev-dependencies] +sqlx = { workspace = true, features = ["mysql"] } + +[lints] +workspace = true diff --git a/src-tauri/vendor/sqlx-mysql/LICENSE-APACHE b/src-tauri/vendor/sqlx-mysql/LICENSE-APACHE new file mode 100644 index 00000000..c79147e8 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/LICENSE-APACHE @@ -0,0 +1,201 @@ +Apache License +Version 2.0, January 2004 +http://www.apache.org/licenses/ + +TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + +1. Definitions. + +"License" shall mean the terms and conditions for use, reproduction, +and distribution as defined by Sections 1 through 9 of this document. + +"Licensor" shall mean the copyright owner or entity authorized by +the copyright owner that is granting the License. + +"Legal Entity" shall mean the union of the acting entity and all +other entities that control, are controlled by, or are under common +control with that entity. For the purposes of this definition, +"control" means (i) the power, direct or indirect, to cause the +direction or management of such entity, whether by contract or +otherwise, or (ii) ownership of fifty percent (50%) or more of the +outstanding shares, or (iii) beneficial ownership of such entity. + +"You" (or "Your") shall mean an individual or Legal Entity +exercising permissions granted by this License. + +"Source" form shall mean the preferred form for making modifications, +including but not limited to software source code, documentation +source, and configuration files. + +"Object" form shall mean any form resulting from mechanical +transformation or translation of a Source form, including but +not limited to compiled object code, generated documentation, +and conversions to other media types. + +"Work" shall mean the work of authorship, whether in Source or +Object form, made available under the License, as indicated by a +copyright notice that is included in or attached to the work +(an example is provided in the Appendix below). + +"Derivative Works" shall mean any work, whether in Source or Object +form, that is based on (or derived from) the Work and for which the +editorial revisions, annotations, elaborations, or other modifications +represent, as a whole, an original work of authorship. For the purposes +of this License, Derivative Works shall not include works that remain +separable from, or merely link (or bind by name) to the interfaces of, +the Work and Derivative Works thereof. + +"Contribution" shall mean any work of authorship, including +the original version of the Work and any modifications or additions +to that Work or Derivative Works thereof, that is intentionally +submitted to Licensor for inclusion in the Work by the copyright owner +or by an individual or Legal Entity authorized to submit on behalf of +the copyright owner. For the purposes of this definition, "submitted" +means any form of electronic, verbal, or written communication sent +to the Licensor or its representatives, including but not limited to +communication on electronic mailing lists, source code control systems, +and issue tracking systems that are managed by, or on behalf of, the +Licensor for the purpose of discussing and improving the Work, but +excluding communication that is conspicuously marked or otherwise +designated in writing by the copyright owner as "Not a Contribution." + +"Contributor" shall mean Licensor and any individual or Legal Entity +on behalf of whom a Contribution has been received by Licensor and +subsequently incorporated within the Work. + +2. Grant of Copyright License. Subject to the terms and conditions of +this License, each Contributor hereby grants to You a perpetual, +worldwide, non-exclusive, no-charge, royalty-free, irrevocable +copyright license to reproduce, prepare Derivative Works of, +publicly display, publicly perform, sublicense, and distribute the +Work and such Derivative Works in Source or Object form. + +3. Grant of Patent License. Subject to the terms and conditions of +this License, each Contributor hereby grants to You a perpetual, +worldwide, non-exclusive, no-charge, royalty-free, irrevocable +(except as stated in this section) patent license to make, have made, +use, offer to sell, sell, import, and otherwise transfer the Work, +where such license applies only to those patent claims licensable +by such Contributor that are necessarily infringed by their +Contribution(s) alone or by combination of their Contribution(s) +with the Work to which such Contribution(s) was submitted. If You +institute patent litigation against any entity (including a +cross-claim or counterclaim in a lawsuit) alleging that the Work +or a Contribution incorporated within the Work constitutes direct +or contributory patent infringement, then any patent licenses +granted to You under this License for that Work shall terminate +as of the date such litigation is filed. + +4. Redistribution. You may reproduce and distribute copies of the +Work or Derivative Works thereof in any medium, with or without +modifications, and in Source or Object form, provided that You +meet the following conditions: + +(a) You must give any other recipients of the Work or +Derivative Works a copy of this License; and + +(b) You must cause any modified files to carry prominent notices +stating that You changed the files; and + +(c) You must retain, in the Source form of any Derivative Works +that You distribute, all copyright, patent, trademark, and +attribution notices from the Source form of the Work, +excluding those notices that do not pertain to any part of +the Derivative Works; and + +(d) If the Work includes a "NOTICE" text file as part of its +distribution, then any Derivative Works that You distribute must +include a readable copy of the attribution notices contained +within such NOTICE file, excluding those notices that do not +pertain to any part of the Derivative Works, in at least one +of the following places: within a NOTICE text file distributed +as part of the Derivative Works; within the Source form or +documentation, if provided along with the Derivative Works; or, +within a display generated by the Derivative Works, if and +wherever such third-party notices normally appear. The contents +of the NOTICE file are for informational purposes only and +do not modify the License. You may add Your own attribution +notices within Derivative Works that You distribute, alongside +or as an addendum to the NOTICE text from the Work, provided +that such additional attribution notices cannot be construed +as modifying the License. + +You may add Your own copyright statement to Your modifications and +may provide additional or different license terms and conditions +for use, reproduction, or distribution of Your modifications, or +for any such Derivative Works as a whole, provided Your use, +reproduction, and distribution of the Work otherwise complies with +the conditions stated in this License. + +5. Submission of Contributions. Unless You explicitly state otherwise, +any Contribution intentionally submitted for inclusion in the Work +by You to the Licensor shall be under the terms and conditions of +this License, without any additional terms or conditions. +Notwithstanding the above, nothing herein shall supersede or modify +the terms of any separate license agreement you may have executed +with Licensor regarding such Contributions. + +6. Trademarks. This License does not grant permission to use the trade +names, trademarks, service marks, or product names of the Licensor, +except as required for reasonable and customary use in describing the +origin of the Work and reproducing the content of the NOTICE file. + +7. Disclaimer of Warranty. Unless required by applicable law or +agreed to in writing, Licensor provides the Work (and each +Contributor provides its Contributions) on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or +implied, including, without limitation, any warranties or conditions +of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A +PARTICULAR PURPOSE. You are solely responsible for determining the +appropriateness of using or redistributing the Work and assume any +risks associated with Your exercise of permissions under this License. + +8. Limitation of Liability. In no event and under no legal theory, +whether in tort (including negligence), contract, or otherwise, +unless required by applicable law (such as deliberate and grossly +negligent acts) or agreed to in writing, shall any Contributor be +liable to You for damages, including any direct, indirect, special, +incidental, or consequential damages of any character arising as a +result of this License or out of the use or inability to use the +Work (including but not limited to damages for loss of goodwill, +work stoppage, computer failure or malfunction, or any and all +other commercial damages or losses), even if such Contributor +has been advised of the possibility of such damages. + +9. Accepting Warranty or Additional Liability. While redistributing +the Work or Derivative Works thereof, You may choose to offer, +and charge a fee for, acceptance of support, warranty, indemnity, +or other liability obligations and/or rights consistent with this +License. However, in accepting such obligations, You may act only +on Your own behalf and on Your sole responsibility, not on behalf +of any other Contributor, and only if You agree to indemnify, +defend, and hold each Contributor harmless for any liability +incurred by, or claims asserted against, such Contributor by reason +of your accepting any such warranty or additional liability. + +END OF TERMS AND CONDITIONS + +APPENDIX: How to apply the Apache License to your work. + +To apply the Apache License to your work, attach the following +boilerplate notice, with the fields enclosed by brackets "[]" +replaced with your own identifying information. (Don't include +the brackets!) The text should be enclosed in the appropriate +comment syntax for the file format. We also recommend that a +file or class name and description of purpose be included on the +same "printed page" as the copyright notice for easier +identification within third-party archives. + +Copyright 2020 LaunchBadge, LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + +http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. \ No newline at end of file diff --git a/src-tauri/vendor/sqlx-mysql/LICENSE-MIT b/src-tauri/vendor/sqlx-mysql/LICENSE-MIT new file mode 100644 index 00000000..861bf608 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/LICENSE-MIT @@ -0,0 +1,25 @@ +Copyright (c) 2020 LaunchBadge, LLC + +Permission is hereby granted, free of charge, to any +person obtaining a copy of this software and associated +documentation files (the "Software"), to deal in the +Software without restriction, including without +limitation the rights to use, copy, modify, merge, +publish, distribute, sublicense, and/or sell copies of +the Software, and to permit persons to whom the Software +is furnished to do so, subject to the following +conditions: + +The above copyright notice and this permission notice +shall be included in all copies or substantial portions +of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF +ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED +TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A +PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT +SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY +CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION +OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR +IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. diff --git a/src-tauri/vendor/sqlx-mysql/src/any.rs b/src-tauri/vendor/sqlx-mysql/src/any.rs new file mode 100644 index 00000000..19b3a6f2 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/any.rs @@ -0,0 +1,227 @@ +use crate::protocol::text::ColumnType; +use crate::{ + MySql, MySqlColumn, MySqlConnectOptions, MySqlConnection, MySqlQueryResult, MySqlRow, + MySqlTransactionManager, MySqlTypeInfo, +}; +use either::Either; +use futures_core::future::BoxFuture; +use futures_core::stream::BoxStream; +use futures_util::{stream, StreamExt, TryFutureExt, TryStreamExt}; +use sqlx_core::any::{ + Any, AnyArguments, AnyColumn, AnyConnectOptions, AnyConnectionBackend, AnyQueryResult, AnyRow, + AnyStatement, AnyTypeInfo, AnyTypeInfoKind, +}; +use sqlx_core::connection::Connection; +use sqlx_core::database::Database; +use sqlx_core::describe::Describe; +use sqlx_core::executor::Executor; +use sqlx_core::transaction::TransactionManager; +use std::borrow::Cow; +use std::{future, pin::pin}; + +sqlx_core::declare_driver_with_optional_migrate!(DRIVER = MySql); + +impl AnyConnectionBackend for MySqlConnection { + fn name(&self) -> &str { + ::NAME + } + + fn close(self: Box) -> BoxFuture<'static, sqlx_core::Result<()>> { + Connection::close(*self) + } + + fn close_hard(self: Box) -> BoxFuture<'static, sqlx_core::Result<()>> { + Connection::close_hard(*self) + } + + fn ping(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + Connection::ping(self) + } + + fn begin( + &mut self, + statement: Option>, + ) -> BoxFuture<'_, sqlx_core::Result<()>> { + MySqlTransactionManager::begin(self, statement) + } + + fn commit(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + MySqlTransactionManager::commit(self) + } + + fn rollback(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + MySqlTransactionManager::rollback(self) + } + + fn start_rollback(&mut self) { + MySqlTransactionManager::start_rollback(self) + } + + fn get_transaction_depth(&self) -> usize { + MySqlTransactionManager::get_transaction_depth(self) + } + + fn shrink_buffers(&mut self) { + Connection::shrink_buffers(self); + } + + fn flush(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + Connection::flush(self) + } + + fn should_flush(&self) -> bool { + Connection::should_flush(self) + } + + #[cfg(feature = "migrate")] + fn as_migrate( + &mut self, + ) -> sqlx_core::Result<&mut (dyn sqlx_core::migrate::Migrate + Send + 'static)> { + Ok(self) + } + + fn fetch_many<'q>( + &'q mut self, + query: &'q str, + persistent: bool, + arguments: Option>, + ) -> BoxStream<'q, sqlx_core::Result>> { + let persistent = persistent && arguments.is_some(); + let arguments = match arguments.as_ref().map(AnyArguments::convert_to).transpose() { + Ok(arguments) => arguments, + Err(error) => { + return stream::once(future::ready(Err(sqlx_core::Error::Encode(error)))).boxed() + } + }; + + Box::pin( + self.run(query, arguments, persistent) + .try_flatten_stream() + .map(|res| { + Ok(match res? { + Either::Left(result) => Either::Left(map_result(result)), + Either::Right(row) => Either::Right(AnyRow::try_from(&row)?), + }) + }), + ) + } + + fn fetch_optional<'q>( + &'q mut self, + query: &'q str, + persistent: bool, + arguments: Option>, + ) -> BoxFuture<'q, sqlx_core::Result>> { + let persistent = persistent && arguments.is_some(); + let arguments = arguments + .as_ref() + .map(AnyArguments::convert_to) + .transpose() + .map_err(sqlx_core::Error::Encode); + + Box::pin(async move { + let arguments = arguments?; + let mut stream = pin!(self.run(query, arguments, persistent).await?); + + while let Some(result) = stream.try_next().await? { + if let Either::Right(row) = result { + return Ok(Some(AnyRow::try_from(&row)?)); + } + } + + Ok(None) + }) + } + + fn prepare_with<'c, 'q: 'c>( + &'c mut self, + sql: &'q str, + _parameters: &[AnyTypeInfo], + ) -> BoxFuture<'c, sqlx_core::Result>> { + Box::pin(async move { + let statement = Executor::prepare_with(self, sql, &[]).await?; + AnyStatement::try_from_statement( + sql, + &statement, + statement.metadata.column_names.clone(), + ) + }) + } + + fn describe<'q>(&'q mut self, sql: &'q str) -> BoxFuture<'q, sqlx_core::Result>> { + Box::pin(async move { + let describe = Executor::describe(self, sql).await?; + describe.try_into_any() + }) + } +} + +impl<'a> TryFrom<&'a MySqlTypeInfo> for AnyTypeInfo { + type Error = sqlx_core::Error; + + fn try_from(type_info: &'a MySqlTypeInfo) -> Result { + Ok(AnyTypeInfo { + kind: match &type_info.r#type { + ColumnType::Null => AnyTypeInfoKind::Null, + ColumnType::Short => AnyTypeInfoKind::SmallInt, + ColumnType::Long => AnyTypeInfoKind::Integer, + ColumnType::LongLong => AnyTypeInfoKind::BigInt, + ColumnType::Float => AnyTypeInfoKind::Real, + ColumnType::Double => AnyTypeInfoKind::Double, + ColumnType::Blob + | ColumnType::TinyBlob + | ColumnType::MediumBlob + | ColumnType::LongBlob => AnyTypeInfoKind::Blob, + ColumnType::String | ColumnType::VarString | ColumnType::VarChar => { + AnyTypeInfoKind::Text + } + _ => { + return Err(sqlx_core::Error::AnyDriverError( + format!("Any driver does not support MySql type {type_info:?}").into(), + )) + } + }, + }) + } +} + +impl<'a> TryFrom<&'a MySqlColumn> for AnyColumn { + type Error = sqlx_core::Error; + + fn try_from(column: &'a MySqlColumn) -> Result { + let type_info = AnyTypeInfo::try_from(&column.type_info)?; + + Ok(AnyColumn { + ordinal: column.ordinal, + name: column.name.clone(), + type_info, + }) + } +} + +impl<'a> TryFrom<&'a MySqlRow> for AnyRow { + type Error = sqlx_core::Error; + + fn try_from(row: &'a MySqlRow) -> Result { + AnyRow::map_from(row, row.column_names.clone()) + } +} + +impl<'a> TryFrom<&'a AnyConnectOptions> for MySqlConnectOptions { + type Error = sqlx_core::Error; + + fn try_from(any_opts: &'a AnyConnectOptions) -> Result { + let mut opts = Self::parse_from_url(&any_opts.database_url)?; + opts.log_settings = any_opts.log_settings.clone(); + Ok(opts) + } +} + +fn map_result(result: MySqlQueryResult) -> AnyQueryResult { + AnyQueryResult { + rows_affected: result.rows_affected, + // Don't expect this to be a problem + #[allow(clippy::cast_possible_wrap)] + last_insert_id: Some(result.last_insert_id as i64), + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/arguments.rs b/src-tauri/vendor/sqlx-mysql/src/arguments.rs new file mode 100644 index 00000000..464529cb --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/arguments.rs @@ -0,0 +1,108 @@ +use crate::encode::{Encode, IsNull}; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo}; +pub(crate) use sqlx_core::arguments::*; +use sqlx_core::error::BoxDynError; +use std::ops::Deref; + +/// Implementation of [`Arguments`] for MySQL. +#[derive(Debug, Default, Clone)] +pub struct MySqlArguments { + pub(crate) values: Vec, + pub(crate) types: Vec, + pub(crate) null_bitmap: NullBitMap, +} + +impl MySqlArguments { + pub(crate) fn add<'q, T>(&mut self, value: T) -> Result<(), BoxDynError> + where + T: Encode<'q, MySql> + Type, + { + let ty = value.produces().unwrap_or_else(T::type_info); + + let value_length_before_encoding = self.values.len(); + let is_null = match value.encode(&mut self.values) { + Ok(is_null) => is_null, + Err(error) => { + // reset the value buffer to its previous value if encoding failed so we don't leave a half-encoded value behind + self.values.truncate(value_length_before_encoding); + return Err(error); + } + }; + + self.types.push(ty); + self.null_bitmap.push(is_null); + + Ok(()) + } +} + +impl<'q> Arguments<'q> for MySqlArguments { + type Database = MySql; + + fn reserve(&mut self, len: usize, size: usize) { + self.types.reserve(len); + self.values.reserve(size); + } + + fn add(&mut self, value: T) -> Result<(), BoxDynError> + where + T: Encode<'q, Self::Database> + Type, + { + self.add(value) + } + + fn len(&self) -> usize { + self.types.len() + } +} + +#[derive(Debug, Default, Clone)] +pub(crate) struct NullBitMap { + bytes: Vec, + length: usize, +} + +impl NullBitMap { + fn push(&mut self, is_null: IsNull) { + let byte_index = self.length / (u8::BITS as usize); + let bit_offset = self.length % (u8::BITS as usize); + + if bit_offset == 0 { + self.bytes.push(0); + } + + self.bytes[byte_index] |= u8::from(is_null.is_null()) << bit_offset; + self.length += 1; + } +} + +impl Deref for NullBitMap { + type Target = [u8]; + + fn deref(&self) -> &Self::Target { + &self.bytes + } +} + +#[cfg(test)] +mod test { + use super::*; + + #[test] + fn null_bit_map_should_push_is_null() { + let mut bit_map = NullBitMap::default(); + + bit_map.push(IsNull::Yes); + bit_map.push(IsNull::No); + bit_map.push(IsNull::Yes); + bit_map.push(IsNull::No); + bit_map.push(IsNull::Yes); + bit_map.push(IsNull::No); + bit_map.push(IsNull::Yes); + bit_map.push(IsNull::No); + bit_map.push(IsNull::Yes); + + assert_eq!([0b01010101, 0b1].as_slice(), bit_map.deref()); + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/collation.rs b/src-tauri/vendor/sqlx-mysql/src/collation.rs new file mode 100644 index 00000000..46a3a3f9 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/collation.rs @@ -0,0 +1,900 @@ +use crate::error::Error; +use std::str::FromStr; + +#[allow(non_camel_case_types)] +#[derive(Copy, Clone)] +pub(crate) enum CharSet { + armscii8, + ascii, + big5, + binary, + cp1250, + cp1251, + cp1256, + cp1257, + cp850, + cp852, + cp866, + cp932, + dec8, + eucjpms, + euckr, + gb18030, + gb2312, + gbk, + geostd8, + greek, + hebrew, + hp8, + keybcs2, + koi8r, + koi8u, + latin1, + latin2, + latin5, + latin7, + macce, + macroman, + sjis, + swe7, + tis620, + ucs2, + ujis, + utf16, + utf16le, + utf32, + utf8, + utf8mb4, +} + +impl CharSet { + pub(crate) fn as_str(&self) -> &'static str { + match self { + CharSet::armscii8 => "armscii8", + CharSet::ascii => "ascii", + CharSet::big5 => "big5", + CharSet::binary => "binary", + CharSet::cp1250 => "cp1250", + CharSet::cp1251 => "cp1251", + CharSet::cp1256 => "cp1256", + CharSet::cp1257 => "cp1257", + CharSet::cp850 => "cp850", + CharSet::cp852 => "cp852", + CharSet::cp866 => "cp866", + CharSet::cp932 => "cp932", + CharSet::dec8 => "dec8", + CharSet::eucjpms => "eucjpms", + CharSet::euckr => "euckr", + CharSet::gb18030 => "gb18030", + CharSet::gb2312 => "gb2312", + CharSet::gbk => "gbk", + CharSet::geostd8 => "geostd8", + CharSet::greek => "greek", + CharSet::hebrew => "hebrew", + CharSet::hp8 => "hp8", + CharSet::keybcs2 => "keybcs2", + CharSet::koi8r => "koi8r", + CharSet::koi8u => "koi8u", + CharSet::latin1 => "latin1", + CharSet::latin2 => "latin2", + CharSet::latin5 => "latin5", + CharSet::latin7 => "latin7", + CharSet::macce => "macce", + CharSet::macroman => "macroman", + CharSet::sjis => "sjis", + CharSet::swe7 => "swe7", + CharSet::tis620 => "tis620", + CharSet::ucs2 => "ucs2", + CharSet::ujis => "ujis", + CharSet::utf16 => "utf16", + CharSet::utf16le => "utf16le", + CharSet::utf32 => "utf32", + CharSet::utf8 => "utf8", + CharSet::utf8mb4 => "utf8mb4", + } + } + + pub(crate) fn default_collation(&self) -> Collation { + match self { + CharSet::armscii8 => Collation::armscii8_general_ci, + CharSet::ascii => Collation::ascii_general_ci, + CharSet::big5 => Collation::big5_chinese_ci, + CharSet::binary => Collation::binary, + CharSet::cp1250 => Collation::cp1250_general_ci, + CharSet::cp1251 => Collation::cp1251_general_ci, + CharSet::cp1256 => Collation::cp1256_general_ci, + CharSet::cp1257 => Collation::cp1257_general_ci, + CharSet::cp850 => Collation::cp850_general_ci, + CharSet::cp852 => Collation::cp852_general_ci, + CharSet::cp866 => Collation::cp866_general_ci, + CharSet::cp932 => Collation::cp932_japanese_ci, + CharSet::dec8 => Collation::dec8_swedish_ci, + CharSet::eucjpms => Collation::eucjpms_japanese_ci, + CharSet::euckr => Collation::euckr_korean_ci, + CharSet::gb18030 => Collation::gb18030_chinese_ci, + CharSet::gb2312 => Collation::gb2312_chinese_ci, + CharSet::gbk => Collation::gbk_chinese_ci, + CharSet::geostd8 => Collation::geostd8_general_ci, + CharSet::greek => Collation::greek_general_ci, + CharSet::hebrew => Collation::hebrew_general_ci, + CharSet::hp8 => Collation::hp8_english_ci, + CharSet::keybcs2 => Collation::keybcs2_general_ci, + CharSet::koi8r => Collation::koi8r_general_ci, + CharSet::koi8u => Collation::koi8u_general_ci, + CharSet::latin1 => Collation::latin1_swedish_ci, + CharSet::latin2 => Collation::latin2_general_ci, + CharSet::latin5 => Collation::latin5_turkish_ci, + CharSet::latin7 => Collation::latin7_general_ci, + CharSet::macce => Collation::macce_general_ci, + CharSet::macroman => Collation::macroman_general_ci, + CharSet::sjis => Collation::sjis_japanese_ci, + CharSet::swe7 => Collation::swe7_swedish_ci, + CharSet::tis620 => Collation::tis620_thai_ci, + CharSet::ucs2 => Collation::ucs2_general_ci, + CharSet::ujis => Collation::ujis_japanese_ci, + CharSet::utf16 => Collation::utf16_general_ci, + CharSet::utf16le => Collation::utf16le_general_ci, + CharSet::utf32 => Collation::utf32_general_ci, + CharSet::utf8 => Collation::utf8_unicode_ci, + CharSet::utf8mb4 => Collation::utf8mb4_unicode_ci, + } + } +} + +impl FromStr for CharSet { + type Err = Error; + + fn from_str(char_set: &str) -> Result { + Ok(match char_set { + "armscii8" => CharSet::armscii8, + "ascii" => CharSet::ascii, + "big5" => CharSet::big5, + "binary" => CharSet::binary, + "cp1250" => CharSet::cp1250, + "cp1251" => CharSet::cp1251, + "cp1256" => CharSet::cp1256, + "cp1257" => CharSet::cp1257, + "cp850" => CharSet::cp850, + "cp852" => CharSet::cp852, + "cp866" => CharSet::cp866, + "cp932" => CharSet::cp932, + "dec8" => CharSet::dec8, + "eucjpms" => CharSet::eucjpms, + "euckr" => CharSet::euckr, + "gb18030" => CharSet::gb18030, + "gb2312" => CharSet::gb2312, + "gbk" => CharSet::gbk, + "geostd8" => CharSet::geostd8, + "greek" => CharSet::greek, + "hebrew" => CharSet::hebrew, + "hp8" => CharSet::hp8, + "keybcs2" => CharSet::keybcs2, + "koi8r" => CharSet::koi8r, + "koi8u" => CharSet::koi8u, + "latin1" => CharSet::latin1, + "latin2" => CharSet::latin2, + "latin5" => CharSet::latin5, + "latin7" => CharSet::latin7, + "macce" => CharSet::macce, + "macroman" => CharSet::macroman, + "sjis" => CharSet::sjis, + "swe7" => CharSet::swe7, + "tis620" => CharSet::tis620, + "ucs2" => CharSet::ucs2, + "ujis" => CharSet::ujis, + "utf16" => CharSet::utf16, + "utf16le" => CharSet::utf16le, + "utf32" => CharSet::utf32, + "utf8" => CharSet::utf8, + "utf8mb4" => CharSet::utf8mb4, + + _ => { + return Err(Error::Configuration( + format!("unsupported MySQL charset: {char_set}").into(), + )); + } + }) + } +} + +#[derive(Copy, Clone)] +#[allow(non_camel_case_types)] +#[repr(u8)] +pub(crate) enum Collation { + armscii8_bin = 64, + armscii8_general_ci = 32, + ascii_bin = 65, + ascii_general_ci = 11, + big5_bin = 84, + big5_chinese_ci = 1, + binary = 63, + cp1250_bin = 66, + cp1250_croatian_ci = 44, + cp1250_czech_cs = 34, + cp1250_general_ci = 26, + cp1250_polish_ci = 99, + cp1251_bin = 50, + cp1251_bulgarian_ci = 14, + cp1251_general_ci = 51, + cp1251_general_cs = 52, + cp1251_ukrainian_ci = 23, + cp1256_bin = 67, + cp1256_general_ci = 57, + cp1257_bin = 58, + cp1257_general_ci = 59, + cp1257_lithuanian_ci = 29, + cp850_bin = 80, + cp850_general_ci = 4, + cp852_bin = 81, + cp852_general_ci = 40, + cp866_bin = 68, + cp866_general_ci = 36, + cp932_bin = 96, + cp932_japanese_ci = 95, + dec8_bin = 69, + dec8_swedish_ci = 3, + eucjpms_bin = 98, + eucjpms_japanese_ci = 97, + euckr_bin = 85, + euckr_korean_ci = 19, + gb18030_bin = 249, + gb18030_chinese_ci = 248, + gb18030_unicode_520_ci = 250, + gb2312_bin = 86, + gb2312_chinese_ci = 24, + gbk_bin = 87, + gbk_chinese_ci = 28, + geostd8_bin = 93, + geostd8_general_ci = 92, + greek_bin = 70, + greek_general_ci = 25, + hebrew_bin = 71, + hebrew_general_ci = 16, + hp8_bin = 72, + hp8_english_ci = 6, + keybcs2_bin = 73, + keybcs2_general_ci = 37, + koi8r_bin = 74, + koi8r_general_ci = 7, + koi8u_bin = 75, + koi8u_general_ci = 22, + latin1_bin = 47, + latin1_danish_ci = 15, + latin1_general_ci = 48, + latin1_general_cs = 49, + latin1_german1_ci = 5, + latin1_german2_ci = 31, + latin1_spanish_ci = 94, + latin1_swedish_ci = 8, + latin2_bin = 77, + latin2_croatian_ci = 27, + latin2_czech_cs = 2, + latin2_general_ci = 9, + latin2_hungarian_ci = 21, + latin5_bin = 78, + latin5_turkish_ci = 30, + latin7_bin = 79, + latin7_estonian_cs = 20, + latin7_general_ci = 41, + latin7_general_cs = 42, + macce_bin = 43, + macce_general_ci = 38, + macroman_bin = 53, + macroman_general_ci = 39, + sjis_bin = 88, + sjis_japanese_ci = 13, + swe7_bin = 82, + swe7_swedish_ci = 10, + tis620_bin = 89, + tis620_thai_ci = 18, + ucs2_bin = 90, + ucs2_croatian_ci = 149, + ucs2_czech_ci = 138, + ucs2_danish_ci = 139, + ucs2_esperanto_ci = 145, + ucs2_estonian_ci = 134, + ucs2_general_ci = 35, + ucs2_general_mysql500_ci = 159, + ucs2_german2_ci = 148, + ucs2_hungarian_ci = 146, + ucs2_icelandic_ci = 129, + ucs2_latvian_ci = 130, + ucs2_lithuanian_ci = 140, + ucs2_persian_ci = 144, + ucs2_polish_ci = 133, + ucs2_roman_ci = 143, + ucs2_romanian_ci = 131, + ucs2_sinhala_ci = 147, + ucs2_slovak_ci = 141, + ucs2_slovenian_ci = 132, + ucs2_spanish_ci = 135, + ucs2_spanish2_ci = 142, + ucs2_swedish_ci = 136, + ucs2_turkish_ci = 137, + ucs2_unicode_520_ci = 150, + ucs2_unicode_ci = 128, + ucs2_vietnamese_ci = 151, + ujis_bin = 91, + ujis_japanese_ci = 12, + utf16_bin = 55, + utf16_croatian_ci = 122, + utf16_czech_ci = 111, + utf16_danish_ci = 112, + utf16_esperanto_ci = 118, + utf16_estonian_ci = 107, + utf16_general_ci = 54, + utf16_german2_ci = 121, + utf16_hungarian_ci = 119, + utf16_icelandic_ci = 102, + utf16_latvian_ci = 103, + utf16_lithuanian_ci = 113, + utf16_persian_ci = 117, + utf16_polish_ci = 106, + utf16_roman_ci = 116, + utf16_romanian_ci = 104, + utf16_sinhala_ci = 120, + utf16_slovak_ci = 114, + utf16_slovenian_ci = 105, + utf16_spanish_ci = 108, + utf16_spanish2_ci = 115, + utf16_swedish_ci = 109, + utf16_turkish_ci = 110, + utf16_unicode_520_ci = 123, + utf16_unicode_ci = 101, + utf16_vietnamese_ci = 124, + utf16le_bin = 62, + utf16le_general_ci = 56, + utf32_bin = 61, + utf32_croatian_ci = 181, + utf32_czech_ci = 170, + utf32_danish_ci = 171, + utf32_esperanto_ci = 177, + utf32_estonian_ci = 166, + utf32_general_ci = 60, + utf32_german2_ci = 180, + utf32_hungarian_ci = 178, + utf32_icelandic_ci = 161, + utf32_latvian_ci = 162, + utf32_lithuanian_ci = 172, + utf32_persian_ci = 176, + utf32_polish_ci = 165, + utf32_roman_ci = 175, + utf32_romanian_ci = 163, + utf32_sinhala_ci = 179, + utf32_slovak_ci = 173, + utf32_slovenian_ci = 164, + utf32_spanish_ci = 167, + utf32_spanish2_ci = 174, + utf32_swedish_ci = 168, + utf32_turkish_ci = 169, + utf32_unicode_520_ci = 182, + utf32_unicode_ci = 160, + utf32_vietnamese_ci = 183, + utf8_bin = 83, + utf8_croatian_ci = 213, + utf8_czech_ci = 202, + utf8_danish_ci = 203, + utf8_esperanto_ci = 209, + utf8_estonian_ci = 198, + utf8_general_ci = 33, + utf8_general_mysql500_ci = 223, + utf8_german2_ci = 212, + utf8_hungarian_ci = 210, + utf8_icelandic_ci = 193, + utf8_latvian_ci = 194, + utf8_lithuanian_ci = 204, + utf8_persian_ci = 208, + utf8_polish_ci = 197, + utf8_roman_ci = 207, + utf8_romanian_ci = 195, + utf8_sinhala_ci = 211, + utf8_slovak_ci = 205, + utf8_slovenian_ci = 196, + utf8_spanish_ci = 199, + utf8_spanish2_ci = 206, + utf8_swedish_ci = 200, + utf8_tolower_ci = 76, + utf8_turkish_ci = 201, + utf8_unicode_520_ci = 214, + utf8_unicode_ci = 192, + utf8_vietnamese_ci = 215, + utf8mb4_0900_ai_ci = 255, + utf8mb4_bin = 46, + utf8mb4_croatian_ci = 245, + utf8mb4_czech_ci = 234, + utf8mb4_danish_ci = 235, + utf8mb4_esperanto_ci = 241, + utf8mb4_estonian_ci = 230, + utf8mb4_general_ci = 45, + utf8mb4_german2_ci = 244, + utf8mb4_hungarian_ci = 242, + utf8mb4_icelandic_ci = 225, + utf8mb4_latvian_ci = 226, + utf8mb4_lithuanian_ci = 236, + utf8mb4_persian_ci = 240, + utf8mb4_polish_ci = 229, + utf8mb4_roman_ci = 239, + utf8mb4_romanian_ci = 227, + utf8mb4_sinhala_ci = 243, + utf8mb4_slovak_ci = 237, + utf8mb4_slovenian_ci = 228, + utf8mb4_spanish_ci = 231, + utf8mb4_spanish2_ci = 238, + utf8mb4_swedish_ci = 232, + utf8mb4_turkish_ci = 233, + utf8mb4_unicode_520_ci = 246, + utf8mb4_unicode_ci = 224, + utf8mb4_vietnamese_ci = 247, +} + +impl Collation { + pub(crate) fn as_str(&self) -> &'static str { + match self { + Collation::armscii8_bin => "armscii8_bin", + Collation::armscii8_general_ci => "armscii8_general_ci", + Collation::ascii_bin => "ascii_bin", + Collation::ascii_general_ci => "ascii_general_ci", + Collation::big5_bin => "big5_bin", + Collation::big5_chinese_ci => "big5_chinese_ci", + Collation::binary => "binary", + Collation::cp1250_bin => "cp1250_bin", + Collation::cp1250_croatian_ci => "cp1250_croatian_ci", + Collation::cp1250_czech_cs => "cp1250_czech_cs", + Collation::cp1250_general_ci => "cp1250_general_ci", + Collation::cp1250_polish_ci => "cp1250_polish_ci", + Collation::cp1251_bin => "cp1251_bin", + Collation::cp1251_bulgarian_ci => "cp1251_bulgarian_ci", + Collation::cp1251_general_ci => "cp1251_general_ci", + Collation::cp1251_general_cs => "cp1251_general_cs", + Collation::cp1251_ukrainian_ci => "cp1251_ukrainian_ci", + Collation::cp1256_bin => "cp1256_bin", + Collation::cp1256_general_ci => "cp1256_general_ci", + Collation::cp1257_bin => "cp1257_bin", + Collation::cp1257_general_ci => "cp1257_general_ci", + Collation::cp1257_lithuanian_ci => "cp1257_lithuanian_ci", + Collation::cp850_bin => "cp850_bin", + Collation::cp850_general_ci => "cp850_general_ci", + Collation::cp852_bin => "cp852_bin", + Collation::cp852_general_ci => "cp852_general_ci", + Collation::cp866_bin => "cp866_bin", + Collation::cp866_general_ci => "cp866_general_ci", + Collation::cp932_bin => "cp932_bin", + Collation::cp932_japanese_ci => "cp932_japanese_ci", + Collation::dec8_bin => "dec8_bin", + Collation::dec8_swedish_ci => "dec8_swedish_ci", + Collation::eucjpms_bin => "eucjpms_bin", + Collation::eucjpms_japanese_ci => "eucjpms_japanese_ci", + Collation::euckr_bin => "euckr_bin", + Collation::euckr_korean_ci => "euckr_korean_ci", + Collation::gb18030_bin => "gb18030_bin", + Collation::gb18030_chinese_ci => "gb18030_chinese_ci", + Collation::gb18030_unicode_520_ci => "gb18030_unicode_520_ci", + Collation::gb2312_bin => "gb2312_bin", + Collation::gb2312_chinese_ci => "gb2312_chinese_ci", + Collation::gbk_bin => "gbk_bin", + Collation::gbk_chinese_ci => "gbk_chinese_ci", + Collation::geostd8_bin => "geostd8_bin", + Collation::geostd8_general_ci => "geostd8_general_ci", + Collation::greek_bin => "greek_bin", + Collation::greek_general_ci => "greek_general_ci", + Collation::hebrew_bin => "hebrew_bin", + Collation::hebrew_general_ci => "hebrew_general_ci", + Collation::hp8_bin => "hp8_bin", + Collation::hp8_english_ci => "hp8_english_ci", + Collation::keybcs2_bin => "keybcs2_bin", + Collation::keybcs2_general_ci => "keybcs2_general_ci", + Collation::koi8r_bin => "koi8r_bin", + Collation::koi8r_general_ci => "koi8r_general_ci", + Collation::koi8u_bin => "koi8u_bin", + Collation::koi8u_general_ci => "koi8u_general_ci", + Collation::latin1_bin => "latin1_bin", + Collation::latin1_danish_ci => "latin1_danish_ci", + Collation::latin1_general_ci => "latin1_general_ci", + Collation::latin1_general_cs => "latin1_general_cs", + Collation::latin1_german1_ci => "latin1_german1_ci", + Collation::latin1_german2_ci => "latin1_german2_ci", + Collation::latin1_spanish_ci => "latin1_spanish_ci", + Collation::latin1_swedish_ci => "latin1_swedish_ci", + Collation::latin2_bin => "latin2_bin", + Collation::latin2_croatian_ci => "latin2_croatian_ci", + Collation::latin2_czech_cs => "latin2_czech_cs", + Collation::latin2_general_ci => "latin2_general_ci", + Collation::latin2_hungarian_ci => "latin2_hungarian_ci", + Collation::latin5_bin => "latin5_bin", + Collation::latin5_turkish_ci => "latin5_turkish_ci", + Collation::latin7_bin => "latin7_bin", + Collation::latin7_estonian_cs => "latin7_estonian_cs", + Collation::latin7_general_ci => "latin7_general_ci", + Collation::latin7_general_cs => "latin7_general_cs", + Collation::macce_bin => "macce_bin", + Collation::macce_general_ci => "macce_general_ci", + Collation::macroman_bin => "macroman_bin", + Collation::macroman_general_ci => "macroman_general_ci", + Collation::sjis_bin => "sjis_bin", + Collation::sjis_japanese_ci => "sjis_japanese_ci", + Collation::swe7_bin => "swe7_bin", + Collation::swe7_swedish_ci => "swe7_swedish_ci", + Collation::tis620_bin => "tis620_bin", + Collation::tis620_thai_ci => "tis620_thai_ci", + Collation::ucs2_bin => "ucs2_bin", + Collation::ucs2_croatian_ci => "ucs2_croatian_ci", + Collation::ucs2_czech_ci => "ucs2_czech_ci", + Collation::ucs2_danish_ci => "ucs2_danish_ci", + Collation::ucs2_esperanto_ci => "ucs2_esperanto_ci", + Collation::ucs2_estonian_ci => "ucs2_estonian_ci", + Collation::ucs2_general_ci => "ucs2_general_ci", + Collation::ucs2_general_mysql500_ci => "ucs2_general_mysql500_ci", + Collation::ucs2_german2_ci => "ucs2_german2_ci", + Collation::ucs2_hungarian_ci => "ucs2_hungarian_ci", + Collation::ucs2_icelandic_ci => "ucs2_icelandic_ci", + Collation::ucs2_latvian_ci => "ucs2_latvian_ci", + Collation::ucs2_lithuanian_ci => "ucs2_lithuanian_ci", + Collation::ucs2_persian_ci => "ucs2_persian_ci", + Collation::ucs2_polish_ci => "ucs2_polish_ci", + Collation::ucs2_roman_ci => "ucs2_roman_ci", + Collation::ucs2_romanian_ci => "ucs2_romanian_ci", + Collation::ucs2_sinhala_ci => "ucs2_sinhala_ci", + Collation::ucs2_slovak_ci => "ucs2_slovak_ci", + Collation::ucs2_slovenian_ci => "ucs2_slovenian_ci", + Collation::ucs2_spanish_ci => "ucs2_spanish_ci", + Collation::ucs2_spanish2_ci => "ucs2_spanish2_ci", + Collation::ucs2_swedish_ci => "ucs2_swedish_ci", + Collation::ucs2_turkish_ci => "ucs2_turkish_ci", + Collation::ucs2_unicode_520_ci => "ucs2_unicode_520_ci", + Collation::ucs2_unicode_ci => "ucs2_unicode_ci", + Collation::ucs2_vietnamese_ci => "ucs2_vietnamese_ci", + Collation::ujis_bin => "ujis_bin", + Collation::ujis_japanese_ci => "ujis_japanese_ci", + Collation::utf16_bin => "utf16_bin", + Collation::utf16_croatian_ci => "utf16_croatian_ci", + Collation::utf16_czech_ci => "utf16_czech_ci", + Collation::utf16_danish_ci => "utf16_danish_ci", + Collation::utf16_esperanto_ci => "utf16_esperanto_ci", + Collation::utf16_estonian_ci => "utf16_estonian_ci", + Collation::utf16_general_ci => "utf16_general_ci", + Collation::utf16_german2_ci => "utf16_german2_ci", + Collation::utf16_hungarian_ci => "utf16_hungarian_ci", + Collation::utf16_icelandic_ci => "utf16_icelandic_ci", + Collation::utf16_latvian_ci => "utf16_latvian_ci", + Collation::utf16_lithuanian_ci => "utf16_lithuanian_ci", + Collation::utf16_persian_ci => "utf16_persian_ci", + Collation::utf16_polish_ci => "utf16_polish_ci", + Collation::utf16_roman_ci => "utf16_roman_ci", + Collation::utf16_romanian_ci => "utf16_romanian_ci", + Collation::utf16_sinhala_ci => "utf16_sinhala_ci", + Collation::utf16_slovak_ci => "utf16_slovak_ci", + Collation::utf16_slovenian_ci => "utf16_slovenian_ci", + Collation::utf16_spanish_ci => "utf16_spanish_ci", + Collation::utf16_spanish2_ci => "utf16_spanish2_ci", + Collation::utf16_swedish_ci => "utf16_swedish_ci", + Collation::utf16_turkish_ci => "utf16_turkish_ci", + Collation::utf16_unicode_520_ci => "utf16_unicode_520_ci", + Collation::utf16_unicode_ci => "utf16_unicode_ci", + Collation::utf16_vietnamese_ci => "utf16_vietnamese_ci", + Collation::utf16le_bin => "utf16le_bin", + Collation::utf16le_general_ci => "utf16le_general_ci", + Collation::utf32_bin => "utf32_bin", + Collation::utf32_croatian_ci => "utf32_croatian_ci", + Collation::utf32_czech_ci => "utf32_czech_ci", + Collation::utf32_danish_ci => "utf32_danish_ci", + Collation::utf32_esperanto_ci => "utf32_esperanto_ci", + Collation::utf32_estonian_ci => "utf32_estonian_ci", + Collation::utf32_general_ci => "utf32_general_ci", + Collation::utf32_german2_ci => "utf32_german2_ci", + Collation::utf32_hungarian_ci => "utf32_hungarian_ci", + Collation::utf32_icelandic_ci => "utf32_icelandic_ci", + Collation::utf32_latvian_ci => "utf32_latvian_ci", + Collation::utf32_lithuanian_ci => "utf32_lithuanian_ci", + Collation::utf32_persian_ci => "utf32_persian_ci", + Collation::utf32_polish_ci => "utf32_polish_ci", + Collation::utf32_roman_ci => "utf32_roman_ci", + Collation::utf32_romanian_ci => "utf32_romanian_ci", + Collation::utf32_sinhala_ci => "utf32_sinhala_ci", + Collation::utf32_slovak_ci => "utf32_slovak_ci", + Collation::utf32_slovenian_ci => "utf32_slovenian_ci", + Collation::utf32_spanish_ci => "utf32_spanish_ci", + Collation::utf32_spanish2_ci => "utf32_spanish2_ci", + Collation::utf32_swedish_ci => "utf32_swedish_ci", + Collation::utf32_turkish_ci => "utf32_turkish_ci", + Collation::utf32_unicode_520_ci => "utf32_unicode_520_ci", + Collation::utf32_unicode_ci => "utf32_unicode_ci", + Collation::utf32_vietnamese_ci => "utf32_vietnamese_ci", + Collation::utf8_bin => "utf8_bin", + Collation::utf8_croatian_ci => "utf8_croatian_ci", + Collation::utf8_czech_ci => "utf8_czech_ci", + Collation::utf8_danish_ci => "utf8_danish_ci", + Collation::utf8_esperanto_ci => "utf8_esperanto_ci", + Collation::utf8_estonian_ci => "utf8_estonian_ci", + Collation::utf8_general_ci => "utf8_general_ci", + Collation::utf8_general_mysql500_ci => "utf8_general_mysql500_ci", + Collation::utf8_german2_ci => "utf8_german2_ci", + Collation::utf8_hungarian_ci => "utf8_hungarian_ci", + Collation::utf8_icelandic_ci => "utf8_icelandic_ci", + Collation::utf8_latvian_ci => "utf8_latvian_ci", + Collation::utf8_lithuanian_ci => "utf8_lithuanian_ci", + Collation::utf8_persian_ci => "utf8_persian_ci", + Collation::utf8_polish_ci => "utf8_polish_ci", + Collation::utf8_roman_ci => "utf8_roman_ci", + Collation::utf8_romanian_ci => "utf8_romanian_ci", + Collation::utf8_sinhala_ci => "utf8_sinhala_ci", + Collation::utf8_slovak_ci => "utf8_slovak_ci", + Collation::utf8_slovenian_ci => "utf8_slovenian_ci", + Collation::utf8_spanish_ci => "utf8_spanish_ci", + Collation::utf8_spanish2_ci => "utf8_spanish2_ci", + Collation::utf8_swedish_ci => "utf8_swedish_ci", + Collation::utf8_tolower_ci => "utf8_tolower_ci", + Collation::utf8_turkish_ci => "utf8_turkish_ci", + Collation::utf8_unicode_520_ci => "utf8_unicode_520_ci", + Collation::utf8_unicode_ci => "utf8_unicode_ci", + Collation::utf8_vietnamese_ci => "utf8_vietnamese_ci", + Collation::utf8mb4_0900_ai_ci => "utf8mb4_0900_ai_ci", + Collation::utf8mb4_bin => "utf8mb4_bin", + Collation::utf8mb4_croatian_ci => "utf8mb4_croatian_ci", + Collation::utf8mb4_czech_ci => "utf8mb4_czech_ci", + Collation::utf8mb4_danish_ci => "utf8mb4_danish_ci", + Collation::utf8mb4_esperanto_ci => "utf8mb4_esperanto_ci", + Collation::utf8mb4_estonian_ci => "utf8mb4_estonian_ci", + Collation::utf8mb4_general_ci => "utf8mb4_general_ci", + Collation::utf8mb4_german2_ci => "utf8mb4_german2_ci", + Collation::utf8mb4_hungarian_ci => "utf8mb4_hungarian_ci", + Collation::utf8mb4_icelandic_ci => "utf8mb4_icelandic_ci", + Collation::utf8mb4_latvian_ci => "utf8mb4_latvian_ci", + Collation::utf8mb4_lithuanian_ci => "utf8mb4_lithuanian_ci", + Collation::utf8mb4_persian_ci => "utf8mb4_persian_ci", + Collation::utf8mb4_polish_ci => "utf8mb4_polish_ci", + Collation::utf8mb4_roman_ci => "utf8mb4_roman_ci", + Collation::utf8mb4_romanian_ci => "utf8mb4_romanian_ci", + Collation::utf8mb4_sinhala_ci => "utf8mb4_sinhala_ci", + Collation::utf8mb4_slovak_ci => "utf8mb4_slovak_ci", + Collation::utf8mb4_slovenian_ci => "utf8mb4_slovenian_ci", + Collation::utf8mb4_spanish_ci => "utf8mb4_spanish_ci", + Collation::utf8mb4_spanish2_ci => "utf8mb4_spanish2_ci", + Collation::utf8mb4_swedish_ci => "utf8mb4_swedish_ci", + Collation::utf8mb4_turkish_ci => "utf8mb4_turkish_ci", + Collation::utf8mb4_unicode_520_ci => "utf8mb4_unicode_520_ci", + Collation::utf8mb4_unicode_ci => "utf8mb4_unicode_ci", + Collation::utf8mb4_vietnamese_ci => "utf8mb4_vietnamese_ci", + } + } +} + +// Handshake packet have only 1 byte for collation_id. +// So we can't use collations with ID > 255. +impl FromStr for Collation { + type Err = Error; + + fn from_str(collation: &str) -> Result { + Ok(match collation { + "big5_chinese_ci" => Collation::big5_chinese_ci, + "swe7_swedish_ci" => Collation::swe7_swedish_ci, + "utf16_unicode_ci" => Collation::utf16_unicode_ci, + "utf16_icelandic_ci" => Collation::utf16_icelandic_ci, + "utf16_latvian_ci" => Collation::utf16_latvian_ci, + "utf16_romanian_ci" => Collation::utf16_romanian_ci, + "utf16_slovenian_ci" => Collation::utf16_slovenian_ci, + "utf16_polish_ci" => Collation::utf16_polish_ci, + "utf16_estonian_ci" => Collation::utf16_estonian_ci, + "utf16_spanish_ci" => Collation::utf16_spanish_ci, + "utf16_swedish_ci" => Collation::utf16_swedish_ci, + "ascii_general_ci" => Collation::ascii_general_ci, + "utf16_turkish_ci" => Collation::utf16_turkish_ci, + "utf16_czech_ci" => Collation::utf16_czech_ci, + "utf16_danish_ci" => Collation::utf16_danish_ci, + "utf16_lithuanian_ci" => Collation::utf16_lithuanian_ci, + "utf16_slovak_ci" => Collation::utf16_slovak_ci, + "utf16_spanish2_ci" => Collation::utf16_spanish2_ci, + "utf16_roman_ci" => Collation::utf16_roman_ci, + "utf16_persian_ci" => Collation::utf16_persian_ci, + "utf16_esperanto_ci" => Collation::utf16_esperanto_ci, + "utf16_hungarian_ci" => Collation::utf16_hungarian_ci, + "ujis_japanese_ci" => Collation::ujis_japanese_ci, + "utf16_sinhala_ci" => Collation::utf16_sinhala_ci, + "utf16_german2_ci" => Collation::utf16_german2_ci, + "utf16_croatian_ci" => Collation::utf16_croatian_ci, + "utf16_unicode_520_ci" => Collation::utf16_unicode_520_ci, + "utf16_vietnamese_ci" => Collation::utf16_vietnamese_ci, + "ucs2_unicode_ci" => Collation::ucs2_unicode_ci, + "ucs2_icelandic_ci" => Collation::ucs2_icelandic_ci, + "sjis_japanese_ci" => Collation::sjis_japanese_ci, + "ucs2_latvian_ci" => Collation::ucs2_latvian_ci, + "ucs2_romanian_ci" => Collation::ucs2_romanian_ci, + "ucs2_slovenian_ci" => Collation::ucs2_slovenian_ci, + "ucs2_polish_ci" => Collation::ucs2_polish_ci, + "ucs2_estonian_ci" => Collation::ucs2_estonian_ci, + "ucs2_spanish_ci" => Collation::ucs2_spanish_ci, + "ucs2_swedish_ci" => Collation::ucs2_swedish_ci, + "ucs2_turkish_ci" => Collation::ucs2_turkish_ci, + "ucs2_czech_ci" => Collation::ucs2_czech_ci, + "ucs2_danish_ci" => Collation::ucs2_danish_ci, + "cp1251_bulgarian_ci" => Collation::cp1251_bulgarian_ci, + "ucs2_lithuanian_ci" => Collation::ucs2_lithuanian_ci, + "ucs2_slovak_ci" => Collation::ucs2_slovak_ci, + "ucs2_spanish2_ci" => Collation::ucs2_spanish2_ci, + "ucs2_roman_ci" => Collation::ucs2_roman_ci, + "ucs2_persian_ci" => Collation::ucs2_persian_ci, + "ucs2_esperanto_ci" => Collation::ucs2_esperanto_ci, + "ucs2_hungarian_ci" => Collation::ucs2_hungarian_ci, + "ucs2_sinhala_ci" => Collation::ucs2_sinhala_ci, + "ucs2_german2_ci" => Collation::ucs2_german2_ci, + "ucs2_croatian_ci" => Collation::ucs2_croatian_ci, + "latin1_danish_ci" => Collation::latin1_danish_ci, + "ucs2_unicode_520_ci" => Collation::ucs2_unicode_520_ci, + "ucs2_vietnamese_ci" => Collation::ucs2_vietnamese_ci, + "ucs2_general_mysql500_ci" => Collation::ucs2_general_mysql500_ci, + "hebrew_general_ci" => Collation::hebrew_general_ci, + "utf32_unicode_ci" => Collation::utf32_unicode_ci, + "utf32_icelandic_ci" => Collation::utf32_icelandic_ci, + "utf32_latvian_ci" => Collation::utf32_latvian_ci, + "utf32_romanian_ci" => Collation::utf32_romanian_ci, + "utf32_slovenian_ci" => Collation::utf32_slovenian_ci, + "utf32_polish_ci" => Collation::utf32_polish_ci, + "utf32_estonian_ci" => Collation::utf32_estonian_ci, + "utf32_spanish_ci" => Collation::utf32_spanish_ci, + "utf32_swedish_ci" => Collation::utf32_swedish_ci, + "utf32_turkish_ci" => Collation::utf32_turkish_ci, + "utf32_czech_ci" => Collation::utf32_czech_ci, + "utf32_danish_ci" => Collation::utf32_danish_ci, + "utf32_lithuanian_ci" => Collation::utf32_lithuanian_ci, + "utf32_slovak_ci" => Collation::utf32_slovak_ci, + "utf32_spanish2_ci" => Collation::utf32_spanish2_ci, + "utf32_roman_ci" => Collation::utf32_roman_ci, + "utf32_persian_ci" => Collation::utf32_persian_ci, + "utf32_esperanto_ci" => Collation::utf32_esperanto_ci, + "utf32_hungarian_ci" => Collation::utf32_hungarian_ci, + "utf32_sinhala_ci" => Collation::utf32_sinhala_ci, + "tis620_thai_ci" => Collation::tis620_thai_ci, + "utf32_german2_ci" => Collation::utf32_german2_ci, + "utf32_croatian_ci" => Collation::utf32_croatian_ci, + "utf32_unicode_520_ci" => Collation::utf32_unicode_520_ci, + "utf32_vietnamese_ci" => Collation::utf32_vietnamese_ci, + "euckr_korean_ci" => Collation::euckr_korean_ci, + "utf8_unicode_ci" => Collation::utf8_unicode_ci, + "utf8_icelandic_ci" => Collation::utf8_icelandic_ci, + "utf8_latvian_ci" => Collation::utf8_latvian_ci, + "utf8_romanian_ci" => Collation::utf8_romanian_ci, + "utf8_slovenian_ci" => Collation::utf8_slovenian_ci, + "utf8_polish_ci" => Collation::utf8_polish_ci, + "utf8_estonian_ci" => Collation::utf8_estonian_ci, + "utf8_spanish_ci" => Collation::utf8_spanish_ci, + "latin2_czech_cs" => Collation::latin2_czech_cs, + "latin7_estonian_cs" => Collation::latin7_estonian_cs, + "utf8_swedish_ci" => Collation::utf8_swedish_ci, + "utf8_turkish_ci" => Collation::utf8_turkish_ci, + "utf8_czech_ci" => Collation::utf8_czech_ci, + "utf8_danish_ci" => Collation::utf8_danish_ci, + "utf8_lithuanian_ci" => Collation::utf8_lithuanian_ci, + "utf8_slovak_ci" => Collation::utf8_slovak_ci, + "utf8_spanish2_ci" => Collation::utf8_spanish2_ci, + "utf8_roman_ci" => Collation::utf8_roman_ci, + "utf8_persian_ci" => Collation::utf8_persian_ci, + "utf8_esperanto_ci" => Collation::utf8_esperanto_ci, + "latin2_hungarian_ci" => Collation::latin2_hungarian_ci, + "utf8_hungarian_ci" => Collation::utf8_hungarian_ci, + "utf8_sinhala_ci" => Collation::utf8_sinhala_ci, + "utf8_german2_ci" => Collation::utf8_german2_ci, + "utf8_croatian_ci" => Collation::utf8_croatian_ci, + "utf8_unicode_520_ci" => Collation::utf8_unicode_520_ci, + "utf8_vietnamese_ci" => Collation::utf8_vietnamese_ci, + "koi8u_general_ci" => Collation::koi8u_general_ci, + "utf8_general_mysql500_ci" => Collation::utf8_general_mysql500_ci, + "utf8mb4_unicode_ci" => Collation::utf8mb4_unicode_ci, + "utf8mb4_icelandic_ci" => Collation::utf8mb4_icelandic_ci, + "utf8mb4_latvian_ci" => Collation::utf8mb4_latvian_ci, + "utf8mb4_romanian_ci" => Collation::utf8mb4_romanian_ci, + "utf8mb4_slovenian_ci" => Collation::utf8mb4_slovenian_ci, + "utf8mb4_polish_ci" => Collation::utf8mb4_polish_ci, + "cp1251_ukrainian_ci" => Collation::cp1251_ukrainian_ci, + "utf8mb4_estonian_ci" => Collation::utf8mb4_estonian_ci, + "utf8mb4_spanish_ci" => Collation::utf8mb4_spanish_ci, + "utf8mb4_swedish_ci" => Collation::utf8mb4_swedish_ci, + "utf8mb4_turkish_ci" => Collation::utf8mb4_turkish_ci, + "utf8mb4_czech_ci" => Collation::utf8mb4_czech_ci, + "utf8mb4_danish_ci" => Collation::utf8mb4_danish_ci, + "utf8mb4_lithuanian_ci" => Collation::utf8mb4_lithuanian_ci, + "utf8mb4_slovak_ci" => Collation::utf8mb4_slovak_ci, + "utf8mb4_spanish2_ci" => Collation::utf8mb4_spanish2_ci, + "utf8mb4_roman_ci" => Collation::utf8mb4_roman_ci, + "gb2312_chinese_ci" => Collation::gb2312_chinese_ci, + "utf8mb4_persian_ci" => Collation::utf8mb4_persian_ci, + "utf8mb4_esperanto_ci" => Collation::utf8mb4_esperanto_ci, + "utf8mb4_hungarian_ci" => Collation::utf8mb4_hungarian_ci, + "utf8mb4_sinhala_ci" => Collation::utf8mb4_sinhala_ci, + "utf8mb4_german2_ci" => Collation::utf8mb4_german2_ci, + "utf8mb4_croatian_ci" => Collation::utf8mb4_croatian_ci, + "utf8mb4_unicode_520_ci" => Collation::utf8mb4_unicode_520_ci, + "utf8mb4_vietnamese_ci" => Collation::utf8mb4_vietnamese_ci, + "gb18030_chinese_ci" => Collation::gb18030_chinese_ci, + "gb18030_bin" => Collation::gb18030_bin, + "greek_general_ci" => Collation::greek_general_ci, + "gb18030_unicode_520_ci" => Collation::gb18030_unicode_520_ci, + "utf8mb4_0900_ai_ci" => Collation::utf8mb4_0900_ai_ci, + "cp1250_general_ci" => Collation::cp1250_general_ci, + "latin2_croatian_ci" => Collation::latin2_croatian_ci, + "gbk_chinese_ci" => Collation::gbk_chinese_ci, + "cp1257_lithuanian_ci" => Collation::cp1257_lithuanian_ci, + "dec8_swedish_ci" => Collation::dec8_swedish_ci, + "latin5_turkish_ci" => Collation::latin5_turkish_ci, + "latin1_german2_ci" => Collation::latin1_german2_ci, + "armscii8_general_ci" => Collation::armscii8_general_ci, + "utf8_general_ci" => Collation::utf8_general_ci, + "cp1250_czech_cs" => Collation::cp1250_czech_cs, + "ucs2_general_ci" => Collation::ucs2_general_ci, + "cp866_general_ci" => Collation::cp866_general_ci, + "keybcs2_general_ci" => Collation::keybcs2_general_ci, + "macce_general_ci" => Collation::macce_general_ci, + "macroman_general_ci" => Collation::macroman_general_ci, + "cp850_general_ci" => Collation::cp850_general_ci, + "cp852_general_ci" => Collation::cp852_general_ci, + "latin7_general_ci" => Collation::latin7_general_ci, + "latin7_general_cs" => Collation::latin7_general_cs, + "macce_bin" => Collation::macce_bin, + "cp1250_croatian_ci" => Collation::cp1250_croatian_ci, + "utf8mb4_general_ci" => Collation::utf8mb4_general_ci, + "utf8mb4_bin" => Collation::utf8mb4_bin, + "latin1_bin" => Collation::latin1_bin, + "latin1_general_ci" => Collation::latin1_general_ci, + "latin1_general_cs" => Collation::latin1_general_cs, + "latin1_german1_ci" => Collation::latin1_german1_ci, + "cp1251_bin" => Collation::cp1251_bin, + "cp1251_general_ci" => Collation::cp1251_general_ci, + "cp1251_general_cs" => Collation::cp1251_general_cs, + "macroman_bin" => Collation::macroman_bin, + "utf16_general_ci" => Collation::utf16_general_ci, + "utf16_bin" => Collation::utf16_bin, + "utf16le_general_ci" => Collation::utf16le_general_ci, + "cp1256_general_ci" => Collation::cp1256_general_ci, + "cp1257_bin" => Collation::cp1257_bin, + "cp1257_general_ci" => Collation::cp1257_general_ci, + "hp8_english_ci" => Collation::hp8_english_ci, + "utf32_general_ci" => Collation::utf32_general_ci, + "utf32_bin" => Collation::utf32_bin, + "utf16le_bin" => Collation::utf16le_bin, + "binary" => Collation::binary, + "armscii8_bin" => Collation::armscii8_bin, + "ascii_bin" => Collation::ascii_bin, + "cp1250_bin" => Collation::cp1250_bin, + "cp1256_bin" => Collation::cp1256_bin, + "cp866_bin" => Collation::cp866_bin, + "dec8_bin" => Collation::dec8_bin, + "koi8r_general_ci" => Collation::koi8r_general_ci, + "greek_bin" => Collation::greek_bin, + "hebrew_bin" => Collation::hebrew_bin, + "hp8_bin" => Collation::hp8_bin, + "keybcs2_bin" => Collation::keybcs2_bin, + "koi8r_bin" => Collation::koi8r_bin, + "koi8u_bin" => Collation::koi8u_bin, + "utf8_tolower_ci" => Collation::utf8_tolower_ci, + "latin2_bin" => Collation::latin2_bin, + "latin5_bin" => Collation::latin5_bin, + "latin7_bin" => Collation::latin7_bin, + "latin1_swedish_ci" => Collation::latin1_swedish_ci, + "cp850_bin" => Collation::cp850_bin, + "cp852_bin" => Collation::cp852_bin, + "swe7_bin" => Collation::swe7_bin, + "utf8_bin" => Collation::utf8_bin, + "big5_bin" => Collation::big5_bin, + "euckr_bin" => Collation::euckr_bin, + "gb2312_bin" => Collation::gb2312_bin, + "gbk_bin" => Collation::gbk_bin, + "sjis_bin" => Collation::sjis_bin, + "tis620_bin" => Collation::tis620_bin, + "latin2_general_ci" => Collation::latin2_general_ci, + "ucs2_bin" => Collation::ucs2_bin, + "ujis_bin" => Collation::ujis_bin, + "geostd8_general_ci" => Collation::geostd8_general_ci, + "geostd8_bin" => Collation::geostd8_bin, + "latin1_spanish_ci" => Collation::latin1_spanish_ci, + "cp932_japanese_ci" => Collation::cp932_japanese_ci, + "cp932_bin" => Collation::cp932_bin, + "eucjpms_japanese_ci" => Collation::eucjpms_japanese_ci, + "eucjpms_bin" => Collation::eucjpms_bin, + "cp1250_polish_ci" => Collation::cp1250_polish_ci, + + _ => { + return Err(Error::Configuration( + format!("unsupported MySQL collation: {collation}").into(), + )); + } + }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/column.rs b/src-tauri/vendor/sqlx-mysql/src/column.rs new file mode 100644 index 00000000..1bb841b9 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/column.rs @@ -0,0 +1,31 @@ +use crate::ext::ustr::UStr; +use crate::protocol::text::ColumnFlags; +use crate::{MySql, MySqlTypeInfo}; +pub(crate) use sqlx_core::column::*; + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct MySqlColumn { + pub(crate) ordinal: usize, + pub(crate) name: UStr, + pub(crate) type_info: MySqlTypeInfo, + + #[cfg_attr(feature = "offline", serde(skip))] + pub(crate) flags: Option, +} + +impl Column for MySqlColumn { + type Database = MySql; + + fn ordinal(&self) -> usize { + self.ordinal + } + + fn name(&self) -> &str { + &self.name + } + + fn type_info(&self) -> &MySqlTypeInfo { + &self.type_info + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/auth.rs b/src-tauri/vendor/sqlx-mysql/src/connection/auth.rs new file mode 100644 index 00000000..613f8e70 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/auth.rs @@ -0,0 +1,197 @@ +use bytes::buf::Chain; +use bytes::Bytes; +use digest::{Digest, OutputSizeUser}; +use generic_array::GenericArray; +use rand::thread_rng; +use rsa::{pkcs8::DecodePublicKey, Oaep, RsaPublicKey}; +use sha1::Sha1; +use sha2::Sha256; + +use crate::connection::stream::MySqlStream; +use crate::error::Error; +use crate::protocol::auth::AuthPlugin; +use crate::protocol::Packet; + +impl AuthPlugin { + pub(super) async fn scramble( + self, + stream: &mut MySqlStream, + password: &str, + nonce: &Chain, + ) -> Result, Error> { + match self { + // https://mariadb.com/kb/en/caching_sha2_password-authentication-plugin/ + AuthPlugin::CachingSha2Password => Ok(scramble_sha256(password, nonce).to_vec()), + + AuthPlugin::MySqlNativePassword => Ok(scramble_sha1(password, nonce).to_vec()), + + // https://mariadb.com/kb/en/sha256_password-plugin/ + AuthPlugin::Sha256Password => encrypt_rsa(stream, 0x01, password, nonce).await, + + AuthPlugin::MySqlClearPassword => { + let mut pw_bytes = password.as_bytes().to_owned(); + pw_bytes.push(0); // null terminate + Ok(pw_bytes) + } + } + } + + pub(super) async fn handle( + self, + stream: &mut MySqlStream, + packet: Packet, + password: &str, + nonce: &Chain, + ) -> Result { + match self { + AuthPlugin::CachingSha2Password if packet[0] == 0x01 => { + match packet[1] { + // AUTH_OK + 0x03 => Ok(true), + + // AUTH_CONTINUE + 0x04 => { + let payload = encrypt_rsa(stream, 0x02, password, nonce).await?; + + stream.write_packet(&*payload)?; + stream.flush().await?; + + Ok(false) + } + + v => { + Err(err_protocol!("unexpected result from fast authentication 0x{:x} when expecting 0x03 (AUTH_OK) or 0x04 (AUTH_CONTINUE)", v)) + } + } + } + + _ => Err(err_protocol!( + "unexpected packet 0x{:02x} for auth plugin '{}' during authentication", + packet[0], + self.name() + )), + } + } +} + +fn scramble_sha1( + password: &str, + nonce: &Chain, +) -> GenericArray::OutputSize> { + // SHA1( password ) ^ SHA1( seed + SHA1( SHA1( password ) ) ) + // https://mariadb.com/kb/en/connection/#mysql_native_password-plugin + + let mut ctx = Sha1::new(); + + ctx.update(password); + + let mut pw_hash = ctx.finalize_reset(); + + ctx.update(pw_hash); + + let pw_hash_hash = ctx.finalize_reset(); + + ctx.update(nonce.first_ref()); + ctx.update(nonce.last_ref()); + ctx.update(pw_hash_hash); + + let pw_seed_hash_hash = ctx.finalize(); + + xor_eq(&mut pw_hash, &pw_seed_hash_hash); + + pw_hash +} + +fn scramble_sha256( + password: &str, + nonce: &Chain, +) -> GenericArray::OutputSize> { + // XOR(SHA256(password), SHA256(seed, SHA256(SHA256(password)))) + // https://mariadb.com/kb/en/caching_sha2_password-authentication-plugin/#sha-2-encrypted-password + let mut ctx = Sha256::new(); + + ctx.update(password); + + let mut pw_hash = ctx.finalize_reset(); + + ctx.update(pw_hash); + + let pw_hash_hash = ctx.finalize_reset(); + + ctx.update(nonce.first_ref()); + ctx.update(nonce.last_ref()); + ctx.update(pw_hash_hash); + + let pw_seed_hash_hash = ctx.finalize(); + + xor_eq(&mut pw_hash, &pw_seed_hash_hash); + + pw_hash +} + +async fn encrypt_rsa<'s>( + stream: &'s mut MySqlStream, + public_key_request_id: u8, + password: &'s str, + nonce: &'s Chain, +) -> Result, Error> { + // https://mariadb.com/kb/en/caching_sha2_password-authentication-plugin/ + + if stream.is_tls { + // If in a TLS stream, send the password directly in clear text + return Ok(to_asciz(password)); + } + + // client sends a public key request + stream.write_packet(&[public_key_request_id][..])?; + stream.flush().await?; + + // server sends a public key response + let packet = stream.recv_packet().await?; + let rsa_pub_key = &packet[1..]; + + // xor the password with the given nonce + let mut pass = to_asciz(password); + + let (a, b) = (nonce.first_ref(), nonce.last_ref()); + let mut nonce = Vec::with_capacity(a.len() + b.len()); + nonce.extend_from_slice(a); + nonce.extend_from_slice(b); + + xor_eq(&mut pass, &nonce); + + // client sends an RSA encrypted password + let pkey = parse_rsa_pub_key(rsa_pub_key)?; + let padding = Oaep::new::(); + pkey.encrypt(&mut thread_rng(), padding, &pass[..]) + .map_err(Error::protocol) +} + +// XOR(x, y) +// If len(y) < len(x), wrap around inside y +fn xor_eq(x: &mut [u8], y: &[u8]) { + let y_len = y.len(); + + for i in 0..x.len() { + x[i] ^= y[i % y_len]; + } +} + +fn to_asciz(s: &str) -> Vec { + let mut z = String::with_capacity(s.len() + 1); + z.push_str(s); + z.push('\0'); + + z.into_bytes() +} + +// https://docs.rs/rsa/0.3.0/rsa/struct.RSAPublicKey.html?search=#example-1 +fn parse_rsa_pub_key(key: &[u8]) -> Result { + let pem = std::str::from_utf8(key).map_err(Error::protocol)?; + + // This takes advantage of the knowledge that we know + // we are receiving a PKCS#8 RSA Public Key at all + // times from MySQL + + RsaPublicKey::from_public_key_pem(pem).map_err(Error::protocol) +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/establish.rs b/src-tauri/vendor/sqlx-mysql/src/connection/establish.rs new file mode 100644 index 00000000..85a9d84f --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/establish.rs @@ -0,0 +1,195 @@ +use bytes::buf::Buf; +use bytes::Bytes; + +use crate::collation::{CharSet, Collation}; +use crate::common::StatementCache; +use crate::connection::{tls, MySqlConnectionInner, MySqlStream, MAX_PACKET_SIZE}; +use crate::error::Error; +use crate::net::{Socket, WithSocket}; +use crate::protocol::connect::{ + AuthSwitchRequest, AuthSwitchResponse, Handshake, HandshakeResponse, +}; +use crate::protocol::Capabilities; +use crate::{MySqlConnectOptions, MySqlConnection, MySqlSslMode}; + +impl MySqlConnection { + pub(crate) async fn establish(options: &MySqlConnectOptions) -> Result { + let do_handshake = DoHandshake::new(options)?; + + let handshake = match &options.socket { + Some(path) => crate::net::connect_uds(path, do_handshake).await?, + None => crate::net::connect_tcp(&options.host, options.port, do_handshake).await?, + }; + + let stream = handshake?; + + Ok(Self { + inner: Box::new(MySqlConnectionInner { + stream, + transaction_depth: 0, + status_flags: Default::default(), + cache_statement: StatementCache::new(options.statement_cache_capacity), + log_settings: options.log_settings.clone(), + }), + }) + } +} + +struct DoHandshake<'a> { + options: &'a MySqlConnectOptions, + charset: CharSet, + collation: Collation, +} + +impl<'a> DoHandshake<'a> { + fn new(options: &'a MySqlConnectOptions) -> Result { + let charset: CharSet = options.charset.parse()?; + let collation: Collation = options + .collation + .as_deref() + .map(|collation| collation.parse()) + .transpose()? + .unwrap_or_else(|| charset.default_collation()); + + if options.enable_cleartext_plugin + && matches!( + options.ssl_mode, + MySqlSslMode::Disabled | MySqlSslMode::Preferred + ) + { + log::warn!("Security warning: sending cleartext passwords without requiring SSL"); + } + + Ok(Self { + options, + charset, + collation, + }) + } + + async fn do_handshake(self, socket: S) -> Result { + let DoHandshake { + options, + charset, + collation, + } = self; + + let mut stream = MySqlStream::with_socket(charset, collation, options, socket); + + // https://dev.mysql.com/doc/internals/en/connection-phase.html + // https://mariadb.com/kb/en/connection/ + + let handshake: Handshake = stream.recv_packet().await?.decode()?; + + let mut plugin = handshake.auth_plugin; + let nonce = handshake.auth_plugin_data; + + // FIXME: server version parse is a bit ugly + // expecting MAJOR.MINOR.PATCH + + let mut server_version = handshake.server_version.split('.'); + + let server_version_major: u16 = server_version + .next() + .unwrap_or_default() + .parse() + .unwrap_or(0); + + let server_version_minor: u16 = server_version + .next() + .unwrap_or_default() + .parse() + .unwrap_or(0); + + let server_version_patch: u16 = server_version + .next() + .unwrap_or_default() + .parse() + .unwrap_or(0); + + stream.server_version = ( + server_version_major, + server_version_minor, + server_version_patch, + ); + + stream.capabilities &= handshake.server_capabilities; + stream.capabilities |= Capabilities::PROTOCOL_41; + + let mut stream = tls::maybe_upgrade(stream, self.options).await?; + + let auth_response = if let (Some(plugin), Some(password)) = (plugin, &options.password) { + Some(plugin.scramble(&mut stream, password, &nonce).await?) + } else { + None + }; + + stream.write_packet(HandshakeResponse { + collation: stream.collation as u8, + max_packet_size: MAX_PACKET_SIZE, + username: &options.username, + database: options.database.as_deref(), + auth_plugin: plugin, + auth_response: auth_response.as_deref(), + })?; + + stream.flush().await?; + + loop { + let packet = stream.recv_packet().await?; + match packet[0] { + 0x00 => { + let _ok = packet.ok()?; + + break; + } + + 0xfe => { + let switch: AuthSwitchRequest = + packet.decode_with(self.options.enable_cleartext_plugin)?; + + plugin = Some(switch.plugin); + let nonce = switch.data.chain(Bytes::new()); + + let response = switch + .plugin + .scramble( + &mut stream, + options.password.as_deref().unwrap_or_default(), + &nonce, + ) + .await?; + + stream.write_packet(AuthSwitchResponse(response))?; + stream.flush().await?; + } + + id => { + if let (Some(plugin), Some(password)) = (plugin, &options.password) { + if plugin.handle(&mut stream, packet, password, &nonce).await? { + // plugin signaled authentication is ok + break; + } + + // plugin signaled to continue authentication + } else { + return Err(err_protocol!( + "unexpected packet 0x{:02x} during authentication", + id + )); + } + } + } + } + + Ok(stream) + } +} + +impl<'a> WithSocket for DoHandshake<'a> { + type Output = Result; + + async fn with_socket(self, socket: S) -> Self::Output { + self.do_handshake(socket).await + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/executor.rs b/src-tauri/vendor/sqlx-mysql/src/connection/executor.rs new file mode 100644 index 00000000..4f5af4bf --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/executor.rs @@ -0,0 +1,428 @@ +use super::MySqlStream; +use crate::connection::stream::Waiting; +use crate::describe::Describe; +use crate::error::Error; +use crate::executor::{Execute, Executor}; +use crate::ext::ustr::UStr; +use crate::io::MySqlBufExt; +use crate::logger::QueryLogger; +use crate::protocol::response::Status; +use crate::protocol::statement::{ + BinaryRow, Execute as StatementExecute, Prepare, PrepareOk, StmtClose, +}; +use crate::protocol::text::{ColumnDefinition, ColumnFlags, Query, TextRow}; +use crate::statement::{MySqlStatement, MySqlStatementMetadata}; +use crate::HashMap; +use crate::{ + MySql, MySqlArguments, MySqlColumn, MySqlConnection, MySqlQueryResult, MySqlRow, MySqlTypeInfo, + MySqlValueFormat, +}; +use either::Either; +use futures_core::future::BoxFuture; +use futures_core::stream::BoxStream; +use futures_core::Stream; +use futures_util::TryStreamExt; +use std::{borrow::Cow, pin::pin, sync::Arc}; + +impl MySqlConnection { + async fn prepare_statement<'c>( + &mut self, + sql: &str, + ) -> Result<(u32, MySqlStatementMetadata), Error> { + // https://dev.mysql.com/doc/internals/en/com-stmt-prepare.html + // https://dev.mysql.com/doc/internals/en/com-stmt-prepare-response.html#packet-COM_STMT_PREPARE_OK + + self.inner + .stream + .send_packet(Prepare { query: sql }) + .await?; + + let ok: PrepareOk = self.inner.stream.recv().await?; + + // the parameter definitions are very unreliable so we skip over them + // as we have little use + + if ok.params > 0 { + for _ in 0..ok.params { + let _def: ColumnDefinition = self.inner.stream.recv().await?; + } + + self.inner.stream.maybe_recv_eof().await?; + } + + // the column definitions are berefit the type information from the + // to-be-bound parameters; we will receive the output column definitions + // once more on execute so we wait for that + + let mut columns = Vec::new(); + + let column_names = if ok.columns > 0 { + recv_result_metadata(&mut self.inner.stream, ok.columns as usize, &mut columns).await? + } else { + Default::default() + }; + + let id = ok.statement_id; + let metadata = MySqlStatementMetadata { + parameters: ok.params as usize, + columns: Arc::new(columns), + column_names: Arc::new(column_names), + }; + + Ok((id, metadata)) + } + + async fn get_or_prepare_statement<'c>( + &mut self, + sql: &str, + ) -> Result<(u32, MySqlStatementMetadata), Error> { + if let Some(statement) = self.inner.cache_statement.get_mut(sql) { + // is internally reference-counted + return Ok((*statement).clone()); + } + + let (id, metadata) = self.prepare_statement(sql).await?; + + // in case of the cache being full, close the least recently used statement + if let Some((id, _)) = self + .inner + .cache_statement + .insert(sql, (id, metadata.clone())) + { + self.inner + .stream + .send_packet(StmtClose { statement: id }) + .await?; + } + + Ok((id, metadata)) + } + + #[allow(clippy::needless_lifetimes)] + pub(crate) async fn run<'e, 'c: 'e, 'q: 'e>( + &'c mut self, + sql: &'q str, + arguments: Option, + persistent: bool, + ) -> Result, Error>> + 'e, Error> + { + let mut logger = QueryLogger::new(sql, self.inner.log_settings.clone()); + + self.inner.stream.wait_until_ready().await?; + self.inner.stream.waiting.push_back(Waiting::Result); + + Ok(try_stream! { + // make a slot for the shared column data + // as long as a reference to a row is not held past one iteration, this enables us + // to re-use this memory freely between result sets + let mut columns = Arc::new(Vec::new()); + + let (mut column_names, format, mut needs_metadata) = if let Some(arguments) = arguments { + if persistent && self.inner.cache_statement.is_enabled() { + let (id, metadata) = self + .get_or_prepare_statement(sql) + .await?; + + // https://dev.mysql.com/doc/internals/en/com-stmt-execute.html + self.inner.stream + .send_packet(StatementExecute { + statement: id, + arguments: &arguments, + }) + .await?; + + (metadata.column_names, MySqlValueFormat::Binary, false) + } else { + let (id, metadata) = self + .prepare_statement(sql) + .await?; + + // https://dev.mysql.com/doc/internals/en/com-stmt-execute.html + self.inner.stream + .send_packet(StatementExecute { + statement: id, + arguments: &arguments, + }) + .await?; + + self.inner.stream.send_packet(StmtClose { statement: id }).await?; + + (metadata.column_names, MySqlValueFormat::Binary, false) + } + } else { + // https://dev.mysql.com/doc/internals/en/com-query.html + self.inner.stream.send_packet(Query(sql)).await?; + + (Arc::default(), MySqlValueFormat::Text, true) + }; + + loop { + // query response is a meta-packet which may be one of: + // Ok, Err, ResultSet, or (unhandled) LocalInfileRequest + let mut packet = self.inner.stream.recv_packet().await?; + + if packet[0] == 0x00 || packet[0] == 0xff { + // first packet in a query response is OK or ERR + // this indicates either a successful query with no rows at all or a failed query + let ok = packet.ok()?; + + self.inner.status_flags = ok.status; + + let rows_affected = ok.affected_rows; + logger.increase_rows_affected(rows_affected); + let done = MySqlQueryResult { + rows_affected, + last_insert_id: ok.last_insert_id, + }; + + r#yield!(Either::Left(done)); + + if ok.status.contains(Status::SERVER_MORE_RESULTS_EXISTS) { + // more result sets exist, continue to the next one + continue; + } + + self.inner.stream.waiting.pop_front(); + return Ok(()); + } + + // otherwise, this first packet is the start of the result-set metadata, + *self.inner.stream.waiting.front_mut().unwrap() = Waiting::Row; + + let num_columns = packet.get_uint_lenenc(); // column count + let num_columns = usize::try_from(num_columns) + .map_err(|_| err_protocol!("column count overflows usize: {num_columns}"))?; + + if needs_metadata { + column_names = Arc::new(recv_result_metadata(&mut self.inner.stream, num_columns, Arc::make_mut(&mut columns)).await?); + } else { + // next time we hit here, it'll be a new result set and we'll need the + // full metadata + needs_metadata = true; + + recv_result_columns(&mut self.inner.stream, num_columns, Arc::make_mut(&mut columns)).await?; + } + + // finally, there will be none or many result-rows + loop { + let packet = self.inner.stream.recv_packet().await?; + + if packet[0] == 0xfe && packet.len() < 9 { + let eof = packet.eof(self.inner.stream.capabilities)?; + + self.inner.status_flags = eof.status; + + r#yield!(Either::Left(MySqlQueryResult { + rows_affected: 0, + last_insert_id: 0, + })); + + if eof.status.contains(Status::SERVER_MORE_RESULTS_EXISTS) { + // more result sets exist, continue to the next one + *self.inner.stream.waiting.front_mut().unwrap() = Waiting::Result; + break; + } + + self.inner.stream.waiting.pop_front(); + return Ok(()); + } + + let row = match format { + MySqlValueFormat::Binary => packet.decode_with::(&columns)?.0, + MySqlValueFormat::Text => packet.decode_with::(&columns)?.0, + }; + + let v = Either::Right(MySqlRow { + row, + format, + columns: Arc::clone(&columns), + column_names: Arc::clone(&column_names), + }); + + logger.increment_rows_returned(); + + r#yield!(v); + } + } + }) + } +} + +impl<'c> Executor<'c> for &'c mut MySqlConnection { + type Database = MySql; + + fn fetch_many<'e, 'q, E>( + self, + mut query: E, + ) -> BoxStream<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + let sql = query.sql(); + let arguments = query.take_arguments().map_err(Error::Encode); + let persistent = query.persistent(); + + Box::pin(try_stream! { + let arguments = arguments?; + let mut s = pin!(self.run(sql, arguments, persistent).await?); + + while let Some(v) = s.try_next().await? { + r#yield!(v); + } + + Ok(()) + }) + } + + fn fetch_optional<'e, 'q, E>(self, query: E) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + let mut s = self.fetch_many(query); + + Box::pin(async move { + while let Some(v) = s.try_next().await? { + if let Either::Right(r) = v { + return Ok(Some(r)); + } + } + + Ok(None) + }) + } + + fn prepare_with<'e, 'q: 'e>( + self, + sql: &'q str, + _parameters: &'e [MySqlTypeInfo], + ) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + Box::pin(async move { + self.inner.stream.wait_until_ready().await?; + + let metadata = if self.inner.cache_statement.is_enabled() { + self.get_or_prepare_statement(sql).await?.1 + } else { + let (id, metadata) = self.prepare_statement(sql).await?; + + self.inner + .stream + .send_packet(StmtClose { statement: id }) + .await?; + + metadata + }; + + Ok(MySqlStatement { + sql: Cow::Borrowed(sql), + // metadata has internal Arcs for expensive data structures + metadata: metadata.clone(), + }) + }) + } + + #[doc(hidden)] + fn describe<'e, 'q: 'e>(self, sql: &'q str) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + Box::pin(async move { + self.inner.stream.wait_until_ready().await?; + + let (id, metadata) = self.prepare_statement(sql).await?; + + self.inner + .stream + .send_packet(StmtClose { statement: id }) + .await?; + + let columns = (*metadata.columns).clone(); + + let nullable = columns + .iter() + .map(|col| { + col.flags + .map(|flags| !flags.contains(ColumnFlags::NOT_NULL)) + }) + .collect(); + + Ok(Describe { + parameters: Some(Either::Right(metadata.parameters)), + columns, + nullable, + }) + }) + } +} + +async fn recv_result_columns( + stream: &mut MySqlStream, + num_columns: usize, + columns: &mut Vec, +) -> Result<(), Error> { + columns.clear(); + columns.reserve(num_columns); + + for ordinal in 0..num_columns { + columns.push(recv_next_result_column(&stream.recv().await?, ordinal)?); + } + + if num_columns > 0 { + stream.maybe_recv_eof().await?; + } + + Ok(()) +} + +fn recv_next_result_column(def: &ColumnDefinition, ordinal: usize) -> Result { + // if the alias is empty, use the alias + // only then use the name + let name = match (def.name()?, def.alias()?) { + (_, alias) if !alias.is_empty() => UStr::new(alias), + (name, _) => UStr::new(name), + }; + + let type_info = MySqlTypeInfo::from_column(def); + + Ok(MySqlColumn { + name, + type_info, + ordinal, + flags: Some(def.flags), + }) +} + +async fn recv_result_metadata( + stream: &mut MySqlStream, + num_columns: usize, + columns: &mut Vec, +) -> Result, Error> { + // the result-set metadata is primarily a listing of each output + // column in the result-set + + let mut column_names = HashMap::with_capacity(num_columns); + + columns.clear(); + columns.reserve(num_columns); + + for ordinal in 0..num_columns { + let def: ColumnDefinition = stream.recv().await?; + + let column = recv_next_result_column(&def, ordinal)?; + + column_names.insert(column.name.clone(), ordinal); + columns.push(column); + } + + stream.maybe_recv_eof().await?; + + Ok(column_names) +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/mod.rs b/src-tauri/vendor/sqlx-mysql/src/connection/mod.rs new file mode 100644 index 00000000..ee5a483d --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/mod.rs @@ -0,0 +1,148 @@ +use std::borrow::Cow; +use std::fmt::{self, Debug, Formatter}; + +use futures_core::future::BoxFuture; +use futures_util::FutureExt; +pub(crate) use sqlx_core::connection::*; +pub(crate) use stream::{MySqlStream, Waiting}; + +use crate::common::StatementCache; +use crate::error::Error; +use crate::protocol::response::Status; +use crate::protocol::statement::StmtClose; +use crate::protocol::text::{Ping, Quit}; +use crate::statement::MySqlStatementMetadata; +use crate::transaction::Transaction; +use crate::{MySql, MySqlConnectOptions}; + +mod auth; +mod establish; +mod executor; +mod stream; +mod tls; + +const MAX_PACKET_SIZE: u32 = 1024; + +/// A connection to a MySQL database. +pub struct MySqlConnection { + pub(crate) inner: Box, +} + +pub(crate) struct MySqlConnectionInner { + // underlying TCP stream, + // wrapped in a potentially TLS stream, + // wrapped in a buffered stream + pub(crate) stream: MySqlStream, + + // transaction status + pub(crate) transaction_depth: usize, + status_flags: Status, + + // cache by query string to the statement id and metadata + cache_statement: StatementCache<(u32, MySqlStatementMetadata)>, + + log_settings: LogSettings, +} + +impl MySqlConnection { + pub(crate) fn in_transaction(&self) -> bool { + self.inner + .status_flags + .intersects(Status::SERVER_STATUS_IN_TRANS) + } +} + +impl Debug for MySqlConnection { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.debug_struct("MySqlConnection").finish() + } +} + +impl Connection for MySqlConnection { + type Database = MySql; + + type Options = MySqlConnectOptions; + + fn close(mut self) -> BoxFuture<'static, Result<(), Error>> { + Box::pin(async move { + self.inner.stream.send_packet(Quit).await?; + self.inner.stream.shutdown().await?; + + Ok(()) + }) + } + + fn close_hard(mut self) -> BoxFuture<'static, Result<(), Error>> { + Box::pin(async move { + self.inner.stream.shutdown().await?; + Ok(()) + }) + } + + fn ping(&mut self) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + // Stroke patch, same as sqlx-postgres: skip the release ping (a + // round trip to TiDB, PlanetScale or Railway after every query) + // when nothing is queued or unread. A dropped stream or a queued + // ROLLBACK leaves `waiting` non-empty and still takes the full path. + if self.inner.stream.waiting.is_empty() && self.inner.stream.write_buffer_is_empty() { + return Ok(()); + } + self.inner.stream.wait_until_ready().await?; + self.inner.stream.send_packet(Ping).await?; + self.inner.stream.recv_ok().await?; + + Ok(()) + }) + } + + #[doc(hidden)] + fn flush(&mut self) -> BoxFuture<'_, Result<(), Error>> { + self.inner.stream.wait_until_ready().boxed() + } + + fn cached_statements_size(&self) -> usize { + self.inner.cache_statement.len() + } + + fn clear_cached_statements(&mut self) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + while let Some((statement_id, _)) = self.inner.cache_statement.remove_lru() { + self.inner + .stream + .send_packet(StmtClose { + statement: statement_id, + }) + .await?; + } + + Ok(()) + }) + } + + #[doc(hidden)] + fn should_flush(&self) -> bool { + !self.inner.stream.write_buffer().is_empty() + } + + fn begin(&mut self) -> BoxFuture<'_, Result, Error>> + where + Self: Sized, + { + Transaction::begin(self, None) + } + + fn begin_with( + &mut self, + statement: impl Into>, + ) -> BoxFuture<'_, Result, Error>> + where + Self: Sized, + { + Transaction::begin(self, Some(statement.into())) + } + + fn shrink_buffers(&mut self) { + self.inner.stream.shrink_buffers(); + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/stream.rs b/src-tauri/vendor/sqlx-mysql/src/connection/stream.rs new file mode 100644 index 00000000..b0e3ff2f --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/stream.rs @@ -0,0 +1,248 @@ +use std::collections::VecDeque; +use std::ops::{Deref, DerefMut}; + +use bytes::{Buf, Bytes, BytesMut}; + +use crate::collation::{CharSet, Collation}; +use crate::error::Error; +use crate::io::MySqlBufExt; +use crate::io::{ProtocolDecode, ProtocolEncode}; +use crate::net::{BufferedSocket, Socket}; +use crate::protocol::response::{EofPacket, ErrPacket, OkPacket, Status}; +use crate::protocol::{Capabilities, Packet}; +use crate::{MySqlConnectOptions, MySqlDatabaseError}; + +pub struct MySqlStream> { + // Wrapping the socket in `Box` allows us to unsize in-place. + pub(crate) socket: BufferedSocket, + pub(crate) server_version: (u16, u16, u16), + pub(super) capabilities: Capabilities, + pub(crate) sequence_id: u8, + pub(crate) waiting: VecDeque, + pub(crate) charset: CharSet, + pub(crate) collation: Collation, + pub(crate) is_tls: bool, +} + +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum Waiting { + // waiting for a result set + Result, + + // waiting for a row within a result set + Row, +} + +impl MySqlStream { + pub(crate) fn with_socket( + charset: CharSet, + collation: Collation, + options: &MySqlConnectOptions, + socket: S, + ) -> Self { + let mut capabilities = Capabilities::PROTOCOL_41 + | Capabilities::IGNORE_SPACE + | Capabilities::DEPRECATE_EOF + | Capabilities::FOUND_ROWS + | Capabilities::TRANSACTIONS + | Capabilities::SECURE_CONNECTION + | Capabilities::PLUGIN_AUTH_LENENC_DATA + | Capabilities::MULTI_STATEMENTS + | Capabilities::MULTI_RESULTS + | Capabilities::PLUGIN_AUTH + | Capabilities::PS_MULTI_RESULTS + | Capabilities::SSL; + + if options.database.is_some() { + capabilities |= Capabilities::CONNECT_WITH_DB; + } + + Self { + waiting: VecDeque::new(), + capabilities, + server_version: (0, 0, 0), + sequence_id: 0, + collation, + charset, + socket: BufferedSocket::new(socket), + is_tls: false, + } + } + + pub(crate) fn write_buffer_is_empty(&self) -> bool { + self.socket.write_buffer().is_empty() + } + + pub(crate) async fn wait_until_ready(&mut self) -> Result<(), Error> { + if !self.socket.write_buffer().is_empty() { + self.socket.flush().await?; + } + + while !self.waiting.is_empty() { + while self.waiting.front() == Some(&Waiting::Row) { + let packet = self.recv_packet().await?; + + if !packet.is_empty() && packet[0] == 0xfe && packet.len() < 9 { + let eof = packet.eof(self.capabilities)?; + + if eof.status.contains(Status::SERVER_MORE_RESULTS_EXISTS) { + *self.waiting.front_mut().unwrap() = Waiting::Result; + } else { + self.waiting.pop_front(); + }; + } + } + + while self.waiting.front() == Some(&Waiting::Result) { + let packet = self.recv_packet().await?; + + if !packet.is_empty() && (packet[0] == 0x00 || packet[0] == 0xff) { + let ok = packet.ok()?; + + if !ok.status.contains(Status::SERVER_MORE_RESULTS_EXISTS) { + self.waiting.pop_front(); + } + } else { + *self.waiting.front_mut().unwrap() = Waiting::Row; + self.skip_result_metadata(packet).await?; + } + } + } + + Ok(()) + } + + pub(crate) async fn send_packet<'en, T>(&mut self, payload: T) -> Result<(), Error> + where + T: ProtocolEncode<'en, Capabilities>, + { + self.sequence_id = 0; + self.write_packet(payload)?; + self.flush().await?; + Ok(()) + } + + pub(crate) fn write_packet<'en, T>(&mut self, payload: T) -> Result<(), Error> + where + T: ProtocolEncode<'en, Capabilities>, + { + self.socket + .write_with(Packet(payload), (self.capabilities, &mut self.sequence_id)) + } + + async fn recv_packet_part(&mut self) -> Result { + // https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_basic_packets.html + // https://mariadb.com/kb/en/library/0-packet/#standard-packet + + let mut header: Bytes = self.socket.read(4).await?; + + // cannot overflow + #[allow(clippy::cast_possible_truncation)] + let packet_size = header.get_uint_le(3) as usize; + let sequence_id = header.get_u8(); + + self.sequence_id = sequence_id.wrapping_add(1); + + let payload: Bytes = self.socket.read(packet_size).await?; + + // TODO: packet compression + + Ok(payload) + } + + // receive the next packet from the database server + // may block (async) on more data from the server + pub(crate) async fn recv_packet(&mut self) -> Result, Error> { + let payload = self.recv_packet_part().await?; + let payload = if payload.len() < 0xFF_FF_FF { + payload + } else { + let mut final_payload = BytesMut::with_capacity(0xFF_FF_FF * 2); + final_payload.extend_from_slice(&payload); + + drop(payload); // we don't need the allocation anymore + + let mut last_read = 0xFF_FF_FF; + while last_read == 0xFF_FF_FF { + let part = self.recv_packet_part().await?; + last_read = part.len(); + final_payload.extend_from_slice(&part); + } + final_payload.into() + }; + + if payload + .first() + .ok_or(err_protocol!("Packet empty"))? + .eq(&0xff) + { + self.waiting.pop_front(); + + // instead of letting this packet be looked at everywhere, we check here + // and emit a proper Error + return Err( + MySqlDatabaseError(ErrPacket::decode_with(payload, self.capabilities)?).into(), + ); + } + + Ok(Packet(payload)) + } + + pub(crate) async fn recv<'de, T>(&mut self) -> Result + where + T: ProtocolDecode<'de, Capabilities>, + { + self.recv_packet().await?.decode_with(self.capabilities) + } + + pub(crate) async fn recv_ok(&mut self) -> Result { + self.recv_packet().await?.ok() + } + + pub(crate) async fn maybe_recv_eof(&mut self) -> Result, Error> { + if self.capabilities.contains(Capabilities::DEPRECATE_EOF) { + Ok(None) + } else { + self.recv().await.map(Some) + } + } + + async fn skip_result_metadata(&mut self, mut packet: Packet) -> Result<(), Error> { + let num_columns: u64 = packet.get_uint_lenenc(); // column count + + for _ in 0..num_columns { + let _ = self.recv_packet().await?; + } + + self.maybe_recv_eof().await?; + + Ok(()) + } + + pub fn boxed_socket(self) -> MySqlStream { + MySqlStream { + socket: self.socket.boxed(), + server_version: self.server_version, + capabilities: self.capabilities, + sequence_id: self.sequence_id, + waiting: self.waiting, + charset: self.charset, + collation: self.collation, + is_tls: self.is_tls, + } + } +} + +impl Deref for MySqlStream { + type Target = BufferedSocket; + + fn deref(&self) -> &Self::Target { + &self.socket + } +} + +impl DerefMut for MySqlStream { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.socket + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/connection/tls.rs b/src-tauri/vendor/sqlx-mysql/src/connection/tls.rs new file mode 100644 index 00000000..eb077c62 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/connection/tls.rs @@ -0,0 +1,109 @@ +use crate::collation::{CharSet, Collation}; +use crate::connection::{MySqlStream, Waiting}; +use crate::error::Error; +use crate::net::tls::TlsConfig; +use crate::net::{tls, BufferedSocket, Socket, WithSocket}; +use crate::protocol::connect::SslRequest; +use crate::protocol::Capabilities; +use crate::{MySqlConnectOptions, MySqlSslMode}; +use std::collections::VecDeque; + +struct MapStream { + server_version: (u16, u16, u16), + capabilities: Capabilities, + sequence_id: u8, + waiting: VecDeque, + charset: CharSet, + collation: Collation, +} + +pub(super) async fn maybe_upgrade( + mut stream: MySqlStream, + options: &MySqlConnectOptions, +) -> Result { + let server_supports_tls = stream.capabilities.contains(Capabilities::SSL); + + if matches!(options.ssl_mode, MySqlSslMode::Disabled) || !tls::available() { + // remove the SSL capability if SSL has been explicitly disabled + stream.capabilities.remove(Capabilities::SSL); + } + + // https://www.postgresql.org/docs/12/libpq-ssl.html#LIBPQ-SSL-SSLMODE-STATEMENTS + match options.ssl_mode { + MySqlSslMode::Disabled => return Ok(stream.boxed_socket()), + + MySqlSslMode::Preferred => { + if !tls::available() { + // Client doesn't support TLS + tracing::debug!("not performing TLS upgrade: TLS support not compiled in"); + return Ok(stream.boxed_socket()); + } + + if !server_supports_tls { + // Server doesn't support TLS + tracing::debug!("not performing TLS upgrade: unsupported by server"); + return Ok(stream.boxed_socket()); + } + } + + MySqlSslMode::Required | MySqlSslMode::VerifyIdentity | MySqlSslMode::VerifyCa => { + tls::error_if_unavailable()?; + + if !server_supports_tls { + // upgrade failed, die + return Err(Error::Tls("server does not support TLS".into())); + } + } + } + + let tls_config = TlsConfig { + accept_invalid_certs: !matches!( + options.ssl_mode, + MySqlSslMode::VerifyCa | MySqlSslMode::VerifyIdentity + ), + accept_invalid_hostnames: !matches!(options.ssl_mode, MySqlSslMode::VerifyIdentity), + hostname: &options.host, + root_cert_path: options.ssl_ca.as_ref(), + client_cert_path: options.ssl_client_cert.as_ref(), + client_key_path: options.ssl_client_key.as_ref(), + }; + + // Request TLS upgrade + stream.write_packet(SslRequest { + max_packet_size: super::MAX_PACKET_SIZE, + collation: stream.collation as u8, + })?; + + stream.flush().await?; + + tls::handshake( + stream.socket.into_inner(), + tls_config, + MapStream { + server_version: stream.server_version, + capabilities: stream.capabilities, + sequence_id: stream.sequence_id, + waiting: stream.waiting, + charset: stream.charset, + collation: stream.collation, + }, + ) + .await +} + +impl WithSocket for MapStream { + type Output = MySqlStream; + + async fn with_socket(self, socket: S) -> Self::Output { + MySqlStream { + socket: BufferedSocket::new(Box::new(socket)), + server_version: self.server_version, + capabilities: self.capabilities, + sequence_id: self.sequence_id, + waiting: self.waiting, + charset: self.charset, + collation: self.collation, + is_tls: true, + } + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/database.rs b/src-tauri/vendor/sqlx-mysql/src/database.rs new file mode 100644 index 00000000..d03a5672 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/database.rs @@ -0,0 +1,38 @@ +use crate::value::{MySqlValue, MySqlValueRef}; +use crate::{ + MySqlArguments, MySqlColumn, MySqlConnection, MySqlQueryResult, MySqlRow, MySqlStatement, + MySqlTransactionManager, MySqlTypeInfo, +}; +pub(crate) use sqlx_core::database::{Database, HasStatementCache}; + +/// MySQL database driver. +#[derive(Debug)] +pub struct MySql; + +impl Database for MySql { + type Connection = MySqlConnection; + + type TransactionManager = MySqlTransactionManager; + + type Row = MySqlRow; + + type QueryResult = MySqlQueryResult; + + type Column = MySqlColumn; + + type TypeInfo = MySqlTypeInfo; + + type Value = MySqlValue; + type ValueRef<'r> = MySqlValueRef<'r>; + + type Arguments<'q> = MySqlArguments; + type ArgumentBuffer<'q> = Vec; + + type Statement<'q> = MySqlStatement<'q>; + + const NAME: &'static str = "MySQL"; + + const URL_SCHEMES: &'static [&'static str] = &["mysql", "mariadb"]; +} + +impl HasStatementCache for MySql {} diff --git a/src-tauri/vendor/sqlx-mysql/src/error.rs b/src-tauri/vendor/sqlx-mysql/src/error.rs new file mode 100644 index 00000000..e7363399 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/error.rs @@ -0,0 +1,177 @@ +use std::error::Error as StdError; +use std::fmt::{self, Debug, Display, Formatter}; + +use crate::protocol::response::ErrPacket; + +use std::borrow::Cow; + +pub(crate) use sqlx_core::error::*; + +/// An error returned from the MySQL database. +pub struct MySqlDatabaseError(pub(super) ErrPacket); + +impl MySqlDatabaseError { + /// The [SQLSTATE](https://dev.mysql.com/doc/mysql-errors/8.0/en/server-error-reference.html) code for this error. + pub fn code(&self) -> Option<&str> { + self.0.sql_state.as_deref() + } + + /// The [number](https://dev.mysql.com/doc/mysql-errors/8.0/en/server-error-reference.html) + /// for this error. + /// + /// MySQL tends to use SQLSTATE as a general error category, and the error number as a more + /// granular indication of the error. + pub fn number(&self) -> u16 { + self.0.error_code + } + + /// The human-readable error message. + pub fn message(&self) -> &str { + &self.0.error_message + } +} + +impl Debug for MySqlDatabaseError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.debug_struct("MySqlDatabaseError") + .field("code", &self.code()) + .field("number", &self.number()) + .field("message", &self.message()) + .finish() + } +} + +impl Display for MySqlDatabaseError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + if let Some(code) = &self.code() { + write!(f, "{} ({}): {}", self.number(), code, self.message()) + } else { + write!(f, "{}: {}", self.number(), self.message()) + } + } +} + +impl StdError for MySqlDatabaseError {} + +impl DatabaseError for MySqlDatabaseError { + #[inline] + fn message(&self) -> &str { + self.message() + } + + #[inline] + fn code(&self) -> Option> { + self.code().map(Cow::Borrowed) + } + + #[doc(hidden)] + fn as_error(&self) -> &(dyn StdError + Send + Sync + 'static) { + self + } + + #[doc(hidden)] + fn as_error_mut(&mut self) -> &mut (dyn StdError + Send + Sync + 'static) { + self + } + + #[doc(hidden)] + fn into_error(self: Box) -> Box { + self + } + + fn kind(&self) -> ErrorKind { + match self.number() { + error_codes::ER_DUP_KEY + | error_codes::ER_DUP_ENTRY + | error_codes::ER_DUP_UNIQUE + | error_codes::ER_DUP_ENTRY_WITH_KEY_NAME + | error_codes::ER_DUP_UNKNOWN_IN_INDEX => ErrorKind::UniqueViolation, + + error_codes::ER_NO_REFERENCED_ROW + | error_codes::ER_NO_REFERENCED_ROW_2 + | error_codes::ER_ROW_IS_REFERENCED + | error_codes::ER_ROW_IS_REFERENCED_2 + | error_codes::ER_FK_COLUMN_NOT_NULL + | error_codes::ER_FK_CANNOT_DELETE_PARENT => ErrorKind::ForeignKeyViolation, + + error_codes::ER_BAD_NULL_ERROR | error_codes::ER_NO_DEFAULT_FOR_FIELD => { + ErrorKind::NotNullViolation + } + + error_codes::ER_CHECK_CONSTRAINT_VIOLATED => ErrorKind::CheckViolation, + + // https://mariadb.com/kb/en/e4025/ + error_codes::mariadb::ER_CONSTRAINT_FAILED + // MySQL uses this code for a completely different error, + // but we can differentiate by SQLSTATE: + // + { + ErrorKind::CheckViolation + } + + _ => ErrorKind::Other, + } + } +} + +/// The MySQL server uses SQLSTATEs as a generic error category, +/// and returns a `error_code` instead within the error packet. +/// +/// For reference: . +pub(crate) mod error_codes { + /// Caused when a DDL operation creates duplicated keys. + pub const ER_DUP_KEY: u16 = 1022; + /// Caused when a DML operation tries create a duplicated entry for a key, + /// be it a unique or primary one. + pub const ER_DUP_ENTRY: u16 = 1062; + /// Similar to `ER_DUP_ENTRY`, but only present in NDB clusters. + /// + /// See: . + pub const ER_DUP_UNIQUE: u16 = 1169; + /// Similar to `ER_DUP_ENTRY`, but with a formatted string message. + /// + /// See: . + pub const ER_DUP_ENTRY_WITH_KEY_NAME: u16 = 1586; + /// Caused when a DDL operation to add a unique index fails, + /// because duplicate items were created by concurrent DML operations. + /// When this happens, the key is unknown, so the server can't use `ER_DUP_KEY`. + /// + /// For example: an `INSERT` operation creates duplicate `name` fields when `ALTER`ing a table and making `name` unique. + pub const ER_DUP_UNKNOWN_IN_INDEX: u16 = 1859; + + /// Caused when inserting an entry with a column with a value that does not reference a foreign row. + pub const ER_NO_REFERENCED_ROW: u16 = 1216; + /// Caused when deleting a row that is referenced in other tables. + pub const ER_ROW_IS_REFERENCED: u16 = 1217; + /// Caused when deleting a row that is referenced in other tables. + /// This differs from `ER_ROW_IS_REFERENCED` in that the error message contains the affected constraint. + pub const ER_ROW_IS_REFERENCED_2: u16 = 1451; + /// Caused when inserting an entry with a column with a value that does not reference a foreign row. + /// This differs from `ER_NO_REFERENCED_ROW` in that the error message contains the affected constraint. + pub const ER_NO_REFERENCED_ROW_2: u16 = 1452; + /// Caused when creating a FK with `ON DELETE SET NULL` or `ON UPDATE SET NULL` to a column that is `NOT NULL`, or vice-versa. + pub const ER_FK_COLUMN_NOT_NULL: u16 = 1830; + /// Removed in 5.7.3. + pub const ER_FK_CANNOT_DELETE_PARENT: u16 = 1834; + + /// Caused when inserting a NULL value to a column marked as NOT NULL. + pub const ER_BAD_NULL_ERROR: u16 = 1048; + /// Caused when inserting a DEFAULT value to a column marked as NOT NULL, which also doesn't have a default value set. + pub const ER_NO_DEFAULT_FOR_FIELD: u16 = 1364; + + /// Caused when a check constraint is violated. + /// + /// Only available after 8.0.16. + pub const ER_CHECK_CONSTRAINT_VIOLATED: u16 = 3819; + + pub(crate) mod mariadb { + /// Error code emitted by MariaDB for constraint errors: + /// + /// MySQL emits this code for a completely different error: + /// + /// + /// You also check that SQLSTATE is `23000`. + pub const ER_CONSTRAINT_FAILED: u16 = 4025; + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/io/buf.rs b/src-tauri/vendor/sqlx-mysql/src/io/buf.rs new file mode 100644 index 00000000..685d5bfd --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/io/buf.rs @@ -0,0 +1,47 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::BufExt; + +pub trait MySqlBufExt: Buf { + // Read a length-encoded integer. + // NOTE: 0xfb or NULL is only returned for binary value encoding to indicate NULL. + // NOTE: 0xff is only returned during a result set to indicate ERR. + // + fn get_uint_lenenc(&mut self) -> u64; + + // Read a length-encoded string. + #[allow(dead_code)] + fn get_str_lenenc(&mut self) -> Result; + + // Read a length-encoded byte sequence. + fn get_bytes_lenenc(&mut self) -> Result; +} + +impl MySqlBufExt for Bytes { + fn get_uint_lenenc(&mut self) -> u64 { + match self.get_u8() { + 0xfc => u64::from(self.get_u16_le()), + 0xfd => self.get_uint_le(3), + 0xfe => self.get_u64_le(), + + v => u64::from(v), + } + } + + fn get_str_lenenc(&mut self) -> Result { + let size = self.get_uint_lenenc(); + let size = usize::try_from(size) + .map_err(|_| err_protocol!("string length overflows usize: {size}"))?; + + self.get_str(size) + } + + fn get_bytes_lenenc(&mut self) -> Result { + let size = self.get_uint_lenenc(); + let size = usize::try_from(size) + .map_err(|_| err_protocol!("string length overflows usize: {size}"))?; + + Ok(self.split_to(size)) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/io/buf_mut.rs b/src-tauri/vendor/sqlx-mysql/src/io/buf_mut.rs new file mode 100644 index 00000000..e40148e7 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/io/buf_mut.rs @@ -0,0 +1,131 @@ +use bytes::BufMut; + +pub trait MySqlBufMutExt: BufMut { + fn put_uint_lenenc(&mut self, v: u64); + + fn put_str_lenenc(&mut self, v: &str); + + fn put_bytes_lenenc(&mut self, v: &[u8]); +} + +impl MySqlBufMutExt for Vec { + fn put_uint_lenenc(&mut self, v: u64) { + // https://dev.mysql.com/doc/internals/en/integer.html + // https://mariadb.com/kb/en/library/protocol-data-types/#length-encoded-integers + + let encoded_le = v.to_le_bytes(); + + match v { + 0..=250 => self.push(encoded_le[0]), + 251..=0xFF_FF => { + self.push(0xfc); + self.extend_from_slice(&encoded_le[..2]); + } + 0x1_00_00..=0xFF_FF_FF => { + self.push(0xfd); + self.extend_from_slice(&encoded_le[..3]); + } + _ => { + self.push(0xfe); + self.extend_from_slice(&encoded_le); + } + } + } + + fn put_str_lenenc(&mut self, v: &str) { + self.put_bytes_lenenc(v.as_bytes()); + } + + fn put_bytes_lenenc(&mut self, v: &[u8]) { + self.put_uint_lenenc(v.len() as u64); + self.extend(v); + } +} + +#[test] +fn test_encodes_int_lenenc_u8() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFA as u64); + + assert_eq!(&buf[..], b"\xFA"); +} + +#[test] +fn test_encodes_int_lenenc_u16() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(std::u16::MAX as u64); + + assert_eq!(&buf[..], b"\xFC\xFF\xFF"); +} + +#[test] +fn test_encodes_int_lenenc_u24() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFF_FF_FF as u64); + + assert_eq!(&buf[..], b"\xFD\xFF\xFF\xFF"); +} + +#[test] +fn test_encodes_int_lenenc_u64() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(std::u64::MAX); + + assert_eq!(&buf[..], b"\xFE\xFF\xFF\xFF\xFF\xFF\xFF\xFF\xFF"); +} + +#[test] +fn test_encodes_int_lenenc_fb() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFB as u64); + + assert_eq!(&buf[..], b"\xFC\xFB\x00"); +} + +#[test] +fn test_encodes_int_lenenc_fc() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFC as u64); + + assert_eq!(&buf[..], b"\xFC\xFC\x00"); +} + +#[test] +fn test_encodes_int_lenenc_fd() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFD as u64); + + assert_eq!(&buf[..], b"\xFC\xFD\x00"); +} + +#[test] +fn test_encodes_int_lenenc_fe() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFE as u64); + + assert_eq!(&buf[..], b"\xFC\xFE\x00"); +} + +#[test] +fn test_encodes_int_lenenc_ff() { + let mut buf = Vec::with_capacity(1024); + buf.put_uint_lenenc(0xFF as u64); + + assert_eq!(&buf[..], b"\xFC\xFF\x00"); +} + +#[test] +fn test_encodes_string_lenenc() { + let mut buf = Vec::with_capacity(1024); + buf.put_str_lenenc("random_string"); + + assert_eq!(&buf[..], b"\x0Drandom_string"); +} + +#[test] +fn test_encodes_byte_lenenc() { + let mut buf = Vec::with_capacity(1024); + buf.put_bytes_lenenc(b"random_string"); + + assert_eq!(&buf[..], b"\x0Drandom_string"); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/io/mod.rs b/src-tauri/vendor/sqlx-mysql/src/io/mod.rs new file mode 100644 index 00000000..cf48e8ab --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/io/mod.rs @@ -0,0 +1,7 @@ +mod buf; +mod buf_mut; + +pub use buf::MySqlBufExt; +pub use buf_mut::MySqlBufMutExt; + +pub(crate) use sqlx_core::io::*; diff --git a/src-tauri/vendor/sqlx-mysql/src/lib.rs b/src-tauri/vendor/sqlx-mysql/src/lib.rs new file mode 100644 index 00000000..7aa14256 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/lib.rs @@ -0,0 +1,73 @@ +//! **MySQL** database driver. +#![deny(clippy::cast_possible_truncation)] +#![deny(clippy::cast_possible_wrap)] +#![deny(clippy::cast_sign_loss)] + +#[macro_use] +extern crate sqlx_core; + +use crate::executor::Executor; + +pub(crate) use sqlx_core::driver_prelude::*; + +#[cfg(feature = "any")] +pub mod any; + +mod arguments; +mod collation; +mod column; +mod connection; +mod database; +mod error; +mod io; +mod options; +mod protocol; +mod query_result; +mod row; +mod statement; +mod transaction; +mod type_checking; +mod type_info; +pub mod types; +mod value; + +#[cfg(feature = "migrate")] +mod migrate; + +#[cfg(feature = "migrate")] +mod testing; + +pub use arguments::MySqlArguments; +pub use column::MySqlColumn; +pub use connection::MySqlConnection; +pub use database::MySql; +pub use error::MySqlDatabaseError; +pub use options::{MySqlConnectOptions, MySqlSslMode}; +pub use query_result::MySqlQueryResult; +pub use row::MySqlRow; +pub use statement::MySqlStatement; +pub use transaction::MySqlTransactionManager; +pub use type_info::MySqlTypeInfo; +pub use value::{MySqlValue, MySqlValueFormat, MySqlValueRef}; + +/// An alias for [`Pool`][crate::pool::Pool], specialized for MySQL. +pub type MySqlPool = crate::pool::Pool; + +/// An alias for [`PoolOptions`][crate::pool::PoolOptions], specialized for MySQL. +pub type MySqlPoolOptions = crate::pool::PoolOptions; + +/// An alias for [`Executor<'_, Database = MySql>`][Executor]. +pub trait MySqlExecutor<'c>: Executor<'c, Database = MySql> {} +impl<'c, T: Executor<'c, Database = MySql>> MySqlExecutor<'c> for T {} + +/// An alias for [`Transaction`][crate::transaction::Transaction], specialized for MySQL. +pub type MySqlTransaction<'c> = crate::transaction::Transaction<'c, MySql>; + +// NOTE: required due to the lack of lazy normalization +impl_into_arguments_for_arguments!(MySqlArguments); +impl_acquire!(MySql, MySqlConnection); +impl_column_index_for_row!(MySqlRow); +impl_column_index_for_statement!(MySqlStatement); + +// required because some databases have a different handling of NULL +impl_encode_for_option!(MySql); diff --git a/src-tauri/vendor/sqlx-mysql/src/migrate.rs b/src-tauri/vendor/sqlx-mysql/src/migrate.rs new file mode 100644 index 00000000..79b55ace --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/migrate.rs @@ -0,0 +1,302 @@ +use std::str::FromStr; +use std::time::Duration; +use std::time::Instant; + +use futures_core::future::BoxFuture; +pub(crate) use sqlx_core::migrate::*; + +use crate::connection::{ConnectOptions, Connection}; +use crate::error::Error; +use crate::executor::Executor; +use crate::query::query; +use crate::query_as::query_as; +use crate::query_scalar::query_scalar; +use crate::{MySql, MySqlConnectOptions, MySqlConnection}; + +fn parse_for_maintenance(url: &str) -> Result<(MySqlConnectOptions, String), Error> { + let mut options = MySqlConnectOptions::from_str(url)?; + + let database = if let Some(database) = &options.database { + database.to_owned() + } else { + return Err(Error::Configuration( + "DATABASE_URL does not specify a database".into(), + )); + }; + + // switch us to database for create/drop commands + options.database = None; + + Ok((options, database)) +} + +impl MigrateDatabase for MySql { + fn create_database(url: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let _ = conn + .execute(&*format!("CREATE DATABASE `{database}`")) + .await?; + + Ok(()) + }) + } + + fn database_exists(url: &str) -> BoxFuture<'_, Result> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let exists: bool = query_scalar( + "select exists(SELECT 1 from INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME = ?)", + ) + .bind(database) + .fetch_one(&mut conn) + .await?; + + Ok(exists) + }) + } + + fn drop_database(url: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let _ = conn + .execute(&*format!("DROP DATABASE IF EXISTS `{database}`")) + .await?; + + Ok(()) + }) + } +} + +impl Migrate for MySqlConnection { + fn ensure_migrations_table(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + // language=MySQL + self.execute( + r#" +CREATE TABLE IF NOT EXISTS _sqlx_migrations ( + version BIGINT PRIMARY KEY, + description TEXT NOT NULL, + installed_on TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + success BOOLEAN NOT NULL, + checksum BLOB NOT NULL, + execution_time BIGINT NOT NULL +); + "#, + ) + .await?; + + Ok(()) + }) + } + + fn dirty_version(&mut self) -> BoxFuture<'_, Result, MigrateError>> { + Box::pin(async move { + // language=SQL + let row: Option<(i64,)> = query_as( + "SELECT version FROM _sqlx_migrations WHERE success = false ORDER BY version LIMIT 1", + ) + .fetch_optional(self) + .await?; + + Ok(row.map(|r| r.0)) + }) + } + + fn list_applied_migrations( + &mut self, + ) -> BoxFuture<'_, Result, MigrateError>> { + Box::pin(async move { + // language=SQL + let rows: Vec<(i64, Vec)> = + query_as("SELECT version, checksum FROM _sqlx_migrations ORDER BY version") + .fetch_all(self) + .await?; + + let migrations = rows + .into_iter() + .map(|(version, checksum)| AppliedMigration { + version, + checksum: checksum.into(), + }) + .collect(); + + Ok(migrations) + }) + } + + fn lock(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + let database_name = current_database(self).await?; + let lock_id = generate_lock_id(&database_name); + + // create an application lock over the database + // this function will not return until the lock is acquired + + // https://www.postgresql.org/docs/current/explicit-locking.html#ADVISORY-LOCKS + // https://www.postgresql.org/docs/current/functions-admin.html#FUNCTIONS-ADVISORY-LOCKS-TABLE + + // language=MySQL + let _ = query("SELECT GET_LOCK(?, -1)") + .bind(lock_id) + .execute(self) + .await?; + + Ok(()) + }) + } + + fn unlock(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + let database_name = current_database(self).await?; + let lock_id = generate_lock_id(&database_name); + + // language=MySQL + let _ = query("SELECT RELEASE_LOCK(?)") + .bind(lock_id) + .execute(self) + .await?; + + Ok(()) + }) + } + + fn apply<'e: 'm, 'm>( + &'e mut self, + migration: &'m Migration, + ) -> BoxFuture<'m, Result> { + Box::pin(async move { + // Use a single transaction for the actual migration script and the essential bookeeping so we never + // execute migrations twice. See https://github.com/launchbadge/sqlx/issues/1966. + // The `execution_time` however can only be measured for the whole transaction. This value _only_ exists for + // data lineage and debugging reasons, so it is not super important if it is lost. So we initialize it to -1 + // and update it once the actual transaction completed. + let mut tx = self.begin().await?; + let start = Instant::now(); + + // For MySQL we cannot really isolate migrations due to implicit commits caused by table modification, see + // https://dev.mysql.com/doc/refman/8.0/en/implicit-commit.html + // + // To somewhat try to detect this, we first insert the migration into the migration table with + // `success=FALSE` and later modify the flag. + // + // language=MySQL + let _ = query( + r#" + INSERT INTO _sqlx_migrations ( version, description, success, checksum, execution_time ) + VALUES ( ?, ?, FALSE, ?, -1 ) + "#, + ) + .bind(migration.version) + .bind(&*migration.description) + .bind(&*migration.checksum) + .execute(&mut *tx) + .await?; + + let _ = tx + .execute(&*migration.sql) + .await + .map_err(|e| MigrateError::ExecuteMigration(e, migration.version))?; + + // language=MySQL + let _ = query( + r#" + UPDATE _sqlx_migrations + SET success = TRUE + WHERE version = ? + "#, + ) + .bind(migration.version) + .execute(&mut *tx) + .await?; + + tx.commit().await?; + + // Update `elapsed_time`. + // NOTE: The process may disconnect/die at this point, so the elapsed time value might be lost. We accept + // this small risk since this value is not super important. + + let elapsed = start.elapsed(); + + #[allow(clippy::cast_possible_truncation)] + let _ = query( + r#" + UPDATE _sqlx_migrations + SET execution_time = ? + WHERE version = ? + "#, + ) + .bind(elapsed.as_nanos() as i64) + .bind(migration.version) + .execute(self) + .await?; + + Ok(elapsed) + }) + } + + fn revert<'e: 'm, 'm>( + &'e mut self, + migration: &'m Migration, + ) -> BoxFuture<'m, Result> { + Box::pin(async move { + // Use a single transaction for the actual migration script and the essential bookeeping so we never + // execute migrations twice. See https://github.com/launchbadge/sqlx/issues/1966. + let mut tx = self.begin().await?; + let start = Instant::now(); + + // For MySQL we cannot really isolate migrations due to implicit commits caused by table modification, see + // https://dev.mysql.com/doc/refman/8.0/en/implicit-commit.html + // + // To somewhat try to detect this, we first insert the migration into the migration table with + // `success=FALSE` and later remove the migration altogether. + // + // language=MySQL + let _ = query( + r#" + UPDATE _sqlx_migrations + SET success = FALSE + WHERE version = ? + "#, + ) + .bind(migration.version) + .execute(&mut *tx) + .await?; + + tx.execute(&*migration.sql).await?; + + // language=SQL + let _ = query(r#"DELETE FROM _sqlx_migrations WHERE version = ?"#) + .bind(migration.version) + .execute(&mut *tx) + .await?; + + tx.commit().await?; + + let elapsed = start.elapsed(); + + Ok(elapsed) + }) + } +} + +async fn current_database(conn: &mut MySqlConnection) -> Result { + // language=MySQL + Ok(query_scalar("SELECT DATABASE()").fetch_one(conn).await?) +} + +// inspired from rails: https://github.com/rails/rails/blob/6e49cc77ab3d16c06e12f93158eaf3e507d4120e/activerecord/lib/active_record/migration.rb#L1308 +fn generate_lock_id(database_name: &str) -> String { + const CRC_IEEE: crc::Crc = crc::Crc::::new(&crc::CRC_32_ISO_HDLC); + // 0x3d32ad9e chosen by fair dice roll + format!( + "{:x}", + 0x3d32ad9e * (CRC_IEEE.checksum(database_name.as_bytes()) as i64) + ) +} diff --git a/src-tauri/vendor/sqlx-mysql/src/options/connect.rs b/src-tauri/vendor/sqlx-mysql/src/options/connect.rs new file mode 100644 index 00000000..116a49cc --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/options/connect.rs @@ -0,0 +1,97 @@ +use crate::connection::ConnectOptions; +use crate::error::Error; +use crate::executor::Executor; +use crate::{MySqlConnectOptions, MySqlConnection}; +use futures_core::future::BoxFuture; +use log::LevelFilter; +use sqlx_core::Url; +use std::time::Duration; + +impl ConnectOptions for MySqlConnectOptions { + type Connection = MySqlConnection; + + fn from_url(url: &Url) -> Result { + Self::parse_from_url(url) + } + + fn to_url_lossy(&self) -> Url { + self.build_url() + } + + fn connect(&self) -> BoxFuture<'_, Result> + where + Self::Connection: Sized, + { + Box::pin(async move { + let mut conn = MySqlConnection::establish(self).await?; + + // After the connection is established, we initialize by configuring a few + // connection parameters + + // https://mariadb.com/kb/en/sql-mode/ + + // PIPES_AS_CONCAT - Allows using the pipe character (ASCII 124) as string concatenation operator. + // This means that "A" || "B" can be used in place of CONCAT("A", "B"). + + // NO_ENGINE_SUBSTITUTION - If not set, if the available storage engine specified by a CREATE TABLE is + // not available, a warning is given and the default storage + // engine is used instead. + + // NO_ZERO_DATE - Don't allow '0000-00-00'. This is invalid in Rust. + + // NO_ZERO_IN_DATE - Don't allow 'YYYY-00-00'. This is invalid in Rust. + + // -- + + // Setting the time zone allows us to assume that the output + // from a TIMESTAMP field is UTC + + // -- + + // https://mathiasbynens.be/notes/mysql-utf8mb4 + + let mut sql_mode = Vec::new(); + if self.pipes_as_concat { + sql_mode.push(r#"PIPES_AS_CONCAT"#); + } + if self.no_engine_substitution { + sql_mode.push(r#"NO_ENGINE_SUBSTITUTION"#); + } + + let mut options = Vec::new(); + if !sql_mode.is_empty() { + options.push(format!( + r#"sql_mode=(SELECT CONCAT(@@sql_mode, ',{}'))"#, + sql_mode.join(",") + )); + } + if let Some(timezone) = &self.timezone { + options.push(format!(r#"time_zone='{}'"#, timezone)); + } + if self.set_names { + options.push(format!( + r#"NAMES {} COLLATE {}"#, + conn.inner.stream.charset.as_str(), + conn.inner.stream.collation.as_str() + )) + } + + if !options.is_empty() { + conn.execute(&*format!(r#"SET {};"#, options.join(","))) + .await?; + } + + Ok(conn) + }) + } + + fn log_statements(mut self, level: LevelFilter) -> Self { + self.log_settings.log_statements(level); + self + } + + fn log_slow_statements(mut self, level: LevelFilter, duration: Duration) -> Self { + self.log_settings.log_slow_statements(level, duration); + self + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/options/mod.rs b/src-tauri/vendor/sqlx-mysql/src/options/mod.rs new file mode 100644 index 00000000..87732cb4 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/options/mod.rs @@ -0,0 +1,518 @@ +use std::path::{Path, PathBuf}; + +mod connect; +mod parse; +mod ssl_mode; + +use crate::{connection::LogSettings, net::tls::CertificateInput}; +pub use ssl_mode::MySqlSslMode; + +/// Options and flags which can be used to configure a MySQL connection. +/// +/// A value of `MySqlConnectOptions` can be parsed from a connection URL, +/// as described by [MySQL](https://dev.mysql.com/doc/connector-j/8.0/en/connector-j-reference-jdbc-url-format.html). +/// +/// The generic format of the connection URL: +/// +/// ```text +/// mysql://[host][/database][?properties] +/// ``` +/// +/// This type also implements [`FromStr`][std::str::FromStr] so you can parse it from a string +/// containing a connection URL and then further adjust options if necessary (see example below). +/// +/// ## Properties +/// +/// |Parameter|Default|Description| +/// |---------|-------|-----------| +/// | `ssl-mode` | `PREFERRED` | Determines whether or with what priority a secure SSL TCP/IP connection will be negotiated. See [`MySqlSslMode`]. | +/// | `ssl-ca` | `None` | Sets the name of a file containing a list of trusted SSL Certificate Authorities. | +/// | `statement-cache-capacity` | `100` | The maximum number of prepared statements stored in the cache. Set to `0` to disable. | +/// | `socket` | `None` | Path to the unix domain socket, which will be used instead of TCP if set. | +/// +/// # Example +/// +/// ```rust,no_run +/// # async fn example() -> sqlx::Result<()> { +/// use sqlx::{Connection, ConnectOptions}; +/// use sqlx::mysql::{MySqlConnectOptions, MySqlConnection, MySqlPool, MySqlSslMode}; +/// +/// // URL connection string +/// let conn = MySqlConnection::connect("mysql://root:password@localhost/db").await?; +/// +/// // Manually-constructed options +/// let conn = MySqlConnectOptions::new() +/// .host("localhost") +/// .username("root") +/// .password("password") +/// .database("db") +/// .connect().await?; +/// +/// // Modifying options parsed from a string +/// let mut opts: MySqlConnectOptions = "mysql://root:password@localhost/db".parse()?; +/// +/// // Change the log verbosity level for queries. +/// // Information about SQL queries is logged at `DEBUG` level by default. +/// opts = opts.log_statements(log::LevelFilter::Trace); +/// +/// let pool = MySqlPool::connect_with(opts).await?; +/// # Ok(()) +/// # } +/// ``` +#[derive(Debug, Clone)] +pub struct MySqlConnectOptions { + pub(crate) host: String, + pub(crate) port: u16, + pub(crate) socket: Option, + pub(crate) username: String, + pub(crate) password: Option, + pub(crate) database: Option, + pub(crate) ssl_mode: MySqlSslMode, + pub(crate) ssl_ca: Option, + pub(crate) ssl_client_cert: Option, + pub(crate) ssl_client_key: Option, + pub(crate) statement_cache_capacity: usize, + pub(crate) charset: String, + pub(crate) collation: Option, + pub(crate) log_settings: LogSettings, + pub(crate) pipes_as_concat: bool, + pub(crate) enable_cleartext_plugin: bool, + pub(crate) no_engine_substitution: bool, + pub(crate) timezone: Option, + pub(crate) set_names: bool, +} + +impl Default for MySqlConnectOptions { + fn default() -> Self { + Self::new() + } +} + +impl MySqlConnectOptions { + /// Creates a new, default set of options ready for configuration + pub fn new() -> Self { + Self { + port: 3306, + host: String::from("localhost"), + socket: None, + username: String::from("root"), + password: None, + database: None, + charset: String::from("utf8mb4"), + collation: None, + ssl_mode: MySqlSslMode::Preferred, + ssl_ca: None, + ssl_client_cert: None, + ssl_client_key: None, + statement_cache_capacity: 100, + log_settings: Default::default(), + pipes_as_concat: true, + enable_cleartext_plugin: false, + no_engine_substitution: true, + timezone: Some(String::from("+00:00")), + set_names: true, + } + } + + /// Sets the name of the host to connect to. + /// + /// The default behavior when the host is not specified, + /// is to connect to localhost. + pub fn host(mut self, host: &str) -> Self { + host.clone_into(&mut self.host); + self + } + + /// Sets the port to connect to at the server host. + /// + /// The default port for MySQL is `3306`. + pub fn port(mut self, port: u16) -> Self { + self.port = port; + self + } + + /// Pass a path to a Unix socket. This changes the connection stream from + /// TCP to UDS. + /// + /// By default set to `None`. + pub fn socket(mut self, path: impl AsRef) -> Self { + self.socket = Some(path.as_ref().to_path_buf()); + self + } + + /// Sets the username to connect as. + pub fn username(mut self, username: &str) -> Self { + username.clone_into(&mut self.username); + self + } + + /// Sets the password to connect with. + pub fn password(mut self, password: &str) -> Self { + self.password = Some(password.to_owned()); + self + } + + /// Sets the database name. + pub fn database(mut self, database: &str) -> Self { + self.database = Some(database.to_owned()); + self + } + + /// Sets whether or with what priority a secure SSL TCP/IP connection will be negotiated + /// with the server. + /// + /// By default, the SSL mode is [`Preferred`](MySqlSslMode::Preferred), and the client will + /// first attempt an SSL connection but fallback to a non-SSL connection on failure. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::Required); + /// ``` + pub fn ssl_mode(mut self, mode: MySqlSslMode) -> Self { + self.ssl_mode = mode; + self + } + + /// Sets the name of a file containing a list of trusted SSL Certificate Authorities. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_ca("path/to/ca.crt"); + /// ``` + pub fn ssl_ca(mut self, file_name: impl AsRef) -> Self { + self.ssl_ca = Some(CertificateInput::File(file_name.as_ref().to_owned())); + self + } + + /// Sets PEM encoded list of trusted SSL Certificate Authorities. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_ca_from_pem(vec![]); + /// ``` + pub fn ssl_ca_from_pem(mut self, pem_certificate: Vec) -> Self { + self.ssl_ca = Some(CertificateInput::Inline(pem_certificate)); + self + } + + /// Sets the name of a file containing SSL client certificate. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_client_cert("path/to/client.crt"); + /// ``` + pub fn ssl_client_cert(mut self, cert: impl AsRef) -> Self { + self.ssl_client_cert = Some(CertificateInput::File(cert.as_ref().to_path_buf())); + self + } + + /// Sets the SSL client certificate as a PEM-encoded byte slice. + /// + /// This should be an ASCII-encoded blob that starts with `-----BEGIN CERTIFICATE-----`. + /// + /// # Example + /// Note: embedding SSL certificates and keys in the binary is not advised. + /// This is for illustration purposes only. + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// + /// const CERT: &[u8] = b"\ + /// -----BEGIN CERTIFICATE----- + /// + /// -----END CERTIFICATE-----"; + /// + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_client_cert_from_pem(CERT); + /// ``` + pub fn ssl_client_cert_from_pem(mut self, cert: impl AsRef<[u8]>) -> Self { + self.ssl_client_cert = Some(CertificateInput::Inline(cert.as_ref().to_vec())); + self + } + + /// Sets the name of a file containing SSL client key. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_client_key("path/to/client.key"); + /// ``` + pub fn ssl_client_key(mut self, key: impl AsRef) -> Self { + self.ssl_client_key = Some(CertificateInput::File(key.as_ref().to_path_buf())); + self + } + + /// Sets the SSL client key as a PEM-encoded byte slice. + /// + /// This should be an ASCII-encoded blob that starts with `-----BEGIN PRIVATE KEY-----`. + /// + /// # Example + /// Note: embedding SSL certificates and keys in the binary is not advised. + /// This is for illustration purposes only. + /// + /// ```rust + /// # use sqlx_mysql::{MySqlSslMode, MySqlConnectOptions}; + /// + /// const KEY: &[u8] = b"\ + /// -----BEGIN PRIVATE KEY----- + /// + /// -----END PRIVATE KEY-----"; + /// + /// let options = MySqlConnectOptions::new() + /// .ssl_mode(MySqlSslMode::VerifyCa) + /// .ssl_client_key_from_pem(KEY); + /// ``` + pub fn ssl_client_key_from_pem(mut self, key: impl AsRef<[u8]>) -> Self { + self.ssl_client_key = Some(CertificateInput::Inline(key.as_ref().to_vec())); + self + } + + /// Sets the capacity of the connection's statement cache in a number of stored + /// distinct statements. Caching is handled using LRU, meaning when the + /// amount of queries hits the defined limit, the oldest statement will get + /// dropped. + /// + /// The default cache capacity is 100 statements. + pub fn statement_cache_capacity(mut self, capacity: usize) -> Self { + self.statement_cache_capacity = capacity; + self + } + + /// Sets the character set for the connection. + /// + /// The default character set is `utf8mb4`. This is supported from MySQL 5.5.3. + /// If you need to connect to an older version, we recommend you to change this to `utf8`. + pub fn charset(mut self, charset: &str) -> Self { + charset.clone_into(&mut self.charset); + self + } + + /// Sets the collation for the connection. + /// + /// The default collation is derived from the `charset`. Normally, you should only have to set + /// the `charset`. + pub fn collation(mut self, collation: &str) -> Self { + self.collation = Some(collation.to_owned()); + self + } + + /// Sets the flag that enables or disables the `PIPES_AS_CONCAT` connection setting + /// + /// The default value is set to true, but some MySql databases such as PlanetScale + /// error out with this connection setting so it needs to be set false in such + /// cases. + pub fn pipes_as_concat(mut self, flag_val: bool) -> Self { + self.pipes_as_concat = flag_val; + self + } + + /// Enables mysql_clear_password plugin support. + /// + /// Security Note: + /// Sending passwords as cleartext may be a security problem in some + /// configurations. Without additional defensive configuration like + /// ssl-mode=VERIFY_IDENTITY, an attacker can compromise a router + /// and trick the application into divulging its credentials. + /// + /// It is strongly recommended to set `.ssl_mode` to `Required`, + /// `VerifyCa`, or `VerifyIdentity` when enabling cleartext plugin. + pub fn enable_cleartext_plugin(mut self, flag_val: bool) -> Self { + self.enable_cleartext_plugin = flag_val; + self + } + + #[deprecated = "renamed to .no_engine_substitution()"] + pub fn no_engine_subsitution(self, flag_val: bool) -> Self { + self.no_engine_substitution(flag_val) + } + + /// Flag that enables or disables the `NO_ENGINE_SUBSTITUTION` sql_mode setting after + /// connection. + /// + /// If not set, if the available storage engine specified by a `CREATE TABLE` is not available, + /// a warning is given and the default storage engine is used instead. + /// + /// By default, this is `true` (`NO_ENGINE_SUBSTITUTION` is passed, forbidding engine + /// substitution). + /// + /// + pub fn no_engine_substitution(mut self, flag_val: bool) -> Self { + self.no_engine_substitution = flag_val; + self + } + + /// If `Some`, sets the `time_zone` option to the given string after connecting to the database. + /// + /// If `None`, no `time_zone` parameter is sent; the server timezone will be used instead. + /// + /// Defaults to `Some(String::from("+00:00"))` to ensure all timestamps are in UTC. + /// + /// ### Warning + /// Changing this setting from its default will apply an unexpected skew to any + /// `time::OffsetDateTime` or `chrono::DateTime` value, whether passed as a parameter or + /// decoded as a result. `TIMESTAMP` values are not encoded with their UTC offset in the MySQL + /// protocol, so encoding and decoding of these types assumes the server timezone is *always* + /// UTC. + /// + /// If you are changing this option, ensure your application only uses + /// `time::PrimitiveDateTime` or `chrono::NaiveDateTime` and that it does not assume these + /// timestamps can be placed on a real timeline without applying the proper offset. + pub fn timezone(mut self, value: impl Into>) -> Self { + self.timezone = value.into(); + self + } + + /// If enabled, `SET NAMES '{charset}' COLLATE '{collation}'` is passed with the values of + /// [`.charset()`] and [`.collation()`] after connecting to the database. + /// + /// This ensures the connection uses the specified character set and collation. + /// + /// Enabled by default. + /// + /// ### Warning + /// If this is disabled and the default charset is not binary-compatible with UTF-8, query + /// strings, column names and string values will likely not decode (or encode) correctly, which + /// may result in unexpected errors or garbage outputs at runtime. + /// + /// For proper functioning, you *must* ensure the server is using a binary-compatible charset, + /// such as ASCII or Latin-1 (ISO 8859-1), and that you do not pass any strings containing + /// codepoints not supported by said charset. + /// + /// Instead of disabling this, you may also consider setting [`.charset()`] to a charset that + /// is supported by your MySQL or MariaDB server version and compatible with UTF-8. + pub fn set_names(mut self, flag_val: bool) -> Self { + self.set_names = flag_val; + self + } +} + +impl MySqlConnectOptions { + /// Get the current host. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .host("127.0.0.1"); + /// assert_eq!(options.get_host(), "127.0.0.1"); + /// ``` + pub fn get_host(&self) -> &str { + &self.host + } + + /// Get the server's port. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .port(6543); + /// assert_eq!(options.get_port(), 6543); + /// ``` + pub fn get_port(&self) -> u16 { + self.port + } + + /// Get the socket path. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .socket("/tmp"); + /// assert!(options.get_socket().is_some()); + /// ``` + pub fn get_socket(&self) -> Option<&PathBuf> { + self.socket.as_ref() + } + + /// Get the current username. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .username("foo"); + /// assert_eq!(options.get_username(), "foo"); + /// ``` + pub fn get_username(&self) -> &str { + &self.username + } + + /// Get the current database name. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .database("postgres"); + /// assert!(options.get_database().is_some()); + /// ``` + pub fn get_database(&self) -> Option<&str> { + self.database.as_deref() + } + + /// Get the SSL mode. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::{MySqlConnectOptions, MySqlSslMode}; + /// let options = MySqlConnectOptions::new(); + /// assert!(matches!(options.get_ssl_mode(), MySqlSslMode::Preferred)); + /// ``` + pub fn get_ssl_mode(&self) -> MySqlSslMode { + self.ssl_mode + } + + /// Get the server charset. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new(); + /// assert_eq!(options.get_charset(), "utf8mb4"); + /// ``` + pub fn get_charset(&self) -> &str { + &self.charset + } + + /// Get the server collation. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let options = MySqlConnectOptions::new() + /// .collation("collation"); + /// assert!(options.get_collation().is_some()); + /// ``` + pub fn get_collation(&self) -> Option<&str> { + self.collation.as_deref() + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/options/parse.rs b/src-tauri/vendor/sqlx-mysql/src/options/parse.rs new file mode 100644 index 00000000..e31ddc46 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/options/parse.rs @@ -0,0 +1,199 @@ +use std::str::FromStr; + +use percent_encoding::{percent_decode_str, utf8_percent_encode, NON_ALPHANUMERIC}; +use sqlx_core::Url; + +use crate::{error::Error, MySqlSslMode}; + +use super::MySqlConnectOptions; + +impl MySqlConnectOptions { + pub(crate) fn parse_from_url(url: &Url) -> Result { + let mut options = Self::new(); + + if let Some(host) = url.host_str() { + options = options.host(host); + } + + if let Some(port) = url.port() { + options = options.port(port); + } + + let username = url.username(); + if !username.is_empty() { + options = options.username( + &percent_decode_str(username) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + if let Some(password) = url.password() { + options = options.password( + &percent_decode_str(password) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + let path = url.path().trim_start_matches('/'); + if !path.is_empty() { + options = options.database( + &percent_decode_str(path) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + for (key, value) in url.query_pairs().into_iter() { + match &*key { + "sslmode" | "ssl-mode" => { + options = options.ssl_mode(value.parse().map_err(Error::config)?); + } + + "sslca" | "ssl-ca" => { + options = options.ssl_ca(&*value); + } + + "charset" => { + options = options.charset(&value); + } + + "collation" => { + options = options.collation(&value); + } + + "sslcert" | "ssl-cert" => options = options.ssl_client_cert(&*value), + + "sslkey" | "ssl-key" => options = options.ssl_client_key(&*value), + + "statement-cache-capacity" => { + options = + options.statement_cache_capacity(value.parse().map_err(Error::config)?); + } + + "socket" => { + options = options.socket(&*value); + } + + "timezone" | "time-zone" => { + options = options.timezone(Some(value.to_string())); + } + + _ => {} + } + } + + Ok(options) + } + + pub(crate) fn build_url(&self) -> Url { + let mut url = Url::parse(&format!( + "mysql://{}@{}:{}", + self.username, self.host, self.port + )) + .expect("BUG: generated un-parseable URL"); + + if let Some(password) = &self.password { + let password = utf8_percent_encode(password, NON_ALPHANUMERIC).to_string(); + let _ = url.set_password(Some(&password)); + } + + if let Some(database) = &self.database { + url.set_path(database); + } + + let ssl_mode = match self.ssl_mode { + MySqlSslMode::Disabled => "DISABLED", + MySqlSslMode::Preferred => "PREFERRED", + MySqlSslMode::Required => "REQUIRED", + MySqlSslMode::VerifyCa => "VERIFY_CA", + MySqlSslMode::VerifyIdentity => "VERIFY_IDENTITY", + }; + url.query_pairs_mut().append_pair("ssl-mode", ssl_mode); + + if let Some(ssl_ca) = &self.ssl_ca { + url.query_pairs_mut() + .append_pair("ssl-ca", &ssl_ca.to_string()); + } + + url.query_pairs_mut().append_pair("charset", &self.charset); + + if let Some(collation) = &self.collation { + url.query_pairs_mut().append_pair("charset", collation); + } + + if let Some(ssl_client_cert) = &self.ssl_client_cert { + url.query_pairs_mut() + .append_pair("ssl-cert", &ssl_client_cert.to_string()); + } + + if let Some(ssl_client_key) = &self.ssl_client_key { + url.query_pairs_mut() + .append_pair("ssl-key", &ssl_client_key.to_string()); + } + + url.query_pairs_mut().append_pair( + "statement-cache-capacity", + &self.statement_cache_capacity.to_string(), + ); + + if let Some(socket) = &self.socket { + url.query_pairs_mut() + .append_pair("socket", &socket.to_string_lossy()); + } + + url + } +} + +impl FromStr for MySqlConnectOptions { + type Err = Error; + + fn from_str(s: &str) -> Result { + let url: Url = s.parse().map_err(Error::config)?; + Self::parse_from_url(&url) + } +} + +#[test] +fn it_parses_username_with_at_sign_correctly() { + let url = "mysql://user@hostname:password@hostname:5432/database"; + let opts = MySqlConnectOptions::from_str(url).unwrap(); + + assert_eq!("user@hostname", &opts.username); +} + +#[test] +fn it_parses_password_with_non_ascii_chars_correctly() { + let url = "mysql://username:p@ssw0rd@hostname:5432/database"; + let opts = MySqlConnectOptions::from_str(url).unwrap(); + + assert_eq!(Some("p@ssw0rd".into()), opts.password); +} + +#[test] +fn it_returns_the_parsed_url() { + let url = "mysql://username:p@ssw0rd@hostname:3306/database"; + let opts = MySqlConnectOptions::from_str(url).unwrap(); + + let mut expected_url = Url::parse(url).unwrap(); + // MySqlConnectOptions defaults + let query_string = "ssl-mode=PREFERRED&charset=utf8mb4&statement-cache-capacity=100"; + expected_url.set_query(Some(query_string)); + + assert_eq!(expected_url, opts.build_url()); +} + +#[test] +fn it_parses_timezone() { + let opts: MySqlConnectOptions = "mysql://user:password@hostname/database?timezone=%2B08:00" + .parse() + .unwrap(); + assert_eq!(opts.timezone.as_deref(), Some("+08:00")); + + let opts: MySqlConnectOptions = "mysql://user:password@hostname/database?time-zone=%2B08:00" + .parse() + .unwrap(); + assert_eq!(opts.timezone.as_deref(), Some("+08:00")); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/options/ssl_mode.rs b/src-tauri/vendor/sqlx-mysql/src/options/ssl_mode.rs new file mode 100644 index 00000000..1238a3e8 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/options/ssl_mode.rs @@ -0,0 +1,52 @@ +use crate::error::Error; +use std::str::FromStr; + +/// Options for controlling the desired security state of the connection to the MySQL server. +/// +/// It is used by the [`ssl_mode`](super::MySqlConnectOptions::ssl_mode) method. +#[derive(Debug, Clone, Copy, Default)] +pub enum MySqlSslMode { + /// Establish an unencrypted connection. + Disabled, + + /// Establish an encrypted connection if the server supports encrypted connections, falling + /// back to an unencrypted connection if an encrypted connection cannot be established. + /// + /// This is the default if `ssl_mode` is not specified. + #[default] + Preferred, + + /// Establish an encrypted connection if the server supports encrypted connections. + /// The connection attempt fails if an encrypted connection cannot be established. + Required, + + /// Like `Required`, but additionally verify the server Certificate Authority (CA) + /// certificate against the configured CA certificates. The connection attempt fails + /// if no valid matching CA certificates are found. + VerifyCa, + + /// Like `VerifyCa`, but additionally perform host name identity verification by + /// checking the host name the client uses for connecting to the server against the + /// identity in the certificate that the server sends to the client. + VerifyIdentity, +} + +impl FromStr for MySqlSslMode { + type Err = Error; + + fn from_str(s: &str) -> Result { + Ok(match &*s.to_ascii_lowercase() { + "disabled" => MySqlSslMode::Disabled, + "preferred" => MySqlSslMode::Preferred, + "required" => MySqlSslMode::Required, + "verify_ca" => MySqlSslMode::VerifyCa, + "verify_identity" => MySqlSslMode::VerifyIdentity, + + _ => { + return Err(Error::Configuration( + format!("unknown value {s:?} for `ssl_mode`").into(), + )); + } + }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/auth.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/auth.rs new file mode 100644 index 00000000..ef1ce4b7 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/auth.rs @@ -0,0 +1,39 @@ +use std::str::FromStr; + +use crate::error::Error; + +#[derive(Debug, Copy, Clone)] +// These have all the same suffix but they match the auth plugin names. +#[allow(clippy::enum_variant_names)] +pub enum AuthPlugin { + MySqlNativePassword, + CachingSha2Password, + Sha256Password, + MySqlClearPassword, +} + +impl AuthPlugin { + pub(crate) fn name(self) -> &'static str { + match self { + AuthPlugin::MySqlNativePassword => "mysql_native_password", + AuthPlugin::CachingSha2Password => "caching_sha2_password", + AuthPlugin::Sha256Password => "sha256_password", + AuthPlugin::MySqlClearPassword => "mysql_clear_password", + } + } +} + +impl FromStr for AuthPlugin { + type Err = Error; + + fn from_str(s: &str) -> Result { + match s { + "mysql_native_password" => Ok(AuthPlugin::MySqlNativePassword), + "caching_sha2_password" => Ok(AuthPlugin::CachingSha2Password), + "sha256_password" => Ok(AuthPlugin::Sha256Password), + "mysql_clear_password" => Ok(AuthPlugin::MySqlClearPassword), + + _ => Err(err_protocol!("unknown authentication plugin: {}", s)), + } + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/capabilities.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/capabilities.rs new file mode 100644 index 00000000..a9c5cc58 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/capabilities.rs @@ -0,0 +1,111 @@ +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/group__group__cs__capabilities__flags.html +// https://mariadb.com/kb/en/library/connection/#capabilities +// +// MySQL defines the capabilities flags as fitting in an `int<4>` but MariaDB +// extends this with more bits sent in a separate field. +// For simplicity, we've chosen to combine these into one type. +bitflags::bitflags! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + pub struct Capabilities: u64 { + // [MariaDB] MySQL compatibility + const MYSQL = 1; + + // [*] Send found rows instead of affected rows in EOF_Packet. + const FOUND_ROWS = 2; + + // Get all column flags. + const LONG_FLAG = 4; + + // [*] Database (schema) name can be specified on connect in Handshake Response Packet. + const CONNECT_WITH_DB = 8; + + // Don't allow database.table.column + const NO_SCHEMA = 16; + + // [*] Compression protocol supported + const COMPRESS = 32; + + // Special handling of ODBC behavior. + const ODBC = 64; + + // Can use LOAD DATA LOCAL + const LOCAL_FILES = 128; + + // [*] Ignore spaces before '(' + const IGNORE_SPACE = 256; + + // [*] New 4.1+ protocol + const PROTOCOL_41 = 512; + + // This is an interactive client + const INTERACTIVE = 1024; + + // Use SSL encryption for this session + const SSL = 2048; + + // Client knows about transactions + const TRANSACTIONS = 8192; + + // 4.1+ authentication + const SECURE_CONNECTION = 1 << 15; + + // Enable/disable multi-statement support for COM_QUERY *and* COM_STMT_PREPARE + const MULTI_STATEMENTS = 1 << 16; + + // Enable/disable multi-results for COM_QUERY + const MULTI_RESULTS = 1 << 17; + + // Enable/disable multi-results for COM_STMT_PREPARE + const PS_MULTI_RESULTS = 1 << 18; + + // Client supports plugin authentication + const PLUGIN_AUTH = 1 << 19; + + // Client supports connection attributes + const CONNECT_ATTRS = 1 << 20; + + // Enable authentication response packet to be larger than 255 bytes. + const PLUGIN_AUTH_LENENC_DATA = 1 << 21; + + // Don't close the connection for a user account with expired password. + const CAN_HANDLE_EXPIRED_PASSWORDS = 1 << 22; + + // Capable of handling server state change information. + const SESSION_TRACK = 1 << 23; + + // Client no longer needs EOF_Packet and will use OK_Packet instead. + const DEPRECATE_EOF = 1 << 24; + + // Support ZSTD protocol compression + const ZSTD_COMPRESSION_ALGORITHM = 1 << 26; + + // Verify server certificate + const SSL_VERIFY_SERVER_CERT = 1 << 30; + + // The client can handle optional metadata information in the resultset + const OPTIONAL_RESULTSET_METADATA = 1 << 25; + + // Don't reset the options after an unsuccessful connect + const REMEMBER_OPTIONS = 1 << 31; + + // Extended capabilities (MariaDB only, as of writing) + // Client support progress indicator (since 10.2) + const MARIADB_CLIENT_PROGRESS = 1 << 32; + + // Permit COM_MULTI protocol + const MARIADB_CLIENT_MULTI = 1 << 33; + + // Permit bulk insert + const MARIADB_CLIENT_STMT_BULK_OPERATIONS = 1 << 34; + + // Add extended metadata information + const MARIADB_CLIENT_EXTENDED_TYPE_INFO = 1 << 35; + + // Permit skipping metadata + const MARIADB_CLIENT_CACHE_METADATA = 1 << 36; + + // when enabled, indicate that Bulk command can use STMT_BULK_FLAG_SEND_UNIT_RESULTS flag + // that permit to return a result-set of all affected rows and auto-increment values + const MARIADB_CLIENT_BULK_UNIT_RESULTS = 1 << 37; + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/connect/auth_switch.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/auth_switch.rs new file mode 100644 index 00000000..e61d26d7 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/auth_switch.rs @@ -0,0 +1,103 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::ProtocolEncode; +use crate::io::{BufExt, ProtocolDecode}; +use crate::protocol::auth::AuthPlugin; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_connection_phase_packets_protocol_auth_switch_request.html + +#[derive(Debug)] +pub struct AuthSwitchRequest { + pub plugin: AuthPlugin, + pub data: Bytes, +} + +impl ProtocolDecode<'_, bool> for AuthSwitchRequest { + fn decode_with(mut buf: Bytes, enable_cleartext_plugin: bool) -> Result { + let header = buf.get_u8(); + if header != 0xfe { + return Err(err_protocol!( + "expected 0xfe (AUTH_SWITCH) but found 0x{:x}", + header + )); + } + + let plugin = buf.get_str_nul()?.parse()?; + + if matches!(plugin, AuthPlugin::MySqlClearPassword) && !enable_cleartext_plugin { + return Err(err_protocol!("mysql_cleartext_plugin disabled")); + } + + if matches!(plugin, AuthPlugin::MySqlClearPassword) && buf.is_empty() { + // Contrary to the MySQL protocol, AWS Aurora with IAM sends + // no data. That is fine because the mysql_clear_password says to + // ignore any data sent. + // See: https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_connection_phase_authentication_methods_clear_text_password.html + return Ok(Self { + plugin, + data: Bytes::new(), + }); + } + + // See: https://github.com/mysql/mysql-server/blob/ea7d2e2d16ac03afdd9cb72a972a95981107bf51/sql/auth/sha2_password.cc#L942 + if buf.len() != 21 { + return Err(err_protocol!( + "expected 21 bytes but found {} bytes", + buf.len() + )); + } + let data = buf.get_bytes(20); + buf.advance(1); // NUL-terminator + + Ok(Self { plugin, data }) + } +} + +#[derive(Debug)] +pub struct AuthSwitchResponse(pub Vec); + +impl ProtocolEncode<'_, Capabilities> for AuthSwitchResponse { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), Error> { + buf.extend_from_slice(&self.0); + Ok(()) + } +} + +#[test] +fn test_decode_auth_switch_packet_data() { + const AUTH_SWITCH_NO_DATA: &[u8] = b"\xfecaching_sha2_password\x00abcdefghijabcdefghij\x00"; + + let p = AuthSwitchRequest::decode_with(AUTH_SWITCH_NO_DATA.into(), true).unwrap(); + + assert!(matches!(p.plugin, AuthPlugin::CachingSha2Password)); + assert_eq!(p.data, &b"abcdefghijabcdefghij"[..]); +} + +#[test] +fn test_decode_auth_switch_cleartext_disabled() { + const AUTH_SWITCH_CLEARTEXT: &[u8] = b"\xfemysql_clear_password\x00abcdefghijabcdefghij\x00"; + + let e = AuthSwitchRequest::decode_with(AUTH_SWITCH_CLEARTEXT.into(), false).unwrap_err(); + + let e_str = e.to_string(); + + let expected = "encountered unexpected or invalid data: mysql_cleartext_plugin disabled"; + + assert!( + // Don't want to assert the full string since it contains the module path now. + e_str.starts_with(expected), + "expected error string to start with {expected:?}, got {e_str:?}" + ); +} + +#[test] +fn test_decode_auth_switch_packet_no_data() { + const AUTH_SWITCH_NO_DATA: &[u8] = b"\xfemysql_clear_password\x00"; + + let p = AuthSwitchRequest::decode_with(AUTH_SWITCH_NO_DATA.into(), true).unwrap(); + + assert!(matches!(p.plugin, AuthPlugin::MySqlClearPassword)); + assert_eq!(p.data, Bytes::new()); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake.rs new file mode 100644 index 00000000..3fef5216 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake.rs @@ -0,0 +1,199 @@ +use bytes::buf::Chain; +use bytes::{Buf, Bytes}; +use std::cmp; + +use crate::error::Error; +use crate::io::{BufExt, ProtocolDecode}; +use crate::protocol::auth::AuthPlugin; +use crate::protocol::response::Status; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::Handshake +// https://mariadb.com/kb/en/connection/#initial-handshake-packet + +#[derive(Debug)] +pub(crate) struct Handshake { + #[allow(unused)] + pub(crate) protocol_version: u8, + pub(crate) server_version: String, + #[allow(unused)] + pub(crate) connection_id: u32, + pub(crate) server_capabilities: Capabilities, + #[allow(unused)] + pub(crate) server_default_collation: u8, + #[allow(unused)] + pub(crate) status: Status, + pub(crate) auth_plugin: Option, + pub(crate) auth_plugin_data: Chain, +} + +impl ProtocolDecode<'_> for Handshake { + fn decode_with(mut buf: Bytes, _: ()) -> Result { + let protocol_version = buf.get_u8(); // int<1> + let server_version = buf.get_str_nul()?; // string + let connection_id = buf.get_u32_le(); // int<4> + let auth_plugin_data_1 = buf.get_bytes(8); // string<8> + + buf.advance(1); // reserved: string<1> + + let capabilities_1 = buf.get_u16_le(); // int<2> + let mut capabilities = Capabilities::from_bits_truncate(capabilities_1.into()); + + let collation = buf.get_u8(); // int<1> + let status = Status::from_bits_truncate(buf.get_u16_le()); + + let capabilities_2 = buf.get_u16_le(); // int<2> + capabilities |= Capabilities::from_bits_truncate(((capabilities_2 as u32) << 16).into()); + + let auth_plugin_data_len = if capabilities.contains(Capabilities::PLUGIN_AUTH) { + buf.get_u8() + } else { + buf.advance(1); // int<1> + 0 + }; + + buf.advance(6); // reserved: string<6> + + if capabilities.contains(Capabilities::MYSQL) { + buf.advance(4); // reserved: string<4> + } else { + let capabilities_3 = buf.get_u32_le(); // int<4> + capabilities |= Capabilities::from_bits_truncate((capabilities_3 as u64) << 32); + } + + let auth_plugin_data_2 = if capabilities.contains(Capabilities::SECURE_CONNECTION) { + let len = cmp::max(auth_plugin_data_len.saturating_sub(9), 12); + let v = buf.get_bytes(len as usize); + buf.advance(1); // NUL-terminator + + v + } else { + Bytes::new() + }; + + let auth_plugin = if capabilities.contains(Capabilities::PLUGIN_AUTH) { + Some(buf.get_str_nul()?.parse()?) + } else { + None + }; + + Ok(Self { + protocol_version, + server_version, + connection_id, + server_default_collation: collation, + status, + server_capabilities: capabilities, + auth_plugin, + auth_plugin_data: auth_plugin_data_1.chain(auth_plugin_data_2), + }) + } +} + +#[test] +fn test_decode_handshake_mysql_8_0_18() { + const HANDSHAKE_MYSQL_8_0_18: &[u8] = b"\n8.0.18\x00\x19\x00\x00\x00\x114aB0c\x06g\x00\xff\xff\xff\x02\x00\xff\xc7\x15\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00tL\x03s\x0f[4\rl4. \x00caching_sha2_password\x00"; + + let p = Handshake::decode(HANDSHAKE_MYSQL_8_0_18.into()).unwrap(); + + assert_eq!(p.protocol_version, 10); + + assert_eq!( + p.server_capabilities, + Capabilities::MYSQL + | Capabilities::FOUND_ROWS + | Capabilities::LONG_FLAG + | Capabilities::CONNECT_WITH_DB + | Capabilities::NO_SCHEMA + | Capabilities::COMPRESS + | Capabilities::ODBC + | Capabilities::LOCAL_FILES + | Capabilities::IGNORE_SPACE + | Capabilities::PROTOCOL_41 + | Capabilities::INTERACTIVE + | Capabilities::SSL + | Capabilities::TRANSACTIONS + | Capabilities::SECURE_CONNECTION + | Capabilities::MULTI_STATEMENTS + | Capabilities::MULTI_RESULTS + | Capabilities::PS_MULTI_RESULTS + | Capabilities::PLUGIN_AUTH + | Capabilities::CONNECT_ATTRS + | Capabilities::PLUGIN_AUTH_LENENC_DATA + | Capabilities::CAN_HANDLE_EXPIRED_PASSWORDS + | Capabilities::SESSION_TRACK + | Capabilities::DEPRECATE_EOF + | Capabilities::ZSTD_COMPRESSION_ALGORITHM + | Capabilities::SSL_VERIFY_SERVER_CERT + | Capabilities::OPTIONAL_RESULTSET_METADATA + | Capabilities::REMEMBER_OPTIONS, + ); + + assert_eq!(p.server_default_collation, 255); + assert!(p.status.contains(Status::SERVER_STATUS_AUTOCOMMIT)); + + assert!(matches!( + p.auth_plugin, + Some(AuthPlugin::CachingSha2Password) + )); + + assert_eq!( + &*p.auth_plugin_data.into_iter().collect::>(), + &[17, 52, 97, 66, 48, 99, 6, 103, 116, 76, 3, 115, 15, 91, 52, 13, 108, 52, 46, 32,] + ); +} + +#[test] +fn test_decode_handshake_mariadb_10_4_7() { + const HANDSHAKE_MARIA_DB_10_4_7: &[u8] = b"\n5.5.5-10.4.7-MariaDB-1:10.4.7+maria~bionic\x00\x0b\x00\x00\x00t6L\\j\"dS\x00\xfe\xf7\x08\x02\x00\xff\x81\x15\x00\x00\x00\x00\x00\x00\x07\x00\x00\x00U14Oph9\">(), + &[116, 54, 76, 92, 106, 34, 100, 83, 85, 49, 52, 79, 112, 104, 57, 34, 60, 72, 53, 110,] + ); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake_response.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake_response.rs new file mode 100644 index 00000000..2e6fec1c --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/handshake_response.rs @@ -0,0 +1,82 @@ +use crate::io::MySqlBufMutExt; +use crate::io::{BufMutExt, ProtocolEncode}; +use crate::protocol::auth::AuthPlugin; +use crate::protocol::connect::ssl_request::SslRequest; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::HandshakeResponse +// https://mariadb.com/kb/en/connection/#client-handshake-response + +#[derive(Debug)] +pub struct HandshakeResponse<'a> { + pub database: Option<&'a str>, + + /// Max size of a command packet that the client wants to send to the server + pub max_packet_size: u32, + + /// Default collation for the connection + pub collation: u8, + + /// Name of the SQL account which client wants to log in + pub username: &'a str, + + /// Authentication method used by the client + pub auth_plugin: Option, + + /// Opaque authentication response + pub auth_response: Option<&'a [u8]>, +} + +impl ProtocolEncode<'_, Capabilities> for HandshakeResponse<'_> { + fn encode_with( + &self, + buf: &mut Vec, + mut context: Capabilities, + ) -> Result<(), crate::Error> { + if self.auth_plugin.is_none() { + // ensure PLUGIN_AUTH is set *only* if we have a defined plugin + context.remove(Capabilities::PLUGIN_AUTH); + } + + // NOTE: Half of this packet is identical to the SSL Request packet + SslRequest { + max_packet_size: self.max_packet_size, + collation: self.collation, + } + .encode_with(buf, context)?; + + buf.put_str_nul(self.username); + + if context.contains(Capabilities::PLUGIN_AUTH_LENENC_DATA) { + buf.put_bytes_lenenc(self.auth_response.unwrap_or_default()); + } else if context.contains(Capabilities::SECURE_CONNECTION) { + let response = self.auth_response.unwrap_or_default(); + + let response_len = u8::try_from(response.len()) + .map_err(|_| err_protocol!("auth_response.len() too long: {}", response.len()))?; + + buf.push(response_len); + buf.extend(response); + } else { + buf.push(0); + } + + if context.contains(Capabilities::CONNECT_WITH_DB) { + if let Some(database) = &self.database { + buf.put_str_nul(database); + } else { + buf.push(0); + } + } + + if context.contains(Capabilities::PLUGIN_AUTH) { + if let Some(plugin) = &self.auth_plugin { + buf.put_str_nul(plugin.name()); + } else { + buf.push(0); + } + } + + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/connect/mod.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/mod.rs new file mode 100644 index 00000000..0222ee89 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/mod.rs @@ -0,0 +1,13 @@ +//! Connection Phase +//! +//! + +mod auth_switch; +mod handshake; +mod handshake_response; +mod ssl_request; + +pub(crate) use auth_switch::{AuthSwitchRequest, AuthSwitchResponse}; +pub(crate) use handshake::Handshake; +pub(crate) use handshake_response::HandshakeResponse; +pub(crate) use ssl_request::SslRequest; diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/connect/ssl_request.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/ssl_request.rs new file mode 100644 index 00000000..cdfc9e51 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/connect/ssl_request.rs @@ -0,0 +1,34 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_connection_phase_packets_protocol_handshake_response.html +// https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::SSLRequest + +#[derive(Debug)] +pub struct SslRequest { + pub max_packet_size: u32, + pub collation: u8, +} + +impl ProtocolEncode<'_, Capabilities> for SslRequest { + fn encode_with(&self, buf: &mut Vec, context: Capabilities) -> Result<(), crate::Error> { + // truncation is intended + #[allow(clippy::cast_possible_truncation)] + buf.extend(&(context.bits() as u32).to_le_bytes()); + buf.extend(&self.max_packet_size.to_le_bytes()); + buf.push(self.collation); + + // reserved: string<19> + buf.extend(&[0_u8; 19]); + + if context.contains(Capabilities::MYSQL) { + // reserved: string<4> + buf.extend(&[0_u8; 4]); + } else { + // extended client capabilities (MariaDB-specified): int<4> + buf.extend(&((context.bits() >> 32) as u32).to_le_bytes()); + } + + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/mod.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/mod.rs new file mode 100644 index 00000000..d1860f5c --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/mod.rs @@ -0,0 +1,12 @@ +pub(crate) mod auth; +mod capabilities; +pub(crate) mod connect; +mod packet; +pub(crate) mod response; +mod row; +pub(crate) mod statement; +pub(crate) mod text; + +pub(crate) use capabilities::Capabilities; +pub(crate) use packet::Packet; +pub(crate) use row::Row; diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/packet.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/packet.rs new file mode 100644 index 00000000..d43338dc --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/packet.rs @@ -0,0 +1,120 @@ +use std::cmp::min; +use std::ops::{Deref, DerefMut}; + +use bytes::Bytes; + +use crate::error::Error; +use crate::io::{ProtocolDecode, ProtocolEncode}; +use crate::protocol::response::{EofPacket, OkPacket}; +use crate::protocol::Capabilities; + +#[derive(Debug)] +pub struct Packet(pub(crate) T); + +impl<'en, 'stream, T> ProtocolEncode<'stream, (Capabilities, &'stream mut u8)> for Packet +where + T: ProtocolEncode<'en, Capabilities>, +{ + fn encode_with( + &self, + buf: &mut Vec, + (capabilities, sequence_id): (Capabilities, &'stream mut u8), + ) -> Result<(), Error> { + let mut next_header = |len: u32| { + let mut buf = len.to_le_bytes(); + buf[3] = *sequence_id; + *sequence_id = sequence_id.wrapping_add(1); + + buf + }; + + // reserve space to write the prefixed length + let offset = buf.len(); + buf.extend(&[0_u8; 4]); + + // encode the payload + self.0.encode_with(buf, capabilities)?; + + // determine the length of the encoded payload + // and write to our reserved space + let len = buf.len() - offset - 4; + let header = &mut buf[offset..]; + + // // `min(.., 0xFF_FF_FF)` cannot overflow + #[allow(clippy::cast_possible_truncation)] + header[..4].copy_from_slice(&next_header(min(len, 0xFF_FF_FF) as u32)); + + // add more packets if we need to split the data + if len >= 0xFF_FF_FF { + let rest = buf.split_off(offset + 4 + 0xFF_FF_FF); + let mut chunks = rest.chunks_exact(0xFF_FF_FF); + + for chunk in chunks.by_ref() { + buf.reserve(chunk.len() + 4); + + // `chunk.len() == 0xFF_FF_FF` + #[allow(clippy::cast_possible_truncation)] + buf.extend(&next_header(chunk.len() as u32)); + buf.extend(chunk); + } + + // this will also handle adding a zero sized packet if the data size is a multiple of 0xFF_FF_FF + let remainder = chunks.remainder(); + buf.reserve(remainder.len() + 4); + + // `remainder.len() < 0xFF_FF_FF` + #[allow(clippy::cast_possible_truncation)] + buf.extend(&next_header(remainder.len() as u32)); + buf.extend(remainder); + } + + Ok(()) + } +} + +impl Packet { + pub(crate) fn decode<'de, T>(self) -> Result + where + T: ProtocolDecode<'de, ()>, + { + self.decode_with(()) + } + + pub(crate) fn decode_with<'de, T, C>(self, context: C) -> Result + where + T: ProtocolDecode<'de, C>, + { + T::decode_with(self.0, context) + } + + pub(crate) fn ok(self) -> Result { + self.decode() + } + + pub(crate) fn eof(self, capabilities: Capabilities) -> Result { + if capabilities.contains(Capabilities::DEPRECATE_EOF) { + let ok = self.ok()?; + + Ok(EofPacket { + warnings: ok.warnings, + status: ok.status, + }) + } else { + self.decode_with(capabilities) + } + } +} + +impl Deref for Packet { + type Target = Bytes; + + fn deref(&self) -> &Bytes { + &self.0 + } +} + +impl DerefMut for Packet { + fn deref_mut(&mut self) -> &mut Bytes { + &mut self.0 + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/response/eof.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/response/eof.rs new file mode 100644 index 00000000..89de9a32 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/response/eof.rs @@ -0,0 +1,36 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::ProtocolDecode; +use crate::protocol::response::Status; +use crate::protocol::Capabilities; + +/// Marks the end of a result set, returning status and warnings. +/// +/// # Note +/// +/// The EOF packet is deprecated as of MySQL 5.7.5. SQLx only uses this packet for MySQL +/// prior MySQL versions. +#[derive(Debug)] +pub struct EofPacket { + #[allow(dead_code)] + pub warnings: u16, + pub status: Status, +} + +impl ProtocolDecode<'_, Capabilities> for EofPacket { + fn decode_with(mut buf: Bytes, _: Capabilities) -> Result { + let header = buf.get_u8(); + if header != 0xfe { + return Err(err_protocol!( + "expected 0xfe (EOF_Packet) but found 0x{:x}", + header + )); + } + + let warnings = buf.get_u16_le(); + let status = Status::from_bits_truncate(buf.get_u16_le()); + + Ok(Self { status, warnings }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/response/err.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/response/err.rs new file mode 100644 index 00000000..085d24f4 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/response/err.rs @@ -0,0 +1,71 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::{BufExt, ProtocolDecode}; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_basic_err_packet.html +// https://mariadb.com/kb/en/err_packet/ + +/// Indicates that an error occurred. +#[derive(Debug)] +pub struct ErrPacket { + pub error_code: u16, + pub sql_state: Option, + pub error_message: String, +} + +impl ProtocolDecode<'_, Capabilities> for ErrPacket { + fn decode_with(mut buf: Bytes, capabilities: Capabilities) -> Result { + let header = buf.get_u8(); + if header != 0xff { + return Err(err_protocol!( + "expected 0xff (ERR_Packet) but found 0x{:x}", + header + )); + } + + let error_code = buf.get_u16_le(); + let mut sql_state = None; + + if capabilities.contains(Capabilities::PROTOCOL_41) { + // If the next byte is '#' then we have a SQL STATE + if buf.starts_with(b"#") { + buf.advance(1); + sql_state = Some(buf.get_str(5)?); + } + } + + let error_message = buf.get_str(buf.len())?; + + Ok(Self { + error_code, + sql_state, + error_message, + }) + } +} + +#[test] +fn test_decode_err_packet_out_of_order() { + const ERR_PACKETS_OUT_OF_ORDER: &[u8] = b"\xff\x84\x04Got packets out of order"; + + let p = + ErrPacket::decode_with(ERR_PACKETS_OUT_OF_ORDER.into(), Capabilities::PROTOCOL_41).unwrap(); + + assert_eq!(&p.error_message, "Got packets out of order"); + assert_eq!(p.error_code, 1156); + assert_eq!(p.sql_state, None); +} + +#[test] +fn test_decode_err_packet_unknown_database() { + const ERR_HANDSHAKE_UNKNOWN_DB: &[u8] = b"\xff\x19\x04#42000Unknown database \'unknown\'"; + + let p = + ErrPacket::decode_with(ERR_HANDSHAKE_UNKNOWN_DB.into(), Capabilities::PROTOCOL_41).unwrap(); + + assert_eq!(p.error_code, 1049); + assert_eq!(p.sql_state.as_deref(), Some("42000")); + assert_eq!(&p.error_message, "Unknown database \'unknown\'"); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/response/mod.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/response/mod.rs new file mode 100644 index 00000000..79767dc6 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/response/mod.rs @@ -0,0 +1,14 @@ +//! Generic Response Packets +//! +//! +//! + +mod eof; +mod err; +mod ok; +mod status; + +pub use eof::EofPacket; +pub use err::ErrPacket; +pub use ok::OkPacket; +pub use status::Status; diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/response/ok.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/response/ok.rs new file mode 100644 index 00000000..d16127d5 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/response/ok.rs @@ -0,0 +1,52 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::MySqlBufExt; +use crate::io::ProtocolDecode; +use crate::protocol::response::Status; + +/// Indicates successful completion of a previous command sent by the client. +#[derive(Debug)] +pub struct OkPacket { + pub affected_rows: u64, + pub last_insert_id: u64, + pub status: Status, + pub warnings: u16, +} + +impl ProtocolDecode<'_> for OkPacket { + fn decode_with(mut buf: Bytes, _: ()) -> Result { + let header = buf.get_u8(); + if header != 0 && header != 0xfe { + return Err(err_protocol!( + "expected 0x00 or 0xfe (OK_Packet) but found 0x{:02x}", + header + )); + } + + let affected_rows = buf.get_uint_lenenc(); + let last_insert_id = buf.get_uint_lenenc(); + let status = Status::from_bits_truncate(buf.get_u16_le()); + let warnings = buf.get_u16_le(); + + Ok(Self { + affected_rows, + last_insert_id, + status, + warnings, + }) + } +} + +#[test] +fn test_decode_ok_packet() { + const DATA: &[u8] = b"\x00\x00\x00\x02@\x00\x00"; + + let p = OkPacket::decode(DATA.into()).unwrap(); + + assert_eq!(p.affected_rows, 0); + assert_eq!(p.last_insert_id, 0); + assert_eq!(p.warnings, 0); + assert!(p.status.contains(Status::SERVER_STATUS_AUTOCOMMIT)); + assert!(p.status.contains(Status::SERVER_SESSION_STATE_CHANGED)); +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/response/status.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/response/status.rs new file mode 100644 index 00000000..4a8bb037 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/response/status.rs @@ -0,0 +1,50 @@ +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/mysql__com_8h.html#a1d854e841086925be1883e4d7b4e8cad +// https://mariadb.com/kb/en/library/mariadb-connectorc-types-and-definitions/#server-status +bitflags::bitflags! { + #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)] + pub struct Status: u16 { + // Is raised when a multi-statement transaction has been started, either explicitly, + // by means of BEGIN or COMMIT AND CHAIN, or implicitly, by the first + // transactional statement, when autocommit=off. + const SERVER_STATUS_IN_TRANS = 1; + + // Autocommit mode is set + const SERVER_STATUS_AUTOCOMMIT = 2; + + // Multi query - next query exists. + const SERVER_MORE_RESULTS_EXISTS = 8; + + const SERVER_QUERY_NO_GOOD_INDEX_USED = 16; + const SERVER_QUERY_NO_INDEX_USED = 32; + + // When using COM_STMT_FETCH, indicate that current cursor still has result + const SERVER_STATUS_CURSOR_EXISTS = 64; + + // When using COM_STMT_FETCH, indicate that current cursor has finished to send results + const SERVER_STATUS_LAST_ROW_SENT = 128; + + // Database has been dropped + const SERVER_STATUS_DB_DROPPED = (1 << 8); + + // Current escape mode is "no backslash escape" + const SERVER_STATUS_NO_BACKSLASH_ESCAPES = (1 << 9); + + // A DDL change did have an impact on an existing PREPARE (an automatic + // re-prepare has been executed) + const SERVER_STATUS_METADATA_CHANGED = (1 << 10); + + // Last statement took more than the time value specified + // in server variable long_query_time. + const SERVER_QUERY_WAS_SLOW = (1 << 11); + + // This result-set contain stored procedure output parameter. + const SERVER_PS_OUT_PARAMS = (1 << 12); + + // Current transaction is a read-only transaction. + const SERVER_STATUS_IN_TRANS_READONLY = (1 << 13); + + // This status flag, when on, implies that one of the state information has changed + // on the server because of the execution of the last statement. + const SERVER_SESSION_STATE_CHANGED = (1 << 14); + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/row.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/row.rs new file mode 100644 index 00000000..327ca216 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/row.rs @@ -0,0 +1,15 @@ +use std::ops::Range; + +use bytes::Bytes; + +#[derive(Debug)] +pub(crate) struct Row { + pub(crate) storage: Bytes, + pub(crate) values: Vec>>, +} + +impl Row { + pub(crate) fn get(&self, index: usize) -> Option<&[u8]> { + self.values[index].clone().map(|col| &self.storage[col]) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/execute.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/execute.rs new file mode 100644 index 00000000..6e51e7b5 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/execute.rs @@ -0,0 +1,40 @@ +use crate::io::ProtocolEncode; +use crate::protocol::text::ColumnFlags; +use crate::protocol::Capabilities; +use crate::MySqlArguments; + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_com_stmt_execute.html + +#[derive(Debug)] +pub struct Execute<'q> { + pub statement: u32, + pub arguments: &'q MySqlArguments, +} + +impl<'q> ProtocolEncode<'_, Capabilities> for Execute<'q> { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x17); // COM_STMT_EXECUTE + buf.extend(&self.statement.to_le_bytes()); + buf.push(0); // NO_CURSOR + buf.extend(&1_u32.to_le_bytes()); // iterations (always 1): int<4> + + if !self.arguments.types.is_empty() { + buf.extend_from_slice(&self.arguments.null_bitmap); + buf.push(1); // send type to server + + for ty in &self.arguments.types { + buf.push(ty.r#type as u8); + + buf.push(if ty.flags.contains(ColumnFlags::UNSIGNED) { + 0x80 + } else { + 0 + }); + } + + buf.extend(&*self.arguments.values); + } + + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/mod.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/mod.rs new file mode 100644 index 00000000..9ae6b3c9 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/mod.rs @@ -0,0 +1,11 @@ +mod execute; +mod prepare; +mod prepare_ok; +mod row; +mod stmt_close; + +pub(crate) use execute::Execute; +pub(crate) use prepare::Prepare; +pub(crate) use prepare_ok::PrepareOk; +pub(crate) use row::BinaryRow; +pub(crate) use stmt_close::StmtClose; diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare.rs new file mode 100644 index 00000000..6012b119 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare.rs @@ -0,0 +1,16 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-stmt-prepare.html#packet-COM_STMT_PREPARE + +pub struct Prepare<'a> { + pub query: &'a str, +} + +impl ProtocolEncode<'_, Capabilities> for Prepare<'_> { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x16); // COM_STMT_PREPARE + buf.extend(self.query.as_bytes()); + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare_ok.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare_ok.rs new file mode 100644 index 00000000..da25047a --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/prepare_ok.rs @@ -0,0 +1,49 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::ProtocolDecode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-stmt-prepare-response.html#packet-COM_STMT_PREPARE_OK + +#[derive(Debug)] +pub(crate) struct PrepareOk { + pub(crate) statement_id: u32, + pub(crate) columns: u16, + pub(crate) params: u16, + #[allow(unused)] + pub(crate) warnings: u16, +} + +impl ProtocolDecode<'_, Capabilities> for PrepareOk { + fn decode_with(buf: Bytes, _: Capabilities) -> Result { + const SIZE: usize = 12; + + let mut slice = buf.get(..SIZE).ok_or_else(|| { + err_protocol!("PrepareOk expected 12 bytes but got {} bytes", buf.len()) + })?; + + let status = slice.get_u8(); + if status != 0x00 { + return Err(err_protocol!( + "expected 0x00 (COM_STMT_PREPARE_OK) but found 0x{:02x}", + status + )); + } + + let statement_id = slice.get_u32_le(); + let columns = slice.get_u16_le(); + let params = slice.get_u16_le(); + + slice.advance(1); // reserved: string<1> + + let warnings = slice.get_u16_le(); + + Ok(Self { + statement_id, + columns, + params, + warnings, + }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/row.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/row.rs new file mode 100644 index 00000000..3007884c --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/row.rs @@ -0,0 +1,107 @@ +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::MySqlBufExt; +use crate::io::{BufExt, ProtocolDecode}; +use crate::protocol::text::ColumnType; +use crate::protocol::Row; +use crate::MySqlColumn; + +// https://dev.mysql.com/doc/internals/en/binary-protocol-resultset-row.html#packet-ProtocolBinary::ResultsetRow +// https://dev.mysql.com/doc/internals/en/binary-protocol-value.html + +#[derive(Debug)] +pub(crate) struct BinaryRow(pub(crate) Row); + +impl<'de> ProtocolDecode<'de, &'de [MySqlColumn]> for BinaryRow { + fn decode_with(mut buf: Bytes, columns: &'de [MySqlColumn]) -> Result { + let header = buf.get_u8(); + if header != 0 { + return Err(err_protocol!( + "exepcted 0x00 (ROW) but found 0x{:02x}", + header + )); + } + + let storage = buf.clone(); + let offset = buf.len(); + + let null_bitmap_len = (columns.len() + 9) / 8; + let null_bitmap = buf.get_bytes(null_bitmap_len); + + let mut values = Vec::with_capacity(columns.len()); + + for (column_idx, column) in columns.iter().enumerate() { + // NOTE: the column index starts at the 3rd bit + let column_null_idx = column_idx + 2; + + let byte_idx = column_null_idx / 8; + let bit_idx = column_null_idx % 8; + + let is_null = null_bitmap[byte_idx] & (1u8 << bit_idx) != 0; + + if is_null { + values.push(None); + continue; + } + + // NOTE: MySQL will never generate NULL types for non-NULL values + let type_info = &column.type_info; + + // Unlike Postgres, MySQL does not length-prefix every value in a binary row. + // Values are *either* fixed-length or length-prefixed, + // so we need to inspect the type code to be sure. + let size: usize = match type_info.r#type { + // All fixed-length types. + ColumnType::LongLong => 8, + ColumnType::Long | ColumnType::Int24 => 4, + ColumnType::Short | ColumnType::Year => 2, + ColumnType::Tiny => 1, + ColumnType::Float => 4, + ColumnType::Double => 8, + + // Blobs and strings are prefixed with their length, + // which is itself a length-encoded integer. + ColumnType::String + | ColumnType::VarChar + | ColumnType::VarString + | ColumnType::Enum + | ColumnType::Set + | ColumnType::LongBlob + | ColumnType::MediumBlob + | ColumnType::Blob + | ColumnType::TinyBlob + | ColumnType::Geometry + | ColumnType::Bit + | ColumnType::Decimal + | ColumnType::Json + | ColumnType::NewDecimal => { + let size = buf.get_uint_lenenc(); + usize::try_from(size) + .map_err(|_| err_protocol!("BLOB length out of range: {size}"))? + } + + // Like strings and blobs, these values are variable-length. + // Unlike strings and blobs, however, they exclusively use one byte for length. + ColumnType::Time + | ColumnType::Timestamp + | ColumnType::Date + | ColumnType::Datetime => { + // Leave the length byte on the front of the value because decoding uses it. + buf[0] as usize + 1 + } + + // NOTE: MySQL will never generate NULL types for non-NULL values + ColumnType::Null => unreachable!(), + }; + + let offset = offset - buf.len(); + + values.push(Some(offset..(offset + size))); + + buf.advance(size); + } + + Ok(BinaryRow(Row { values, storage })) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/statement/stmt_close.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/stmt_close.rs new file mode 100644 index 00000000..a92f0310 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/statement/stmt_close.rs @@ -0,0 +1,17 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-stmt-close.html + +#[derive(Debug)] +pub struct StmtClose { + pub statement: u32, +} + +impl ProtocolEncode<'_, Capabilities> for StmtClose { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x19); // COM_STMT_CLOSE + buf.extend(&self.statement.to_le_bytes()); + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/column.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/column.rs new file mode 100644 index 00000000..425a5cdc --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/column.rs @@ -0,0 +1,261 @@ +use std::str::from_utf8; + +use bitflags::bitflags; +use bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::MySqlBufExt; +use crate::io::ProtocolDecode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/group__group__cs__column__definition__flags.html + +bitflags! { + #[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + pub(crate) struct ColumnFlags: u16 { + /// Field can't be `NULL`. + const NOT_NULL = 1; + + /// Field is part of a primary key. + const PRIMARY_KEY = 2; + + /// Field is part of a unique key. + const UNIQUE_KEY = 4; + + /// Field is part of a multi-part unique or primary key. + const MULTIPLE_KEY = 8; + + /// Field is a blob. + const BLOB = 16; + + /// Field is unsigned. + const UNSIGNED = 32; + + /// Field is zero filled. + const ZEROFILL = 64; + + /// Field is binary. + const BINARY = 128; + + /// Field is an enumeration. + const ENUM = 256; + + /// Field is an auto-incement field. + const AUTO_INCREMENT = 512; + + /// Field is a timestamp. + const TIMESTAMP = 1024; + + /// Field is a set. + const SET = 2048; + + /// Field does not have a default value. + const NO_DEFAULT_VALUE = 4096; + + /// Field is set to NOW on UPDATE. + const ON_UPDATE_NOW = 8192; + + /// Field is a number. + const NUM = 32768; + } +} + +// https://dev.mysql.com/doc/internals/en/com-query-response.html#column-type + +#[derive(Debug, Copy, Clone, PartialEq)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +#[repr(u8)] +pub enum ColumnType { + Decimal = 0x00, + Tiny = 0x01, + Short = 0x02, + Long = 0x03, + Float = 0x04, + Double = 0x05, + Null = 0x06, + Timestamp = 0x07, + LongLong = 0x08, + Int24 = 0x09, + Date = 0x0a, + Time = 0x0b, + Datetime = 0x0c, + Year = 0x0d, + VarChar = 0x0f, + Bit = 0x10, + Json = 0xf5, + NewDecimal = 0xf6, + Enum = 0xf7, + Set = 0xf8, + TinyBlob = 0xf9, + MediumBlob = 0xfa, + LongBlob = 0xfb, + Blob = 0xfc, + VarString = 0xfd, + String = 0xfe, + Geometry = 0xff, +} + +// https://dev.mysql.com/doc/dev/mysql-server/8.0.12/page_protocol_com_query_response_text_resultset_column_definition.html +// https://mariadb.com/kb/en/resultset/#column-definition-packet +// https://dev.mysql.com/doc/internals/en/com-query-response.html#packet-Protocol::ColumnDefinition41 + +#[derive(Debug)] +pub(crate) struct ColumnDefinition { + #[allow(unused)] + catalog: Bytes, + #[allow(unused)] + schema: Bytes, + #[allow(unused)] + table_alias: Bytes, + #[allow(unused)] + table: Bytes, + alias: Bytes, + name: Bytes, + #[allow(unused)] + pub(crate) collation: u16, + pub(crate) max_size: u32, + pub(crate) r#type: ColumnType, + pub(crate) flags: ColumnFlags, + #[allow(unused)] + decimals: u8, +} + +impl ColumnDefinition { + // NOTE: strings in-protocol are transmitted according to the client character set + // as this is UTF-8, all these strings should be UTF-8 + + pub(crate) fn name(&self) -> Result<&str, Error> { + from_utf8(&self.name).map_err(Error::protocol) + } + + pub(crate) fn alias(&self) -> Result<&str, Error> { + from_utf8(&self.alias).map_err(Error::protocol) + } +} + +impl ProtocolDecode<'_, Capabilities> for ColumnDefinition { + fn decode_with(mut buf: Bytes, _: Capabilities) -> Result { + let catalog = buf.get_bytes_lenenc()?; + let schema = buf.get_bytes_lenenc()?; + let table_alias = buf.get_bytes_lenenc()?; + let table = buf.get_bytes_lenenc()?; + let alias = buf.get_bytes_lenenc()?; + let name = buf.get_bytes_lenenc()?; + let _next_len = buf.get_uint_lenenc(); // always 0x0c + let collation = buf.get_u16_le(); + let max_size = buf.get_u32_le(); + let type_id = buf.get_u8(); + let flags = buf.get_u16_le(); + let decimals = buf.get_u8(); + + Ok(Self { + catalog, + schema, + table_alias, + table, + alias, + name, + collation, + max_size, + r#type: ColumnType::try_from_u16(type_id)?, + flags: ColumnFlags::from_bits_truncate(flags), + decimals, + }) + } +} + +impl ColumnType { + pub(crate) fn name(self, flags: ColumnFlags, max_size: Option) -> &'static str { + let is_binary = flags.contains(ColumnFlags::BINARY); + let is_unsigned = flags.contains(ColumnFlags::UNSIGNED); + let is_enum = flags.contains(ColumnFlags::ENUM); + + match self { + ColumnType::Tiny if max_size == Some(1) => "BOOLEAN", + ColumnType::Tiny if is_unsigned => "TINYINT UNSIGNED", + ColumnType::Short if is_unsigned => "SMALLINT UNSIGNED", + ColumnType::Long if is_unsigned => "INT UNSIGNED", + ColumnType::Int24 if is_unsigned => "MEDIUMINT UNSIGNED", + ColumnType::LongLong if is_unsigned => "BIGINT UNSIGNED", + ColumnType::Tiny => "TINYINT", + ColumnType::Short => "SMALLINT", + ColumnType::Long => "INT", + ColumnType::Int24 => "MEDIUMINT", + ColumnType::LongLong => "BIGINT", + ColumnType::Float => "FLOAT", + ColumnType::Double => "DOUBLE", + ColumnType::Null => "NULL", + ColumnType::Timestamp => "TIMESTAMP", + ColumnType::Date => "DATE", + ColumnType::Time => "TIME", + ColumnType::Datetime => "DATETIME", + ColumnType::Year => "YEAR", + ColumnType::Bit => "BIT", + ColumnType::Enum => "ENUM", + ColumnType::Set => "SET", + ColumnType::Decimal | ColumnType::NewDecimal => "DECIMAL", + ColumnType::Geometry => "GEOMETRY", + ColumnType::Json => "JSON", + + ColumnType::String if is_binary => "BINARY", + ColumnType::String if is_enum => "ENUM", + ColumnType::VarChar | ColumnType::VarString if is_binary => "VARBINARY", + + ColumnType::String => "CHAR", + ColumnType::VarChar | ColumnType::VarString => "VARCHAR", + + ColumnType::TinyBlob if is_binary => "TINYBLOB", + ColumnType::TinyBlob => "TINYTEXT", + + ColumnType::Blob if is_binary => "BLOB", + ColumnType::Blob => "TEXT", + + ColumnType::MediumBlob if is_binary => "MEDIUMBLOB", + ColumnType::MediumBlob => "MEDIUMTEXT", + + ColumnType::LongBlob if is_binary => "LONGBLOB", + ColumnType::LongBlob => "LONGTEXT", + } + } + + pub(crate) fn try_from_u16(id: u8) -> Result { + Ok(match id { + 0x00 => ColumnType::Decimal, + 0x01 => ColumnType::Tiny, + 0x02 => ColumnType::Short, + 0x03 => ColumnType::Long, + 0x04 => ColumnType::Float, + 0x05 => ColumnType::Double, + 0x06 => ColumnType::Null, + 0x07 => ColumnType::Timestamp, + 0x08 => ColumnType::LongLong, + 0x09 => ColumnType::Int24, + 0x0a => ColumnType::Date, + 0x0b => ColumnType::Time, + 0x0c => ColumnType::Datetime, + 0x0d => ColumnType::Year, + // [internal] 0x0e => ColumnType::NewDate, + 0x0f => ColumnType::VarChar, + 0x10 => ColumnType::Bit, + // [internal] 0x11 => ColumnType::Timestamp2, + // [internal] 0x12 => ColumnType::Datetime2, + // [internal] 0x13 => ColumnType::Time2, + 0xf5 => ColumnType::Json, + 0xf6 => ColumnType::NewDecimal, + 0xf7 => ColumnType::Enum, + 0xf8 => ColumnType::Set, + 0xf9 => ColumnType::TinyBlob, + 0xfa => ColumnType::MediumBlob, + 0xfb => ColumnType::LongBlob, + 0xfc => ColumnType::Blob, + 0xfd => ColumnType::VarString, + 0xfe => ColumnType::String, + 0xff => ColumnType::Geometry, + + _ => { + return Err(err_protocol!("unknown column type 0x{:02x}", id)); + } + }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/mod.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/mod.rs new file mode 100644 index 00000000..2286ee89 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/mod.rs @@ -0,0 +1,11 @@ +mod column; +mod ping; +mod query; +mod quit; +mod row; + +pub(crate) use column::{ColumnDefinition, ColumnFlags, ColumnType}; +pub(crate) use ping::Ping; +pub(crate) use query::Query; +pub(crate) use quit::Quit; +pub(crate) use row::TextRow; diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/ping.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/ping.rs new file mode 100644 index 00000000..4eb8ab2e --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/ping.rs @@ -0,0 +1,14 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-ping.html + +#[derive(Debug)] +pub(crate) struct Ping; + +impl ProtocolEncode<'_, Capabilities> for Ping { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x0e); // COM_PING + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/query.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/query.rs new file mode 100644 index 00000000..b3533adb --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/query.rs @@ -0,0 +1,15 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-query.html + +#[derive(Debug)] +pub(crate) struct Query<'q>(pub(crate) &'q str); + +impl ProtocolEncode<'_, Capabilities> for Query<'_> { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x03); // COM_QUERY + buf.extend(self.0.as_bytes()); + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/quit.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/quit.rs new file mode 100644 index 00000000..c0d8729e --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/quit.rs @@ -0,0 +1,14 @@ +use crate::io::ProtocolEncode; +use crate::protocol::Capabilities; + +// https://dev.mysql.com/doc/internals/en/com-quit.html + +#[derive(Debug)] +pub(crate) struct Quit; + +impl ProtocolEncode<'_, Capabilities> for Quit { + fn encode_with(&self, buf: &mut Vec, _: Capabilities) -> Result<(), crate::Error> { + buf.push(0x01); // COM_QUIT + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/protocol/text/row.rs b/src-tauri/vendor/sqlx-mysql/src/protocol/text/row.rs new file mode 100644 index 00000000..e5f820c6 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/protocol/text/row.rs @@ -0,0 +1,46 @@ +use bytes::{Buf, Bytes}; + +use crate::column::MySqlColumn; +use crate::error::Error; +use crate::io::MySqlBufExt; +use crate::io::ProtocolDecode; +use crate::protocol::Row; + +#[derive(Debug)] +pub(crate) struct TextRow(pub(crate) Row); + +impl<'de> ProtocolDecode<'de, &'de [MySqlColumn]> for TextRow { + fn decode_with(mut buf: Bytes, columns: &'de [MySqlColumn]) -> Result { + let storage = buf.clone(); + let offset = buf.len(); + + let mut values = Vec::with_capacity(columns.len()); + + for c in columns { + if buf[0] == 0xfb { + // NULL is sent as 0xfb + values.push(None); + buf.advance(1); + } else { + let size = buf.get_uint_lenenc(); + if (buf.remaining() as u64) < size { + return Err(err_protocol!( + "buffer exhausted when reading data for column {:?}; decoded length is {}, but only {} bytes remain in buffer. Malformed packet or protocol error?", + c, + size, + buf.remaining())); + } + let size = usize::try_from(size) + .map_err(|_| err_protocol!("TextRow length out of range: {size}"))?; + + let offset = offset - buf.len(); + + values.push(Some(offset..(offset + size))); + + buf.advance(size); + } + } + + Ok(TextRow(Row { values, storage })) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/query_result.rs b/src-tauri/vendor/sqlx-mysql/src/query_result.rs new file mode 100644 index 00000000..f008db06 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/query_result.rs @@ -0,0 +1,36 @@ +use std::iter::{Extend, IntoIterator}; + +#[derive(Debug, Default)] +pub struct MySqlQueryResult { + pub(super) rows_affected: u64, + pub(super) last_insert_id: u64, +} + +impl MySqlQueryResult { + pub fn last_insert_id(&self) -> u64 { + self.last_insert_id + } + + pub fn rows_affected(&self) -> u64 { + self.rows_affected + } +} + +impl Extend for MySqlQueryResult { + fn extend>(&mut self, iter: T) { + for elem in iter { + self.rows_affected += elem.rows_affected; + self.last_insert_id = elem.last_insert_id; + } + } +} +#[cfg(feature = "any")] +/// This conversion attempts to save last_insert_id by converting to i64. +impl From for sqlx_core::any::AnyQueryResult { + fn from(done: MySqlQueryResult) -> Self { + sqlx_core::any::AnyQueryResult { + rows_affected: done.rows_affected(), + last_insert_id: done.last_insert_id().try_into().ok(), + } + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/row.rs b/src-tauri/vendor/sqlx-mysql/src/row.rs new file mode 100644 index 00000000..e7191366 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/row.rs @@ -0,0 +1,51 @@ +use std::sync::Arc; + +pub(crate) use sqlx_core::row::*; + +use crate::column::ColumnIndex; +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::HashMap; +use crate::{protocol, MySql, MySqlColumn, MySqlValueFormat, MySqlValueRef}; + +/// Implementation of [`Row`] for MySQL. +#[derive(Debug)] +pub struct MySqlRow { + pub(crate) row: protocol::Row, + pub(crate) format: MySqlValueFormat, + pub(crate) columns: Arc>, + pub(crate) column_names: Arc>, +} + +impl Row for MySqlRow { + type Database = MySql; + + fn columns(&self) -> &[MySqlColumn] { + &self.columns + } + + fn try_get_raw(&self, index: I) -> Result, Error> + where + I: ColumnIndex, + { + let index = index.index(self)?; + let column = &self.columns[index]; + let value = self.row.get(index); + + Ok(MySqlValueRef { + format: self.format, + row: Some(&self.row.storage), + type_info: column.type_info.clone(), + value, + }) + } +} + +impl ColumnIndex for &'_ str { + fn index(&self, row: &MySqlRow) -> Result { + row.column_names + .get(*self) + .ok_or_else(|| Error::ColumnNotFound((*self).into())) + .copied() + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/statement.rs b/src-tauri/vendor/sqlx-mysql/src/statement.rs new file mode 100644 index 00000000..e9578403 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/statement.rs @@ -0,0 +1,60 @@ +use super::MySqlColumn; +use crate::column::ColumnIndex; +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::HashMap; +use crate::{MySql, MySqlArguments, MySqlTypeInfo}; +use either::Either; +use std::borrow::Cow; +use std::sync::Arc; + +pub(crate) use sqlx_core::statement::*; + +#[derive(Debug, Clone)] +pub struct MySqlStatement<'q> { + pub(crate) sql: Cow<'q, str>, + pub(crate) metadata: MySqlStatementMetadata, +} + +#[derive(Debug, Default, Clone)] +pub(crate) struct MySqlStatementMetadata { + pub(crate) columns: Arc>, + pub(crate) column_names: Arc>, + pub(crate) parameters: usize, +} + +impl<'q> Statement<'q> for MySqlStatement<'q> { + type Database = MySql; + + fn to_owned(&self) -> MySqlStatement<'static> { + MySqlStatement::<'static> { + sql: Cow::Owned(self.sql.clone().into_owned()), + metadata: self.metadata.clone(), + } + } + + fn sql(&self) -> &str { + &self.sql + } + + fn parameters(&self) -> Option> { + Some(Either::Right(self.metadata.parameters)) + } + + fn columns(&self) -> &[MySqlColumn] { + &self.metadata.columns + } + + impl_statement_query!(MySqlArguments); +} + +impl ColumnIndex> for &'_ str { + fn index(&self, statement: &MySqlStatement<'_>) -> Result { + statement + .metadata + .column_names + .get(*self) + .ok_or_else(|| Error::ColumnNotFound((*self).into())) + .copied() + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/testing/mod.rs b/src-tauri/vendor/sqlx-mysql/src/testing/mod.rs new file mode 100644 index 00000000..2b6d4671 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/testing/mod.rs @@ -0,0 +1,253 @@ +use std::ops::Deref; +use std::str::FromStr; +use std::time::Duration; + +use futures_core::future::BoxFuture; + +use crate::error::Error; +use crate::executor::Executor; +use crate::pool::{Pool, PoolOptions}; +use crate::query::query; +use crate::{MySql, MySqlConnectOptions, MySqlConnection, MySqlDatabaseError}; +use once_cell::sync::OnceCell; +use sqlx_core::connection::Connection; +use sqlx_core::query_builder::QueryBuilder; +use sqlx_core::query_scalar::query_scalar; +use std::fmt::Write; + +pub(crate) use sqlx_core::testing::*; + +// Using a blocking `OnceCell` here because the critical sections are short. +static MASTER_POOL: OnceCell> = OnceCell::new(); + +impl TestSupport for MySql { + fn test_context(args: &TestArgs) -> BoxFuture<'_, Result, Error>> { + Box::pin(async move { test_context(args).await }) + } + + fn cleanup_test(db_name: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let mut conn = MASTER_POOL + .get() + .expect("cleanup_test() invoked outside `#[sqlx::test]`") + .acquire() + .await?; + + do_cleanup(&mut conn, db_name).await + }) + } + + fn cleanup_test_dbs() -> BoxFuture<'static, Result, Error>> { + Box::pin(async move { + let url = dotenvy::var("DATABASE_URL").expect("DATABASE_URL must be set"); + + let mut conn = MySqlConnection::connect(&url).await?; + + let delete_db_names: Vec = + query_scalar("select db_name from _sqlx_test_databases") + .fetch_all(&mut conn) + .await?; + + if delete_db_names.is_empty() { + return Ok(None); + } + + let mut deleted_db_names = Vec::with_capacity(delete_db_names.len()); + + let mut command = String::new(); + + for db_name in &delete_db_names { + command.clear(); + + let db_name = format!("_sqlx_test_database_{db_name}"); + + writeln!(command, "drop database if exists {db_name};").ok(); + match conn.execute(&*command).await { + Ok(_deleted) => { + deleted_db_names.push(db_name); + } + // Assume a database error just means the DB is still in use. + Err(Error::Database(dbe)) => { + eprintln!("could not clean test database {db_name:?}: {dbe}") + } + // Bubble up other errors + Err(e) => return Err(e), + } + } + + if deleted_db_names.is_empty() { + return Ok(None); + } + + let mut query = + QueryBuilder::new("delete from _sqlx_test_databases where db_name in ("); + + let mut separated = query.separated(","); + + for db_name in &deleted_db_names { + separated.push_bind(db_name); + } + + query.push(")").build().execute(&mut conn).await?; + + let _ = conn.close().await; + Ok(Some(delete_db_names.len())) + }) + } + + fn snapshot( + _conn: &mut Self::Connection, + ) -> BoxFuture<'_, Result, Error>> { + // TODO: I want to get the testing feature out the door so this will have to wait, + // but I'm keeping the code around for now because I plan to come back to it. + todo!() + } +} + +async fn test_context(args: &TestArgs) -> Result, Error> { + let url = dotenvy::var("DATABASE_URL").expect("DATABASE_URL must be set"); + + let master_opts = MySqlConnectOptions::from_str(&url).expect("failed to parse DATABASE_URL"); + + let pool = PoolOptions::new() + // MySql's normal connection limit is 150 plus 1 superuser connection + // We don't want to use the whole cap and there may be fuzziness here due to + // concurrently running tests anyway. + .max_connections(20) + // Immediately close master connections. Tokio's I/O streams don't like hopping runtimes. + .after_release(|_conn, _| Box::pin(async move { Ok(false) })) + .connect_lazy_with(master_opts); + + let master_pool = match MASTER_POOL.try_insert(pool) { + Ok(inserted) => inserted, + Err((existing, pool)) => { + // Sanity checks. + assert_eq!( + existing.connect_options().host, + pool.connect_options().host, + "DATABASE_URL changed at runtime, host differs" + ); + + assert_eq!( + existing.connect_options().database, + pool.connect_options().database, + "DATABASE_URL changed at runtime, database differs" + ); + + existing + } + }; + + let mut conn = master_pool.acquire().await?; + + cleanup_old_dbs(&mut conn).await?; + + // language=MySQL + conn.execute( + r#" + create table if not exists _sqlx_test_databases ( + db_name text not null, + test_path text not null, + created_at timestamp not null default current_timestamp, + -- BLOB/TEXT columns can only be used as index keys with a prefix length: + -- https://dev.mysql.com/doc/refman/8.4/en/column-indexes.html#column-indexes-prefix + primary key(db_name(63)) + ); + "#, + ) + .await?; + + let db_name = MySql::db_name(args); + do_cleanup(&mut conn, &db_name).await?; + + query("insert into _sqlx_test_databases(db_name, test_path) values (?, ?)") + .bind(&db_name) + .bind(args.test_path) + .execute(&mut *conn) + .await?; + + conn.execute(&format!("create database {db_name}")[..]) + .await?; + + eprintln!("created database {db_name}"); + + Ok(TestContext { + pool_opts: PoolOptions::new() + // Don't allow a single test to take all the connections. + // Most tests shouldn't require more than 5 connections concurrently, + // or else they're likely doing too much in one test. + .max_connections(5) + // Close connections ASAP if left in the idle queue. + .idle_timeout(Some(Duration::from_secs(1))) + .parent(master_pool.clone()), + connect_opts: master_pool + .connect_options() + .deref() + .clone() + .database(&db_name), + db_name, + }) +} + +async fn do_cleanup(conn: &mut MySqlConnection, db_name: &str) -> Result<(), Error> { + let delete_db_command = format!("drop database if exists {db_name};"); + conn.execute(&*delete_db_command).await?; + query("delete from _sqlx_test_databases where db_name = ?") + .bind(db_name) + .execute(&mut *conn) + .await?; + + Ok(()) +} + +/// Pre <0.8.4, test databases were stored by integer ID. +async fn cleanup_old_dbs(conn: &mut MySqlConnection) -> Result<(), Error> { + let res: Result, Error> = query_scalar("select db_id from _sqlx_test_databases") + .fetch_all(&mut *conn) + .await; + + let db_ids = match res { + Ok(db_ids) => db_ids, + Err(e) => { + if let Some(dbe) = e.as_database_error() { + match dbe.downcast_ref::().number() { + // Column `db_id` does not exist: + // https://dev.mysql.com/doc/mysql-errors/8.0/en/server-error-reference.html#error_er_bad_field_error + // + // The table has already been migrated. + 1054 => return Ok(()), + // Table `_sqlx_test_databases` does not exist. + // No cleanup needed. + // https://dev.mysql.com/doc/mysql-errors/8.0/en/server-error-reference.html#error_er_no_such_table + 1146 => return Ok(()), + _ => (), + } + } + + return Err(e); + } + }; + + // Drop old-style test databases. + for id in db_ids { + match conn + .execute(&*format!( + "drop database if exists _sqlx_test_database_{id}" + )) + .await + { + Ok(_deleted) => (), + // Assume a database error just means the DB is still in use. + Err(Error::Database(dbe)) => { + eprintln!("could not clean old test database _sqlx_test_database_{id}: {dbe}"); + } + // Bubble up other errors + Err(e) => return Err(e), + } + } + + conn.execute("drop table if exists _sqlx_test_databases") + .await?; + + Ok(()) +} diff --git a/src-tauri/vendor/sqlx-mysql/src/transaction.rs b/src-tauri/vendor/sqlx-mysql/src/transaction.rs new file mode 100644 index 00000000..545cb5f4 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/transaction.rs @@ -0,0 +1,86 @@ +use std::borrow::Cow; + +use futures_core::future::BoxFuture; + +use crate::connection::Waiting; +use crate::error::Error; +use crate::executor::Executor; +use crate::protocol::text::Query; +use crate::{MySql, MySqlConnection}; + +pub(crate) use sqlx_core::transaction::*; + +/// Implementation of [`TransactionManager`] for MySQL. +pub struct MySqlTransactionManager; + +impl TransactionManager for MySqlTransactionManager { + type Database = MySql; + + fn begin<'conn>( + conn: &'conn mut MySqlConnection, + statement: Option>, + ) -> BoxFuture<'conn, Result<(), Error>> { + Box::pin(async move { + let depth = conn.inner.transaction_depth; + let statement = match statement { + // custom `BEGIN` statements are not allowed if we're already in a transaction + // (we need to issue a `SAVEPOINT` instead) + Some(_) if depth > 0 => return Err(Error::InvalidSavePointStatement), + Some(statement) => statement, + None => begin_ansi_transaction_sql(depth), + }; + conn.execute(&*statement).await?; + if !conn.in_transaction() { + return Err(Error::BeginFailed); + } + conn.inner.transaction_depth += 1; + + Ok(()) + }) + } + + fn commit(conn: &mut MySqlConnection) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let depth = conn.inner.transaction_depth; + + if depth > 0 { + conn.execute(&*commit_ansi_transaction_sql(depth)).await?; + conn.inner.transaction_depth = depth - 1; + } + + Ok(()) + }) + } + + fn rollback(conn: &mut MySqlConnection) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let depth = conn.inner.transaction_depth; + + if depth > 0 { + conn.execute(&*rollback_ansi_transaction_sql(depth)).await?; + conn.inner.transaction_depth = depth - 1; + } + + Ok(()) + }) + } + + fn start_rollback(conn: &mut MySqlConnection) { + let depth = conn.inner.transaction_depth; + + if depth > 0 { + conn.inner.stream.waiting.push_back(Waiting::Result); + conn.inner.stream.sequence_id = 0; + conn.inner + .stream + .write_packet(Query(&rollback_ansi_transaction_sql(depth))) + .expect("BUG: unexpected error queueing ROLLBACK"); + + conn.inner.transaction_depth = depth - 1; + } + } + + fn get_transaction_depth(conn: &MySqlConnection) -> usize { + conn.inner.transaction_depth + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/type_checking.rs b/src-tauri/vendor/sqlx-mysql/src/type_checking.rs new file mode 100644 index 00000000..3f3ce583 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/type_checking.rs @@ -0,0 +1,65 @@ +// Type mappings used by the macros and `Debug` impls. + +#[allow(unused_imports)] +use sqlx_core as sqlx; + +use crate::MySql; + +impl_type_checking!( + MySql { + u8, + u16, + u32, + u64, + i8, + i16, + i32, + i64, + f32, + f64, + + // ordering is important here as otherwise we might infer strings to be binary + // CHAR, VAR_CHAR, TEXT + String, + + // BINARY, VAR_BINARY, BLOB + Vec, + + // Types from third-party crates need to be referenced at a known path + // for the macros to work, but we don't want to require the user to add extra dependencies. + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveTime, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveDate, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveDateTime, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::DateTime, + + #[cfg(feature = "time")] + sqlx::types::time::Time, + + #[cfg(feature = "time")] + sqlx::types::time::Date, + + #[cfg(feature = "time")] + sqlx::types::time::PrimitiveDateTime, + + #[cfg(feature = "time")] + sqlx::types::time::OffsetDateTime, + + #[cfg(feature = "bigdecimal")] + sqlx::types::BigDecimal, + + #[cfg(feature = "rust_decimal")] + sqlx::types::Decimal, + + #[cfg(feature = "json")] + sqlx::types::JsonValue, + }, + ParamChecking::Weak, + feature-types: info => info.__type_feature_gate(), +); diff --git a/src-tauri/vendor/sqlx-mysql/src/type_info.rs b/src-tauri/vendor/sqlx-mysql/src/type_info.rs new file mode 100644 index 00000000..a80b233f --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/type_info.rs @@ -0,0 +1,118 @@ +use std::fmt::{self, Display, Formatter}; + +pub(crate) use sqlx_core::type_info::*; + +use crate::protocol::text::{ColumnDefinition, ColumnFlags, ColumnType}; + +/// Type information for a MySql type. +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct MySqlTypeInfo { + pub(crate) r#type: ColumnType, + pub(crate) flags: ColumnFlags, + + // [max_size] for integer types, this is (M) in BIT(M) or TINYINT(M) + #[cfg_attr(feature = "offline", serde(default))] + pub(crate) max_size: Option, +} + +impl MySqlTypeInfo { + pub(crate) const fn binary(ty: ColumnType) -> Self { + Self { + r#type: ty, + flags: ColumnFlags::BINARY, + max_size: None, + } + } + + #[doc(hidden)] + pub const fn __enum() -> Self { + // Newer versions of MySQL seem to expect that a parameter binding of `MYSQL_TYPE_ENUM` + // means that the value is encoded as an integer. + // + // For "strong" enums inputted as strings, we need to specify this type instead + // for wider compatibility. This works on all covered versions of MySQL and MariaDB. + // + // Annoyingly, MySQL's developer documentation doesn't really explain this anywhere; + // this had to be determined experimentally. + Self { + r#type: ColumnType::String, + flags: ColumnFlags::ENUM, + max_size: None, + } + } + + #[doc(hidden)] + pub fn __type_feature_gate(&self) -> Option<&'static str> { + match self.r#type { + ColumnType::Date | ColumnType::Time | ColumnType::Timestamp | ColumnType::Datetime => { + Some("time") + } + + ColumnType::Json => Some("json"), + ColumnType::NewDecimal => Some("bigdecimal"), + + _ => None, + } + } + + pub(crate) fn from_column(column: &ColumnDefinition) -> Self { + Self { + r#type: column.r#type, + flags: column.flags, + max_size: Some(column.max_size), + } + } +} + +impl Display for MySqlTypeInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.pad(self.name()) + } +} + +impl TypeInfo for MySqlTypeInfo { + fn is_null(&self) -> bool { + matches!(self.r#type, ColumnType::Null) + } + + fn name(&self) -> &str { + self.r#type.name(self.flags, self.max_size) + } +} + +impl PartialEq for MySqlTypeInfo { + fn eq(&self, other: &MySqlTypeInfo) -> bool { + if self.r#type != other.r#type { + return false; + } + + match self.r#type { + ColumnType::Tiny + | ColumnType::Short + | ColumnType::Long + | ColumnType::Int24 + | ColumnType::LongLong => { + return self.flags.contains(ColumnFlags::UNSIGNED) + == other.flags.contains(ColumnFlags::UNSIGNED); + } + + // for string types, check that our charset matches + ColumnType::VarChar + | ColumnType::Blob + | ColumnType::TinyBlob + | ColumnType::MediumBlob + | ColumnType::LongBlob + | ColumnType::String + | ColumnType::VarString + | ColumnType::Enum => { + return self.flags == other.flags; + } + _ => {} + } + + true + } +} + +impl Eq for MySqlTypeInfo {} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/bigdecimal.rs b/src-tauri/vendor/sqlx-mysql/src/types/bigdecimal.rs new file mode 100644 index 00000000..11bca048 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/bigdecimal.rs @@ -0,0 +1,33 @@ +use bigdecimal::BigDecimal; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::protocol::text::ColumnType; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for BigDecimal { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::NewDecimal) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!(ty.r#type, ColumnType::Decimal | ColumnType::NewDecimal) + } +} + +impl Encode<'_, MySql> for BigDecimal { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for BigDecimal { + fn decode(value: MySqlValueRef<'_>) -> Result { + Ok(value.as_str()?.parse()?) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/bool.rs b/src-tauri/vendor/sqlx-mysql/src/types/bool.rs new file mode 100644 index 00000000..7d8c243d --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/bool.rs @@ -0,0 +1,43 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{ + protocol::text::{ColumnFlags, ColumnType}, + MySql, MySqlTypeInfo, MySqlValueRef, +}; + +impl Type for bool { + fn type_info() -> MySqlTypeInfo { + // MySQL has no actual `BOOLEAN` type, the type is an alias of `TINYINT(1)` + MySqlTypeInfo { + flags: ColumnFlags::BINARY | ColumnFlags::UNSIGNED, + max_size: Some(1), + r#type: ColumnType::Tiny, + } + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!( + ty.r#type, + ColumnType::Tiny + | ColumnType::Short + | ColumnType::Long + | ColumnType::Int24 + | ColumnType::LongLong + | ColumnType::Bit + ) + } +} + +impl Encode<'_, MySql> for bool { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + >::encode(*self as i8, buf) + } +} + +impl Decode<'_, MySql> for bool { + fn decode(value: MySqlValueRef<'_>) -> Result { + Ok(>::decode(value)? != 0) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/bytes.rs b/src-tauri/vendor/sqlx-mysql/src/types/bytes.rs new file mode 100644 index 00000000..ade079ad --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/bytes.rs @@ -0,0 +1,85 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::protocol::text::ColumnType; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for [u8] { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Blob) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!( + ty.r#type, + ColumnType::VarChar + | ColumnType::Blob + | ColumnType::TinyBlob + | ColumnType::MediumBlob + | ColumnType::LongBlob + | ColumnType::String + | ColumnType::VarString + | ColumnType::Enum + ) + } +} + +impl Encode<'_, MySql> for &'_ [u8] { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_bytes_lenenc(self); + + Ok(IsNull::No) + } +} + +impl<'r> Decode<'r, MySql> for &'r [u8] { + fn decode(value: MySqlValueRef<'r>) -> Result { + value.as_bytes() + } +} + +impl Type for Box<[u8]> { + fn type_info() -> MySqlTypeInfo { + <&[u8] as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&[u8] as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Box<[u8]> { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + <&[u8] as Encode>::encode(self.as_ref(), buf) + } +} + +impl<'r> Decode<'r, MySql> for Box<[u8]> { + fn decode(value: MySqlValueRef<'r>) -> Result { + <&[u8] as Decode>::decode(value).map(Box::from) + } +} + +impl Type for Vec { + fn type_info() -> MySqlTypeInfo { + <[u8] as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&[u8] as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Vec { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + <&[u8] as Encode>::encode(&**self, buf) + } +} + +impl Decode<'_, MySql> for Vec { + fn decode(value: MySqlValueRef<'_>) -> Result { + <&[u8] as Decode>::decode(value).map(ToOwned::to_owned) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/chrono.rs b/src-tauri/vendor/sqlx-mysql/src/types/chrono.rs new file mode 100644 index 00000000..39e215be --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/chrono.rs @@ -0,0 +1,362 @@ +use bytes::Buf; +use chrono::{ + DateTime, Datelike, Local, NaiveDate, NaiveDateTime, NaiveTime, TimeZone, Timelike, Utc, +}; +use sqlx_core::database::Database; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::{BoxDynError, UnexpectedNullError}; +use crate::protocol::text::ColumnType; +use crate::type_info::MySqlTypeInfo; +use crate::types::{MySqlTime, MySqlTimeSign, Type}; +use crate::{MySql, MySqlValueFormat, MySqlValueRef}; + +impl Type for DateTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Timestamp) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!(ty.r#type, ColumnType::Datetime | ColumnType::Timestamp) + } +} + +/// Note: assumes the connection's `time_zone` is set to `+00:00` (UTC). +impl Encode<'_, MySql> for DateTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + Encode::::encode(self.naive_utc(), buf) + } +} + +/// Note: assumes the connection's `time_zone` is set to `+00:00` (UTC). +impl<'r> Decode<'r, MySql> for DateTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + let naive: NaiveDateTime = Decode::::decode(value)?; + + Ok(Utc.from_utc_datetime(&naive)) + } +} + +impl Type for DateTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Timestamp) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!(ty.r#type, ColumnType::Datetime | ColumnType::Timestamp) + } +} + +/// Note: assumes the connection's `time_zone` is set to `+00:00` (UTC). +impl Encode<'_, MySql> for DateTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + Encode::::encode(self.naive_utc(), buf) + } +} + +/// Note: assumes the connection's `time_zone` is set to `+00:00` (UTC). +impl<'r> Decode<'r, MySql> for DateTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + Ok( as Decode<'r, MySql>>::decode(value)?.with_timezone(&Local)) + } +} + +impl Type for NaiveTime { + fn type_info() -> MySqlTypeInfo { + MySqlTime::type_info() + } +} + +impl Encode<'_, MySql> for NaiveTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + let len = naive_time_encoded_len(self); + buf.push(len); + + // NaiveTime is not negative + buf.push(0); + + // Number of days in the interval; always 0 for time-of-day values. + // https://mariadb.com/kb/en/resultset-row/#teimstamp-binary-encoding + buf.extend_from_slice(&[0_u8; 4]); + + encode_time(self, len > 8, buf); + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + naive_time_encoded_len(self) as usize + 1 // plus length byte + } +} + +/// Decode from a `TIME` value. +/// +/// ### Errors +/// Returns an error if the `TIME` value is negative or exceeds `23:59:59.999999`. +impl<'r> Decode<'r, MySql> for NaiveTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + // Covers most possible failure modes. + MySqlTime::decode(value)?.try_into() + } + // Retaining this parsing for now as it allows us to cross-check our impl. + MySqlValueFormat::Text => { + let s = value.as_str()?; + NaiveTime::parse_from_str(s, "%H:%M:%S%.f").map_err(Into::into) + } + } + } +} + +impl TryFrom for NaiveTime { + type Error = BoxDynError; + + fn try_from(time: MySqlTime) -> Result { + NaiveTime::from_hms_micro_opt( + time.hours(), + time.minutes() as u32, + time.seconds() as u32, + time.microseconds(), + ) + .ok_or_else(|| format!("Cannot convert `MySqlTime` value to `NaiveTime`: {time}").into()) + } +} + +impl From for chrono::TimeDelta { + fn from(time: MySqlTime) -> Self { + chrono::TimeDelta::new(time.whole_seconds_signed(), time.subsec_nanos()) + .expect("BUG: chrono::TimeDelta should have a greater range than MySqlTime") + } +} + +impl TryFrom for MySqlTime { + type Error = BoxDynError; + + fn try_from(value: chrono::TimeDelta) -> Result { + let sign = if value < chrono::TimeDelta::zero() { + MySqlTimeSign::Negative + } else { + MySqlTimeSign::Positive + }; + + Ok( + // `std::time::Duration` has a greater positive range than `TimeDelta` + // which makes it a great intermediate if you ignore the sign. + MySqlTime::try_from(value.abs().to_std()?)?.with_sign(sign), + ) + } +} + +impl Type for chrono::TimeDelta { + fn type_info() -> MySqlTypeInfo { + MySqlTime::type_info() + } +} + +impl<'r> Decode<'r, MySql> for chrono::TimeDelta { + fn decode(value: ::ValueRef<'r>) -> Result { + Ok(MySqlTime::decode(value)?.into()) + } +} + +impl Type for NaiveDate { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Date) + } +} + +impl Encode<'_, MySql> for NaiveDate { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.push(4); + + encode_date(self, buf)?; + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + 5 + } +} + +impl<'r> Decode<'r, MySql> for NaiveDate { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + // Row decoding should have left the length prefix. + if buf.is_empty() { + return Err("empty buffer".into()); + } + + decode_date(&buf[1..])?.ok_or_else(|| UnexpectedNullError.into()) + } + + MySqlValueFormat::Text => { + let s = value.as_str()?; + NaiveDate::parse_from_str(s, "%Y-%m-%d").map_err(Into::into) + } + } + } +} + +impl Type for NaiveDateTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Datetime) + } +} + +impl Encode<'_, MySql> for NaiveDateTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + let len = naive_dt_encoded_len(self); + buf.push(len); + + encode_date(&self.date(), buf)?; + + if len > 4 { + encode_time(&self.time(), len > 7, buf); + } + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + naive_dt_encoded_len(self) as usize + 1 // plus length byte + } +} + +impl<'r> Decode<'r, MySql> for NaiveDateTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + if buf.is_empty() { + return Err("empty buffer".into()); + } + + let len = buf[0]; + let date = decode_date(&buf[1..])?.ok_or(UnexpectedNullError)?; + + let dt = if len > 4 { + date.and_time(decode_time(len - 4, &buf[5..])?) + } else { + date.and_hms_opt(0, 0, 0) + .expect("expected `NaiveDate::and_hms_opt(0, 0, 0)` to be valid") + }; + + Ok(dt) + } + + MySqlValueFormat::Text => { + let s = value.as_str()?; + NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f").map_err(Into::into) + } + } + } +} + +fn encode_date(date: &NaiveDate, buf: &mut Vec) -> Result<(), BoxDynError> { + // MySQL supports years 1000 - 9999 + let year = u16::try_from(date.year()) + .map_err(|_| format!("NaiveDateTime out of range for Mysql: {date}"))?; + + buf.extend_from_slice(&year.to_le_bytes()); + + // `NaiveDate` guarantees the ranges of these values + #[allow(clippy::cast_possible_truncation)] + { + buf.push(date.month() as u8); + buf.push(date.day() as u8); + } + + Ok(()) +} + +fn decode_date(mut buf: &[u8]) -> Result, BoxDynError> { + match buf.len() { + // MySQL specifies that if there are no bytes, this is all zeros + 0 => Ok(None), + 4.. => { + let year = buf.get_u16_le() as i32; + let month = buf[0] as u32; + let day = buf[1] as u32; + + let date = NaiveDate::from_ymd_opt(year, month, day) + .ok_or_else(|| format!("server returned invalid date: {year}/{month}/{day}"))?; + + Ok(Some(date)) + } + len => Err(format!("expected at least 4 bytes for date, got {len}").into()), + } +} + +fn encode_time(time: &NaiveTime, include_micros: bool, buf: &mut Vec) { + // `NaiveTime` API guarantees the ranges of these values + #[allow(clippy::cast_possible_truncation)] + { + buf.push(time.hour() as u8); + buf.push(time.minute() as u8); + buf.push(time.second() as u8); + } + + if include_micros { + buf.extend((time.nanosecond() / 1000).to_le_bytes()); + } +} + +fn decode_time(len: u8, mut buf: &[u8]) -> Result { + let hour = buf.get_u8(); + let minute = buf.get_u8(); + let seconds = buf.get_u8(); + + let micros = if len > 3 { + // microseconds : int + buf.get_uint_le(buf.len()) + } else { + 0 + }; + + let micros = u32::try_from(micros) + .map_err(|_| format!("server returned microseconds out of range: {micros}"))?; + + NaiveTime::from_hms_micro_opt(hour as u32, minute as u32, seconds as u32, micros) + .ok_or_else(|| format!("server returned invalid time: {hour:02}:{minute:02}:{seconds:02}; micros: {micros}").into()) +} + +#[inline(always)] +fn naive_dt_encoded_len(time: &NaiveDateTime) -> u8 { + // to save space the packet can be compressed: + match ( + time.hour(), + time.minute(), + time.second(), + #[allow(deprecated)] + time.timestamp_subsec_nanos(), + ) { + // if hour, minutes, seconds and micro_seconds are all 0, + // length is 4 and no other field is sent + (0, 0, 0, 0) => 4, + + // if micro_seconds is 0, length is 7 + // and micro_seconds is not sent + (_, _, _, 0) => 7, + + // otherwise length is 11 + (_, _, _, _) => 11, + } +} + +#[inline(always)] +fn naive_time_encoded_len(time: &NaiveTime) -> u8 { + if time.nanosecond() == 0 { + // if micro_seconds is 0, length is 8 and micro_seconds is not sent + 8 + } else { + // otherwise length is 12 + 12 + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/float.rs b/src-tauri/vendor/sqlx-mysql/src/types/float.rs new file mode 100644 index 00000000..44acb31b --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/float.rs @@ -0,0 +1,108 @@ +use byteorder::{ByteOrder, LittleEndian}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::protocol::text::ColumnType; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueFormat, MySqlValueRef}; + +fn real_compatible(ty: &MySqlTypeInfo) -> bool { + // NOTE: `DECIMAL` is explicitly excluded because floating-point numbers have different semantics. + matches!(ty.r#type, ColumnType::Float | ColumnType::Double) +} + +impl Type for f32 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Float) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + real_compatible(ty) + } +} + +impl Type for f64 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Double) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + real_compatible(ty) + } +} + +impl Encode<'_, MySql> for f32 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for f64 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for f32 { + fn decode(value: MySqlValueRef<'_>) -> Result { + Ok(match value.format() { + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + match buf.len() { + // These functions panic if `buf` is not exactly the right size. + 4 => LittleEndian::read_f32(buf), + // MySQL can return 8-byte DOUBLE values for a FLOAT + // We take and truncate to f32 as that's the same behavior as *in* MySQL, + #[allow(clippy::cast_possible_truncation)] + 8 => LittleEndian::read_f64(buf) as f32, + other => { + // Users may try to decode a DECIMAL as floating point; + // inform them why that's a bad idea. + return Err(format!( + "expected a FLOAT as 4 or 8 bytes, got {other} bytes; \ + note that decoding DECIMAL as `f32` is not supported \ + due to differing semantics" + ) + .into()); + } + } + } + + MySqlValueFormat::Text => value.as_str()?.parse()?, + }) + } +} + +impl Decode<'_, MySql> for f64 { + fn decode(value: MySqlValueRef<'_>) -> Result { + Ok(match value.format() { + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + // The `read_*` functions panic if `buf` is not exactly the right size. + match buf.len() { + // Allow implicit widening here + 4 => LittleEndian::read_f32(buf) as f64, + 8 => LittleEndian::read_f64(buf), + other => { + // Users may try to decode a DECIMAL as floating point; + // inform them why that's a bad idea. + return Err(format!( + "expected a DOUBLE as 4 or 8 bytes, got {other} bytes; \ + note that decoding DECIMAL as `f64` is not supported \ + due to differing semantics" + ) + .into()); + } + } + } + MySqlValueFormat::Text => value.as_str()?.parse()?, + }) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/inet.rs b/src-tauri/vendor/sqlx-mysql/src/types/inet.rs new file mode 100644 index 00000000..19e59028 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/inet.rs @@ -0,0 +1,92 @@ +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for Ipv4Addr { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Ipv4Addr { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Ipv4Addr { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &str type to decode from MySQL + let text = <&str as Decode>::decode(value)?; + + // parse a Ipv4Addr from the text + text.parse().map_err(Into::into) + } +} + +impl Type for Ipv6Addr { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Ipv6Addr { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Ipv6Addr { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &str type to decode from MySQL + let text = <&str as Decode>::decode(value)?; + + // parse a Ipv6Addr from the text + text.parse().map_err(Into::into) + } +} + +impl Type for IpAddr { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for IpAddr { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for IpAddr { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &str type to decode from MySQL + let text = <&str as Decode>::decode(value)?; + + // parse a IpAddr from the text + text.parse().map_err(Into::into) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/int.rs b/src-tauri/vendor/sqlx-mysql/src/types/int.rs new file mode 100644 index 00000000..0e5b1622 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/int.rs @@ -0,0 +1,139 @@ +use byteorder::{ByteOrder, LittleEndian}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::protocol::text::{ColumnFlags, ColumnType}; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueFormat, MySqlValueRef}; + +fn int_compatible(ty: &MySqlTypeInfo) -> bool { + matches!( + ty.r#type, + ColumnType::Tiny + | ColumnType::Short + | ColumnType::Long + | ColumnType::Int24 + | ColumnType::LongLong + ) && !ty.flags.contains(ColumnFlags::UNSIGNED) +} + +impl Type for i8 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Tiny) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + int_compatible(ty) + } +} + +impl Type for i16 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Short) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + int_compatible(ty) + } +} + +impl Type for i32 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Long) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + int_compatible(ty) + } +} + +impl Type for i64 { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::LongLong) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + int_compatible(ty) + } +} + +impl Encode<'_, MySql> for i8 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for i16 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for i32 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for i64 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +fn int_decode(value: MySqlValueRef<'_>) -> Result { + Ok(match value.format() { + MySqlValueFormat::Text => value.as_str()?.parse()?, + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + // Check conditions that could cause `read_int()` to panic. + if buf.is_empty() { + return Err("empty buffer".into()); + } + + if buf.len() > 8 { + return Err(format!( + "expected no more than 8 bytes for integer value, got {}", + buf.len() + ) + .into()); + } + + LittleEndian::read_int(buf, buf.len()) + } + }) +} + +impl Decode<'_, MySql> for i8 { + fn decode(value: MySqlValueRef<'_>) -> Result { + int_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for i16 { + fn decode(value: MySqlValueRef<'_>) -> Result { + int_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for i32 { + fn decode(value: MySqlValueRef<'_>) -> Result { + int_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for i64 { + fn decode(value: MySqlValueRef<'_>) -> Result { + int_decode(value) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/json.rs b/src-tauri/vendor/sqlx-mysql/src/types/json.rs new file mode 100644 index 00000000..b47baa35 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/json.rs @@ -0,0 +1,68 @@ +use serde::{Deserialize, Serialize}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::protocol::text::ColumnType; +use crate::types::{Json, Type}; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for Json { + fn type_info() -> MySqlTypeInfo { + // MySql uses the `CHAR` type to pass JSON data from and to the client + // NOTE: This is forwards-compatible with MySQL v8+ as CHAR is a common transmission format + // and has nothing to do with the native storage ability of MySQL v8+ + MySqlTypeInfo::binary(ColumnType::String) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + ty.r#type == ColumnType::Json + || <&str as Type>::compatible(ty) + || <&[u8] as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Json +where + T: Serialize, +{ + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + // Encode JSON as a length-prefixed string. + // + // The previous implementation encoded into an intermediate buffer to get the final length. + // This is because the length prefix for the string is itself length-encoded, so we have + // to know the length first before we can start encoding in the buffer... or do we? + // + // The docs suggest that the integer length-encoding doesn't actually enforce a range on + // the value itself as long as it fits in the chosen encoding, so why not just choose + // the full length encoding to begin with? Then we can just reserve the space up-front + // and encode directly into the buffer. + // + // If someone is storing a JSON value it's likely large enough that the overhead of using + // the full-length integer encoding doesn't really matter. And if it's so large it overflows + // a `u64` then the process is likely to run OOM during the encoding process first anyway. + + let lenenc_start = buf.len(); + + buf.extend_from_slice(&[0u8; 9]); + + let encode_start = buf.len(); + self.encode_to(buf)?; + let encoded_len = (buf.len() - encode_start) as u64; + + // This prefix indicates that the following 8 bytes are a little-endian integer. + buf[lenenc_start] = 0xFE; + buf[lenenc_start + 1..][..8].copy_from_slice(&encoded_len.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl<'r, T> Decode<'r, MySql> for Json +where + T: 'r + Deserialize<'r>, +{ + fn decode(value: MySqlValueRef<'r>) -> Result { + Json::decode_from_string(value.as_str()?) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/mod.rs b/src-tauri/vendor/sqlx-mysql/src/types/mod.rs new file mode 100644 index 00000000..dc5105fa --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/mod.rs @@ -0,0 +1,191 @@ +//! Conversions between Rust and **MySQL/MariaDB** types. +//! +//! # Types +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `bool` | TINYINT(1), BOOLEAN, BOOL (see below) | +//! | `i8` | TINYINT | +//! | `i16` | SMALLINT | +//! | `i32` | INT | +//! | `i64` | BIGINT | +//! | `u8` | TINYINT UNSIGNED | +//! | `u16` | SMALLINT UNSIGNED | +//! | `u32` | INT UNSIGNED | +//! | `u64` | BIGINT UNSIGNED | +//! | `f32` | FLOAT | +//! | `f64` | DOUBLE | +//! | `&str`, [`String`] | VARCHAR, CHAR, TEXT | +//! | `&[u8]`, `Vec` | VARBINARY, BINARY, BLOB | +//! | `IpAddr` | VARCHAR, TEXT | +//! | `Ipv4Addr` | INET4 (MariaDB-only), VARCHAR, TEXT | +//! | `Ipv6Addr` | INET6 (MariaDB-only), VARCHAR, TEXT | +//! | [`MySqlTime`] | TIME (encode and decode full range) | +//! | [`Duration`][std::time::Duration] | TIME (for decoding positive values only) | +//! +//! ##### Note: `BOOLEAN`/`BOOL` Type +//! MySQL and MariaDB treat `BOOLEAN` as an alias of the `TINYINT` type: +//! +//! * [Using Data Types from Other Database Engines (MySQL)](https://dev.mysql.com/doc/refman/8.0/en/other-vendor-data-types.html) +//! * [BOOLEAN (MariaDB)](https://mariadb.com/kb/en/boolean/) +//! +//! For the most part, you can simply use the Rust type `bool` when encoding or decoding a value +//! using the dynamic query interface, or passing a boolean as a parameter to the query macros +//! (`query!()` _et al._). +//! +//! However, because the MySQL wire protocol does not distinguish between `TINYINT` and `BOOLEAN`, +//! the query macros cannot know that a `TINYINT` column is semantically a boolean. +//! By default, they will map a `TINYINT` column as `i8` instead, as that is the safer assumption. +//! +//! Thus, you must use the type override syntax in the query to tell the macros you are expecting +//! a `bool` column. See the docs for `query!()` and `query_as!()` for details on this syntax. +//! +//! ### NOTE: MySQL's `TIME` type is signed +//! MySQL's `TIME` type can be used as either a time-of-day value, or a signed interval. +//! Thus, it may take on negative values. +//! +//! Decoding a [`std::time::Duration`] returns an error if the `TIME` value is negative. +//! +//! ### [`chrono`](https://crates.io/crates/chrono) +//! +//! Requires the `chrono` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `chrono::DateTime` | TIMESTAMP | +//! | `chrono::DateTime` | TIMESTAMP | +//! | `chrono::NaiveDateTime` | DATETIME | +//! | `chrono::NaiveDate` | DATE | +//! | `chrono::NaiveTime` | TIME (time-of-day only) | +//! | `chrono::TimeDelta` | TIME (decodes full range; see note for encoding) | +//! +//! ### NOTE: MySQL's `TIME` type is dual-purpose +//! MySQL's `TIME` type can be used as either a time-of-day value, or an interval. +//! However, `chrono::NaiveTime` is designed only to represent a time-of-day. +//! +//! Decoding a `TIME` value as `chrono::NaiveTime` will return an error if the value is out of range. +//! +//! The [`MySqlTime`] type supports the full range and it also implements `TryInto`. +//! +//! Decoding a `chrono::TimeDelta` also supports the full range. +//! +//! To encode a `chrono::TimeDelta`, convert it to [`MySqlTime`] first using `TryFrom`/`TryInto`. +//! +//! ### [`time`](https://crates.io/crates/time) +//! +//! Requires the `time` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `time::PrimitiveDateTime` | DATETIME | +//! | `time::OffsetDateTime` | TIMESTAMP | +//! | `time::Date` | DATE | +//! | `time::Time` | TIME (time-of-day only) | +//! | `time::Duration` | TIME (decodes full range; see note for encoding) | +//! +//! ### NOTE: MySQL's `TIME` type is dual-purpose +//! MySQL's `TIME` type can be used as either a time-of-day value, or an interval. +//! However, `time::Time` is designed only to represent a time-of-day. +//! +//! Decoding a `TIME` value as `time::Time` will return an error if the value is out of range. +//! +//! The [`MySqlTime`] type supports the full range, and it also implements `TryInto`. +//! +//! Decoding a `time::Duration` also supports the full range. +//! +//! To encode a `time::Duration`, convert it to [`MySqlTime`] first using `TryFrom`/`TryInto`. +//! +//! ### [`bigdecimal`](https://crates.io/crates/bigdecimal) +//! Requires the `bigdecimal` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `bigdecimal::BigDecimal` | DECIMAL | +//! +//! ### [`decimal`](https://crates.io/crates/rust_decimal) +//! Requires the `decimal` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `rust_decimal::Decimal` | DECIMAL | +//! +//! ### [`uuid`](https://crates.io/crates/uuid) +//! +//! Requires the `uuid` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `uuid::Uuid` | BINARY(16) (see note) | +//! | `uuid::fmt::Hyphenated` | CHAR(36), VARCHAR, TEXT, UUID (MariaDB-only) | +//! | `uuid::fmt::Simple` | CHAR(32), VARCHAR, TEXT | +//! +//! #### Note: `Uuid` uses binary format +//! +//! MySQL does not have a native datatype for UUIDs. +//! The `UUID()` function returns a 36-character `TEXT` value, +//! which encourages storing UUIDs as text. +//! +//! MariaDB's `UUID` type stores and retrieves as text, though it has a better representation +//! for index sorting (see [MariaDB manual: UUID data-type][mariadb-uuid] for details). +//! +//! As an opinionated library, SQLx chose to map `uuid::Uuid` to/from binary format by default +//! (16 bytes, the raw value of a UUID; SQL type `BINARY(16)`). +//! This saves 20 bytes over the text format for each value. +//! +//! The `impl Decode for Uuid` does not support the text format, and will return an error. +//! +//! If you want to use the text format compatible with the `UUID()` function, +//! use [`uuid::fmt::Hyphenated`][::uuid::fmt::Hyphenated] in the place of `Uuid`. +//! +//! The MySQL official blog has an article showing how to support both binary and text format UUIDs +//! by storing the binary and adding a generated column for the text format, though this is rather +//! verbose and fiddly: +//! +//! [mariadb-uuid]: https://mariadb.com/kb/en/uuid-data-type/ +//! +//! ### [`json`](https://crates.io/crates/serde_json) +//! +//! Requires the `json` Cargo feature flag. +//! +//! | Rust type | MySQL/MariaDB type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | [`Json`] | JSON | +//! | `serde_json::JsonValue` | JSON | +//! | `&serde_json::value::RawValue` | JSON | +//! +//! # Nullable +//! +//! In addition, `Option` is supported where `T` implements `Type`. An `Option` represents +//! a potentially `NULL` value from MySQL/MariaDB. + +pub(crate) use sqlx_core::types::*; + +pub use mysql_time::{MySqlTime, MySqlTimeError, MySqlTimeSign}; + +mod bool; +mod bytes; +mod float; +mod inet; +mod int; +mod mysql_time; +mod str; +mod text; +mod uint; + +#[cfg(feature = "json")] +mod json; + +#[cfg(feature = "bigdecimal")] +mod bigdecimal; + +#[cfg(feature = "rust_decimal")] +mod rust_decimal; + +#[cfg(feature = "chrono")] +mod chrono; + +#[cfg(feature = "time")] +mod time; + +#[cfg(feature = "uuid")] +mod uuid; diff --git a/src-tauri/vendor/sqlx-mysql/src/types/mysql_time.rs b/src-tauri/vendor/sqlx-mysql/src/types/mysql_time.rs new file mode 100644 index 00000000..b549af57 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/mysql_time.rs @@ -0,0 +1,712 @@ +//! The [`MysqlTime`] type. + +use crate::protocol::text::ColumnType; +use crate::{MySql, MySqlTypeInfo, MySqlValueFormat}; +use bytes::{Buf, BufMut}; +use sqlx_core::database::Database; +use sqlx_core::decode::Decode; +use sqlx_core::encode::{Encode, IsNull}; +use sqlx_core::error::BoxDynError; +use sqlx_core::types::Type; +use std::cmp::Ordering; +use std::fmt::{Debug, Display, Formatter, Write}; +use std::time::Duration; + +// Similar to `PgInterval` +/// Container for a MySQL `TIME` value, which may be an interval or a time-of-day. +/// +/// Allowed range is `-838:59:59.0` to `838:59:59.0`. +/// +/// If this value is used for a time-of-day, the range should be `00:00:00.0` to `23:59:59.999999`. +/// You can use [`Self::is_valid_time_of_day()`] to check this easily. +/// +/// * [MySQL Manual 13.2.3: The TIME Type](https://dev.mysql.com/doc/refman/8.3/en/time.html) +/// * [MariaDB Manual: TIME](https://mariadb.com/kb/en/time/) +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub struct MySqlTime { + pub(crate) sign: MySqlTimeSign, + pub(crate) magnitude: TimeMagnitude, +} + +// By using a subcontainer for the actual time magnitude, +// we can still use a derived `Ord` implementation and just flip the comparison for negative values. +#[derive(Debug, Copy, Clone, Ord, PartialOrd, Eq, PartialEq)] +pub(crate) struct TimeMagnitude { + pub(crate) hours: u32, + pub(crate) minutes: u8, + pub(crate) seconds: u8, + pub(crate) microseconds: u32, +} + +const MAGNITUDE_ZERO: TimeMagnitude = TimeMagnitude { + hours: 0, + minutes: 0, + seconds: 0, + microseconds: 0, +}; + +/// Maximum magnitude (positive or negative). +const MAGNITUDE_MAX: TimeMagnitude = TimeMagnitude { + hours: MySqlTime::HOURS_MAX, + minutes: 59, + seconds: 59, + // Surprisingly this is not 999_999 which is why `MySqlTimeError::SubsecondExcess`. + microseconds: 0, +}; + +/// The sign for a [`MySqlTime`] type. +#[derive(Debug, Copy, Clone, Ord, PartialOrd, Eq, PartialEq)] +pub enum MySqlTimeSign { + // The protocol actually specifies negative as 1 and positive as 0, + // but by specifying variants this way we can derive `Ord` and it works as expected. + /// The interval is negative (invalid for time-of-day values). + Negative, + /// The interval is positive, or represents a time-of-day. + Positive, +} + +/// Errors returned by [`MySqlTime::new()`]. +#[derive(Debug, thiserror::Error)] +pub enum MySqlTimeError { + /// A field of [`MySqlTime`] exceeded its max range. + #[error("`MySqlTime` field `{field}` cannot exceed {max}, got {value}")] + FieldRange { + field: &'static str, + max: u32, + value: u64, + }, + /// Error returned for time magnitudes (positive or negative) between `838:59:59.0` and `839:00:00.0`. + /// + /// Other range errors should be covered by [`Self::FieldRange`] for the `hours` field. + /// + /// For applications which can tolerate rounding, a valid truncated value is provided. + #[error( + "`MySqlTime` cannot exceed +/-838:59:59.000000; got {sign}838:59:59.{microseconds:06}" + )] + SubsecondExcess { + /// The sign of the magnitude. + sign: MySqlTimeSign, + /// The number of microseconds over the maximum. + microseconds: u32, + /// The truncated value, + /// either [`MySqlTime::MIN`] if negative or [`MySqlTime::MAX`] if positive. + truncated: MySqlTime, + }, + /// MySQL coerces `-00:00:00` to `00:00:00` but this API considers that an error. + /// + /// For applications which can tolerate coercion, you can convert this error to [`MySqlTime::ZERO`]. + #[error("attempted to construct a `MySqlTime` value of negative zero")] + NegativeZero, +} + +impl MySqlTime { + /// The `MySqlTime` value corresponding to `TIME '0:00:00.0'` (zero). + pub const ZERO: Self = MySqlTime { + sign: MySqlTimeSign::Positive, + magnitude: MAGNITUDE_ZERO, + }; + + /// The `MySqlTime` value corresponding to `TIME '838:59:59.0'` (max value). + pub const MAX: Self = MySqlTime { + sign: MySqlTimeSign::Positive, + magnitude: MAGNITUDE_MAX, + }; + + /// The `MySqlTime` value corresponding to `TIME '-838:59:59.0'` (min value). + pub const MIN: Self = MySqlTime { + sign: MySqlTimeSign::Negative, + // Same magnitude, opposite sign. + magnitude: MAGNITUDE_MAX, + }; + + // The maximums for the other values are self-evident, but not necessarily this one. + pub(crate) const HOURS_MAX: u32 = 838; + + /// Construct a [`MySqlTime`] that is valid for use as a `TIME` value. + /// + /// ### Errors + /// * [`MySqlTimeError::NegativeZero`] if all fields are 0 but `sign` is [`MySqlTimeSign::Negative`]. + /// * [`MySqlTimeError::FieldRange`] if any field is out of range: + /// * `hours > 838` + /// * `minutes > 59` + /// * `seconds > 59` + /// * `microseconds > 999_999` + /// * [`MySqlTimeError::SubsecondExcess`] if the magnitude is less than one second over the maximum. + /// * Durations 839 hours or greater are covered by `FieldRange`. + pub fn new( + sign: MySqlTimeSign, + hours: u32, + minutes: u8, + seconds: u8, + microseconds: u32, + ) -> Result { + macro_rules! check_fields { + ($($name:ident: $max:expr),+ $(,)?) => { + $( + if $name > $max { + return Err(MySqlTimeError::FieldRange { + field: stringify!($name), + max: $max as u32, + value: $name as u64 + }) + } + )+ + } + } + + check_fields!( + hours: Self::HOURS_MAX, + minutes: 59, + seconds: 59, + microseconds: 999_999 + ); + + let values = TimeMagnitude { + hours, + minutes, + seconds, + microseconds, + }; + + if sign.is_negative() && values == MAGNITUDE_ZERO { + return Err(MySqlTimeError::NegativeZero); + } + + // This is only `true` if less than 1 second over the maximum magnitude + if values > MAGNITUDE_MAX { + return Err(MySqlTimeError::SubsecondExcess { + sign, + microseconds, + truncated: if sign.is_positive() { + Self::MAX + } else { + Self::MIN + }, + }); + } + + Ok(Self { + sign, + magnitude: values, + }) + } + + /// Update the `sign` of this value. + pub fn with_sign(self, sign: MySqlTimeSign) -> Self { + Self { sign, ..self } + } + + /// Return the sign (positive or negative) for this TIME value. + pub fn sign(&self) -> MySqlTimeSign { + self.sign + } + + /// Returns `true` if `self` is zero (equal to [`Self::ZERO`]). + pub fn is_zero(&self) -> bool { + self == &Self::ZERO + } + + /// Returns `true` if `self` is positive or zero, `false` if negative. + pub fn is_positive(&self) -> bool { + self.sign.is_positive() + } + + /// Returns `true` if `self` is negative, `false` if positive or zero. + pub fn is_negative(&self) -> bool { + self.sign.is_positive() + } + + /// Returns `true` if this interval is a valid time-of-day. + /// + /// If `true`, the sign is positive and `hours` is not greater than 23. + pub fn is_valid_time_of_day(&self) -> bool { + self.sign.is_positive() && self.hours() < 24 + } + + /// Get the total number of hours in this interval, from 0 to 838. + /// + /// If this value represents a time-of-day, the range is 0 to 23. + pub fn hours(&self) -> u32 { + self.magnitude.hours + } + + /// Get the number of minutes in this interval, from 0 to 59. + pub fn minutes(&self) -> u8 { + self.magnitude.minutes + } + + /// Get the number of seconds in this interval, from 0 to 59. + pub fn seconds(&self) -> u8 { + self.magnitude.seconds + } + + /// Get the number of seconds in this interval, from 0 to 999,999. + pub fn microseconds(&self) -> u32 { + self.magnitude.microseconds + } + + /// Convert this TIME value to a [`std::time::Duration`]. + /// + /// Returns `None` if this value is negative (cannot be represented). + pub fn to_duration(&self) -> Option { + self.is_positive() + .then(|| Duration::new(self.whole_seconds() as u64, self.subsec_nanos())) + } + + /// Get the whole number of seconds (`seconds + (minutes * 60) + (hours * 3600)`) in this time. + /// + /// Sign is ignored. + pub(crate) fn whole_seconds(&self) -> u32 { + // If `hours` does not exceed 838 then this cannot overflow. + self.hours() * 3600 + self.minutes() as u32 * 60 + self.seconds() as u32 + } + + #[cfg_attr(not(any(feature = "time", feature = "chrono")), allow(dead_code))] + pub(crate) fn whole_seconds_signed(&self) -> i64 { + self.whole_seconds() as i64 * self.sign.signum() as i64 + } + + pub(crate) fn subsec_nanos(&self) -> u32 { + self.microseconds() * 1000 + } + + fn encoded_len(&self) -> u8 { + if self.is_zero() { + 0 + } else if self.microseconds() == 0 { + 8 + } else { + 12 + } + } +} + +impl PartialOrd for MySqlTime { + fn partial_cmp(&self, other: &MySqlTime) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for MySqlTime { + fn cmp(&self, other: &Self) -> Ordering { + // If the sides have different signs, we just need to compare those. + if self.sign != other.sign { + return self.sign.cmp(&other.sign); + } + + // We've checked that both sides have the same sign + match self.sign { + MySqlTimeSign::Positive => self.magnitude.cmp(&other.magnitude), + // Reverse the comparison for negative values (smaller negative magnitude = greater) + MySqlTimeSign::Negative => other.magnitude.cmp(&self.magnitude), + } + } +} + +impl Display for MySqlTime { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + let TimeMagnitude { + hours, + minutes, + seconds, + microseconds, + } = self.magnitude; + + // Obeys the `+` flag. + Display::fmt(&self.sign(), f)?; + + write!(f, "{hours}:{minutes:02}:{seconds:02}")?; + + // Write microseconds if not zero or a nonzero precision was explicitly requested. + if f.precision().map_or(microseconds != 0, |it| it != 0) { + f.write_char('.')?; + + let mut remaining_precision = f.precision(); + let mut remainder = microseconds; + let mut power_of_10 = 10u32.pow(5); + + // Write digits from most-significant to least, up to the requested precision. + while remainder > 0 && remaining_precision != Some(0) { + let digit = remainder / power_of_10; + // 1 % 1 = 0 + remainder %= power_of_10; + power_of_10 /= 10; + + write!(f, "{digit}")?; + + if let Some(remaining_precision) = &mut remaining_precision { + *remaining_precision = remaining_precision.saturating_sub(1); + } + } + + // If any requested precision remains, pad with zeroes. + if let Some(precision) = remaining_precision.filter(|it| *it != 0) { + write!(f, "{:0precision$}", 0)?; + } + } + + Ok(()) + } +} + +impl Type for MySqlTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Time) + } +} + +impl<'r> Decode<'r, MySql> for MySqlTime { + fn decode(value: ::ValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + let mut buf = value.as_bytes()?; + + // Row decoding should have left the length byte on the front. + if buf.is_empty() { + return Err("empty buffer".into()); + } + + let length = buf.get_u8(); + + // MySQL specifies that if all fields are 0 then the length is 0 and no further data is sent + // https://dev.mysql.com/doc/internals/en/binary-protocol-value.html + if length == 0 { + return Ok(Self::ZERO); + } + + if !matches!(buf.len(), 8 | 12) { + return Err(format!( + "expected 8 or 12 bytes for TIME value, got {}", + buf.len() + ) + .into()); + } + + let sign = MySqlTimeSign::from_byte(buf.get_u8())?; + // The wire protocol includes days but the text format doesn't. Isn't that crazy? + let days = buf.get_u32_le(); + let hours = buf.get_u8(); + let minutes = buf.get_u8(); + let seconds = buf.get_u8(); + + let microseconds = if !buf.is_empty() { buf.get_u32_le() } else { 0 }; + + let whole_hours = days + .checked_mul(24) + .and_then(|days_to_hours| days_to_hours.checked_add(hours as u32)) + .ok_or("overflow calculating whole hours from `days * 24 + hours`")?; + + Ok(Self::new( + sign, + whole_hours, + minutes, + seconds, + microseconds, + )?) + } + MySqlValueFormat::Text => parse(value.as_str()?), + } + } +} + +impl<'q> Encode<'q, MySql> for MySqlTime { + fn encode_by_ref( + &self, + buf: &mut ::ArgumentBuffer<'q>, + ) -> Result { + if self.is_zero() { + buf.put_u8(0); + return Ok(IsNull::No); + } + + buf.put_u8(self.encoded_len()); + buf.put_u8(self.sign.to_byte()); + + let TimeMagnitude { + hours: whole_hours, + minutes, + seconds, + microseconds, + } = self.magnitude; + + let days = whole_hours / 24; + let hours = (whole_hours % 24) as u8; + + buf.put_u32_le(days); + buf.put_u8(hours); + buf.put_u8(minutes); + buf.put_u8(seconds); + + if microseconds != 0 { + buf.put_u32_le(microseconds); + } + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + self.encoded_len() as usize + 1 + } +} + +/// Convert [`MySqlTime`] from [`std::time::Duration`]. +/// +/// ### Note: Precision Truncation +/// [`Duration`] supports nanosecond precision, but MySQL `TIME` values only support microsecond +/// precision. +/// +/// For simplicity, higher precision values are truncated when converting. +/// If you prefer another rounding mode instead, you should apply that to the `Duration` first. +/// +/// See also: [MySQL Manual, section 13.2.6: Fractional Seconds in Time Values](https://dev.mysql.com/doc/refman/8.3/en/fractional-seconds.html) +/// +/// ### Errors: +/// Returns [`MySqlTimeError::FieldRange`] if the given duration is longer than `838:59:59.999999`. +/// +impl TryFrom for MySqlTime { + type Error = MySqlTimeError; + + fn try_from(value: Duration) -> Result { + let hours = value.as_secs() / 3600; + let rem_seconds = value.as_secs() % 3600; + let minutes = (rem_seconds / 60) as u8; + let seconds = (rem_seconds % 60) as u8; + + // Simply divides by 1000 + let microseconds = value.subsec_micros(); + + Self::new( + MySqlTimeSign::Positive, + hours.try_into().map_err(|_| MySqlTimeError::FieldRange { + field: "hours", + max: Self::HOURS_MAX, + value: hours, + })?, + minutes, + seconds, + microseconds, + ) + } +} + +impl MySqlTimeSign { + fn from_byte(b: u8) -> Result { + match b { + 0 => Ok(Self::Positive), + 1 => Ok(Self::Negative), + other => Err(format!("expected 0 or 1 for TIME sign byte, got {other}").into()), + } + } + + fn to_byte(self) -> u8 { + match self { + // We can't use `#[repr(u8)]` because this is opposite of the ordering we want from `Ord` + Self::Negative => 1, + Self::Positive => 0, + } + } + + fn signum(&self) -> i32 { + match self { + Self::Negative => -1, + Self::Positive => 1, + } + } + + /// Returns `true` if positive, `false` if negative. + pub fn is_positive(&self) -> bool { + matches!(self, Self::Positive) + } + + /// Returns `true` if negative, `false` if positive. + pub fn is_negative(&self) -> bool { + matches!(self, Self::Negative) + } +} + +impl Display for MySqlTimeSign { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::Positive if f.sign_plus() => f.write_char('+'), + Self::Negative => f.write_char('-'), + _ => Ok(()), + } + } +} + +impl Type for Duration { + fn type_info() -> MySqlTypeInfo { + MySqlTime::type_info() + } +} + +impl<'r> Decode<'r, MySql> for Duration { + fn decode(value: ::ValueRef<'r>) -> Result { + let time = MySqlTime::decode(value)?; + + time.to_duration().ok_or_else(|| { + format!("`std::time::Duration` can only decode positive TIME values; got {time}").into() + }) + } +} + +// Not exposing this as a `FromStr` impl currently because `MySqlTime` is not designed to be +// a general interchange type. +fn parse(text: &str) -> Result { + let mut segments = text.split(':'); + + let hours = segments + .next() + .ok_or("expected hours segment, got nothing")?; + + let minutes = segments + .next() + .ok_or("expected minutes segment, got nothing")?; + + let seconds = segments + .next() + .ok_or("expected seconds segment, got nothing")?; + + // Include the sign in parsing for convenience; + // the allowed range of whole hours is much smaller than `i32`'s positive range. + let hours: i32 = hours + .parse() + .map_err(|e| format!("error parsing hours from {text:?} (segment {hours:?}): {e}"))?; + + let sign = if hours.is_negative() { + MySqlTimeSign::Negative + } else { + MySqlTimeSign::Positive + }; + + let hours = hours.unsigned_abs(); + + let minutes: u8 = minutes + .parse() + .map_err(|e| format!("error parsing minutes from {text:?} (segment {minutes:?}): {e}"))?; + + let (seconds, microseconds): (u8, u32) = if let Some((seconds, microseconds)) = + seconds.split_once('.') + { + ( + seconds.parse().map_err(|e| { + format!("error parsing seconds from {text:?} (segment {seconds:?}): {e}") + })?, + parse_microseconds(microseconds).map_err(|e| { + format!("error parsing microseconds from {text:?} (segment {microseconds:?}): {e}") + })?, + ) + } else { + ( + seconds.parse().map_err(|e| { + format!("error parsing seconds from {text:?} (segment {seconds:?}): {e}") + })?, + 0, + ) + }; + + Ok(MySqlTime::new(sign, hours, minutes, seconds, microseconds)?) +} + +/// Parse microseconds from a fractional seconds string. +fn parse_microseconds(micros: &str) -> Result { + const EXPECTED_DIGITS: usize = 6; + + match micros.len() { + 0 => Err("empty string".into()), + len @ ..=EXPECTED_DIGITS => { + // Fewer than 6 digits, multiply to the correct magnitude + let micros: u32 = micros.parse()?; + // cast cannot overflow + #[allow(clippy::cast_possible_truncation)] + Ok(micros * 10u32.pow((EXPECTED_DIGITS - len) as u32)) + } + // More digits than expected, truncate + _ => Ok(micros[..EXPECTED_DIGITS].parse()?), + } +} + +#[cfg(test)] +mod tests { + use super::MySqlTime; + use crate::types::MySqlTimeSign; + + use super::parse_microseconds; + + #[test] + fn test_display() { + assert_eq!(MySqlTime::ZERO.to_string(), "0:00:00"); + + assert_eq!(format!("{:.0}", MySqlTime::ZERO), "0:00:00"); + + assert_eq!(format!("{:.3}", MySqlTime::ZERO), "0:00:00.000"); + + assert_eq!(format!("{:.6}", MySqlTime::ZERO), "0:00:00.000000"); + + assert_eq!(format!("{:.9}", MySqlTime::ZERO), "0:00:00.000000000"); + + assert_eq!(format!("{:.0}", MySqlTime::MAX), "838:59:59"); + + assert_eq!(format!("{:.3}", MySqlTime::MAX), "838:59:59.000"); + + assert_eq!(format!("{:.6}", MySqlTime::MAX), "838:59:59.000000"); + + assert_eq!(format!("{:.9}", MySqlTime::MAX), "838:59:59.000000000"); + + assert_eq!(format!("{:+.0}", MySqlTime::MAX), "+838:59:59"); + + assert_eq!(format!("{:+.3}", MySqlTime::MAX), "+838:59:59.000"); + + assert_eq!(format!("{:+.6}", MySqlTime::MAX), "+838:59:59.000000"); + + assert_eq!(format!("{:+.9}", MySqlTime::MAX), "+838:59:59.000000000"); + + assert_eq!(format!("{:.0}", MySqlTime::MIN), "-838:59:59"); + + assert_eq!(format!("{:.3}", MySqlTime::MIN), "-838:59:59.000"); + + assert_eq!(format!("{:.6}", MySqlTime::MIN), "-838:59:59.000000"); + + assert_eq!(format!("{:.9}", MySqlTime::MIN), "-838:59:59.000000000"); + + let positive = MySqlTime::new(MySqlTimeSign::Positive, 123, 45, 56, 890011).unwrap(); + + assert_eq!(positive.to_string(), "123:45:56.890011"); + assert_eq!(format!("{positive:.0}"), "123:45:56"); + assert_eq!(format!("{positive:.3}"), "123:45:56.890"); + assert_eq!(format!("{positive:.6}"), "123:45:56.890011"); + assert_eq!(format!("{positive:.9}"), "123:45:56.890011000"); + + assert_eq!(format!("{positive:+.0}"), "+123:45:56"); + assert_eq!(format!("{positive:+.3}"), "+123:45:56.890"); + assert_eq!(format!("{positive:+.6}"), "+123:45:56.890011"); + assert_eq!(format!("{positive:+.9}"), "+123:45:56.890011000"); + + let negative = MySqlTime::new(MySqlTimeSign::Negative, 123, 45, 56, 890011).unwrap(); + + assert_eq!(negative.to_string(), "-123:45:56.890011"); + assert_eq!(format!("{negative:.0}"), "-123:45:56"); + assert_eq!(format!("{negative:.3}"), "-123:45:56.890"); + assert_eq!(format!("{negative:.6}"), "-123:45:56.890011"); + assert_eq!(format!("{negative:.9}"), "-123:45:56.890011000"); + } + + #[test] + fn test_parse_microseconds() { + assert_eq!(parse_microseconds("010").unwrap(), 10_000); + + assert_eq!(parse_microseconds("0100000000").unwrap(), 10_000); + + assert_eq!(parse_microseconds("890").unwrap(), 890_000); + + assert_eq!(parse_microseconds("0890").unwrap(), 89_000); + + assert_eq!( + // Case in point about not exposing this: + // we always truncate excess precision because it's simpler than rounding + // and MySQL should never return a higher precision. + parse_microseconds("123456789").unwrap(), + 123456, + ); + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/rust_decimal.rs b/src-tauri/vendor/sqlx-mysql/src/types/rust_decimal.rs new file mode 100644 index 00000000..6e78243c --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/rust_decimal.rs @@ -0,0 +1,33 @@ +use rust_decimal::Decimal; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::protocol::text::ColumnType; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for Decimal { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::NewDecimal) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!(ty.r#type, ColumnType::Decimal | ColumnType::NewDecimal) + } +} + +impl Encode<'_, MySql> for Decimal { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Decimal { + fn decode(value: MySqlValueRef<'_>) -> Result { + Ok(value.as_str()?.parse()?) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/str.rs b/src-tauri/vendor/sqlx-mysql/src/types/str.rs new file mode 100644 index 00000000..8233e908 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/str.rs @@ -0,0 +1,116 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::protocol::text::{ColumnFlags, ColumnType}; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; +use std::borrow::Cow; + +impl Type for str { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo { + r#type: ColumnType::VarString, // VARCHAR + flags: ColumnFlags::empty(), + max_size: None, + } + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + // TODO: Support more collations being returned from SQL? + matches!( + ty.r#type, + ColumnType::VarChar + | ColumnType::Blob + | ColumnType::TinyBlob + | ColumnType::MediumBlob + | ColumnType::LongBlob + | ColumnType::String + | ColumnType::VarString + | ColumnType::Enum + ) && !ty.flags.contains(ColumnFlags::BINARY) + } +} + +impl Encode<'_, MySql> for &'_ str { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(self); + + Ok(IsNull::No) + } +} + +impl<'r> Decode<'r, MySql> for &'r str { + fn decode(value: MySqlValueRef<'r>) -> Result { + value.as_str() + } +} + +impl Type for Box { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Box { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + <&str as Encode>::encode(&**self, buf) + } +} + +impl<'r> Decode<'r, MySql> for Box { + fn decode(value: MySqlValueRef<'r>) -> Result { + <&str as Decode>::decode(value).map(Box::from) + } +} + +impl Type for String { + fn type_info() -> MySqlTypeInfo { + >::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + >::compatible(ty) + } +} + +impl Encode<'_, MySql> for String { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + <&str as Encode>::encode(&**self, buf) + } +} + +impl Decode<'_, MySql> for String { + fn decode(value: MySqlValueRef<'_>) -> Result { + <&str as Decode>::decode(value).map(ToOwned::to_owned) + } +} + +impl Type for Cow<'_, str> { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Cow<'_, str> { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + match self { + Cow::Borrowed(str) => <&str as Encode>::encode(*str, buf), + Cow::Owned(str) => <&str as Encode>::encode(&**str, buf), + } + } +} + +impl<'r> Decode<'r, MySql> for Cow<'r, str> { + fn decode(value: MySqlValueRef<'r>) -> Result { + value.as_str().map(Cow::Borrowed) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/text.rs b/src-tauri/vendor/sqlx-mysql/src/types/text.rs new file mode 100644 index 00000000..ad61c1be --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/text.rs @@ -0,0 +1,49 @@ +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; +use sqlx_core::decode::Decode; +use sqlx_core::encode::{Encode, IsNull}; +use sqlx_core::error::BoxDynError; +use sqlx_core::types::{Text, Type}; +use std::fmt::Display; +use std::str::FromStr; + +impl Type for Text { + fn type_info() -> MySqlTypeInfo { + >::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + >::compatible(ty) + } +} + +impl<'q, T> Encode<'q, MySql> for Text +where + T: Display, +{ + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + // We can't really do the trick like with Postgres where we reserve the space for the + // length up-front and then overwrite it later, because MySQL appears to enforce that + // length-encoded integers use the smallest encoding for the value: + // https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_basic_dt_integers.html#sect_protocol_basic_dt_int_le + // + // So we'd have to reserve space for the max-width encoding, format into the buffer, + // then figure out how many bytes our length-encoded integer needs to be and move the + // value bytes down to use up the empty space. + // + // Copying from a completely separate buffer instead is easier. It may or may not be faster + // or slower depending on a ton of different variables, but I don't currently have the time + // to implement both approaches and compare their performance. + Encode::::encode(self.0.to_string(), buf) + } +} + +impl<'r, T> Decode<'r, MySql> for Text +where + T: FromStr, + BoxDynError: From<::Err>, +{ + fn decode(value: MySqlValueRef<'r>) -> Result { + let s: &str = Decode::::decode(value)?; + Ok(Self(s.parse()?)) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/time.rs b/src-tauri/vendor/sqlx-mysql/src/types/time.rs new file mode 100644 index 00000000..e04f8928 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/time.rs @@ -0,0 +1,337 @@ +use byteorder::{ByteOrder, LittleEndian}; +use bytes::Buf; +use sqlx_core::database::Database; +use time::macros::format_description; +use time::{Date, OffsetDateTime, PrimitiveDateTime, Time, UtcOffset}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::{BoxDynError, UnexpectedNullError}; +use crate::protocol::text::ColumnType; +use crate::type_info::MySqlTypeInfo; +use crate::types::{MySqlTime, MySqlTimeSign, Type}; +use crate::{MySql, MySqlValueFormat, MySqlValueRef}; + +impl Type for OffsetDateTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Timestamp) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + matches!(ty.r#type, ColumnType::Datetime | ColumnType::Timestamp) + } +} + +impl Encode<'_, MySql> for OffsetDateTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + let utc_dt = self.to_offset(UtcOffset::UTC); + let primitive_dt = PrimitiveDateTime::new(utc_dt.date(), utc_dt.time()); + + Encode::::encode(primitive_dt, buf) + } +} + +impl<'r> Decode<'r, MySql> for OffsetDateTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + let primitive: PrimitiveDateTime = Decode::::decode(value)?; + + Ok(primitive.assume_utc()) + } +} + +impl Type for Time { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Time) + } +} + +impl Encode<'_, MySql> for Time { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + let len = time_encoded_len(self); + buf.push(len); + + // sign byte: Time is never negative + buf.push(0); + + // Number of days in the interval; always 0 for time-of-day values. + // https://mariadb.com/kb/en/resultset-row/#teimstamp-binary-encoding + buf.extend_from_slice(&[0_u8; 4]); + + encode_time(self, len > 8, buf); + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + time_encoded_len(self) as usize + 1 // plus length byte + } +} + +impl<'r> Decode<'r, MySql> for Time { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + // Should never panic. + MySqlTime::decode(value)?.try_into() + } + + // Retaining this parsing for now as it allows us to cross-check our impl. + MySqlValueFormat::Text => Time::parse( + value.as_str()?, + &format_description!("[hour]:[minute]:[second].[subsecond]"), + ) + .map_err(Into::into), + } + } +} + +impl TryFrom for Time { + type Error = BoxDynError; + + fn try_from(time: MySqlTime) -> Result { + if !time.is_valid_time_of_day() { + return Err(format!("MySqlTime value out of range for `time::Time`: {time}").into()); + } + + #[allow(clippy::cast_possible_truncation)] + Ok(Time::from_hms_micro( + // `is_valid_time_of_day()` ensures this won't overflow + time.hours() as u8, + time.minutes(), + time.seconds(), + time.microseconds(), + )?) + } +} + +impl From for time::Duration { + fn from(time: MySqlTime) -> Self { + // `subsec_nanos()` is guaranteed to be between 0 and 10^9 + #[allow(clippy::cast_possible_wrap)] + time::Duration::new(time.whole_seconds_signed(), time.subsec_nanos() as i32) + } +} + +impl TryFrom for MySqlTime { + type Error = BoxDynError; + + fn try_from(value: time::Duration) -> Result { + let sign = if value.is_negative() { + MySqlTimeSign::Negative + } else { + MySqlTimeSign::Positive + }; + + // Similar to `TryFrom`, use `std::time::Duration` as an intermediate. + Ok(MySqlTime::try_from(std::time::Duration::try_from(value.abs())?)?.with_sign(sign)) + } +} + +impl Type for time::Duration { + fn type_info() -> MySqlTypeInfo { + MySqlTime::type_info() + } +} + +impl<'r> Decode<'r, MySql> for time::Duration { + fn decode(value: ::ValueRef<'r>) -> Result { + Ok(MySqlTime::decode(value)?.into()) + } +} + +impl Type for Date { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Date) + } +} + +impl Encode<'_, MySql> for Date { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.push(4); + + encode_date(self, buf)?; + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + 5 + } +} + +impl<'r> Decode<'r, MySql> for Date { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + // Row decoding should leave the length byte on the front. + if buf.is_empty() { + return Err("empty buffer".into()); + } + + Ok(decode_date(&buf[1..])?.ok_or(UnexpectedNullError)?) + } + MySqlValueFormat::Text => { + let s = value.as_str()?; + Date::parse(s, &format_description!("[year]-[month]-[day]")).map_err(Into::into) + } + } + } +} + +impl Type for PrimitiveDateTime { + fn type_info() -> MySqlTypeInfo { + MySqlTypeInfo::binary(ColumnType::Datetime) + } +} + +impl Encode<'_, MySql> for PrimitiveDateTime { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + let len = primitive_dt_encoded_len(self); + buf.push(len); + + encode_date(&self.date(), buf)?; + + if len > 4 { + encode_time(&self.time(), len > 7, buf); + } + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + primitive_dt_encoded_len(self) as usize + 1 // plus length byte + } +} + +impl<'r> Decode<'r, MySql> for PrimitiveDateTime { + fn decode(value: MySqlValueRef<'r>) -> Result { + match value.format() { + MySqlValueFormat::Binary => { + let mut buf = value.as_bytes()?; + + if buf.is_empty() { + return Err("empty buffer".into()); + } + + let len = buf.get_u8(); + + let date = decode_date(buf)?.ok_or(UnexpectedNullError)?; + + let dt = if len > 4 { + date.with_time(decode_time(&buf[4..])?) + } else { + date.midnight() + }; + + Ok(dt) + } + + MySqlValueFormat::Text => { + let s = value.as_str()?; + + // If there are no nanoseconds parse without them + if s.contains('.') { + PrimitiveDateTime::parse( + s, + &format_description!( + "[year]-[month]-[day] [hour]:[minute]:[second].[subsecond]" + ), + ) + .map_err(Into::into) + } else { + PrimitiveDateTime::parse( + s, + &format_description!("[year]-[month]-[day] [hour]:[minute]:[second]"), + ) + .map_err(Into::into) + } + } + } + } +} + +fn encode_date(date: &Date, buf: &mut Vec) -> Result<(), BoxDynError> { + // MySQL supports years from 1000 - 9999 + let year = + u16::try_from(date.year()).map_err(|_| format!("Date out of range for Mysql: {date}"))?; + + buf.extend_from_slice(&year.to_le_bytes()); + buf.push(date.month().into()); + buf.push(date.day()); + + Ok(()) +} + +fn decode_date(buf: &[u8]) -> Result, BoxDynError> { + if buf.is_empty() { + // zero buffer means a zero date (null) + return Ok(None); + } + + Date::from_calendar_date( + LittleEndian::read_u16(buf) as i32, + time::Month::try_from(buf[2])?, + buf[3], + ) + .map_err(Into::into) + .map(Some) +} + +fn encode_time(time: &Time, include_micros: bool, buf: &mut Vec) { + buf.push(time.hour()); + buf.push(time.minute()); + buf.push(time.second()); + + if include_micros { + buf.extend(&(time.nanosecond() / 1000).to_le_bytes()); + } +} + +fn decode_time(mut buf: &[u8]) -> Result { + let hour = buf.get_u8(); + let minute = buf.get_u8(); + let seconds = buf.get_u8(); + + let micros = if !buf.is_empty() { + // microseconds : int + buf.get_uint_le(buf.len()) + } else { + 0 + }; + + let micros = u32::try_from(micros) + .map_err(|_| format!("MySQL returned microseconds out of range: {micros}"))?; + + Time::from_hms_micro(hour, minute, seconds, micros) + .map_err(|e| format!("Time out of range for MySQL: {e}").into()) +} + +#[inline(always)] +fn primitive_dt_encoded_len(time: &PrimitiveDateTime) -> u8 { + // to save space the packet can be compressed: + match (time.hour(), time.minute(), time.second(), time.nanosecond()) { + // if hour, minutes, seconds and micro_seconds are all 0, + // length is 4 and no other field is sent + (0, 0, 0, 0) => 4, + + // if micro_seconds is 0, length is 7 + // and micro_seconds is not sent + (_, _, _, 0) => 7, + + // otherwise length is 11 + (_, _, _, _) => 11, + } +} + +#[inline(always)] +fn time_encoded_len(time: &Time) -> u8 { + if time.nanosecond() == 0 { + // if micro_seconds is 0, length is 8 and micro_seconds is not sent + 8 + } else { + // otherwise length is 12 + 12 + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/uint.rs b/src-tauri/vendor/sqlx-mysql/src/types/uint.rs new file mode 100644 index 00000000..ca8eb753 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/uint.rs @@ -0,0 +1,162 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::protocol::text::{ColumnFlags, ColumnType}; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueFormat, MySqlValueRef}; +use byteorder::{ByteOrder, LittleEndian}; + +fn uint_type_info(ty: ColumnType) -> MySqlTypeInfo { + MySqlTypeInfo { + r#type: ty, + flags: ColumnFlags::BINARY | ColumnFlags::UNSIGNED, + max_size: None, + } +} + +fn uint_compatible(ty: &MySqlTypeInfo) -> bool { + matches!( + ty.r#type, + ColumnType::Tiny + | ColumnType::Short + | ColumnType::Long + | ColumnType::Int24 + | ColumnType::LongLong + | ColumnType::Year + | ColumnType::Bit + ) && ty.flags.contains(ColumnFlags::UNSIGNED) +} + +impl Type for u8 { + fn type_info() -> MySqlTypeInfo { + uint_type_info(ColumnType::Tiny) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + uint_compatible(ty) + } +} + +impl Type for u16 { + fn type_info() -> MySqlTypeInfo { + uint_type_info(ColumnType::Short) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + uint_compatible(ty) + } +} + +impl Type for u32 { + fn type_info() -> MySqlTypeInfo { + uint_type_info(ColumnType::Long) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + uint_compatible(ty) + } +} + +impl Type for u64 { + fn type_info() -> MySqlTypeInfo { + uint_type_info(ColumnType::LongLong) + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + uint_compatible(ty) + } +} + +impl Encode<'_, MySql> for u8 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for u16 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for u32 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, MySql> for u64 { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.extend(&self.to_le_bytes()); + + Ok(IsNull::No) + } +} + +fn uint_decode(value: MySqlValueRef<'_>) -> Result { + if value.type_info.r#type == ColumnType::Bit { + // NOTE: Regardless of the value format, there is raw binary data here + + let buf = value.as_bytes()?; + let mut value: u64 = 0; + + for b in buf { + value = (*b as u64) | (value << 8); + } + + return Ok(value); + } + + Ok(match value.format() { + MySqlValueFormat::Text => value.as_str()?.parse()?, + + MySqlValueFormat::Binary => { + let buf = value.as_bytes()?; + + // Check conditions that could cause `read_uint()` to panic. + if buf.is_empty() { + return Err("empty buffer".into()); + } + + if buf.len() > 8 { + return Err(format!( + "expected no more than 8 bytes for unsigned integer value, got {}", + buf.len() + ) + .into()); + } + + LittleEndian::read_uint(buf, buf.len()) + } + }) +} + +impl Decode<'_, MySql> for u8 { + fn decode(value: MySqlValueRef<'_>) -> Result { + uint_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for u16 { + fn decode(value: MySqlValueRef<'_>) -> Result { + uint_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for u32 { + fn decode(value: MySqlValueRef<'_>) -> Result { + uint_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Decode<'_, MySql> for u64 { + fn decode(value: MySqlValueRef<'_>) -> Result { + uint_decode(value) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/types/uuid.rs b/src-tauri/vendor/sqlx-mysql/src/types/uuid.rs new file mode 100644 index 00000000..8bd4d37f --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/types/uuid.rs @@ -0,0 +1,108 @@ +use uuid::{ + fmt::{Hyphenated, Simple}, + Uuid, +}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::io::MySqlBufMutExt; +use crate::types::Type; +use crate::{MySql, MySqlTypeInfo, MySqlValueRef}; + +impl Type for Uuid { + fn type_info() -> MySqlTypeInfo { + <&[u8] as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&[u8] as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Uuid { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_bytes_lenenc(self.as_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Uuid { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &[u8] type to decode from MySQL + let bytes = <&[u8] as Decode>::decode(value)?; + + if bytes.len() != 16 { + return Err(format!( + "Expected 16 bytes, got {}; `Uuid` uses binary format for MySQL/MariaDB. \ + For text-formatted UUIDs, use `uuid::fmt::Hyphenated` instead of `Uuid`.", + bytes.len(), + ) + .into()); + } + + // construct a Uuid from the returned bytes + Uuid::from_slice(bytes).map_err(Into::into) + } +} + +impl Type for Hyphenated { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Hyphenated { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Hyphenated { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &str type to decode from MySQL + let text = <&str as Decode>::decode(value)?; + + // parse a UUID from the text + Uuid::parse_str(text) + .map_err(Into::into) + .map(|u| u.hyphenated()) + } +} + +impl Type for Simple { + fn type_info() -> MySqlTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &MySqlTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Encode<'_, MySql> for Simple { + fn encode_by_ref(&self, buf: &mut Vec) -> Result { + buf.put_str_lenenc(&self.to_string()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, MySql> for Simple { + fn decode(value: MySqlValueRef<'_>) -> Result { + // delegate to the &str type to decode from MySQL + let text = <&str as Decode>::decode(value)?; + + // parse a UUID from the text + Uuid::parse_str(text) + .map_err(Into::into) + .map(|u| u.simple()) + } +} diff --git a/src-tauri/vendor/sqlx-mysql/src/value.rs b/src-tauri/vendor/sqlx-mysql/src/value.rs new file mode 100644 index 00000000..fe8b50c6 --- /dev/null +++ b/src-tauri/vendor/sqlx-mysql/src/value.rs @@ -0,0 +1,115 @@ +use std::borrow::Cow; +use std::str::from_utf8; + +use bytes::Bytes; +pub(crate) use sqlx_core::value::*; + +use crate::error::{BoxDynError, UnexpectedNullError}; +use crate::protocol::text::ColumnType; +use crate::{MySql, MySqlTypeInfo}; + +#[derive(Debug, Clone, Copy)] +#[repr(u8)] +pub enum MySqlValueFormat { + Text, + Binary, +} + +/// Implementation of [`Value`] for MySQL. +#[derive(Clone)] +pub struct MySqlValue { + value: Option, + type_info: MySqlTypeInfo, + format: MySqlValueFormat, +} + +/// Implementation of [`ValueRef`] for MySQL. +#[derive(Clone)] +pub struct MySqlValueRef<'r> { + pub(crate) value: Option<&'r [u8]>, + pub(crate) row: Option<&'r Bytes>, + pub(crate) type_info: MySqlTypeInfo, + pub(crate) format: MySqlValueFormat, +} + +impl<'r> MySqlValueRef<'r> { + pub(crate) fn format(&self) -> MySqlValueFormat { + self.format + } + + pub(crate) fn as_bytes(&self) -> Result<&'r [u8], BoxDynError> { + match &self.value { + Some(v) => Ok(v), + None => Err(UnexpectedNullError.into()), + } + } + + pub(crate) fn as_str(&self) -> Result<&'r str, BoxDynError> { + Ok(from_utf8(self.as_bytes()?)?) + } +} + +impl Value for MySqlValue { + type Database = MySql; + + fn as_ref(&self) -> MySqlValueRef<'_> { + MySqlValueRef { + value: self.value.as_deref(), + row: None, + type_info: self.type_info.clone(), + format: self.format, + } + } + + fn type_info(&self) -> Cow<'_, MySqlTypeInfo> { + Cow::Borrowed(&self.type_info) + } + + fn is_null(&self) -> bool { + is_null(self.value.as_deref(), &self.type_info) + } +} + +impl<'r> ValueRef<'r> for MySqlValueRef<'r> { + type Database = MySql; + + fn to_owned(&self) -> MySqlValue { + let value = match (self.row, self.value) { + (Some(row), Some(value)) => Some(row.slice_ref(value)), + + (None, Some(value)) => Some(Bytes::copy_from_slice(value)), + + _ => None, + }; + + MySqlValue { + value, + format: self.format, + type_info: self.type_info.clone(), + } + } + + fn type_info(&self) -> Cow<'_, MySqlTypeInfo> { + Cow::Borrowed(&self.type_info) + } + + #[inline] + fn is_null(&self) -> bool { + is_null(self.value, &self.type_info) + } +} + +fn is_null(value: Option<&[u8]>, ty: &MySqlTypeInfo) -> bool { + if let Some(value) = value { + // zero dates and date times should be treated the same as NULL + if matches!( + ty.r#type, + ColumnType::Date | ColumnType::Timestamp | ColumnType::Datetime + ) && value.starts_with(b"\0") + { + return true; + } + } + + value.is_none() +} diff --git a/src-tauri/vendor/sqlx-postgres/Cargo.toml b/src-tauri/vendor/sqlx-postgres/Cargo.toml new file mode 100644 index 00000000..f8f51d53 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/Cargo.toml @@ -0,0 +1,260 @@ +# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO +# +# When uploading crates to the registry Cargo will automatically +# "normalize" Cargo.toml files for maximal compatibility +# with all versions of Cargo and also rewrite `path` dependencies +# to registry (e.g., crates.io) dependencies. +# +# If you are reading this file be aware that the original Cargo.toml +# will likely look very different (and much more reasonable). +# See Cargo.toml.orig for the original contents. + +[package] +edition = "2021" +name = "sqlx-postgres" +version = "0.8.6" +authors = [ + "Ryan Leckey ", + "Austin Bonander ", + "Chloe Ross ", + "Daniel Akhterov ", +] +description = "PostgreSQL driver implementation for SQLx. Not for direct use; see the `sqlx` crate for details." +documentation = "https://docs.rs/sqlx" +license = "MIT OR Apache-2.0" +repository = "https://github.com/launchbadge/sqlx" + +[dependencies.atoi] +version = "2.0" + +[dependencies.base64] +version = "0.22.0" +features = ["std"] +default-features = false + +[dependencies.bigdecimal] +version = "0.4.0" +optional = true + +[dependencies.bit-vec] +version = "0.6.3" +optional = true + +[dependencies.bitflags] +version = "2" +default-features = false + +[dependencies.byteorder] +version = "1.4.3" +features = ["std"] +default-features = false + +[dependencies.chrono] +version = "0.4.34" +features = [ + "std", + "clock", +] +optional = true +default-features = false + +[dependencies.crc] +version = "3.0.0" + +[dependencies.dotenvy] +version = "0.15.7" +default-features = false + +[dependencies.futures-channel] +version = "0.3.19" +features = [ + "sink", + "alloc", + "std", +] +default-features = false + +[dependencies.futures-core] +version = "0.3.19" +default-features = false + +[dependencies.futures-util] +version = "0.3.19" +features = [ + "alloc", + "sink", + "io", +] +default-features = false + +[dependencies.hex] +version = "0.4.3" + +[dependencies.hkdf] +version = "0.12.0" + +[dependencies.hmac] +version = "0.12.0" +features = ["reset"] +default-features = false + +[dependencies.home] +version = "0.5.5" + +[dependencies.ipnet] +version = "2.3.0" +optional = true + +[dependencies.ipnetwork] +version = "0.20.0" +optional = true + +[dependencies.itoa] +version = "1.0.1" + +[dependencies.log] +version = "0.4.18" + +[dependencies.mac_address] +version = "1.1.5" +optional = true + +[dependencies.md-5] +version = "0.10.0" +default-features = false + +[dependencies.memchr] +version = "2.4.1" +default-features = false + +[dependencies.num-bigint] +version = "0.4.3" +optional = true + +[dependencies.once_cell] +version = "1.9.0" + +[dependencies.rand] +version = "0.8.4" +features = [ + "std", + "std_rng", +] +default-features = false + +[dependencies.rust_decimal] +version = "1.26.1" +features = ["std"] +optional = true +default-features = false + +[dependencies.serde] +version = "1.0.144" +features = ["derive"] + +[dependencies.serde_json] +version = "1.0.85" +features = ["raw_value"] + +[dependencies.sha2] +version = "0.10.0" +default-features = false + +[dependencies.smallvec] +version = "1.7.0" +features = ["serde"] + +[dependencies.sqlx-core] +version = "=0.8.6" +features = ["json"] + +[dependencies.stringprep] +version = "0.1.2" + +[dependencies.thiserror] +version = "2.0.0" + +[dependencies.time] +version = "0.3.36" +features = [ + "formatting", + "parsing", + "macros", +] +optional = true + +[dependencies.tracing] +version = "0.1.37" +features = ["log"] + +[dependencies.uuid] +version = "1.1.2" +optional = true + +[dependencies.whoami] +version = "1.2.1" +default-features = false + +[dev-dependencies.sqlx] +version = "=0.8.6" +features = [ + "postgres", + "derive", +] +default-features = false + +[features] +any = ["sqlx-core/any"] +bigdecimal = [ + "dep:bigdecimal", + "dep:num-bigint", + "sqlx-core/bigdecimal", +] +bit-vec = [ + "dep:bit-vec", + "sqlx-core/bit-vec", +] +chrono = [ + "dep:chrono", + "sqlx-core/chrono", +] +ipnet = [ + "dep:ipnet", + "sqlx-core/ipnet", +] +ipnetwork = [ + "dep:ipnetwork", + "sqlx-core/ipnetwork", +] +json = ["sqlx-core/json"] +mac_address = [ + "dep:mac_address", + "sqlx-core/mac_address", +] +migrate = ["sqlx-core/migrate"] +offline = ["sqlx-core/offline"] +rust_decimal = [ + "dep:rust_decimal", + "rust_decimal/maths", + "sqlx-core/rust_decimal", +] +time = [ + "dep:time", + "sqlx-core/time", +] +uuid = [ + "dep:uuid", + "sqlx-core/uuid", +] + +[target."cfg(target_os = \"windows\")".dependencies.etcetera] +version = "0.8.0" + +[lints.clippy] +cast_possible_truncation = "deny" +cast_possible_wrap = "deny" +cast_sign_loss = "deny" +disallowed_methods = "deny" + +[lints.rust] +warnings = "allow" diff --git a/src-tauri/vendor/sqlx-postgres/Cargo.toml.orig b/src-tauri/vendor/sqlx-postgres/Cargo.toml.orig new file mode 100644 index 00000000..818aadba --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/Cargo.toml.orig @@ -0,0 +1,89 @@ +[package] +name = "sqlx-postgres" +documentation = "https://docs.rs/sqlx" +description = "PostgreSQL driver implementation for SQLx. Not for direct use; see the `sqlx` crate for details." +version.workspace = true +license.workspace = true +edition.workspace = true +authors.workspace = true +repository.workspace = true +# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html + +[features] +any = ["sqlx-core/any"] +json = ["sqlx-core/json"] +migrate = ["sqlx-core/migrate"] +offline = ["sqlx-core/offline"] + +# Type Integration features +bigdecimal = ["dep:bigdecimal", "dep:num-bigint", "sqlx-core/bigdecimal"] +bit-vec = ["dep:bit-vec", "sqlx-core/bit-vec"] +chrono = ["dep:chrono", "sqlx-core/chrono"] +ipnet = ["dep:ipnet", "sqlx-core/ipnet"] +ipnetwork = ["dep:ipnetwork", "sqlx-core/ipnetwork"] +mac_address = ["dep:mac_address", "sqlx-core/mac_address"] +rust_decimal = ["dep:rust_decimal", "rust_decimal/maths", "sqlx-core/rust_decimal"] +time = ["dep:time", "sqlx-core/time"] +uuid = ["dep:uuid", "sqlx-core/uuid"] + +[dependencies] +# Futures crates +futures-channel = { version = "0.3.19", default-features = false, features = ["sink", "alloc", "std"] } +futures-core = { version = "0.3.19", default-features = false } +futures-util = { version = "0.3.19", default-features = false, features = ["alloc", "sink", "io"] } + +# Cryptographic Primitives +crc = "3.0.0" +hkdf = "0.12.0" +hmac = { version = "0.12.0", default-features = false, features = ["reset"]} +md-5 = { version = "0.10.0", default-features = false } +rand = { version = "0.8.4", default-features = false, features = ["std", "std_rng"] } +sha2 = { version = "0.10.0", default-features = false } + +# Type Integrations (versions inherited from `[workspace.dependencies]`) +bigdecimal = { workspace = true, optional = true } +bit-vec = { workspace = true, optional = true } +chrono = { workspace = true, optional = true } +ipnet = { workspace = true, optional = true } +ipnetwork = { workspace = true, optional = true } +mac_address = { workspace = true, optional = true } +rust_decimal = { workspace = true, optional = true } +time = { workspace = true, optional = true } +uuid = { workspace = true, optional = true } + +# Misc +atoi = "2.0" +base64 = { version = "0.22.0", default-features = false, features = ["std"] } +bitflags = { version = "2", default-features = false } +byteorder = { version = "1.4.3", default-features = false, features = ["std"] } +dotenvy = { workspace = true } +hex = "0.4.3" +home = "0.5.5" +itoa = "1.0.1" +log = "0.4.18" +memchr = { version = "2.4.1", default-features = false } +num-bigint = { version = "0.4.3", optional = true } +once_cell = "1.9.0" +smallvec = { version = "1.7.0", features = ["serde"] } +stringprep = "0.1.2" +thiserror = "2.0.0" +tracing = { version = "0.1.37", features = ["log"] } +whoami = { version = "1.2.1", default-features = false } + +serde = { version = "1.0.144", features = ["derive"] } +serde_json = { version = "1.0.85", features = ["raw_value"] } + +[dependencies.sqlx-core] +workspace = true +# We use JSON in the driver implementation itself so there's no reason not to enable it here. +features = ["json"] + +[dev-dependencies.sqlx] +workspace = true +features = ["postgres", "derive"] + +[target.'cfg(target_os = "windows")'.dependencies] +etcetera = "0.8.0" + +[lints] +workspace = true diff --git a/src-tauri/vendor/sqlx-postgres/LICENSE-APACHE b/src-tauri/vendor/sqlx-postgres/LICENSE-APACHE new file mode 100644 index 00000000..c79147e8 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/LICENSE-APACHE @@ -0,0 +1,201 @@ +Apache License +Version 2.0, January 2004 +http://www.apache.org/licenses/ + +TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + +1. Definitions. + +"License" shall mean the terms and conditions for use, reproduction, +and distribution as defined by Sections 1 through 9 of this document. + +"Licensor" shall mean the copyright owner or entity authorized by +the copyright owner that is granting the License. + +"Legal Entity" shall mean the union of the acting entity and all +other entities that control, are controlled by, or are under common +control with that entity. For the purposes of this definition, +"control" means (i) the power, direct or indirect, to cause the +direction or management of such entity, whether by contract or +otherwise, or (ii) ownership of fifty percent (50%) or more of the +outstanding shares, or (iii) beneficial ownership of such entity. + +"You" (or "Your") shall mean an individual or Legal Entity +exercising permissions granted by this License. + +"Source" form shall mean the preferred form for making modifications, +including but not limited to software source code, documentation +source, and configuration files. + +"Object" form shall mean any form resulting from mechanical +transformation or translation of a Source form, including but +not limited to compiled object code, generated documentation, +and conversions to other media types. + +"Work" shall mean the work of authorship, whether in Source or +Object form, made available under the License, as indicated by a +copyright notice that is included in or attached to the work +(an example is provided in the Appendix below). + +"Derivative Works" shall mean any work, whether in Source or Object +form, that is based on (or derived from) the Work and for which the +editorial revisions, annotations, elaborations, or other modifications +represent, as a whole, an original work of authorship. For the purposes +of this License, Derivative Works shall not include works that remain +separable from, or merely link (or bind by name) to the interfaces of, +the Work and Derivative Works thereof. + +"Contribution" shall mean any work of authorship, including +the original version of the Work and any modifications or additions +to that Work or Derivative Works thereof, that is intentionally +submitted to Licensor for inclusion in the Work by the copyright owner +or by an individual or Legal Entity authorized to submit on behalf of +the copyright owner. For the purposes of this definition, "submitted" +means any form of electronic, verbal, or written communication sent +to the Licensor or its representatives, including but not limited to +communication on electronic mailing lists, source code control systems, +and issue tracking systems that are managed by, or on behalf of, the +Licensor for the purpose of discussing and improving the Work, but +excluding communication that is conspicuously marked or otherwise +designated in writing by the copyright owner as "Not a Contribution." + +"Contributor" shall mean Licensor and any individual or Legal Entity +on behalf of whom a Contribution has been received by Licensor and +subsequently incorporated within the Work. + +2. Grant of Copyright License. Subject to the terms and conditions of +this License, each Contributor hereby grants to You a perpetual, +worldwide, non-exclusive, no-charge, royalty-free, irrevocable +copyright license to reproduce, prepare Derivative Works of, +publicly display, publicly perform, sublicense, and distribute the +Work and such Derivative Works in Source or Object form. + +3. Grant of Patent License. Subject to the terms and conditions of +this License, each Contributor hereby grants to You a perpetual, +worldwide, non-exclusive, no-charge, royalty-free, irrevocable +(except as stated in this section) patent license to make, have made, +use, offer to sell, sell, import, and otherwise transfer the Work, +where such license applies only to those patent claims licensable +by such Contributor that are necessarily infringed by their +Contribution(s) alone or by combination of their Contribution(s) +with the Work to which such Contribution(s) was submitted. If You +institute patent litigation against any entity (including a +cross-claim or counterclaim in a lawsuit) alleging that the Work +or a Contribution incorporated within the Work constitutes direct +or contributory patent infringement, then any patent licenses +granted to You under this License for that Work shall terminate +as of the date such litigation is filed. + +4. Redistribution. You may reproduce and distribute copies of the +Work or Derivative Works thereof in any medium, with or without +modifications, and in Source or Object form, provided that You +meet the following conditions: + +(a) You must give any other recipients of the Work or +Derivative Works a copy of this License; and + +(b) You must cause any modified files to carry prominent notices +stating that You changed the files; and + +(c) You must retain, in the Source form of any Derivative Works +that You distribute, all copyright, patent, trademark, and +attribution notices from the Source form of the Work, +excluding those notices that do not pertain to any part of +the Derivative Works; and + +(d) If the Work includes a "NOTICE" text file as part of its +distribution, then any Derivative Works that You distribute must +include a readable copy of the attribution notices contained +within such NOTICE file, excluding those notices that do not +pertain to any part of the Derivative Works, in at least one +of the following places: within a NOTICE text file distributed +as part of the Derivative Works; within the Source form or +documentation, if provided along with the Derivative Works; or, +within a display generated by the Derivative Works, if and +wherever such third-party notices normally appear. The contents +of the NOTICE file are for informational purposes only and +do not modify the License. You may add Your own attribution +notices within Derivative Works that You distribute, alongside +or as an addendum to the NOTICE text from the Work, provided +that such additional attribution notices cannot be construed +as modifying the License. + +You may add Your own copyright statement to Your modifications and +may provide additional or different license terms and conditions +for use, reproduction, or distribution of Your modifications, or +for any such Derivative Works as a whole, provided Your use, +reproduction, and distribution of the Work otherwise complies with +the conditions stated in this License. + +5. Submission of Contributions. Unless You explicitly state otherwise, +any Contribution intentionally submitted for inclusion in the Work +by You to the Licensor shall be under the terms and conditions of +this License, without any additional terms or conditions. +Notwithstanding the above, nothing herein shall supersede or modify +the terms of any separate license agreement you may have executed +with Licensor regarding such Contributions. + +6. Trademarks. This License does not grant permission to use the trade +names, trademarks, service marks, or product names of the Licensor, +except as required for reasonable and customary use in describing the +origin of the Work and reproducing the content of the NOTICE file. + +7. Disclaimer of Warranty. Unless required by applicable law or +agreed to in writing, Licensor provides the Work (and each +Contributor provides its Contributions) on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or +implied, including, without limitation, any warranties or conditions +of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A +PARTICULAR PURPOSE. You are solely responsible for determining the +appropriateness of using or redistributing the Work and assume any +risks associated with Your exercise of permissions under this License. + +8. Limitation of Liability. In no event and under no legal theory, +whether in tort (including negligence), contract, or otherwise, +unless required by applicable law (such as deliberate and grossly +negligent acts) or agreed to in writing, shall any Contributor be +liable to You for damages, including any direct, indirect, special, +incidental, or consequential damages of any character arising as a +result of this License or out of the use or inability to use the +Work (including but not limited to damages for loss of goodwill, +work stoppage, computer failure or malfunction, or any and all +other commercial damages or losses), even if such Contributor +has been advised of the possibility of such damages. + +9. Accepting Warranty or Additional Liability. While redistributing +the Work or Derivative Works thereof, You may choose to offer, +and charge a fee for, acceptance of support, warranty, indemnity, +or other liability obligations and/or rights consistent with this +License. However, in accepting such obligations, You may act only +on Your own behalf and on Your sole responsibility, not on behalf +of any other Contributor, and only if You agree to indemnify, +defend, and hold each Contributor harmless for any liability +incurred by, or claims asserted against, such Contributor by reason +of your accepting any such warranty or additional liability. + +END OF TERMS AND CONDITIONS + +APPENDIX: How to apply the Apache License to your work. + +To apply the Apache License to your work, attach the following +boilerplate notice, with the fields enclosed by brackets "[]" +replaced with your own identifying information. (Don't include +the brackets!) The text should be enclosed in the appropriate +comment syntax for the file format. We also recommend that a +file or class name and description of purpose be included on the +same "printed page" as the copyright notice for easier +identification within third-party archives. + +Copyright 2020 LaunchBadge, LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + +http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. \ No newline at end of file diff --git a/src-tauri/vendor/sqlx-postgres/LICENSE-MIT b/src-tauri/vendor/sqlx-postgres/LICENSE-MIT new file mode 100644 index 00000000..861bf608 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/LICENSE-MIT @@ -0,0 +1,25 @@ +Copyright (c) 2020 LaunchBadge, LLC + +Permission is hereby granted, free of charge, to any +person obtaining a copy of this software and associated +documentation files (the "Software"), to deal in the +Software without restriction, including without +limitation the rights to use, copy, modify, merge, +publish, distribute, sublicense, and/or sell copies of +the Software, and to permit persons to whom the Software +is furnished to do so, subject to the following +conditions: + +The above copyright notice and this permission notice +shall be included in all copies or substantial portions +of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF +ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED +TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A +PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT +SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY +CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION +OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR +IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. diff --git a/src-tauri/vendor/sqlx-postgres/src/advisory_lock.rs b/src-tauri/vendor/sqlx-postgres/src/advisory_lock.rs new file mode 100644 index 00000000..d1aef176 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/advisory_lock.rs @@ -0,0 +1,421 @@ +use crate::error::Result; +use crate::Either; +use crate::PgConnection; +use hkdf::Hkdf; +use once_cell::sync::OnceCell; +use sha2::Sha256; +use std::ops::{Deref, DerefMut}; + +/// A mutex-like type utilizing [Postgres advisory locks]. +/// +/// Advisory locks are a mechanism provided by Postgres to have mutually exclusive or shared +/// locks tracked in the database with application-defined semantics, as opposed to the standard +/// row-level or table-level locks which may not fit all use-cases. +/// +/// This API provides a convenient wrapper for generating and storing the integer keys that +/// advisory locks use, as well as RAII guards for releasing advisory locks when they fall out +/// of scope. +/// +/// This API only handles session-scoped advisory locks (explicitly locked and unlocked, or +/// automatically released when a connection is closed). +/// +/// It is also possible to use transaction-scoped locks but those can be used by beginning a +/// transaction and calling the appropriate lock functions (e.g. `SELECT pg_advisory_xact_lock()`) +/// manually, and cannot be explicitly released, but are automatically released when a transaction +/// ends (is committed or rolled back). +/// +/// Session-level locks can be acquired either inside or outside a transaction and are not +/// tied to transaction semantics; a lock acquired inside a transaction is still held when that +/// transaction is committed or rolled back, until explicitly released or the connection is closed. +/// +/// Locks can be acquired in either shared or exclusive modes, which can be thought of as read locks +/// and write locks, respectively. Multiple shared locks are allowed for the same key, but a single +/// exclusive lock prevents any other lock being taken for a given key until it is released. +/// +/// [Postgres advisory locks]: https://www.postgresql.org/docs/current/explicit-locking.html#ADVISORY-LOCKS +#[derive(Debug, Clone)] +pub struct PgAdvisoryLock { + key: PgAdvisoryLockKey, + /// The query to execute to release this lock. + release_query: OnceCell, +} + +/// A key type natively used by Postgres advisory locks. +/// +/// Currently, Postgres advisory locks have two different key spaces: one keyed by a single +/// 64-bit integer, and one keyed by a pair of two 32-bit integers. The Postgres docs +/// specify that these key spaces "do not overlap": +/// +/// +/// +/// The documentation for the `pg_locks` system view explains further how advisory locks +/// are treated in Postgres: +/// +/// +#[derive(Debug, Clone, PartialEq, Eq)] +#[non_exhaustive] +pub enum PgAdvisoryLockKey { + /// The keyspace designated by a single 64-bit integer. + /// + /// When [PgAdvisoryLock] is constructed with [::new()][PgAdvisoryLock::new()], + /// this is the keyspace used. + BigInt(i64), + /// The keyspace designated by two 32-bit integers. + IntPair(i32, i32), +} + +/// A wrapper for `PgConnection` (or a similar type) that represents a held Postgres advisory lock. +/// +/// Can be acquired by [`PgAdvisoryLock::acquire()`] or [`PgAdvisoryLock::try_acquire()`]. +/// Released on-drop or via [`Self::release_now()`]. +/// +/// ### Note: Release-on-drop is not immediate! +/// On drop, this guard queues a `pg_advisory_unlock()` call on the connection which will be +/// flushed to the server the next time it is used, or when it is returned to +/// a [`PgPool`][crate::PgPool] in the case of +/// [`PoolConnection`][crate::pool::PoolConnection]. +/// +/// This means the lock is not actually released as soon as the guard is dropped. To ensure the +/// lock is eagerly released, you can call [`.release_now().await`][Self::release_now()]. +pub struct PgAdvisoryLockGuard<'lock, C: AsMut> { + lock: &'lock PgAdvisoryLock, + conn: Option, +} + +impl PgAdvisoryLock { + /// Construct a `PgAdvisoryLock` using the given string as a key. + /// + /// This is intended to make it easier to use an advisory lock by using a human-readable string + /// for a key as opposed to manually generating a unique integer key. The generated integer key + /// is guaranteed to be stable and in the single 64-bit integer keyspace + /// (see [`PgAdvisoryLockKey`] for details). + /// + /// This is done by applying the [Hash-based Key Derivation Function (HKDF; IETF RFC 5869)][hkdf] + /// to the bytes of the input string, but in a way that the calculated integer is unlikely + /// to collide with any similar implementations (although we don't currently know of any). + /// See the source of this method for details. + /// + /// [hkdf]: https://datatracker.ietf.org/doc/html/rfc5869 + /// ### Example + /// ```rust + /// use sqlx::postgres::{PgAdvisoryLock, PgAdvisoryLockKey}; + /// + /// let lock = PgAdvisoryLock::new("my first Postgres advisory lock!"); + /// // Negative values are fine because of how Postgres treats advisory lock keys. + /// // See the documentation for the `pg_locks` system view for details. + /// assert_eq!(lock.key(), &PgAdvisoryLockKey::BigInt(-5560419505042474287)); + /// ``` + pub fn new(key_string: impl AsRef) -> Self { + let input_key_material = key_string.as_ref(); + + // HKDF was chosen because it is designed to concentrate the entropy in a variable-length + // input key and produce a higher quality but reduced-length output key with a + // well-specified and reproducible algorithm. + // + // Granted, the input key is usually meant to be pseudorandom and not human readable, + // but we're not trying to produce an unguessable value by any means; just one that's as + // unlikely to already be in use as possible, but still deterministic. + // + // SHA-256 was chosen as the hash function because it's already used in the Postgres driver, + // which should save on codegen and optimization. + + // We don't supply a salt as that is intended to be random, but we want a deterministic key. + let hkdf = Hkdf::::new(None, input_key_material.as_bytes()); + + let mut output_key_material = [0u8; 8]; + + // The first string is the "info" string of the HKDF which is intended to tie the output + // exclusively to SQLx. This should avoid collisions with implementations using a similar + // strategy. If you _want_ this to match some other implementation then you should get + // the calculated integer key from it and use that directly. + // + // Do *not* change this string as it will affect the output! + hkdf.expand( + b"SQLx (Rust) Postgres advisory lock", + &mut output_key_material, + ) + // `Hkdf::expand()` only returns an error if you ask for more than 255 times the digest size. + // This is specified by RFC 5869 but not elaborated upon: + // https://datatracker.ietf.org/doc/html/rfc5869#section-2.3 + // Since we're only asking for 8 bytes, this error shouldn't be returned. + .expect("BUG: `output_key_material` should be of acceptable length"); + + // For ease of use, this method assumes the user doesn't care which keyspace is used. + // + // It doesn't seem likely that someone would care about using the `(int, int)` keyspace + // specifically unless they already had keys to use, in which case they wouldn't + // care about this method. That's why we also provide `with_key()`. + // + // The choice of `from_le_bytes()` is mostly due to x86 being the most popular + // architecture for server software, so it should be a no-op there. + let key = PgAdvisoryLockKey::BigInt(i64::from_le_bytes(output_key_material)); + + tracing::trace!( + ?key, + key_string = ?input_key_material, + "generated key from key string", + ); + + Self::with_key(key) + } + + /// Construct a `PgAdvisoryLock` with a manually supplied key. + pub fn with_key(key: PgAdvisoryLockKey) -> Self { + Self { + key, + release_query: OnceCell::new(), + } + } + + /// Returns the current key. + pub fn key(&self) -> &PgAdvisoryLockKey { + &self.key + } + + // Why doesn't this use `Acquire`? Well, I tried it and got really useless errors + // about "cannot project lifetimes to parent scope". + // + // It has something to do with how lifetimes work on the `Acquire` trait, I couldn't + // be bothered to figure it out. Probably another issue with a lack of `async fn` in traits + // or lazy normalization. + + /// Acquires an exclusive lock using `pg_advisory_lock()`, waiting until the lock is acquired. + /// + /// For a version that returns immediately instead of waiting, see [`Self::try_acquire()`]. + /// + /// A connection-like type is required to execute the call. Allowed types include `PgConnection`, + /// `PoolConnection` and `Transaction`, as well as mutable references to + /// any of these. + /// + /// The returned guard queues a `pg_advisory_unlock()` call on the connection when dropped, + /// which will be executed the next time the connection is used, or when returned to a + /// [`PgPool`][crate::PgPool] in the case of `PoolConnection`. + /// + /// Postgres allows a single connection to acquire a given lock more than once without releasing + /// it first, so in that sense the lock is re-entrant. However, the number of unlock operations + /// must match the number of lock operations for the lock to actually be released. + /// + /// See [Postgres' documentation for the Advisory Lock Functions][advisory-funcs] for details. + /// + /// [advisory-funcs]: https://www.postgresql.org/docs/current/functions-admin.html#FUNCTIONS-ADVISORY-LOCKS + pub async fn acquire>( + &self, + mut conn: C, + ) -> Result> { + match &self.key { + PgAdvisoryLockKey::BigInt(key) => { + crate::query::query("SELECT pg_advisory_lock($1)") + .bind(key) + .execute(conn.as_mut()) + .await?; + } + PgAdvisoryLockKey::IntPair(key1, key2) => { + crate::query::query("SELECT pg_advisory_lock($1, $2)") + .bind(key1) + .bind(key2) + .execute(conn.as_mut()) + .await?; + } + } + + Ok(PgAdvisoryLockGuard::new(self, conn)) + } + + /// Acquires an exclusive lock using `pg_try_advisory_lock()`, returning immediately + /// if the lock could not be acquired. + /// + /// For a version that waits until the lock is acquired, see [`Self::acquire()`]. + /// + /// A connection-like type is required to execute the call. Allowed types include `PgConnection`, + /// `PoolConnection` and `Transaction`, as well as mutable references to + /// any of these. The connection is returned if the lock could not be acquired. + /// + /// The returned guard queues a `pg_advisory_unlock()` call on the connection when dropped, + /// which will be executed the next time the connection is used, or when returned to a + /// [`PgPool`][crate::PgPool] in the case of `PoolConnection`. + /// + /// Postgres allows a single connection to acquire a given lock more than once without releasing + /// it first, so in that sense the lock is re-entrant. However, the number of unlock operations + /// must match the number of lock operations for the lock to actually be released. + /// + /// See [Postgres' documentation for the Advisory Lock Functions][advisory-funcs] for details. + /// + /// [advisory-funcs]: https://www.postgresql.org/docs/current/functions-admin.html#FUNCTIONS-ADVISORY-LOCKS + pub async fn try_acquire>( + &self, + mut conn: C, + ) -> Result, C>> { + let locked: bool = match &self.key { + PgAdvisoryLockKey::BigInt(key) => { + crate::query_scalar::query_scalar("SELECT pg_try_advisory_lock($1)") + .bind(key) + .fetch_one(conn.as_mut()) + .await? + } + PgAdvisoryLockKey::IntPair(key1, key2) => { + crate::query_scalar::query_scalar("SELECT pg_try_advisory_lock($1, $2)") + .bind(key1) + .bind(key2) + .fetch_one(conn.as_mut()) + .await? + } + }; + + if locked { + Ok(Either::Left(PgAdvisoryLockGuard::new(self, conn))) + } else { + Ok(Either::Right(conn)) + } + } + + /// Execute `pg_advisory_unlock()` for this lock's key on the given connection. + /// + /// This is used by [`PgAdvisoryLockGuard::release_now()`] and is also provided for manually + /// releasing the lock from connections returned by [`PgAdvisoryLockGuard::leak()`]. + /// + /// An error should only be returned if there is something wrong with the connection, + /// in which case the lock will be automatically released by the connection closing anyway. + /// + /// The `boolean` value is that returned by `pg_advisory_lock()`. If it is `false`, it + /// indicates that the lock was not actually held by the given connection and that a warning + /// has been logged by the Postgres server. + pub async fn force_release>(&self, mut conn: C) -> Result<(C, bool)> { + let released: bool = match &self.key { + PgAdvisoryLockKey::BigInt(key) => { + crate::query_scalar::query_scalar("SELECT pg_advisory_unlock($1)") + .bind(key) + .fetch_one(conn.as_mut()) + .await? + } + PgAdvisoryLockKey::IntPair(key1, key2) => { + crate::query_scalar::query_scalar("SELECT pg_advisory_unlock($1, $2)") + .bind(key1) + .bind(key2) + .fetch_one(conn.as_mut()) + .await? + } + }; + + Ok((conn, released)) + } + + fn get_release_query(&self) -> &str { + self.release_query.get_or_init(|| match &self.key { + PgAdvisoryLockKey::BigInt(key) => format!("SELECT pg_advisory_unlock({key})"), + PgAdvisoryLockKey::IntPair(key1, key2) => { + format!("SELECT pg_advisory_unlock({key1}, {key2})") + } + }) + } +} + +impl PgAdvisoryLockKey { + /// Converts `Self::Bigint(bigint)` to `Some(bigint)` and all else to `None`. + pub fn as_bigint(&self) -> Option { + if let Self::BigInt(bigint) = self { + Some(*bigint) + } else { + None + } + } +} + +const NONE_ERR: &str = "BUG: PgAdvisoryLockGuard.conn taken"; + +impl<'lock, C: AsMut> PgAdvisoryLockGuard<'lock, C> { + fn new(lock: &'lock PgAdvisoryLock, conn: C) -> Self { + PgAdvisoryLockGuard { + lock, + conn: Some(conn), + } + } + + /// Immediately release the held advisory lock instead of when the connection is next used. + /// + /// An error should only be returned if there is something wrong with the connection, + /// in which case the lock will be automatically released by the connection closing anyway. + /// + /// If `pg_advisory_unlock()` returns `false`, a warning will be logged, both by SQLx as + /// well as the Postgres server. This would only happen if the lock was released without + /// using this guard, or the connection was swapped using [`std::mem::replace()`]. + pub async fn release_now(mut self) -> Result { + let (conn, released) = self + .lock + .force_release(self.conn.take().expect(NONE_ERR)) + .await?; + + if !released { + tracing::warn!( + lock = ?self.lock.key, + "PgAdvisoryLockGuard: advisory lock was not held by the contained connection", + ); + } + + Ok(conn) + } + + /// Cancel the release of the advisory lock, keeping it held until the connection is closed. + /// + /// To manually release the lock later, see [`PgAdvisoryLock::force_release()`]. + pub fn leak(mut self) -> C { + self.conn.take().expect(NONE_ERR) + } +} + +impl<'lock, C: AsMut + AsRef> Deref for PgAdvisoryLockGuard<'lock, C> { + type Target = PgConnection; + + fn deref(&self) -> &Self::Target { + self.conn.as_ref().expect(NONE_ERR).as_ref() + } +} + +/// Mutable access to the underlying connection is provided so it can still be used like normal, +/// even allowing locks to be taken recursively. +/// +/// However, replacing the connection with a different one using, e.g. [`std::mem::replace()`] +/// is a logic error and will cause a warning to be logged by the PostgreSQL server when this +/// guard attempts to release the lock. +impl<'lock, C: AsMut + AsRef> DerefMut + for PgAdvisoryLockGuard<'lock, C> +{ + fn deref_mut(&mut self) -> &mut Self::Target { + self.conn.as_mut().expect(NONE_ERR).as_mut() + } +} + +impl<'lock, C: AsMut + AsRef> AsRef + for PgAdvisoryLockGuard<'lock, C> +{ + fn as_ref(&self) -> &PgConnection { + self.conn.as_ref().expect(NONE_ERR).as_ref() + } +} + +/// Mutable access to the underlying connection is provided so it can still be used like normal, +/// even allowing locks to be taken recursively. +/// +/// However, replacing the connection with a different one using, e.g. [`std::mem::replace()`] +/// is a logic error and will cause a warning to be logged by the PostgreSQL server when this +/// guard attempts to release the lock. +impl<'lock, C: AsMut> AsMut for PgAdvisoryLockGuard<'lock, C> { + fn as_mut(&mut self) -> &mut PgConnection { + self.conn.as_mut().expect(NONE_ERR).as_mut() + } +} + +/// Queues a `pg_advisory_unlock()` call on the wrapped connection which will be flushed +/// to the server the next time it is used, or when it is returned to [`PgPool`][crate::PgPool] +/// in the case of [`PoolConnection`][crate::pool::PoolConnection]. +impl<'lock, C: AsMut> Drop for PgAdvisoryLockGuard<'lock, C> { + fn drop(&mut self) { + if let Some(mut conn) = self.conn.take() { + // Queue a simple query message to execute next time the connection is used. + // The `async fn` versions can safely use the prepared statement protocol, + // but this is the safest way to queue a query to execute on the next opportunity. + conn.as_mut() + .queue_simple_query(self.lock.get_release_query()) + .expect("BUG: PgAdvisoryLock::get_release_query() somehow too long for protocol"); + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/any.rs b/src-tauri/vendor/sqlx-postgres/src/any.rs new file mode 100644 index 00000000..e5b8a366 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/any.rs @@ -0,0 +1,258 @@ +use crate::{ + Either, PgColumn, PgConnectOptions, PgConnection, PgQueryResult, PgRow, PgTransactionManager, + PgTypeInfo, Postgres, +}; +use futures_core::future::BoxFuture; +use futures_core::stream::BoxStream; +use futures_util::{stream, StreamExt, TryFutureExt, TryStreamExt}; +use std::borrow::Cow; +use std::{future, pin::pin}; + +use sqlx_core::any::{ + Any, AnyArguments, AnyColumn, AnyConnectOptions, AnyConnectionBackend, AnyQueryResult, AnyRow, + AnyStatement, AnyTypeInfo, AnyTypeInfoKind, +}; + +use crate::type_info::PgType; +use sqlx_core::connection::Connection; +use sqlx_core::database::Database; +use sqlx_core::describe::Describe; +use sqlx_core::executor::Executor; +use sqlx_core::ext::ustr::UStr; +use sqlx_core::transaction::TransactionManager; + +sqlx_core::declare_driver_with_optional_migrate!(DRIVER = Postgres); + +impl AnyConnectionBackend for PgConnection { + fn name(&self) -> &str { + ::NAME + } + + fn close(self: Box) -> BoxFuture<'static, sqlx_core::Result<()>> { + Connection::close(*self) + } + + fn close_hard(self: Box) -> BoxFuture<'static, sqlx_core::Result<()>> { + Connection::close_hard(*self) + } + + fn ping(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + Connection::ping(self) + } + + fn begin( + &mut self, + statement: Option>, + ) -> BoxFuture<'_, sqlx_core::Result<()>> { + PgTransactionManager::begin(self, statement) + } + + fn commit(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + PgTransactionManager::commit(self) + } + + fn rollback(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + PgTransactionManager::rollback(self) + } + + fn start_rollback(&mut self) { + PgTransactionManager::start_rollback(self) + } + + fn get_transaction_depth(&self) -> usize { + PgTransactionManager::get_transaction_depth(self) + } + + fn shrink_buffers(&mut self) { + Connection::shrink_buffers(self); + } + + fn flush(&mut self) -> BoxFuture<'_, sqlx_core::Result<()>> { + Connection::flush(self) + } + + fn should_flush(&self) -> bool { + Connection::should_flush(self) + } + + #[cfg(feature = "migrate")] + fn as_migrate( + &mut self, + ) -> sqlx_core::Result<&mut (dyn sqlx_core::migrate::Migrate + Send + 'static)> { + Ok(self) + } + + fn fetch_many<'q>( + &'q mut self, + query: &'q str, + persistent: bool, + arguments: Option>, + ) -> BoxStream<'q, sqlx_core::Result>> { + let persistent = persistent && arguments.is_some(); + let arguments = match arguments.as_ref().map(AnyArguments::convert_to).transpose() { + Ok(arguments) => arguments, + Err(error) => { + return stream::once(future::ready(Err(sqlx_core::Error::Encode(error)))).boxed() + } + }; + + Box::pin( + self.run(query, arguments, persistent, None) + .try_flatten_stream() + .map( + move |res: sqlx_core::Result>| match res? { + Either::Left(result) => Ok(Either::Left(map_result(result))), + Either::Right(row) => Ok(Either::Right(AnyRow::try_from(&row)?)), + }, + ), + ) + } + + fn fetch_optional<'q>( + &'q mut self, + query: &'q str, + persistent: bool, + arguments: Option>, + ) -> BoxFuture<'q, sqlx_core::Result>> { + let persistent = persistent && arguments.is_some(); + let arguments = arguments + .as_ref() + .map(AnyArguments::convert_to) + .transpose() + .map_err(sqlx_core::Error::Encode); + + Box::pin(async move { + let arguments = arguments?; + let mut stream = pin!(self.run(query, arguments, persistent, None).await?); + + if let Some(Either::Right(row)) = stream.try_next().await? { + return Ok(Some(AnyRow::try_from(&row)?)); + } + + Ok(None) + }) + } + + fn prepare_with<'c, 'q: 'c>( + &'c mut self, + sql: &'q str, + _parameters: &[AnyTypeInfo], + ) -> BoxFuture<'c, sqlx_core::Result>> { + Box::pin(async move { + let statement = Executor::prepare_with(self, sql, &[]).await?; + AnyStatement::try_from_statement( + sql, + &statement, + statement.metadata.column_names.clone(), + ) + }) + } + + fn describe<'q>(&'q mut self, sql: &'q str) -> BoxFuture<'q, sqlx_core::Result>> { + Box::pin(async move { + let describe = Executor::describe(self, sql).await?; + + let columns = describe + .columns + .iter() + .map(AnyColumn::try_from) + .collect::, _>>()?; + + let parameters = match describe.parameters { + Some(Either::Left(parameters)) => Some(Either::Left( + parameters + .iter() + .enumerate() + .map(|(i, type_info)| { + AnyTypeInfo::try_from(type_info).map_err(|_| { + sqlx_core::Error::AnyDriverError( + format!( + "Any driver does not support type {type_info} of parameter {i}" + ) + .into(), + ) + }) + }) + .collect::, _>>()?, + )), + Some(Either::Right(count)) => Some(Either::Right(count)), + None => None, + }; + + Ok(Describe { + columns, + parameters, + nullable: describe.nullable, + }) + }) + } +} + +impl<'a> TryFrom<&'a PgTypeInfo> for AnyTypeInfo { + type Error = sqlx_core::Error; + + fn try_from(pg_type: &'a PgTypeInfo) -> Result { + Ok(AnyTypeInfo { + kind: match &pg_type.0 { + PgType::Bool => AnyTypeInfoKind::Bool, + PgType::Void => AnyTypeInfoKind::Null, + PgType::Int2 => AnyTypeInfoKind::SmallInt, + PgType::Int4 => AnyTypeInfoKind::Integer, + PgType::Int8 => AnyTypeInfoKind::BigInt, + PgType::Float4 => AnyTypeInfoKind::Real, + PgType::Float8 => AnyTypeInfoKind::Double, + PgType::Bytea => AnyTypeInfoKind::Blob, + PgType::Text | PgType::Varchar => AnyTypeInfoKind::Text, + PgType::DeclareWithName(UStr::Static("citext")) => AnyTypeInfoKind::Text, + _ => { + return Err(sqlx_core::Error::AnyDriverError( + format!("Any driver does not support the Postgres type {pg_type:?}").into(), + )) + } + }, + }) + } +} + +impl<'a> TryFrom<&'a PgColumn> for AnyColumn { + type Error = sqlx_core::Error; + + fn try_from(col: &'a PgColumn) -> Result { + let type_info = + AnyTypeInfo::try_from(&col.type_info).map_err(|e| sqlx_core::Error::ColumnDecode { + index: col.name.to_string(), + source: e.into(), + })?; + + Ok(AnyColumn { + ordinal: col.ordinal, + name: col.name.clone(), + type_info, + }) + } +} + +impl<'a> TryFrom<&'a PgRow> for AnyRow { + type Error = sqlx_core::Error; + + fn try_from(row: &'a PgRow) -> Result { + AnyRow::map_from(row, row.metadata.column_names.clone()) + } +} + +impl<'a> TryFrom<&'a AnyConnectOptions> for PgConnectOptions { + type Error = sqlx_core::Error; + + fn try_from(value: &'a AnyConnectOptions) -> Result { + let mut opts = PgConnectOptions::parse_from_url(&value.database_url)?; + opts.log_settings = value.log_settings.clone(); + Ok(opts) + } +} + +fn map_result(res: PgQueryResult) -> AnyQueryResult { + AnyQueryResult { + rows_affected: res.rows_affected(), + last_insert_id: None, + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/arguments.rs b/src-tauri/vendor/sqlx-postgres/src/arguments.rs new file mode 100644 index 00000000..62a227e5 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/arguments.rs @@ -0,0 +1,296 @@ +use std::fmt::{self, Write}; +use std::ops::{Deref, DerefMut}; +use std::sync::Arc; + +use crate::encode::{Encode, IsNull}; +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::types::Type; +use crate::{PgConnection, PgTypeInfo, Postgres}; + +use crate::type_info::PgArrayOf; +pub(crate) use sqlx_core::arguments::Arguments; +use sqlx_core::error::BoxDynError; + +// TODO: buf.patch(|| ...) is a poor name, can we think of a better name? Maybe `buf.lazy(||)` ? +// TODO: Extend the patch system to support dynamic lengths +// Considerations: +// - The prefixed-len offset needs to be back-tracked and updated +// - message::Bind needs to take a &PgArguments and use a `write` method instead of +// referencing a buffer directly +// - The basic idea is that we write bytes for the buffer until we get somewhere +// that has a patch, we then apply the patch which should write to &mut Vec, +// backtrack and update the prefixed-len, then write until the next patch offset + +#[derive(Default, Debug, Clone)] +pub struct PgArgumentBuffer { + buffer: Vec, + + // Number of arguments + count: usize, + + // Whenever an `Encode` impl needs to defer some work until after we resolve parameter types + // it can use `patch`. + // + // This currently is only setup to be useful if there is a *fixed-size* slot that needs to be + // tweaked from the input type. However, that's the only use case we currently have. + patches: Vec, + + // Whenever an `Encode` impl encounters a `PgTypeInfo` object that does not have an OID + // It pushes a "hole" that must be patched later. + // + // The hole is a `usize` offset into the buffer with the type name that should be resolved + // This is done for Records and Arrays as the OID is needed well before we are in an async + // function and can just ask postgres. + // + type_holes: Vec<(usize, HoleKind)>, // Vec<{ offset, type_name }> +} + +#[derive(Debug, Clone)] +enum HoleKind { + Type { name: UStr }, + Array(Arc), +} + +#[derive(Clone)] +struct Patch { + buf_offset: usize, + arg_index: usize, + #[allow(clippy::type_complexity)] + callback: Arc, +} + +impl fmt::Debug for Patch { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Patch") + .field("buf_offset", &self.buf_offset) + .field("arg_index", &self.arg_index) + .field("callback", &"") + .finish() + } +} + +/// Implementation of [`Arguments`] for PostgreSQL. +#[derive(Default, Debug, Clone)] +pub struct PgArguments { + // Types of each bind parameter + pub(crate) types: Vec, + + // Buffer of encoded bind parameters + pub(crate) buffer: PgArgumentBuffer, +} + +impl PgArguments { + pub(crate) fn add<'q, T>(&mut self, value: T) -> Result<(), BoxDynError> + where + T: Encode<'q, Postgres> + Type, + { + let type_info = value.produces().unwrap_or_else(T::type_info); + + let buffer_snapshot = self.buffer.snapshot(); + + // encode the value into our buffer + if let Err(error) = self.buffer.encode(value) { + // reset the value buffer to its previous value if encoding failed, + // so we don't leave a half-encoded value behind + self.buffer.reset_to_snapshot(buffer_snapshot); + return Err(error); + }; + + // remember the type information for this value + self.types.push(type_info); + // increment the number of arguments we are tracking + self.buffer.count += 1; + + Ok(()) + } + + // Apply patches + // This should only go out and ask postgres if we have not seen the type name yet + pub(crate) async fn apply_patches( + &mut self, + conn: &mut PgConnection, + parameters: &[PgTypeInfo], + ) -> Result<(), Error> { + let PgArgumentBuffer { + ref patches, + ref type_holes, + ref mut buffer, + .. + } = self.buffer; + + for patch in patches { + let buf = &mut buffer[patch.buf_offset..]; + let ty = ¶meters[patch.arg_index]; + + (patch.callback)(buf, ty); + } + + for (offset, kind) in type_holes { + let oid = match kind { + HoleKind::Type { name } => conn.fetch_type_id_by_name(name).await?, + HoleKind::Array(array) => conn.fetch_array_type_id(array).await?, + }; + buffer[*offset..(*offset + 4)].copy_from_slice(&oid.0.to_be_bytes()); + } + + Ok(()) + } +} + +impl<'q> Arguments<'q> for PgArguments { + type Database = Postgres; + + fn reserve(&mut self, additional: usize, size: usize) { + self.types.reserve(additional); + self.buffer.reserve(size); + } + + fn add(&mut self, value: T) -> Result<(), BoxDynError> + where + T: Encode<'q, Self::Database> + Type, + { + self.add(value) + } + + fn format_placeholder(&self, writer: &mut W) -> fmt::Result { + write!(writer, "${}", self.buffer.count) + } + + #[inline(always)] + fn len(&self) -> usize { + self.buffer.count + } +} + +impl PgArgumentBuffer { + pub(crate) fn encode<'q, T>(&mut self, value: T) -> Result<(), BoxDynError> + where + T: Encode<'q, Postgres>, + { + // Won't catch everything but is a good sanity check + value_size_int4_checked(value.size_hint())?; + + // reserve space to write the prefixed length of the value + let offset = self.len(); + + self.extend(&[0; 4]); + + // encode the value into our buffer + let len = if let IsNull::No = value.encode(self)? { + // Ensure that the value size does not overflow i32 + value_size_int4_checked(self.len() - offset - 4)? + } else { + // Write a -1 to indicate NULL + // NOTE: It is illegal for [encode] to write any data + debug_assert_eq!(self.len(), offset + 4); + -1_i32 + }; + + // write the len to the beginning of the value + // (offset + 4) cannot overflow because it would have failed at `self.extend()`. + self[offset..(offset + 4)].copy_from_slice(&len.to_be_bytes()); + + Ok(()) + } + + // Adds a callback to be invoked later when we know the parameter type + #[allow(dead_code)] + pub(crate) fn patch(&mut self, callback: F) + where + F: Fn(&mut [u8], &PgTypeInfo) + 'static + Send + Sync, + { + let offset = self.len(); + let arg_index = self.count; + + self.patches.push(Patch { + buf_offset: offset, + arg_index, + callback: Arc::new(callback), + }); + } + + // Extends the inner buffer by enough space to have an OID + // Remembers where the OID goes and type name for the OID + pub(crate) fn patch_type_by_name(&mut self, type_name: &UStr) { + let offset = self.len(); + + self.extend_from_slice(&0_u32.to_be_bytes()); + self.type_holes.push(( + offset, + HoleKind::Type { + name: type_name.clone(), + }, + )); + } + + pub(crate) fn patch_array_type(&mut self, array: Arc) { + let offset = self.len(); + + self.extend_from_slice(&0_u32.to_be_bytes()); + self.type_holes.push((offset, HoleKind::Array(array))); + } + + fn snapshot(&self) -> PgArgumentBufferSnapshot { + let Self { + buffer, + count, + patches, + type_holes, + } = self; + + PgArgumentBufferSnapshot { + buffer_length: buffer.len(), + count: *count, + patches_length: patches.len(), + type_holes_length: type_holes.len(), + } + } + + fn reset_to_snapshot( + &mut self, + PgArgumentBufferSnapshot { + buffer_length, + count, + patches_length, + type_holes_length, + }: PgArgumentBufferSnapshot, + ) { + self.buffer.truncate(buffer_length); + self.count = count; + self.patches.truncate(patches_length); + self.type_holes.truncate(type_holes_length); + } +} + +struct PgArgumentBufferSnapshot { + buffer_length: usize, + count: usize, + patches_length: usize, + type_holes_length: usize, +} + +impl Deref for PgArgumentBuffer { + type Target = Vec; + + #[inline] + fn deref(&self) -> &Self::Target { + &self.buffer + } +} + +impl DerefMut for PgArgumentBuffer { + #[inline] + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.buffer + } +} + +pub(crate) fn value_size_int4_checked(size: usize) -> Result { + i32::try_from(size).map_err(|_| { + format!( + "value size would overflow in the binary protocol encoding: {size} > {}", + i32::MAX + ) + }) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/column.rs b/src-tauri/vendor/sqlx-postgres/src/column.rs new file mode 100644 index 00000000..a838c27b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/column.rs @@ -0,0 +1,54 @@ +use crate::ext::ustr::UStr; +use crate::{PgTypeInfo, Postgres}; + +pub(crate) use sqlx_core::column::{Column, ColumnIndex}; + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct PgColumn { + pub(crate) ordinal: usize, + pub(crate) name: UStr, + pub(crate) type_info: PgTypeInfo, + #[cfg_attr(feature = "offline", serde(skip))] + pub(crate) relation_id: Option, + #[cfg_attr(feature = "offline", serde(skip))] + pub(crate) relation_attribute_no: Option, +} + +impl PgColumn { + /// Returns the OID of the table this column is from, if applicable. + /// + /// This will be `None` if the column is the result of an expression. + /// + /// Corresponds to column `attrelid` of the `pg_catalog.pg_attribute` table: + /// + pub fn relation_id(&self) -> Option { + self.relation_id + } + + /// Returns the 1-based index of this column in its parent table, if applicable. + /// + /// This will be `None` if the column is the result of an expression. + /// + /// Corresponds to column `attnum` of the `pg_catalog.pg_attribute` table: + /// + pub fn relation_attribute_no(&self) -> Option { + self.relation_attribute_no + } +} + +impl Column for PgColumn { + type Database = Postgres; + + fn ordinal(&self) -> usize { + self.ordinal + } + + fn name(&self) -> &str { + &self.name + } + + fn type_info(&self) -> &PgTypeInfo { + &self.type_info + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/describe.rs b/src-tauri/vendor/sqlx-postgres/src/connection/describe.rs new file mode 100644 index 00000000..a27578c5 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/describe.rs @@ -0,0 +1,681 @@ +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::io::StatementId; +use crate::message::{ParameterDescription, RowDescription}; +use crate::query_as::query_as; +use crate::query_scalar::query_scalar; +use crate::statement::PgStatementMetadata; +use crate::type_info::{PgArrayOf, PgCustomType, PgType, PgTypeKind}; +use crate::types::Json; +use crate::types::Oid; +use crate::HashMap; +use crate::{PgColumn, PgConnection, PgTypeInfo}; +use smallvec::SmallVec; +use sqlx_core::query_builder::QueryBuilder; +use std::sync::Arc; + +/// Describes the type of the `pg_type.typtype` column +/// +/// See +#[derive(Copy, Clone, Debug, Eq, PartialEq)] +enum TypType { + Base, + Composite, + Domain, + Enum, + Pseudo, + Range, +} + +impl TryFrom for TypType { + type Error = (); + + fn try_from(t: i8) -> Result { + let t = u8::try_from(t).or(Err(()))?; + + let t = match t { + b'b' => Self::Base, + b'c' => Self::Composite, + b'd' => Self::Domain, + b'e' => Self::Enum, + b'p' => Self::Pseudo, + b'r' => Self::Range, + _ => return Err(()), + }; + Ok(t) + } +} + +/// Describes the type of the `pg_type.typcategory` column +/// +/// See +#[derive(Copy, Clone, Debug, Eq, PartialEq)] +enum TypCategory { + Array, + Boolean, + Composite, + DateTime, + Enum, + Geometric, + Network, + Numeric, + Pseudo, + Range, + String, + Timespan, + User, + BitString, + Unknown, +} + +impl TryFrom for TypCategory { + type Error = (); + + fn try_from(c: i8) -> Result { + let c = u8::try_from(c).or(Err(()))?; + + let c = match c { + b'A' => Self::Array, + b'B' => Self::Boolean, + b'C' => Self::Composite, + b'D' => Self::DateTime, + b'E' => Self::Enum, + b'G' => Self::Geometric, + b'I' => Self::Network, + b'N' => Self::Numeric, + b'P' => Self::Pseudo, + b'R' => Self::Range, + b'S' => Self::String, + b'T' => Self::Timespan, + b'U' => Self::User, + b'V' => Self::BitString, + b'X' => Self::Unknown, + _ => return Err(()), + }; + Ok(c) + } +} + +impl PgConnection { + pub(super) async fn handle_row_description( + &mut self, + desc: Option, + should_fetch: bool, + ) -> Result<(Vec, HashMap), Error> { + let mut columns = Vec::new(); + let mut column_names = HashMap::new(); + + let desc = if let Some(desc) = desc { + desc + } else { + // no rows + return Ok((columns, column_names)); + }; + + columns.reserve(desc.fields.len()); + column_names.reserve(desc.fields.len()); + + for (index, field) in desc.fields.into_iter().enumerate() { + let name = UStr::from(field.name); + + let type_info = self + .maybe_fetch_type_info_by_oid(field.data_type_id, should_fetch) + .await?; + + let column = PgColumn { + ordinal: index, + name: name.clone(), + type_info, + relation_id: field.relation_id, + relation_attribute_no: field.relation_attribute_no, + }; + + columns.push(column); + column_names.insert(name, index); + } + + Ok((columns, column_names)) + } + + pub(super) async fn handle_parameter_description( + &mut self, + desc: ParameterDescription, + ) -> Result, Error> { + let mut params = Vec::with_capacity(desc.types.len()); + + for ty in desc.types { + params.push(self.maybe_fetch_type_info_by_oid(ty, true).await?); + } + + Ok(params) + } + + async fn maybe_fetch_type_info_by_oid( + &mut self, + oid: Oid, + should_fetch: bool, + ) -> Result { + // first we check if this is a built-in type + // in the average application, the vast majority of checks should flow through this + if let Some(info) = PgTypeInfo::try_from_oid(oid) { + return Ok(info); + } + + // next we check a local cache for user-defined type names <-> object id + if let Some(info) = self.inner.cache_type_info.get(&oid) { + return Ok(info.clone()); + } + + // fallback to asking the database directly for a type name + if should_fetch { + // we're boxing this future here so we can use async recursion + let info = Box::pin(async { self.fetch_type_by_oid(oid).await }).await?; + + // cache the type name <-> oid relationship in a paired hashmap + // so we don't come down this road again + self.inner.cache_type_info.insert(oid, info.clone()); + self.inner + .cache_type_oid + .insert(info.0.name().to_string().into(), oid); + + Ok(info) + } else { + // we are not in a place that *can* run a query + // this generally means we are in the middle of another query + // this _should_ only happen for complex types sent through the TEXT protocol + // we're open to ideas to correct this.. but it'd probably be more efficient to figure + // out a way to "prime" the type cache for connections rather than make this + // fallback work correctly for complex user-defined types for the TEXT protocol + Ok(PgTypeInfo(PgType::DeclareWithOid(oid))) + } + } + + async fn fetch_type_by_oid(&mut self, oid: Oid) -> Result { + let (name, typ_type, category, relation_id, element, base_type): ( + String, + i8, + i8, + Oid, + Oid, + Oid, + ) = query_as( + // Converting the OID to `regtype` and then `text` will give us the name that + // the type will need to be found at by search_path. + "SELECT oid::regtype::text, \ + typtype, \ + typcategory, \ + typrelid, \ + typelem, \ + typbasetype \ + FROM pg_catalog.pg_type \ + WHERE oid = $1", + ) + .bind(oid) + .fetch_one(&mut *self) + .await?; + + let typ_type = TypType::try_from(typ_type); + let category = TypCategory::try_from(category); + + match (typ_type, category) { + (Ok(TypType::Domain), _) => self.fetch_domain_by_oid(oid, base_type, name).await, + + (Ok(TypType::Base), Ok(TypCategory::Array)) => { + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + kind: PgTypeKind::Array( + self.maybe_fetch_type_info_by_oid(element, true).await?, + ), + name: name.into(), + oid, + })))) + } + + (Ok(TypType::Pseudo), Ok(TypCategory::Pseudo)) => { + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + kind: PgTypeKind::Pseudo, + name: name.into(), + oid, + })))) + } + + (Ok(TypType::Range), Ok(TypCategory::Range)) => { + self.fetch_range_by_oid(oid, name).await + } + + (Ok(TypType::Enum), Ok(TypCategory::Enum)) => self.fetch_enum_by_oid(oid, name).await, + + (Ok(TypType::Composite), Ok(TypCategory::Composite)) => { + self.fetch_composite_by_oid(oid, relation_id, name).await + } + + _ => Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + kind: PgTypeKind::Simple, + name: name.into(), + oid, + })))), + } + } + + async fn fetch_enum_by_oid(&mut self, oid: Oid, name: String) -> Result { + let variants: Vec = query_scalar( + r#" +SELECT enumlabel +FROM pg_catalog.pg_enum +WHERE enumtypid = $1 +ORDER BY enumsortorder + "#, + ) + .bind(oid) + .fetch_all(self) + .await?; + + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + oid, + name: name.into(), + kind: PgTypeKind::Enum(Arc::from(variants)), + })))) + } + + async fn fetch_composite_by_oid( + &mut self, + oid: Oid, + relation_id: Oid, + name: String, + ) -> Result { + let raw_fields: Vec<(String, Oid)> = query_as( + r#" +SELECT attname, atttypid +FROM pg_catalog.pg_attribute +WHERE attrelid = $1 +AND NOT attisdropped +AND attnum > 0 +ORDER BY attnum + "#, + ) + .bind(relation_id) + .fetch_all(&mut *self) + .await?; + + let mut fields = Vec::new(); + + for (field_name, field_oid) in raw_fields.into_iter() { + let field_type = self.maybe_fetch_type_info_by_oid(field_oid, true).await?; + + fields.push((field_name, field_type)); + } + + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + oid, + name: name.into(), + kind: PgTypeKind::Composite(Arc::from(fields)), + })))) + } + + async fn fetch_domain_by_oid( + &mut self, + oid: Oid, + base_type: Oid, + name: String, + ) -> Result { + let base_type = self.maybe_fetch_type_info_by_oid(base_type, true).await?; + + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + oid, + name: name.into(), + kind: PgTypeKind::Domain(base_type), + })))) + } + + async fn fetch_range_by_oid(&mut self, oid: Oid, name: String) -> Result { + let element_oid: Oid = query_scalar( + r#" +SELECT rngsubtype +FROM pg_catalog.pg_range +WHERE rngtypid = $1 + "#, + ) + .bind(oid) + .fetch_one(&mut *self) + .await?; + + let element = self.maybe_fetch_type_info_by_oid(element_oid, true).await?; + + Ok(PgTypeInfo(PgType::Custom(Arc::new(PgCustomType { + kind: PgTypeKind::Range(element), + name: name.into(), + oid, + })))) + } + + pub(crate) async fn resolve_type_id(&mut self, ty: &PgType) -> Result { + if let Some(oid) = ty.try_oid() { + return Ok(oid); + } + + match ty { + PgType::DeclareWithName(name) => self.fetch_type_id_by_name(name).await, + PgType::DeclareArrayOf(array) => self.fetch_array_type_id(array).await, + // `.try_oid()` should return `Some()` or it should be covered here + _ => unreachable!("(bug) OID should be resolvable for type {ty:?}"), + } + } + + pub(crate) async fn fetch_type_id_by_name(&mut self, name: &str) -> Result { + if let Some(oid) = self.inner.cache_type_oid.get(name) { + return Ok(*oid); + } + + // language=SQL + let (oid,): (Oid,) = query_as("SELECT $1::regtype::oid") + .bind(name) + .fetch_optional(&mut *self) + .await? + .ok_or_else(|| Error::TypeNotFound { + type_name: name.into(), + })?; + + self.inner + .cache_type_oid + .insert(name.to_string().into(), oid); + Ok(oid) + } + + pub(crate) async fn fetch_array_type_id(&mut self, array: &PgArrayOf) -> Result { + if let Some(oid) = self + .inner + .cache_type_oid + .get(&array.elem_name) + .and_then(|elem_oid| self.inner.cache_elem_type_to_array.get(elem_oid)) + { + return Ok(*oid); + } + + // language=SQL + let (elem_oid, array_oid): (Oid, Oid) = + query_as("SELECT oid, typarray FROM pg_catalog.pg_type WHERE oid = $1::regtype::oid") + .bind(&*array.elem_name) + .fetch_optional(&mut *self) + .await? + .ok_or_else(|| Error::TypeNotFound { + type_name: array.name.to_string(), + })?; + + // Avoids copying `elem_name` until necessary + self.inner + .cache_type_oid + .entry_ref(&array.elem_name) + .insert(elem_oid); + self.inner + .cache_elem_type_to_array + .insert(elem_oid, array_oid); + + Ok(array_oid) + } + + /// Check whether EXPLAIN statements are supported by the current connection + fn is_explain_available(&self) -> bool { + let parameter_statuses = &self.inner.stream.parameter_statuses; + let is_cockroachdb = parameter_statuses.contains_key("crdb_version"); + let is_materialize = parameter_statuses.contains_key("mz_version"); + let is_questdb = parameter_statuses.contains_key("questdb_version"); + !is_cockroachdb && !is_materialize && !is_questdb + } + + pub(crate) async fn get_nullable_for_columns( + &mut self, + stmt_id: StatementId, + meta: &PgStatementMetadata, + ) -> Result>, Error> { + if meta.columns.is_empty() { + return Ok(vec![]); + } + + if meta.columns.len() * 3 > 65535 { + tracing::debug!( + ?stmt_id, + num_columns = meta.columns.len(), + "number of columns in query is too large to pull nullability for" + ); + } + + // Query for NOT NULL constraints for each column in the query. + // + // This will include columns that don't have a `relation_id` (are not from a table); + // assuming those are a minority of columns, it's less code to _not_ work around it + // and just let Postgres return `NULL`. + // + // Use `UNION ALL` syntax instead of `VALUES` due to frequent lack of + // support for `VALUES` in pgwire supported databases. + let mut nullable_query = QueryBuilder::new("SELECT NOT attnotnull FROM ( "); + let mut separated = nullable_query.separated("UNION ALL "); + + let mut column_iter = meta.columns.iter().zip(0i32..); + if let Some((column, i)) = column_iter.next() { + separated.push("( SELECT "); + separated + .push_bind_unseparated(i) + .push_unseparated("::int4 AS idx, "); + separated + .push_bind_unseparated(column.relation_id) + .push_unseparated("::int4 AS table_id, "); + separated + .push_bind_unseparated(column.relation_attribute_no) + .push_unseparated("::int2 AS col_idx ) "); + } + + for (column, i) in column_iter { + separated.push("( SELECT "); + separated + .push_bind_unseparated(i) + .push_unseparated("::int4, "); + separated + .push_bind_unseparated(column.relation_id) + .push_unseparated("::int4, "); + separated + .push_bind_unseparated(column.relation_attribute_no) + .push_unseparated("::int2 ) "); + } + + nullable_query.push( + ") AS col LEFT JOIN pg_catalog.pg_attribute \ + ON table_id IS NOT NULL \ + AND attrelid = table_id \ + AND attnum = col_idx \ + ORDER BY idx", + ); + + let mut nullables: Vec> = nullable_query + .build_query_scalar() + .fetch_all(&mut *self) + .await + .map_err(|e| { + err_protocol!( + "error from nullables query: {e}; query: {:?}", + nullable_query.sql() + ) + })?; + + // If the server doesn't support EXPLAIN statements, skip this step (#1248). + if self.is_explain_available() { + // patch up our null inference with data from EXPLAIN + let nullable_patch = self + .nullables_from_explain(stmt_id, meta.parameters.len()) + .await?; + + for (nullable, patch) in nullables.iter_mut().zip(nullable_patch) { + *nullable = patch.or(*nullable); + } + } + + Ok(nullables) + } + + /// Infer nullability for columns of this statement using EXPLAIN VERBOSE. + /// + /// This currently only marks columns that are on the inner half of an outer join + /// and returns `None` for all others. + async fn nullables_from_explain( + &mut self, + stmt_id: StatementId, + params_len: usize, + ) -> Result>, Error> { + let stmt_id_display = stmt_id + .display() + .ok_or_else(|| err_protocol!("cannot EXPLAIN unnamed statement: {stmt_id:?}"))?; + + let mut explain = format!("EXPLAIN (VERBOSE, FORMAT JSON) EXECUTE {stmt_id_display}"); + let mut comma = false; + + if params_len > 0 { + explain += "("; + + // fill the arguments list with NULL, which should theoretically be valid + for _ in 0..params_len { + if comma { + explain += ", "; + } + + explain += "NULL"; + comma = true; + } + + explain += ")"; + } + + let (Json(explains),): (Json>,) = + query_as(&explain).fetch_one(self).await?; + + let mut nullables = Vec::new(); + + if let Some(Explain::Plan { + plan: + plan @ Plan { + output: Some(ref outputs), + .. + }, + }) = explains.first() + { + nullables.resize(outputs.len(), None); + visit_plan(plan, outputs, &mut nullables); + } + + Ok(nullables) + } +} + +fn visit_plan(plan: &Plan, outputs: &[String], nullables: &mut Vec>) { + if let Some(plan_outputs) = &plan.output { + // all outputs of a Full Join must be marked nullable + // otherwise, all outputs of the inner half of an outer join must be marked nullable + if plan.join_type.as_deref() == Some("Full") + || plan.parent_relation.as_deref() == Some("Inner") + { + for output in plan_outputs { + if let Some(i) = outputs.iter().position(|o| o == output) { + // N.B. this may produce false positives but those don't cause runtime errors + nullables[i] = Some(true); + } + } + } + } + + if let Some(plans) = &plan.plans { + if let Some("Left") | Some("Right") = plan.join_type.as_deref() { + for plan in plans { + visit_plan(plan, outputs, nullables); + } + } + } +} + +#[derive(serde::Deserialize, Debug)] +#[serde(untagged)] +enum Explain { + // NOTE: the returned JSON may not contain a `plan` field, for example, with `CALL` statements: + // https://github.com/launchbadge/sqlx/issues/1449 + // + // In this case, we should just fall back to assuming all is nullable. + // + // It may also contain additional fields we don't care about, which should not break parsing: + // https://github.com/launchbadge/sqlx/issues/2587 + // https://github.com/launchbadge/sqlx/issues/2622 + Plan { + #[serde(rename = "Plan")] + plan: Plan, + }, + + // This ensures that parsing never technically fails. + // + // We don't want to specifically expect `"Utility Statement"` because there might be other cases + // and we don't care unless it contains a query plan anyway. + Other(serde::de::IgnoredAny), +} + +#[derive(serde::Deserialize, Debug)] +struct Plan { + #[serde(rename = "Join Type")] + join_type: Option, + #[serde(rename = "Parent Relationship")] + parent_relation: Option, + #[serde(rename = "Output")] + output: Option>, + #[serde(rename = "Plans")] + plans: Option>, +} + +#[test] +fn explain_parsing() { + let normal_plan = r#"[ + { + "Plan": { + "Node Type": "Result", + "Parallel Aware": false, + "Async Capable": false, + "Startup Cost": 0.00, + "Total Cost": 0.01, + "Plan Rows": 1, + "Plan Width": 4, + "Output": ["1"] + } + } +]"#; + + // https://github.com/launchbadge/sqlx/issues/2622 + let extra_field = r#"[ + { + "Plan": { + "Node Type": "Result", + "Parallel Aware": false, + "Async Capable": false, + "Startup Cost": 0.00, + "Total Cost": 0.01, + "Plan Rows": 1, + "Plan Width": 4, + "Output": ["1"] + }, + "Query Identifier": 1147616880456321454 + } +]"#; + + // https://github.com/launchbadge/sqlx/issues/1449 + let utility_statement = r#"["Utility Statement"]"#; + + let normal_plan_parsed = serde_json::from_str::<[Explain; 1]>(normal_plan).unwrap(); + let extra_field_parsed = serde_json::from_str::<[Explain; 1]>(extra_field).unwrap(); + let utility_statement_parsed = serde_json::from_str::<[Explain; 1]>(utility_statement).unwrap(); + + assert!( + matches!(normal_plan_parsed, [Explain::Plan { plan: Plan { .. } }]), + "unexpected parse from {normal_plan:?}: {normal_plan_parsed:?}" + ); + + assert!( + matches!(extra_field_parsed, [Explain::Plan { plan: Plan { .. } }]), + "unexpected parse from {extra_field:?}: {extra_field_parsed:?}" + ); + + assert!( + matches!(utility_statement_parsed, [Explain::Other(_)]), + "unexpected parse from {utility_statement:?}: {utility_statement_parsed:?}" + ) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/establish.rs b/src-tauri/vendor/sqlx-postgres/src/connection/establish.rs new file mode 100644 index 00000000..1bc4172f --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/establish.rs @@ -0,0 +1,155 @@ +use crate::HashMap; + +use crate::common::StatementCache; +use crate::connection::{sasl, stream::PgStream}; +use crate::error::Error; +use crate::io::StatementId; +use crate::message::{ + Authentication, BackendKeyData, BackendMessageFormat, Password, ReadyForQuery, Startup, +}; +use crate::{PgConnectOptions, PgConnection}; + +use super::PgConnectionInner; + +// https://www.postgresql.org/docs/current/protocol-flow.html#id-1.10.5.7.3 +// https://www.postgresql.org/docs/current/protocol-flow.html#id-1.10.5.7.11 + +impl PgConnection { + pub(crate) async fn establish(options: &PgConnectOptions) -> Result { + // Upgrade to TLS if we were asked to and the server supports it + let mut stream = PgStream::connect(options).await?; + + // To begin a session, a frontend opens a connection to the server + // and sends a startup message. + + let mut params = vec![ + // Sets the display format for date and time values, + // as well as the rules for interpreting ambiguous date input values. + ("DateStyle", "ISO, MDY"), + // Sets the client-side encoding (character set). + // + ("client_encoding", "UTF8"), + // Sets the time zone for displaying and interpreting time stamps. + ("TimeZone", "UTC"), + ]; + + if let Some(ref extra_float_digits) = options.extra_float_digits { + params.push(("extra_float_digits", extra_float_digits)); + } + + if let Some(ref application_name) = options.application_name { + params.push(("application_name", application_name)); + } + + if let Some(ref options) = options.options { + params.push(("options", options)); + } + + stream.write(Startup { + username: Some(&options.username), + database: options.database.as_deref(), + params: ¶ms, + })?; + + stream.flush().await?; + + // The server then uses this information and the contents of + // its configuration files (such as pg_hba.conf) to determine whether the connection is + // provisionally acceptable, and what additional + // authentication is required (if any). + + let mut process_id = 0; + let mut secret_key = 0; + let transaction_status; + + loop { + let message = stream.recv().await?; + match message.format { + BackendMessageFormat::Authentication => match message.decode()? { + Authentication::Ok => { + // the authentication exchange is successfully completed + // do nothing; no more information is required to continue + } + + Authentication::CleartextPassword => { + // The frontend must now send a [PasswordMessage] containing the + // password in clear-text form. + + stream + .send(Password::Cleartext( + options.password.as_deref().unwrap_or_default(), + )) + .await?; + } + + Authentication::Md5Password(body) => { + // The frontend must now send a [PasswordMessage] containing the + // password (with user name) encrypted via MD5, then encrypted again + // using the 4-byte random salt specified in the + // [AuthenticationMD5Password] message. + + stream + .send(Password::Md5 { + username: &options.username, + password: options.password.as_deref().unwrap_or_default(), + salt: body.salt, + }) + .await?; + } + + Authentication::Sasl(body) => { + sasl::authenticate(&mut stream, options, body).await?; + } + + method => { + return Err(err_protocol!( + "unsupported authentication method: {:?}", + method + )); + } + }, + + BackendMessageFormat::BackendKeyData => { + // provides secret-key data that the frontend must save if it wants to be + // able to issue cancel requests later + + let data: BackendKeyData = message.decode()?; + + process_id = data.process_id; + secret_key = data.secret_key; + } + + BackendMessageFormat::ReadyForQuery => { + // start-up is completed. The frontend can now issue commands + transaction_status = message.decode::()?.transaction_status; + + break; + } + + _ => { + return Err(err_protocol!( + "establish: unexpected message: {:?}", + message.format + )) + } + } + } + + Ok(PgConnection { + inner: Box::new(PgConnectionInner { + stream, + process_id, + secret_key, + transaction_status, + transaction_depth: 0, + pending_ready_for_query_count: 0, + next_statement_id: StatementId::NAMED_START, + cache_statement: StatementCache::new(options.statement_cache_capacity), + cache_type_oid: HashMap::new(), + cache_type_info: HashMap::new(), + cache_elem_type_to_array: HashMap::new(), + log_settings: options.log_settings.clone(), + }), + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/executor.rs b/src-tauri/vendor/sqlx-postgres/src/connection/executor.rs new file mode 100644 index 00000000..3fe4f402 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/executor.rs @@ -0,0 +1,486 @@ +use crate::describe::Describe; +use crate::error::Error; +use crate::executor::{Execute, Executor}; +use crate::io::{PortalId, StatementId}; +use crate::logger::QueryLogger; +use crate::message::{ + self, BackendMessageFormat, Bind, Close, CommandComplete, DataRow, ParameterDescription, Parse, + ParseComplete, Query, RowDescription, +}; +use crate::statement::PgStatementMetadata; +use crate::{ + statement::PgStatement, PgArguments, PgConnection, PgQueryResult, PgRow, PgTypeInfo, + PgValueFormat, Postgres, +}; +use futures_core::future::BoxFuture; +use futures_core::stream::BoxStream; +use futures_core::Stream; +use futures_util::TryStreamExt; +use sqlx_core::arguments::Arguments; +use sqlx_core::Either; +use std::{borrow::Cow, pin::pin, sync::Arc}; + +async fn prepare( + conn: &mut PgConnection, + sql: &str, + parameters: &[PgTypeInfo], + metadata: Option>, + persistent: bool, +) -> Result<(StatementId, Arc), Error> { + let id = if persistent { + let id = conn.inner.next_statement_id; + conn.inner.next_statement_id = id.next(); + id + } else { + StatementId::UNNAMED + }; + + // build a list of type OIDs to send to the database in the PARSE command + // we have not yet started the query sequence, so we are *safe* to cleanly make + // additional queries here to get any missing OIDs + + let mut param_types = Vec::with_capacity(parameters.len()); + + for ty in parameters { + param_types.push(conn.resolve_type_id(&ty.0).await?); + } + + // flush and wait until we are re-ready + conn.wait_until_ready().await?; + + // next we send the PARSE command to the server + conn.inner.stream.write_msg(Parse { + param_types: ¶m_types, + query: sql, + statement: id, + })?; + + if metadata.is_none() { + // get the statement columns and parameters + conn.inner + .stream + .write_msg(message::Describe::Statement(id))?; + } + + // we ask for the server to immediately send us the result of the PARSE command + conn.write_sync(); + conn.inner.stream.flush().await?; + + // indicates that the SQL query string is now successfully parsed and has semantic validity + conn.inner.stream.recv_expect::().await?; + + let metadata = if let Some(metadata) = metadata { + // each SYNC produces one READY FOR QUERY + conn.recv_ready_for_query().await?; + + // we already have metadata + metadata + } else { + let parameters = recv_desc_params(conn).await?; + + let rows = recv_desc_rows(conn).await?; + + // each SYNC produces one READY FOR QUERY + conn.recv_ready_for_query().await?; + + let parameters = conn.handle_parameter_description(parameters).await?; + + let (columns, column_names) = conn.handle_row_description(rows, true).await?; + + // ensure that if we did fetch custom data, we wait until we are fully ready before + // continuing + conn.wait_until_ready().await?; + + Arc::new(PgStatementMetadata { + parameters, + columns, + column_names: Arc::new(column_names), + }) + }; + + Ok((id, metadata)) +} + +async fn recv_desc_params(conn: &mut PgConnection) -> Result { + conn.inner.stream.recv_expect().await +} + +async fn recv_desc_rows(conn: &mut PgConnection) -> Result, Error> { + let rows: Option = match conn.inner.stream.recv().await? { + // describes the rows that will be returned when the statement is eventually executed + message if message.format == BackendMessageFormat::RowDescription => { + Some(message.decode()?) + } + + // no data would be returned if this statement was executed + message if message.format == BackendMessageFormat::NoData => None, + + message => { + return Err(err_protocol!( + "expecting RowDescription or NoData but received {:?}", + message.format + )); + } + }; + + Ok(rows) +} + +impl PgConnection { + // wait for CloseComplete to indicate a statement was closed + pub(super) async fn wait_for_close_complete(&mut self, mut count: usize) -> Result<(), Error> { + // we need to wait for the [CloseComplete] to be returned from the server + while count > 0 { + match self.inner.stream.recv().await? { + message if message.format == BackendMessageFormat::PortalSuspended => { + // there was an open portal + // this can happen if the last time a statement was used it was not fully executed + } + + message if message.format == BackendMessageFormat::CloseComplete => { + // successfully closed the statement (and freed up the server resources) + count -= 1; + } + + message => { + return Err(err_protocol!( + "expecting PortalSuspended or CloseComplete but received {:?}", + message.format + )); + } + } + } + + Ok(()) + } + + #[inline(always)] + pub(crate) fn write_sync(&mut self) { + self.inner + .stream + .write_msg(message::Sync) + .expect("BUG: Sync should not be too big for protocol"); + + // all SYNC messages will return a ReadyForQuery + self.inner.pending_ready_for_query_count += 1; + } + + async fn get_or_prepare<'a>( + &mut self, + sql: &str, + parameters: &[PgTypeInfo], + persistent: bool, + // optional metadata that was provided by the user, this means they are reusing + // a statement object + metadata: Option>, + ) -> Result<(StatementId, Arc), Error> { + if let Some(statement) = self.inner.cache_statement.get_mut(sql) { + return Ok((*statement).clone()); + } + + let statement = prepare(self, sql, parameters, metadata, persistent).await?; + + if persistent && self.inner.cache_statement.is_enabled() { + if let Some((id, _)) = self.inner.cache_statement.insert(sql, statement.clone()) { + self.inner.stream.write_msg(Close::Statement(id))?; + self.write_sync(); + + self.inner.stream.flush().await?; + + self.wait_for_close_complete(1).await?; + self.recv_ready_for_query().await?; + } + } + + Ok(statement) + } + + pub(crate) async fn run<'e, 'c: 'e, 'q: 'e>( + &'c mut self, + query: &'q str, + arguments: Option, + persistent: bool, + metadata_opt: Option>, + ) -> Result, Error>> + 'e, Error> { + let mut logger = QueryLogger::new(query, self.inner.log_settings.clone()); + + // before we continue, wait until we are "ready" to accept more queries + self.wait_until_ready().await?; + + let mut metadata: Arc; + + let format = if let Some(mut arguments) = arguments { + // Check this before we write anything to the stream. + // + // Note: Postgres actually interprets this value as unsigned, + // making the max number of parameters 65535, not 32767 + // https://github.com/launchbadge/sqlx/issues/3464 + // https://www.postgresql.org/docs/current/limits.html + let num_params = u16::try_from(arguments.len()).map_err(|_| { + err_protocol!( + "PgConnection::run(): too many arguments for query: {}", + arguments.len() + ) + })?; + + // prepare the statement if this our first time executing it + // always return the statement ID here + let (statement, metadata_) = self + .get_or_prepare(query, &arguments.types, persistent, metadata_opt) + .await?; + + metadata = metadata_; + + // patch holes created during encoding + arguments.apply_patches(self, &metadata.parameters).await?; + + // consume messages till `ReadyForQuery` before bind and execute + self.wait_until_ready().await?; + + // bind to attach the arguments to the statement and create a portal + self.inner.stream.write_msg(Bind { + portal: PortalId::UNNAMED, + statement, + formats: &[PgValueFormat::Binary], + num_params, + params: &arguments.buffer, + result_formats: &[PgValueFormat::Binary], + })?; + + // executes the portal up to the passed limit + // the protocol-level limit acts nearly identically to the `LIMIT` in SQL + self.inner.stream.write_msg(message::Execute { + portal: PortalId::UNNAMED, + // Non-zero limits cause query plan pessimization by disabling parallel workers: + // https://github.com/launchbadge/sqlx/issues/3673 + limit: 0, + })?; + // From https://www.postgresql.org/docs/current/protocol-flow.html: + // + // "An unnamed portal is destroyed at the end of the transaction, or as + // soon as the next Bind statement specifying the unnamed portal as + // destination is issued. (Note that a simple Query message also + // destroys the unnamed portal." + + // we ask the database server to close the unnamed portal and free the associated resources + // earlier - after the execution of the current query. + self.inner + .stream + .write_msg(Close::Portal(PortalId::UNNAMED))?; + + // finally, [Sync] asks postgres to process the messages that we sent and respond with + // a [ReadyForQuery] message when it's completely done. Theoretically, we could send + // dozens of queries before a [Sync] and postgres can handle that. Execution on the server + // is still serial but it would reduce round-trips. Some kind of builder pattern that is + // termed batching might suit this. + self.write_sync(); + + // prepared statements are binary + PgValueFormat::Binary + } else { + // Query will trigger a ReadyForQuery + self.inner.stream.write_msg(Query(query))?; + self.inner.pending_ready_for_query_count += 1; + + // metadata starts out as "nothing" + metadata = Arc::new(PgStatementMetadata::default()); + + // and unprepared statements are text + PgValueFormat::Text + }; + + self.inner.stream.flush().await?; + + Ok(try_stream! { + loop { + let message = self.inner.stream.recv().await?; + + match message.format { + BackendMessageFormat::BindComplete + | BackendMessageFormat::ParseComplete + | BackendMessageFormat::ParameterDescription + | BackendMessageFormat::NoData + // unnamed portal has been closed + | BackendMessageFormat::CloseComplete + => { + // harmless messages to ignore + } + + // "Execute phase is always terminated by the appearance of + // exactly one of these messages: CommandComplete, + // EmptyQueryResponse (if the portal was created from an + // empty query string), ErrorResponse, or PortalSuspended" + BackendMessageFormat::CommandComplete => { + // a SQL command completed normally + let cc: CommandComplete = message.decode()?; + + let rows_affected = cc.rows_affected(); + logger.increase_rows_affected(rows_affected); + r#yield!(Either::Left(PgQueryResult { + rows_affected, + })); + } + + BackendMessageFormat::EmptyQueryResponse => { + // empty query string passed to an unprepared execute + } + + // Message::ErrorResponse is handled in self.stream.recv() + + // incomplete query execution has finished + BackendMessageFormat::PortalSuspended => {} + + BackendMessageFormat::RowDescription => { + // indicates that a *new* set of rows are about to be returned + let (columns, column_names) = self + .handle_row_description(Some(message.decode()?), false) + .await?; + + metadata = Arc::new(PgStatementMetadata { + column_names: Arc::new(column_names), + columns, + parameters: Vec::default(), + }); + } + + BackendMessageFormat::DataRow => { + logger.increment_rows_returned(); + + // one of the set of rows returned by a SELECT, FETCH, etc query + let data: DataRow = message.decode()?; + let row = PgRow { + data, + format, + metadata: Arc::clone(&metadata), + }; + + r#yield!(Either::Right(row)); + } + + BackendMessageFormat::ReadyForQuery => { + // processing of the query string is complete + self.handle_ready_for_query(message)?; + break; + } + + _ => { + return Err(err_protocol!( + "execute: unexpected message: {:?}", + message.format + )); + } + } + } + + Ok(()) + }) + } +} + +impl<'c> Executor<'c> for &'c mut PgConnection { + type Database = Postgres; + + fn fetch_many<'e, 'q, E>( + self, + mut query: E, + ) -> BoxStream<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + let sql = query.sql(); + // False positive: https://github.com/rust-lang/rust-clippy/issues/12560 + #[allow(clippy::map_clone)] + let metadata = query.statement().map(|s| Arc::clone(&s.metadata)); + let arguments = query.take_arguments().map_err(Error::Encode); + let persistent = query.persistent(); + + Box::pin(try_stream! { + let arguments = arguments?; + let mut s = pin!(self.run(sql, arguments, persistent, metadata).await?); + + while let Some(v) = s.try_next().await? { + r#yield!(v); + } + + Ok(()) + }) + } + + fn fetch_optional<'e, 'q, E>(self, mut query: E) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + let sql = query.sql(); + // False positive: https://github.com/rust-lang/rust-clippy/issues/12560 + #[allow(clippy::map_clone)] + let metadata = query.statement().map(|s| Arc::clone(&s.metadata)); + let arguments = query.take_arguments().map_err(Error::Encode); + let persistent = query.persistent(); + + Box::pin(async move { + let arguments = arguments?; + let mut s = pin!(self.run(sql, arguments, persistent, metadata).await?); + + // With deferred constraints we need to check all responses as we + // could get a OK response (with uncommitted data), only to get an + // error response after (when the deferred constraint is actually + // checked). + let mut ret = None; + while let Some(result) = s.try_next().await? { + match result { + Either::Right(r) if ret.is_none() => ret = Some(r), + _ => {} + } + } + Ok(ret) + }) + } + + fn prepare_with<'e, 'q: 'e>( + self, + sql: &'q str, + parameters: &'e [PgTypeInfo], + ) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + Box::pin(async move { + self.wait_until_ready().await?; + + let (_, metadata) = self.get_or_prepare(sql, parameters, true, None).await?; + + Ok(PgStatement { + sql: Cow::Borrowed(sql), + metadata, + }) + }) + } + + fn describe<'e, 'q: 'e>( + self, + sql: &'q str, + ) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + Box::pin(async move { + self.wait_until_ready().await?; + + let (stmt_id, metadata) = self.get_or_prepare(sql, &[], true, None).await?; + + let nullable = self.get_nullable_for_columns(stmt_id, &metadata).await?; + + Ok(Describe { + columns: metadata.columns.clone(), + nullable, + parameters: Some(Either::Left(metadata.parameters.clone())), + }) + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/mod.rs b/src-tauri/vendor/sqlx-postgres/src/connection/mod.rs new file mode 100644 index 00000000..0005f8e5 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/mod.rs @@ -0,0 +1,269 @@ +use std::borrow::Cow; +use std::fmt::{self, Debug, Formatter}; +use std::sync::Arc; + +use crate::HashMap; +use futures_core::future::BoxFuture; +use futures_util::FutureExt; + +use crate::common::StatementCache; +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::io::StatementId; +use crate::message::{ + BackendMessageFormat, Close, Query, ReadyForQuery, ReceivedMessage, Terminate, + TransactionStatus, +}; +use crate::statement::PgStatementMetadata; +use crate::transaction::Transaction; +use crate::types::Oid; +use crate::{PgConnectOptions, PgTypeInfo, Postgres}; + +pub(crate) use sqlx_core::connection::*; + +pub use self::stream::PgStream; + +pub(crate) mod describe; +mod establish; +mod executor; +mod sasl; +mod stream; +mod tls; + +/// A connection to a PostgreSQL database. +/// +/// See [`PgConnectOptions`] for connection URL reference. +pub struct PgConnection { + pub(crate) inner: Box, +} + +pub struct PgConnectionInner { + // underlying TCP or UDS stream, + // wrapped in a potentially TLS stream, + // wrapped in a buffered stream + pub(crate) stream: PgStream, + + // process id of this backend + // used to send cancel requests + #[allow(dead_code)] + process_id: u32, + + // secret key of this backend + // used to send cancel requests + #[allow(dead_code)] + secret_key: u32, + + // sequence of statement IDs for use in preparing statements + // in PostgreSQL, the statement is prepared to a user-supplied identifier + next_statement_id: StatementId, + + // cache statement by query string to the id and columns + cache_statement: StatementCache<(StatementId, Arc)>, + + // cache user-defined types by id <-> info + cache_type_info: HashMap, + cache_type_oid: HashMap, + cache_elem_type_to_array: HashMap, + + // number of ReadyForQuery messages that we are currently expecting + pub(crate) pending_ready_for_query_count: usize, + + // current transaction status + transaction_status: TransactionStatus, + pub(crate) transaction_depth: usize, + + log_settings: LogSettings, +} + +impl PgConnection { + /// the version number of the server in `libpq` format + pub fn server_version_num(&self) -> Option { + self.inner.stream.server_version_num + } + + // will return when the connection is ready for another query + pub(crate) async fn wait_until_ready(&mut self) -> Result<(), Error> { + if !self.inner.stream.write_buffer_mut().is_empty() { + self.inner.stream.flush().await?; + } + + while self.inner.pending_ready_for_query_count > 0 { + let message = self.inner.stream.recv().await?; + + if let BackendMessageFormat::ReadyForQuery = message.format { + self.handle_ready_for_query(message)?; + } + } + + Ok(()) + } + + async fn recv_ready_for_query(&mut self) -> Result<(), Error> { + let r: ReadyForQuery = self.inner.stream.recv_expect().await?; + + self.inner.pending_ready_for_query_count -= 1; + self.inner.transaction_status = r.transaction_status; + + Ok(()) + } + + #[inline(always)] + fn handle_ready_for_query(&mut self, message: ReceivedMessage) -> Result<(), Error> { + self.inner.pending_ready_for_query_count = self + .inner + .pending_ready_for_query_count + .checked_sub(1) + .ok_or_else(|| err_protocol!("received more ReadyForQuery messages than expected"))?; + + self.inner.transaction_status = message.decode::()?.transaction_status; + + Ok(()) + } + + /// Queue a simple query (not prepared) to execute the next time this connection is used. + /// + /// Used for rolling back transactions and releasing advisory locks. + #[inline(always)] + pub(crate) fn queue_simple_query(&mut self, query: &str) -> Result<(), Error> { + self.inner.stream.write_msg(Query(query))?; + self.inner.pending_ready_for_query_count += 1; + + Ok(()) + } + + pub(crate) fn in_transaction(&self) -> bool { + match self.inner.transaction_status { + TransactionStatus::Transaction => true, + TransactionStatus::Error | TransactionStatus::Idle => false, + } + } +} + +impl Debug for PgConnection { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.debug_struct("PgConnection").finish() + } +} + +impl Connection for PgConnection { + type Database = Postgres; + + type Options = PgConnectOptions; + + fn close(mut self) -> BoxFuture<'static, Result<(), Error>> { + // The normal, graceful termination procedure is that the frontend sends a Terminate + // message and immediately closes the connection. + + // On receipt of this message, the backend closes the + // connection and terminates. + + Box::pin(async move { + self.inner.stream.send(Terminate).await?; + self.inner.stream.shutdown().await?; + + Ok(()) + }) + } + + fn close_hard(mut self) -> BoxFuture<'static, Result<(), Error>> { + Box::pin(async move { + self.inner.stream.shutdown().await?; + + Ok(()) + }) + } + + fn ping(&mut self) -> BoxFuture<'_, Result<(), Error>> { + // Users were complaining about this showing up in query statistics on the server. + // By sending a comment we avoid an error if the connection was in the middle of a rowset + // self.execute("/* SQLx ping */").map_ok(|_| ()).boxed() + + Box::pin(async move { + // Stroke patch: the pool pings every connection it takes back, one + // full round trip before the connection can serve the next query + // (265-535ms to Neon, Supabase, Nile, Prisma), and Nile's proxy + // answers a bare Sync with "Internal error", so every query there + // threw its connection away. With nothing queued and nothing + // unread there is nothing to flush, so a clean connection skips it. + // A dropped stream or a queued rollback still takes the full path. + if self.inner.pending_ready_for_query_count == 0 + && self.inner.stream.write_buffer_mut().is_empty() + { + return Ok(()); + } + // The simplest call-and-response that's possible. + self.write_sync(); + self.wait_until_ready().await + }) + } + + fn begin(&mut self) -> BoxFuture<'_, Result, Error>> + where + Self: Sized, + { + Transaction::begin(self, None) + } + + fn begin_with( + &mut self, + statement: impl Into>, + ) -> BoxFuture<'_, Result, Error>> + where + Self: Sized, + { + Transaction::begin(self, Some(statement.into())) + } + + fn cached_statements_size(&self) -> usize { + self.inner.cache_statement.len() + } + + fn clear_cached_statements(&mut self) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + self.inner.cache_type_oid.clear(); + + let mut cleared = 0_usize; + + self.wait_until_ready().await?; + + while let Some((id, _)) = self.inner.cache_statement.remove_lru() { + self.inner.stream.write_msg(Close::Statement(id))?; + cleared += 1; + } + + if cleared > 0 { + self.write_sync(); + self.inner.stream.flush().await?; + + self.wait_for_close_complete(cleared).await?; + self.recv_ready_for_query().await?; + } + + Ok(()) + }) + } + + fn shrink_buffers(&mut self) { + self.inner.stream.shrink_buffers(); + } + + #[doc(hidden)] + fn flush(&mut self) -> BoxFuture<'_, Result<(), Error>> { + self.wait_until_ready().boxed() + } + + #[doc(hidden)] + fn should_flush(&self) -> bool { + !self.inner.stream.write_buffer().is_empty() + } +} + +// Implement `AsMut` so that `PgConnection` can be wrapped in +// a `PgAdvisoryLockGuard`. +// +// See: https://github.com/launchbadge/sqlx/issues/2520 +impl AsMut for PgConnection { + fn as_mut(&mut self) -> &mut PgConnection { + self + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/sasl.rs b/src-tauri/vendor/sqlx-postgres/src/connection/sasl.rs new file mode 100644 index 00000000..729cc1fc --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/sasl.rs @@ -0,0 +1,225 @@ +use crate::connection::stream::PgStream; +use crate::error::Error; +use crate::message::{Authentication, AuthenticationSasl, SaslInitialResponse, SaslResponse}; +use crate::PgConnectOptions; +use hmac::{Hmac, Mac}; +use rand::Rng; +use sha2::{Digest, Sha256}; +use stringprep::saslprep; + +use base64::prelude::{Engine as _, BASE64_STANDARD}; + +const GS2_HEADER: &str = "n,,"; +const CHANNEL_ATTR: &str = "c"; +const USERNAME_ATTR: &str = "n"; +const CLIENT_PROOF_ATTR: &str = "p"; +const NONCE_ATTR: &str = "r"; + +pub(crate) async fn authenticate( + stream: &mut PgStream, + options: &PgConnectOptions, + data: AuthenticationSasl, +) -> Result<(), Error> { + let mut has_sasl = false; + let mut has_sasl_plus = false; + let mut unknown = Vec::new(); + + for mechanism in data.mechanisms() { + match mechanism { + "SCRAM-SHA-256" => { + has_sasl = true; + } + + "SCRAM-SHA-256-PLUS" => { + has_sasl_plus = true; + } + + _ => { + unknown.push(mechanism.to_owned()); + } + } + } + + if !has_sasl_plus && !has_sasl { + return Err(err_protocol!( + "unsupported SASL authentication mechanisms: {}", + unknown.join(", ") + )); + } + + // channel-binding = "c=" base64 + let mut channel_binding = format!("{CHANNEL_ATTR}="); + BASE64_STANDARD.encode_string(GS2_HEADER, &mut channel_binding); + + // "n=" saslname ;; Usernames are prepared using SASLprep. + let username = format!("{}={}", USERNAME_ATTR, options.username); + let username = match saslprep(&username) { + Ok(v) => v, + // TODO(danielakhterov): Remove panic when we have proper support for configuration errors + Err(_) => panic!("Failed to saslprep username"), + }; + + // nonce = "r=" c-nonce [s-nonce] ;; Second part provided by server. + let nonce = gen_nonce(); + + // client-first-message-bare = [reserved-mext ","] username "," nonce ["," extensions] + let client_first_message_bare = format!("{username},{nonce}"); + + let client_first_message = format!("{GS2_HEADER}{client_first_message_bare}"); + + stream + .send(SaslInitialResponse { + response: &client_first_message, + plus: false, + }) + .await?; + + let cont = match stream.recv_expect().await? { + Authentication::SaslContinue(data) => data, + + auth => { + return Err(err_protocol!( + "expected SASLContinue but received {:?}", + auth + )); + } + }; + + // SaltedPassword := Hi(Normalize(password), salt, i) + let salted_password = hi( + options.password.as_deref().unwrap_or_default(), + &cont.salt, + cont.iterations, + )?; + + // ClientKey := HMAC(SaltedPassword, "Client Key") + let mut mac = Hmac::::new_from_slice(&salted_password).map_err(Error::protocol)?; + mac.update(b"Client Key"); + + let client_key = mac.finalize().into_bytes(); + + // StoredKey := H(ClientKey) + let stored_key = Sha256::digest(client_key); + + // client-final-message-without-proof + let client_final_message_wo_proof = format!( + "{channel_binding},r={nonce}", + channel_binding = channel_binding, + nonce = &cont.nonce + ); + + // AuthMessage := client-first-message-bare + "," + server-first-message + "," + client-final-message-without-proof + let auth_message = format!( + "{client_first_message_bare},{server_first_message},{client_final_message_wo_proof}", + client_first_message_bare = client_first_message_bare, + server_first_message = cont.message, + client_final_message_wo_proof = client_final_message_wo_proof + ); + + // ClientSignature := HMAC(StoredKey, AuthMessage) + let mut mac = Hmac::::new_from_slice(&stored_key).map_err(Error::protocol)?; + mac.update(auth_message.as_bytes()); + + let client_signature = mac.finalize().into_bytes(); + + // ClientProof := ClientKey XOR ClientSignature + let client_proof: Vec = client_key + .iter() + .zip(client_signature.iter()) + .map(|(&a, &b)| a ^ b) + .collect(); + + // ServerKey := HMAC(SaltedPassword, "Server Key") + let mut mac = Hmac::::new_from_slice(&salted_password).map_err(Error::protocol)?; + mac.update(b"Server Key"); + + let server_key = mac.finalize().into_bytes(); + + // ServerSignature := HMAC(ServerKey, AuthMessage) + let mut mac = Hmac::::new_from_slice(&server_key).map_err(Error::protocol)?; + mac.update(auth_message.as_bytes()); + + // client-final-message = client-final-message-without-proof "," proof + let mut client_final_message = format!("{client_final_message_wo_proof},{CLIENT_PROOF_ATTR}="); + BASE64_STANDARD.encode_string(client_proof, &mut client_final_message); + + stream.send(SaslResponse(&client_final_message)).await?; + + let data = match stream.recv_expect().await? { + Authentication::SaslFinal(data) => data, + + auth => { + return Err(err_protocol!("expected SASLFinal but received {:?}", auth)); + } + }; + + // authentication is only considered valid if this verification passes + mac.verify_slice(&data.verifier).map_err(Error::protocol)?; + + Ok(()) +} + +// nonce is a sequence of random printable bytes +fn gen_nonce() -> String { + let mut rng = rand::thread_rng(); + let count = rng.gen_range(64..128); + + // printable = %x21-2B / %x2D-7E + // ;; Printable ASCII except ",". + // ;; Note that any "printable" is also + // ;; a valid "value". + let nonce: String = std::iter::repeat(()) + .map(|()| { + let mut c = rng.gen_range(0x21u8..0x7F); + + while c == 0x2C { + c = rng.gen_range(0x21u8..0x7F); + } + + c + }) + .take(count) + .map(|c| c as char) + .collect(); + + rng.gen_range(32..128); + format!("{NONCE_ATTR}={nonce}") +} + +// Hi(str, salt, i): +fn hi<'a>(s: &'a str, salt: &'a [u8], iter_count: u32) -> Result<[u8; 32], Error> { + let mut mac = Hmac::::new_from_slice(s.as_bytes()).map_err(Error::protocol)?; + + mac.update(salt); + mac.update(&1u32.to_be_bytes()); + + let mut u = mac.finalize_reset().into_bytes(); + let mut hi = u; + + for _ in 1..iter_count { + mac.update(u.as_slice()); + u = mac.finalize_reset().into_bytes(); + hi = hi.iter().zip(u.iter()).map(|(&a, &b)| a ^ b).collect(); + } + + Ok(hi.into()) +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_sasl_hi(b: &mut test::Bencher) { + use test::black_box; + + let mut rng = rand::thread_rng(); + let nonce: Vec = std::iter::repeat(()) + .map(|()| rng.sample(rand::distributions::Alphanumeric)) + .take(64) + .collect(); + b.iter(|| { + let _ = hi( + test::black_box("secret_password"), + test::black_box(&nonce), + test::black_box(4096), + ); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/stream.rs b/src-tauri/vendor/sqlx-postgres/src/connection/stream.rs new file mode 100644 index 00000000..e8a1aedc --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/stream.rs @@ -0,0 +1,282 @@ +use std::collections::BTreeMap; +use std::ops::{ControlFlow, Deref, DerefMut}; +use std::str::FromStr; + +use futures_channel::mpsc::UnboundedSender; +use futures_util::SinkExt; +use log::Level; +use sqlx_core::bytes::Buf; + +use crate::connection::tls::MaybeUpgradeTls; +use crate::error::Error; +use crate::message::{ + BackendMessage, BackendMessageFormat, EncodeMessage, FrontendMessage, Notice, Notification, + ParameterStatus, ReceivedMessage, +}; +use crate::net::{self, BufferedSocket, Socket}; +use crate::{PgConnectOptions, PgDatabaseError, PgSeverity}; + +// the stream is a separate type from the connection to uphold the invariant where an instantiated +// [PgConnection] is a **valid** connection to postgres + +// when a new connection is asked for, we work directly on the [PgStream] type until the +// connection is fully established + +// in other words, `self` in any PgConnection method is a live connection to postgres that +// is fully prepared to receive queries + +pub struct PgStream { + // A trait object is okay here as the buffering amortizes the overhead of both the dynamic + // function call as well as the syscall. + inner: BufferedSocket>, + + // buffer of unreceived notification messages from `PUBLISH` + // this is set when creating a PgListener and only written to if that listener is + // re-used for query execution in-between receiving messages + pub(crate) notifications: Option>, + + pub(crate) parameter_statuses: BTreeMap, + + pub(crate) server_version_num: Option, +} + +impl PgStream { + pub(super) async fn connect(options: &PgConnectOptions) -> Result { + let socket_result = match options.fetch_socket() { + Some(ref path) => net::connect_uds(path, MaybeUpgradeTls(options)).await?, + None => net::connect_tcp(&options.host, options.port, MaybeUpgradeTls(options)).await?, + }; + + let socket = socket_result?; + + Ok(Self { + inner: BufferedSocket::new(socket), + notifications: None, + parameter_statuses: BTreeMap::default(), + server_version_num: None, + }) + } + + #[inline(always)] + pub(crate) fn write_msg(&mut self, message: impl FrontendMessage) -> Result<(), Error> { + self.write(EncodeMessage(message)) + } + + pub(crate) async fn send(&mut self, message: T) -> Result<(), Error> + where + T: FrontendMessage, + { + self.write_msg(message)?; + self.flush().await?; + Ok(()) + } + + // Expect a specific type and format + pub(crate) async fn recv_expect(&mut self) -> Result { + self.recv().await?.decode() + } + + pub(crate) async fn recv_unchecked(&mut self) -> Result { + // NOTE: to not break everything, this should be cancel-safe; + // DO NOT modify `buf` unless a full message has been read + self.inner + .try_read(|buf| { + // all packets in postgres start with a 5-byte header + // this header contains the message type and the total length of the message + let Some(mut header) = buf.get(..5) else { + return Ok(ControlFlow::Continue(5)); + }; + + let format = BackendMessageFormat::try_from_u8(header.get_u8())?; + + let message_len = header.get_u32() as usize; + + let expected_len = message_len + .checked_add(1) + // this shouldn't really happen but is mostly a sanity check + .ok_or_else(|| { + err_protocol!("message_len + 1 overflows usize: {message_len}") + })?; + + if buf.len() < expected_len { + return Ok(ControlFlow::Continue(expected_len)); + } + + // `buf` SHOULD NOT be modified ABOVE this line + + // pop off the format code since it's not counted in `message_len` + buf.advance(1); + + // consume the message, including the length prefix + let mut contents = buf.split_to(message_len).freeze(); + + // cut off the length prefix + contents.advance(4); + + Ok(ControlFlow::Break(ReceivedMessage { format, contents })) + }) + .await + } + + // Get the next message from the server + // May wait for more data from the server + pub(crate) async fn recv(&mut self) -> Result { + loop { + let message = self.recv_unchecked().await?; + + match message.format { + BackendMessageFormat::ErrorResponse => { + // An error returned from the database server. + return Err(message.decode::()?.into()); + } + + BackendMessageFormat::NotificationResponse => { + if let Some(buffer) = &mut self.notifications { + let notification: Notification = message.decode()?; + let _ = buffer.send(notification).await; + + continue; + } + } + + BackendMessageFormat::ParameterStatus => { + // informs the frontend about the current (initial) + // setting of backend parameters + + let ParameterStatus { name, value } = message.decode()?; + // TODO: handle `client_encoding`, `DateStyle` change + + match name.as_str() { + "server_version" => { + self.server_version_num = parse_server_version(&value); + } + _ => { + self.parameter_statuses.insert(name, value); + } + } + + continue; + } + + BackendMessageFormat::NoticeResponse => { + // do we need this to be more configurable? + // if you are reading this comment and think so, open an issue + + let notice: Notice = message.decode()?; + + let (log_level, tracing_level) = match notice.severity() { + PgSeverity::Fatal | PgSeverity::Panic | PgSeverity::Error => { + (Level::Error, tracing::Level::ERROR) + } + PgSeverity::Warning => (Level::Warn, tracing::Level::WARN), + PgSeverity::Notice => (Level::Info, tracing::Level::INFO), + PgSeverity::Debug => (Level::Debug, tracing::Level::DEBUG), + PgSeverity::Info | PgSeverity::Log => (Level::Trace, tracing::Level::TRACE), + }; + + let log_is_enabled = log::log_enabled!( + target: "sqlx::postgres::notice", + log_level + ) || sqlx_core::private_tracing_dynamic_enabled!( + target: "sqlx::postgres::notice", + tracing_level + ); + if log_is_enabled { + sqlx_core::private_tracing_dynamic_event!( + target: "sqlx::postgres::notice", + tracing_level, + message = notice.message() + ); + } + + continue; + } + + _ => {} + } + + return Ok(message); + } + } +} + +impl Deref for PgStream { + type Target = BufferedSocket>; + + #[inline] + fn deref(&self) -> &Self::Target { + &self.inner + } +} + +impl DerefMut for PgStream { + #[inline] + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.inner + } +} + +// reference: +// https://github.com/postgres/postgres/blob/6feebcb6b44631c3dc435e971bd80c2dd218a5ab/src/interfaces/libpq/fe-exec.c#L1030-L1065 +fn parse_server_version(s: &str) -> Option { + let mut parts = Vec::::with_capacity(3); + + let mut from = 0; + let mut chs = s.char_indices().peekable(); + while let Some((i, ch)) = chs.next() { + match ch { + '.' => { + if let Ok(num) = u32::from_str(&s[from..i]) { + parts.push(num); + from = i + 1; + } else { + break; + } + } + _ if ch.is_ascii_digit() => { + if chs.peek().is_none() { + if let Ok(num) = u32::from_str(&s[from..]) { + parts.push(num); + } + break; + } + } + _ => { + if let Ok(num) = u32::from_str(&s[from..i]) { + parts.push(num); + } + break; + } + }; + } + + let version_num = match parts.as_slice() { + [major, minor, rev] => (100 * major + minor) * 100 + rev, + [major, minor] if *major >= 10 => 100 * 100 * major + minor, + [major, minor] => (100 * major + minor) * 100, + [major] => 100 * 100 * major, + _ => return None, + }; + + Some(version_num) +} + +#[cfg(test)] +mod tests { + use super::parse_server_version; + + #[test] + fn test_parse_server_version_num() { + // old style + assert_eq!(parse_server_version("9.6.1"), Some(90601)); + // new style + assert_eq!(parse_server_version("10.1"), Some(100001)); + // old style without minor version + assert_eq!(parse_server_version("9.6devel"), Some(90600)); + // new style without minor version, e.g. */ + assert_eq!(parse_server_version("10devel"), Some(100000)); + assert_eq!(parse_server_version("13devel87"), Some(130000)); + // unknown + assert_eq!(parse_server_version("unknown"), None); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/connection/tls.rs b/src-tauri/vendor/sqlx-postgres/src/connection/tls.rs new file mode 100644 index 00000000..16b7333b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/connection/tls.rs @@ -0,0 +1,100 @@ +use crate::error::Error; +use crate::net::tls::{self, TlsConfig}; +use crate::net::{Socket, SocketIntoBox, WithSocket}; + +use crate::message::SslRequest; +use crate::{PgConnectOptions, PgSslMode}; + +pub struct MaybeUpgradeTls<'a>(pub &'a PgConnectOptions); + +impl<'a> WithSocket for MaybeUpgradeTls<'a> { + type Output = crate::Result>; + + async fn with_socket(self, socket: S) -> Self::Output { + maybe_upgrade(socket, self.0).await + } +} + +async fn maybe_upgrade( + mut socket: S, + options: &PgConnectOptions, +) -> Result, Error> { + // https://www.postgresql.org/docs/12/libpq-ssl.html#LIBPQ-SSL-SSLMODE-STATEMENTS + match options.ssl_mode { + // FIXME: Implement ALLOW + PgSslMode::Allow | PgSslMode::Disable => return Ok(Box::new(socket)), + + PgSslMode::Prefer => { + if !tls::available() { + return Ok(Box::new(socket)); + } + + // try upgrade, but its okay if we fail + if !request_upgrade(&mut socket, options).await? { + return Ok(Box::new(socket)); + } + } + + PgSslMode::Require | PgSslMode::VerifyFull | PgSslMode::VerifyCa => { + tls::error_if_unavailable()?; + + if !request_upgrade(&mut socket, options).await? { + // upgrade failed, die + return Err(Error::Tls("server does not support TLS".into())); + } + } + } + + let accept_invalid_certs = !matches!( + options.ssl_mode, + PgSslMode::VerifyCa | PgSslMode::VerifyFull + ); + let accept_invalid_hostnames = !matches!(options.ssl_mode, PgSslMode::VerifyFull); + + let config = TlsConfig { + accept_invalid_certs, + accept_invalid_hostnames, + hostname: &options.host, + root_cert_path: options.ssl_root_cert.as_ref(), + client_cert_path: options.ssl_client_cert.as_ref(), + client_key_path: options.ssl_client_key.as_ref(), + }; + + tls::handshake(socket, config, SocketIntoBox).await +} + +async fn request_upgrade( + socket: &mut impl Socket, + _options: &PgConnectOptions, +) -> Result { + // https://www.postgresql.org/docs/current/protocol-flow.html#id-1.10.5.7.11 + + // To initiate an SSL-encrypted connection, the frontend initially sends an + // SSLRequest message rather than a StartupMessage + + socket.write(SslRequest::BYTES).await?; + + // The server then responds with a single byte containing S or N, indicating that + // it is willing or unwilling to perform SSL, respectively. + + let mut response = [0u8]; + + socket.read(&mut &mut response[..]).await?; + + match response[0] { + b'S' => { + // The server is ready and willing to accept an SSL connection + Ok(true) + } + + b'N' => { + // The server is _unwilling_ to perform SSL + Ok(false) + } + + other => Err(err_protocol!( + "unexpected response from SSLRequest: 0x{:02x}", + other + )), + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/copy.rs b/src-tauri/vendor/sqlx-postgres/src/copy.rs new file mode 100644 index 00000000..1315ea0e --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/copy.rs @@ -0,0 +1,352 @@ +use std::borrow::Cow; +use std::ops::{Deref, DerefMut}; + +use futures_core::future::BoxFuture; +use futures_core::stream::BoxStream; + +use sqlx_core::bytes::{BufMut, Bytes}; + +use crate::connection::PgConnection; +use crate::error::{Error, Result}; +use crate::ext::async_stream::TryAsyncStream; +use crate::io::AsyncRead; +use crate::message::{ + BackendMessageFormat, CommandComplete, CopyData, CopyDone, CopyFail, CopyInResponse, + CopyOutResponse, CopyResponseData, Query, ReadyForQuery, +}; +use crate::pool::{Pool, PoolConnection}; +use crate::Postgres; + +impl PgConnection { + /// Issue a `COPY FROM STDIN` statement and transition the connection to streaming data + /// to Postgres. This is a more efficient way to import data into Postgres as compared to + /// `INSERT` but requires one of a few specific data formats (text/CSV/binary). + /// + /// If `statement` is anything other than a `COPY ... FROM STDIN ...` command, an error is + /// returned. + /// + /// Command examples and accepted formats for `COPY` data are shown here: + /// + /// + /// ### Note + /// [PgCopyIn::finish] or [PgCopyIn::abort] *must* be called when finished or the connection + /// will return an error the next time it is used. + pub async fn copy_in_raw(&mut self, statement: &str) -> Result> { + PgCopyIn::begin(self, statement).await + } + + /// Issue a `COPY TO STDOUT` statement and transition the connection to streaming data + /// from Postgres. This is a more efficient way to export data from Postgres but + /// arrives in chunks of one of a few data formats (text/CSV/binary). + /// + /// If `statement` is anything other than a `COPY ... TO STDOUT ...` command, + /// an error is returned. + /// + /// Note that once this process has begun, unless you read the stream to completion, + /// it can only be canceled in two ways: + /// + /// 1. by closing the connection, or: + /// 2. by using another connection to kill the server process that is sending the data as shown + /// [in this StackOverflow answer](https://stackoverflow.com/a/35319598). + /// + /// If you don't read the stream to completion, the next time the connection is used it will + /// need to read and discard all the remaining queued data, which could take some time. + /// + /// Command examples and accepted formats for `COPY` data are shown here: + /// + #[allow(clippy::needless_lifetimes)] + pub async fn copy_out_raw<'c>( + &'c mut self, + statement: &str, + ) -> Result>> { + pg_begin_copy_out(self, statement).await + } +} + +/// Implements methods for directly executing `COPY FROM/TO STDOUT` on a [`PgPool`][crate::PgPool]. +/// +/// This is a replacement for the inherent methods on `PgPool` which could not exist +/// once the Postgres driver was moved out into its own crate. +pub trait PgPoolCopyExt { + /// Issue a `COPY FROM STDIN` statement and begin streaming data to Postgres. + /// This is a more efficient way to import data into Postgres as compared to + /// `INSERT` but requires one of a few specific data formats (text/CSV/binary). + /// + /// A single connection will be checked out for the duration. + /// + /// If `statement` is anything other than a `COPY ... FROM STDIN ...` command, an error is + /// returned. + /// + /// Command examples and accepted formats for `COPY` data are shown here: + /// + /// + /// ### Note + /// [PgCopyIn::finish] or [PgCopyIn::abort] *must* be called when finished or the connection + /// will return an error the next time it is used. + fn copy_in_raw<'a>( + &'a self, + statement: &'a str, + ) -> BoxFuture<'a, Result>>>; + + /// Issue a `COPY TO STDOUT` statement and begin streaming data + /// from Postgres. This is a more efficient way to export data from Postgres but + /// arrives in chunks of one of a few data formats (text/CSV/binary). + /// + /// If `statement` is anything other than a `COPY ... TO STDOUT ...` command, + /// an error is returned. + /// + /// Note that once this process has begun, unless you read the stream to completion, + /// it can only be canceled in two ways: + /// + /// 1. by closing the connection, or: + /// 2. by using another connection to kill the server process that is sending the data as shown + /// [in this StackOverflow answer](https://stackoverflow.com/a/35319598). + /// + /// If you don't read the stream to completion, the next time the connection is used it will + /// need to read and discard all the remaining queued data, which could take some time. + /// + /// Command examples and accepted formats for `COPY` data are shown here: + /// + fn copy_out_raw<'a>( + &'a self, + statement: &'a str, + ) -> BoxFuture<'a, Result>>>; +} + +impl PgPoolCopyExt for Pool { + fn copy_in_raw<'a>( + &'a self, + statement: &'a str, + ) -> BoxFuture<'a, Result>>> { + Box::pin(async { PgCopyIn::begin(self.acquire().await?, statement).await }) + } + + fn copy_out_raw<'a>( + &'a self, + statement: &'a str, + ) -> BoxFuture<'a, Result>>> { + Box::pin(async { pg_begin_copy_out(self.acquire().await?, statement).await }) + } +} + +// (1 GiB - 1) - 1 - length prefix (4 bytes) +pub const PG_COPY_MAX_DATA_LEN: usize = 0x3fffffff - 1 - 4; + +/// A connection in streaming `COPY FROM STDIN` mode. +/// +/// Created by [PgConnection::copy_in_raw] or [Pool::copy_out_raw]. +/// +/// ### Note +/// [PgCopyIn::finish] or [PgCopyIn::abort] *must* be called when finished or the connection +/// will return an error the next time it is used. +#[must_use = "connection will error on next use if `.finish()` or `.abort()` is not called"] +pub struct PgCopyIn> { + conn: Option, + response: CopyResponseData, +} + +impl> PgCopyIn { + async fn begin(mut conn: C, statement: &str) -> Result { + conn.wait_until_ready().await?; + conn.inner.stream.send(Query(statement)).await?; + + let response = match conn.inner.stream.recv_expect::().await { + Ok(res) => res.0, + Err(e) => { + conn.inner.stream.recv().await?; + return Err(e); + } + }; + + Ok(PgCopyIn { + conn: Some(conn), + response, + }) + } + + /// Returns `true` if Postgres is expecting data in text or CSV format. + pub fn is_textual(&self) -> bool { + self.response.format == 0 + } + + /// Returns the number of columns expected in the input. + pub fn num_columns(&self) -> usize { + assert_eq!( + self.response.num_columns.unsigned_abs() as usize, + self.response.format_codes.len(), + "num_columns does not match format_codes.len()" + ); + self.response.format_codes.len() + } + + /// Check if a column is expecting data in text format (`true`) or binary format (`false`). + /// + /// ### Panics + /// If `column` is out of range according to [`.num_columns()`][Self::num_columns]. + pub fn column_is_textual(&self, column: usize) -> bool { + self.response.format_codes[column] == 0 + } + + /// Send a chunk of `COPY` data. + /// + /// The data is sent in chunks if it exceeds the maximum length of a `CopyData` message (1 GiB - 6 + /// bytes) and may be partially sent if this call is cancelled. + /// + /// If you're copying data from an `AsyncRead`, maybe consider [Self::read_from] instead. + pub async fn send(&mut self, data: impl Deref) -> Result<&mut Self> { + for chunk in data.deref().chunks(PG_COPY_MAX_DATA_LEN) { + self.conn + .as_deref_mut() + .expect("send_data: conn taken") + .inner + .stream + .send(CopyData(chunk)) + .await?; + } + + Ok(self) + } + + /// Copy data directly from `source` to the database without requiring an intermediate buffer. + /// + /// `source` will be read to the end. + /// + /// ### Note: Completion Step Required + /// You must still call either [Self::finish] or [Self::abort] to complete the process. + /// + /// ### Note: Runtime Features + /// This method uses the `AsyncRead` trait which is re-exported from either Tokio or `async-std` + /// depending on which runtime feature is used. + /// + /// The runtime features _used_ to be mutually exclusive, but are no longer. + /// If both `runtime-async-std` and `runtime-tokio` features are enabled, the Tokio version + /// takes precedent. + pub async fn read_from(&mut self, mut source: impl AsyncRead + Unpin) -> Result<&mut Self> { + let conn: &mut PgConnection = self.conn.as_deref_mut().expect("copy_from: conn taken"); + loop { + let buf = conn.inner.stream.write_buffer_mut(); + + // Write the CopyData format code and reserve space for the length. + // This may end up sending an empty `CopyData` packet if, after this point, + // we get canceled or read 0 bytes, but that should be fine. + buf.put_slice(b"d\0\0\0\x04"); + + let read = buf.read_from(&mut source).await?; + + if read == 0 { + break; + } + + // Write the length + let read32 = i32::try_from(read) + .map_err(|_| err_protocol!("number of bytes read exceeds 2^31 - 1: {}", read))?; + + (&mut buf.get_mut()[1..]).put_i32(read32 + 4); + + conn.inner.stream.flush().await?; + } + + Ok(self) + } + + /// Signal that the `COPY` process should be aborted and any data received should be discarded. + /// + /// The given message can be used for indicating the reason for the abort in the database logs. + /// + /// The server is expected to respond with an error, so only _unexpected_ errors are returned. + pub async fn abort(mut self, msg: impl Into) -> Result<()> { + let mut conn = self + .conn + .take() + .expect("PgCopyIn::fail_with: conn taken illegally"); + + conn.inner.stream.send(CopyFail::new(msg)).await?; + + match conn.inner.stream.recv().await { + Ok(msg) => Err(err_protocol!( + "fail_with: expected ErrorResponse, got: {:?}", + msg.format + )), + Err(Error::Database(e)) => { + match e.code() { + Some(Cow::Borrowed("57014")) => { + // postgres abort received error code + conn.inner.stream.recv_expect::().await?; + Ok(()) + } + _ => Err(Error::Database(e)), + } + } + Err(e) => Err(e), + } + } + + /// Signal that the `COPY` process is complete. + /// + /// The number of rows affected is returned. + pub async fn finish(mut self) -> Result { + let mut conn = self + .conn + .take() + .expect("CopyWriter::finish: conn taken illegally"); + + conn.inner.stream.send(CopyDone).await?; + let cc: CommandComplete = match conn.inner.stream.recv_expect().await { + Ok(cc) => cc, + Err(e) => { + conn.inner.stream.recv().await?; + return Err(e); + } + }; + + conn.inner.stream.recv_expect::().await?; + + Ok(cc.rows_affected()) + } +} + +impl> Drop for PgCopyIn { + fn drop(&mut self) { + if let Some(mut conn) = self.conn.take() { + conn.inner + .stream + .write_msg(CopyFail::new( + "PgCopyIn dropped without calling finish() or fail()", + )) + .expect("BUG: PgCopyIn abort message should not be too large"); + } + } +} + +async fn pg_begin_copy_out<'c, C: DerefMut + Send + 'c>( + mut conn: C, + statement: &str, +) -> Result>> { + conn.wait_until_ready().await?; + conn.inner.stream.send(Query(statement)).await?; + + let _: CopyOutResponse = conn.inner.stream.recv_expect().await?; + + let stream: TryAsyncStream<'c, Bytes> = try_stream! { + loop { + match conn.inner.stream.recv().await { + Err(e) => { + conn.inner.stream.recv_expect::().await?; + return Err(e); + }, + Ok(msg) => match msg.format { + BackendMessageFormat::CopyData => r#yield!(msg.decode::>()?.0), + BackendMessageFormat::CopyDone => { + let _ = msg.decode::()?; + conn.inner.stream.recv_expect::().await?; + conn.inner.stream.recv_expect::().await?; + return Ok(()) + }, + _ => return Err(err_protocol!("unexpected message format during copy out: {:?}", msg.format)) + } + } + } + }; + + Ok(Box::pin(stream)) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/database.rs b/src-tauri/vendor/sqlx-postgres/src/database.rs new file mode 100644 index 00000000..876e2958 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/database.rs @@ -0,0 +1,40 @@ +use crate::arguments::PgArgumentBuffer; +use crate::value::{PgValue, PgValueRef}; +use crate::{ + PgArguments, PgColumn, PgConnection, PgQueryResult, PgRow, PgStatement, PgTransactionManager, + PgTypeInfo, +}; + +pub(crate) use sqlx_core::database::{Database, HasStatementCache}; + +/// PostgreSQL database driver. +#[derive(Debug)] +pub struct Postgres; + +impl Database for Postgres { + type Connection = PgConnection; + + type TransactionManager = PgTransactionManager; + + type Row = PgRow; + + type QueryResult = PgQueryResult; + + type Column = PgColumn; + + type TypeInfo = PgTypeInfo; + + type Value = PgValue; + type ValueRef<'r> = PgValueRef<'r>; + + type Arguments<'q> = PgArguments; + type ArgumentBuffer<'q> = PgArgumentBuffer; + + type Statement<'q> = PgStatement<'q>; + + const NAME: &'static str = "PostgreSQL"; + + const URL_SCHEMES: &'static [&'static str] = &["postgres", "postgresql"]; +} + +impl HasStatementCache for Postgres {} diff --git a/src-tauri/vendor/sqlx-postgres/src/error.rs b/src-tauri/vendor/sqlx-postgres/src/error.rs new file mode 100644 index 00000000..db8bcc8a --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/error.rs @@ -0,0 +1,242 @@ +use std::error::Error as StdError; +use std::fmt::{self, Debug, Display, Formatter}; + +use atoi::atoi; +use smallvec::alloc::borrow::Cow; +use sqlx_core::bytes::Bytes; +pub(crate) use sqlx_core::error::*; + +use crate::message::{BackendMessage, BackendMessageFormat, Notice, PgSeverity}; + +/// An error returned from the PostgreSQL database. +pub struct PgDatabaseError(pub(crate) Notice); + +// Error message fields are documented: +// https://www.postgresql.org/docs/current/protocol-error-fields.html + +impl PgDatabaseError { + #[inline] + pub fn severity(&self) -> PgSeverity { + self.0.severity() + } + + /// The [SQLSTATE](https://www.postgresql.org/docs/current/errcodes-appendix.html) code for + /// this error. + #[inline] + pub fn code(&self) -> &str { + self.0.code() + } + + /// The primary human-readable error message. This should be accurate but + /// terse (typically one line). + #[inline] + pub fn message(&self) -> &str { + self.0.message() + } + + /// An optional secondary error message carrying more detail about the problem. + /// Might run to multiple lines. + #[inline] + pub fn detail(&self) -> Option<&str> { + self.0.get(b'D') + } + + /// An optional suggestion what to do about the problem. This is intended to differ from + /// `detail` in that it offers advice (potentially inappropriate) rather than hard facts. + /// Might run to multiple lines. + #[inline] + pub fn hint(&self) -> Option<&str> { + self.0.get(b'H') + } + + /// Indicates an error cursor position as an index into the original query string; or, + /// a position into an internally generated query. + #[inline] + pub fn position(&self) -> Option> { + self.0 + .get_raw(b'P') + .and_then(atoi) + .map(PgErrorPosition::Original) + .or_else(|| { + let position = self.0.get_raw(b'p').and_then(atoi)?; + let query = self.0.get(b'q')?; + + Some(PgErrorPosition::Internal { position, query }) + }) + } + + /// An indication of the context in which the error occurred. Presently this includes a call + /// stack traceback of active procedural language functions and internally-generated queries. + /// The trace is one entry per line, most recent first. + pub fn r#where(&self) -> Option<&str> { + self.0.get(b'W') + } + + /// If this error is with a specific database object, the + /// name of the schema containing that object, if any. + pub fn schema(&self) -> Option<&str> { + self.0.get(b's') + } + + /// If this error is with a specific table, the name of the table. + pub fn table(&self) -> Option<&str> { + self.0.get(b't') + } + + /// If the error is with a specific table column, the name of the column. + pub fn column(&self) -> Option<&str> { + self.0.get(b'c') + } + + /// If the error is with a specific data type, the name of the data type. + pub fn data_type(&self) -> Option<&str> { + self.0.get(b'd') + } + + /// If the error is with a specific constraint, the name of the constraint. + /// For this purpose, indexes are constraints, even if they weren't created + /// with constraint syntax. + pub fn constraint(&self) -> Option<&str> { + self.0.get(b'n') + } + + /// The file name of the source-code location where this error was reported. + pub fn file(&self) -> Option<&str> { + self.0.get(b'F') + } + + /// The line number of the source-code location where this error was reported. + pub fn line(&self) -> Option { + self.0.get_raw(b'L').and_then(atoi) + } + + /// The name of the source-code routine reporting this error. + pub fn routine(&self) -> Option<&str> { + self.0.get(b'R') + } +} + +#[derive(Debug, Eq, PartialEq)] +pub enum PgErrorPosition<'a> { + /// A position (in characters) into the original query. + Original(usize), + + /// A position into the internally-generated query. + Internal { + /// The position in characters. + position: usize, + + /// The text of a failed internally-generated command. This could be, for example, + /// the SQL query issued by a PL/pgSQL function. + query: &'a str, + }, +} + +impl Debug for PgDatabaseError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.debug_struct("PgDatabaseError") + .field("severity", &self.severity()) + .field("code", &self.code()) + .field("message", &self.message()) + .field("detail", &self.detail()) + .field("hint", &self.hint()) + .field("position", &self.position()) + .field("where", &self.r#where()) + .field("schema", &self.schema()) + .field("table", &self.table()) + .field("column", &self.column()) + .field("data_type", &self.data_type()) + .field("constraint", &self.constraint()) + .field("file", &self.file()) + .field("line", &self.line()) + .field("routine", &self.routine()) + .finish() + } +} + +impl Display for PgDatabaseError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.write_str(self.message()) + } +} + +impl StdError for PgDatabaseError {} + +impl DatabaseError for PgDatabaseError { + fn message(&self) -> &str { + self.message() + } + + fn code(&self) -> Option> { + Some(Cow::Borrowed(self.code())) + } + + #[doc(hidden)] + fn as_error(&self) -> &(dyn StdError + Send + Sync + 'static) { + self + } + + #[doc(hidden)] + fn as_error_mut(&mut self) -> &mut (dyn StdError + Send + Sync + 'static) { + self + } + + #[doc(hidden)] + fn into_error(self: Box) -> BoxDynError { + self + } + + fn is_transient_in_connect_phase(&self) -> bool { + // https://www.postgresql.org/docs/current/errcodes-appendix.html + [ + // too_many_connections + // This may be returned if we just un-gracefully closed a connection, + // give the database a chance to notice it and clean it up. + "53300", + // cannot_connect_now + // Returned if the database is still starting up. + "57P03", + ] + .contains(&self.code()) + } + + fn constraint(&self) -> Option<&str> { + self.constraint() + } + + fn table(&self) -> Option<&str> { + self.table() + } + + fn kind(&self) -> ErrorKind { + match self.code() { + error_codes::UNIQUE_VIOLATION => ErrorKind::UniqueViolation, + error_codes::FOREIGN_KEY_VIOLATION => ErrorKind::ForeignKeyViolation, + error_codes::NOT_NULL_VIOLATION => ErrorKind::NotNullViolation, + error_codes::CHECK_VIOLATION => ErrorKind::CheckViolation, + _ => ErrorKind::Other, + } + } +} + +// ErrorResponse is the same structure as NoticeResponse but a different format code. +impl BackendMessage for PgDatabaseError { + const FORMAT: BackendMessageFormat = BackendMessageFormat::ErrorResponse; + + #[inline(always)] + fn decode_body(buf: Bytes) -> std::result::Result { + Ok(Self(Notice::decode_body(buf)?)) + } +} + +/// For reference: +pub(crate) mod error_codes { + /// Caused when a unique or primary key is violated. + pub const UNIQUE_VIOLATION: &str = "23505"; + /// Caused when a foreign key is violated. + pub const FOREIGN_KEY_VIOLATION: &str = "23503"; + /// Caused when a column marked as NOT NULL received a null value. + pub const NOT_NULL_VIOLATION: &str = "23502"; + /// Caused when a check constraint is violated. + pub const CHECK_VIOLATION: &str = "23514"; +} diff --git a/src-tauri/vendor/sqlx-postgres/src/io/buf_mut.rs b/src-tauri/vendor/sqlx-postgres/src/io/buf_mut.rs new file mode 100644 index 00000000..0fe3809b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/io/buf_mut.rs @@ -0,0 +1,58 @@ +use crate::io::{PortalId, StatementId}; + +pub trait PgBufMutExt { + fn put_length_prefixed(&mut self, f: F) -> Result<(), crate::Error> + where + F: FnOnce(&mut Vec) -> Result<(), crate::Error>; + + fn put_statement_name(&mut self, id: StatementId); + + fn put_portal_name(&mut self, id: PortalId); +} + +impl PgBufMutExt for Vec { + // writes a length-prefixed message, this is used when encoding nearly all messages as postgres + // wants us to send the length of the often-variable-sized messages up front + fn put_length_prefixed(&mut self, write_contents: F) -> Result<(), crate::Error> + where + F: FnOnce(&mut Vec) -> Result<(), crate::Error>, + { + // reserve space to write the prefixed length + let offset = self.len(); + self.extend(&[0; 4]); + + // write the main body of the message + let write_result = write_contents(self); + + let size_result = write_result.and_then(|_| { + let size = self.len() - offset; + i32::try_from(size) + .map_err(|_| err_protocol!("message size out of range for protocol: {size}")) + }); + + match size_result { + Ok(size) => { + // now calculate the size of what we wrote and set the length value + self[offset..(offset + 4)].copy_from_slice(&size.to_be_bytes()); + Ok(()) + } + Err(e) => { + // Put the buffer back to where it was. + self.truncate(offset); + Err(e) + } + } + } + + // writes a statement name by ID + #[inline] + fn put_statement_name(&mut self, id: StatementId) { + id.put_name_with_nul(self); + } + + // writes a portal name by ID + #[inline] + fn put_portal_name(&mut self, id: PortalId) { + id.put_name_with_nul(self); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/io/mod.rs b/src-tauri/vendor/sqlx-postgres/src/io/mod.rs new file mode 100644 index 00000000..72f2a978 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/io/mod.rs @@ -0,0 +1,156 @@ +mod buf_mut; + +pub use buf_mut::PgBufMutExt; +use std::fmt; +use std::fmt::{Display, Formatter}; +use std::num::{NonZeroU32, Saturating}; + +pub(crate) use sqlx_core::io::*; + +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub(crate) struct StatementId(IdInner); + +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub(crate) struct PortalId(IdInner); + +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +struct IdInner(Option); + +pub(crate) struct DisplayId { + prefix: &'static str, + id: NonZeroU32, +} + +impl StatementId { + #[allow(dead_code)] + pub const UNNAMED: Self = Self(IdInner::UNNAMED); + + pub const NAMED_START: Self = Self(IdInner::NAMED_START); + + #[cfg(test)] + pub const TEST_VAL: Self = Self(IdInner::TEST_VAL); + + const NAME_PREFIX: &'static str = "sqlx_s_"; + + pub fn next(&self) -> Self { + Self(self.0.next()) + } + + pub fn name_len(&self) -> Saturating { + self.0.name_len(Self::NAME_PREFIX) + } + + /// Get a type to format this statement ID with [`Display`]. + /// + /// Returns `None` if this is the unnamed statement. + #[inline(always)] + pub fn display(&self) -> Option { + self.0.display(Self::NAME_PREFIX) + } + + pub fn put_name_with_nul(&self, buf: &mut Vec) { + self.0.put_name_with_nul(Self::NAME_PREFIX, buf) + } +} + +impl Display for DisplayId { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + write!(f, "{}{}", self.prefix, self.id) + } +} + +#[allow(dead_code)] +impl PortalId { + // None selects the unnamed portal + pub const UNNAMED: Self = PortalId(IdInner::UNNAMED); + + pub const NAMED_START: Self = PortalId(IdInner::NAMED_START); + + #[cfg(test)] + pub const TEST_VAL: Self = Self(IdInner::TEST_VAL); + + const NAME_PREFIX: &'static str = "sqlx_p_"; + + /// If ID represents a named portal, return the next ID, wrapping on overflow. + /// + /// If this ID represents the unnamed portal, return the same. + pub fn next(&self) -> Self { + Self(self.0.next()) + } + + /// Calculate the number of bytes that will be written by [`Self::put_name_with_nul()`]. + pub fn name_len(&self) -> Saturating { + self.0.name_len(Self::NAME_PREFIX) + } + + pub fn put_name_with_nul(&self, buf: &mut Vec) { + self.0.put_name_with_nul(Self::NAME_PREFIX, buf) + } +} + +impl IdInner { + const UNNAMED: Self = Self(None); + + const NAMED_START: Self = Self(Some(NonZeroU32::MIN)); + + #[cfg(test)] + pub const TEST_VAL: Self = Self(NonZeroU32::new(1234567890)); + + #[inline(always)] + fn next(&self) -> Self { + Self( + self.0 + .map(|id| id.checked_add(1).unwrap_or(NonZeroU32::MIN)), + ) + } + + #[inline(always)] + fn display(&self, prefix: &'static str) -> Option { + self.0.map(|id| DisplayId { prefix, id }) + } + + #[inline(always)] + fn name_len(&self, name_prefix: &str) -> Saturating { + let mut len = Saturating(0); + + if let Some(id) = self.0 { + len += name_prefix.len(); + // estimate the length of the ID in decimal + // `.ilog10()` can't panic since the value is never zero + len += id.get().ilog10() as usize; + // add one to compensate for `ilog10()` rounding down. + len += 1; + } + + // count the NUL terminator + len += 1; + + len + } + + #[inline(always)] + fn put_name_with_nul(&self, name_prefix: &str, buf: &mut Vec) { + if let Some(id) = self.0 { + buf.extend_from_slice(name_prefix.as_bytes()); + buf.extend_from_slice(itoa::Buffer::new().format(id.get()).as_bytes()); + } + + buf.push(0); + } +} + +#[test] +fn statement_id_display_matches_encoding() { + const EXPECTED_STR: &str = "sqlx_s_1234567890"; + const EXPECTED_BYTES: &[u8] = b"sqlx_s_1234567890\0"; + + let mut bytes = Vec::new(); + + StatementId::TEST_VAL.put_name_with_nul(&mut bytes); + + assert_eq!(bytes, EXPECTED_BYTES); + + let str = StatementId::TEST_VAL.display().unwrap().to_string(); + + assert_eq!(str, EXPECTED_STR); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/lib.rs b/src-tauri/vendor/sqlx-postgres/src/lib.rs new file mode 100644 index 00000000..bded7549 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/lib.rs @@ -0,0 +1,83 @@ +//! **PostgreSQL** database driver. + +#[macro_use] +extern crate sqlx_core; + +use crate::executor::Executor; + +mod advisory_lock; +mod arguments; +mod column; +mod connection; +mod copy; +mod database; +mod error; +mod io; +mod listener; +mod message; +mod options; +mod query_result; +mod row; +mod statement; +mod transaction; +mod type_checking; +mod type_info; +pub mod types; +mod value; + +#[cfg(feature = "any")] +// We are hiding the any module with its AnyConnectionBackend trait +// so that IDEs don't show it in the autocompletion list +// and end users don't accidentally use it. This can result in +// nested transactions not behaving as expected. +// For more information, see https://github.com/launchbadge/sqlx/pull/3254#issuecomment-2144043823 +#[doc(hidden)] +pub mod any; + +#[doc(hidden)] +pub use copy::PG_COPY_MAX_DATA_LEN; + +#[cfg(feature = "migrate")] +mod migrate; + +#[cfg(feature = "migrate")] +mod testing; + +pub(crate) use sqlx_core::driver_prelude::*; + +pub use advisory_lock::{PgAdvisoryLock, PgAdvisoryLockGuard, PgAdvisoryLockKey}; +pub use arguments::{PgArgumentBuffer, PgArguments}; +pub use column::PgColumn; +pub use connection::PgConnection; +pub use copy::{PgCopyIn, PgPoolCopyExt}; +pub use database::Postgres; +pub use error::{PgDatabaseError, PgErrorPosition}; +pub use listener::{PgListener, PgNotification}; +pub use message::PgSeverity; +pub use options::{PgConnectOptions, PgSslMode}; +pub use query_result::PgQueryResult; +pub use row::PgRow; +pub use statement::PgStatement; +pub use transaction::PgTransactionManager; +pub use type_info::{PgTypeInfo, PgTypeKind}; +pub use types::PgHasArrayType; +pub use value::{PgValue, PgValueFormat, PgValueRef}; + +/// An alias for [`Pool`][crate::pool::Pool], specialized for Postgres. +pub type PgPool = crate::pool::Pool; + +/// An alias for [`PoolOptions`][crate::pool::PoolOptions], specialized for Postgres. +pub type PgPoolOptions = crate::pool::PoolOptions; + +/// An alias for [`Executor<'_, Database = Postgres>`][Executor]. +pub trait PgExecutor<'c>: Executor<'c, Database = Postgres> {} +impl<'c, T: Executor<'c, Database = Postgres>> PgExecutor<'c> for T {} + +/// An alias for [`Transaction`][crate::transaction::Transaction], specialized for Postgres. +pub type PgTransaction<'c> = crate::transaction::Transaction<'c, Postgres>; + +impl_into_arguments_for_arguments!(PgArguments); +impl_acquire!(Postgres, PgConnection); +impl_column_index_for_row!(PgRow); +impl_column_index_for_statement!(PgStatement); +impl_encode_for_option!(Postgres); diff --git a/src-tauri/vendor/sqlx-postgres/src/listener.rs b/src-tauri/vendor/sqlx-postgres/src/listener.rs new file mode 100644 index 00000000..17a46a91 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/listener.rs @@ -0,0 +1,517 @@ +use std::fmt::{self, Debug}; +use std::io; +use std::str::from_utf8; + +use futures_channel::mpsc; +use futures_core::future::BoxFuture; +use futures_core::stream::{BoxStream, Stream}; +use futures_util::{FutureExt, StreamExt, TryFutureExt, TryStreamExt}; +use sqlx_core::acquire::Acquire; +use sqlx_core::transaction::Transaction; +use sqlx_core::Either; +use tracing::Instrument; + +use crate::describe::Describe; +use crate::error::Error; +use crate::executor::{Execute, Executor}; +use crate::message::{BackendMessageFormat, Notification}; +use crate::pool::PoolOptions; +use crate::pool::{Pool, PoolConnection}; +use crate::{PgConnection, PgQueryResult, PgRow, PgStatement, PgTypeInfo, Postgres}; + +/// A stream of asynchronous notifications from Postgres. +/// +/// This listener will auto-reconnect. If the active +/// connection being used ever dies, this listener will detect that event, create a +/// new connection, will re-subscribe to all of the originally specified channels, and will resume +/// operations as normal. +pub struct PgListener { + pool: Pool, + connection: Option>, + buffer_rx: mpsc::UnboundedReceiver, + buffer_tx: Option>, + channels: Vec, + ignore_close_event: bool, + eager_reconnect: bool, +} + +/// An asynchronous notification from Postgres. +pub struct PgNotification(Notification); + +impl PgListener { + pub async fn connect(url: &str) -> Result { + // Create a pool of 1 without timeouts (as they don't apply here) + // We only use the pool to handle re-connections + let pool = PoolOptions::::new() + .max_connections(1) + .max_lifetime(None) + .idle_timeout(None) + .connect(url) + .await?; + + let mut this = Self::connect_with(&pool).await?; + // We don't need to handle close events + this.ignore_close_event = true; + + Ok(this) + } + + pub async fn connect_with(pool: &Pool) -> Result { + // Pull out an initial connection + let mut connection = pool.acquire().await?; + + // Setup a notification buffer + let (sender, receiver) = mpsc::unbounded(); + connection.inner.stream.notifications = Some(sender); + + Ok(Self { + pool: pool.clone(), + connection: Some(connection), + buffer_rx: receiver, + buffer_tx: None, + channels: Vec::new(), + ignore_close_event: false, + eager_reconnect: true, + }) + } + + /// Set whether or not to ignore [`Pool::close_event()`]. Defaults to `false`. + /// + /// By default, when [`Pool::close()`] is called on the pool this listener is using + /// while [`Self::recv()`] or [`Self::try_recv()`] are waiting for a message, the wait is + /// cancelled and `Err(PoolClosed)` is returned. + /// + /// This is because `Pool::close()` will wait until _all_ connections are returned and closed, + /// including the one being used by this listener. + /// + /// Otherwise, `pool.close().await` would have to wait until `PgListener` encountered a + /// need to acquire a new connection (timeout, error, etc.) and dropped the one it was + /// currently holding, at which point `.recv()` or `.try_recv()` would return `Err(PoolClosed)` + /// on the attempt to acquire a new connection anyway. + /// + /// However, if you want `PgListener` to ignore the close event and continue waiting for a + /// message as long as it can, set this to `true`. + /// + /// Does nothing if this was constructed with [`PgListener::connect()`], as that creates an + /// internal pool just for the new instance of `PgListener` which cannot be closed manually. + pub fn ignore_pool_close_event(&mut self, val: bool) { + self.ignore_close_event = val; + } + + /// Set whether a lost connection in `try_recv()` should be re-established before it returns + /// `Ok(None)`, or on the next call to `try_recv()`. + /// + /// By default, this is `true` and the connection is re-established before returning `Ok(None)`. + /// + /// If this is set to `false` then notifications will continue to be lost until the next call + /// to `try_recv()`. If your recovery logic uses a different database connection then + /// notifications that occur after it completes may be lost without any way to tell that they + /// have been. + pub fn eager_reconnect(&mut self, val: bool) { + self.eager_reconnect = val; + } + + /// Starts listening for notifications on a channel. + /// The channel name is quoted here to ensure case sensitivity. + pub async fn listen(&mut self, channel: &str) -> Result<(), Error> { + self.connection() + .await? + .execute(&*format!(r#"LISTEN "{}""#, ident(channel))) + .await?; + + self.channels.push(channel.to_owned()); + + Ok(()) + } + + /// Starts listening for notifications on all channels. + pub async fn listen_all( + &mut self, + channels: impl IntoIterator, + ) -> Result<(), Error> { + let beg = self.channels.len(); + self.channels.extend(channels.into_iter().map(|s| s.into())); + + let query = build_listen_all_query(&self.channels[beg..]); + self.connection().await?.execute(&*query).await?; + + Ok(()) + } + + /// Stops listening for notifications on a channel. + /// The channel name is quoted here to ensure case sensitivity. + pub async fn unlisten(&mut self, channel: &str) -> Result<(), Error> { + // use RAW connection and do NOT re-connect automatically, since this is not required for + // UNLISTEN (we've disconnected anyways) + if let Some(connection) = self.connection.as_mut() { + connection + .execute(&*format!(r#"UNLISTEN "{}""#, ident(channel))) + .await?; + } + + if let Some(pos) = self.channels.iter().position(|s| s == channel) { + self.channels.remove(pos); + } + + Ok(()) + } + + /// Stops listening for notifications on all channels. + pub async fn unlisten_all(&mut self) -> Result<(), Error> { + // use RAW connection and do NOT re-connect automatically, since this is not required for + // UNLISTEN (we've disconnected anyways) + if let Some(connection) = self.connection.as_mut() { + connection.execute("UNLISTEN *").await?; + } + + self.channels.clear(); + + Ok(()) + } + + #[inline] + async fn connect_if_needed(&mut self) -> Result<(), Error> { + if self.connection.is_none() { + let mut connection = self.pool.acquire().await?; + connection.inner.stream.notifications = self.buffer_tx.take(); + + connection + .execute(&*build_listen_all_query(&self.channels)) + .await?; + + self.connection = Some(connection); + } + + Ok(()) + } + + #[inline] + async fn connection(&mut self) -> Result<&mut PgConnection, Error> { + // Ensure we have an active connection to work with. + self.connect_if_needed().await?; + + Ok(self.connection.as_mut().unwrap()) + } + + /// Receives the next notification available from any of the subscribed channels. + /// + /// If the connection to PostgreSQL is lost, it is automatically reconnected on the next + /// call to `recv()`, and should be entirely transparent (as long as it was just an + /// intermittent network failure or long-lived connection reaper). + /// + /// As notifications are transient, any received while the connection was lost, will not + /// be returned. If you'd prefer the reconnection to be explicit and have a chance to + /// do something before, please see [`try_recv`](Self::try_recv). + /// + /// # Example + /// + /// ```rust,no_run + /// # use sqlx::postgres::PgListener; + /// # + /// # sqlx::__rt::test_block_on(async move { + /// let mut listener = PgListener::connect("postgres:// ...").await?; + /// loop { + /// // ask for next notification, re-connecting (transparently) if needed + /// let notification = listener.recv().await?; + /// + /// // handle notification, do something interesting + /// } + /// # Result::<(), sqlx::Error>::Ok(()) + /// # }).unwrap(); + /// ``` + pub async fn recv(&mut self) -> Result { + loop { + if let Some(notification) = self.try_recv().await? { + return Ok(notification); + } + } + } + + /// Receives the next notification available from any of the subscribed channels. + /// + /// If the connection to PostgreSQL is lost, `None` is returned, and the connection is + /// reconnected either immediately, or on the next call to `try_recv()`, depending on + /// the value of [`eager_reconnect`]. + /// + /// # Example + /// + /// ```rust,no_run + /// # use sqlx::postgres::PgListener; + /// # + /// # sqlx::__rt::test_block_on(async move { + /// # let mut listener = PgListener::connect("postgres:// ...").await?; + /// loop { + /// // start handling notifications, connecting if needed + /// while let Some(notification) = listener.try_recv().await? { + /// // handle notification + /// } + /// + /// // connection lost, do something interesting + /// } + /// # Result::<(), sqlx::Error>::Ok(()) + /// # }).unwrap(); + /// ``` + /// + /// [`eager_reconnect`]: PgListener::eager_reconnect + pub async fn try_recv(&mut self) -> Result, Error> { + // Flush the buffer first, if anything + // This would only fill up if this listener is used as a connection + if let Some(notification) = self.next_buffered() { + return Ok(Some(notification)); + } + + // Fetch our `CloseEvent` listener, if applicable. + let mut close_event = (!self.ignore_close_event).then(|| self.pool.close_event()); + + loop { + let next_message = self.connection().await?.inner.stream.recv_unchecked(); + + let res = if let Some(ref mut close_event) = close_event { + // cancels the wait and returns `Err(PoolClosed)` if the pool is closed + // before `next_message` returns, or if the pool was already closed + close_event.do_until(next_message).await? + } else { + next_message.await + }; + + let message = match res { + Ok(message) => message, + + // The connection is dead, ensure that it is dropped, + // update self state, and loop to try again. + Err(Error::Io(err)) + if matches!( + err.kind(), + io::ErrorKind::ConnectionAborted | + io::ErrorKind::UnexpectedEof | + // see ERRORS section in tcp(7) man page (https://man7.org/linux/man-pages/man7/tcp.7.html) + io::ErrorKind::TimedOut | + io::ErrorKind::BrokenPipe + ) => + { + if let Some(mut conn) = self.connection.take() { + self.buffer_tx = conn.inner.stream.notifications.take(); + // Close the connection in a background task, so we can continue. + conn.close_on_drop(); + } + + if self.eager_reconnect { + self.connect_if_needed().await?; + } + + // lost connection + return Ok(None); + } + + // Forward other errors + Err(error) => { + return Err(error); + } + }; + + match message.format { + // We've received an async notification, return it. + BackendMessageFormat::NotificationResponse => { + return Ok(Some(PgNotification(message.decode()?))); + } + + // Mark the connection as ready for another query + BackendMessageFormat::ReadyForQuery => { + self.connection().await?.inner.pending_ready_for_query_count -= 1; + } + + // Ignore unexpected messages + _ => {} + } + } + } + + /// Receives the next notification that already exists in the connection buffer, if any. + /// + /// This is similar to `try_recv`, except it will not wait if the connection has not yet received a notification. + /// + /// This is helpful if you want to retrieve all buffered notifications and process them in batches. + pub fn next_buffered(&mut self) -> Option { + if let Ok(Some(notification)) = self.buffer_rx.try_next() { + Some(PgNotification(notification)) + } else { + None + } + } + + /// Consume this listener, returning a `Stream` of notifications. + /// + /// The backing connection will be automatically reconnected should it be lost. + /// + /// This has the same potential drawbacks as [`recv`](PgListener::recv). + /// + pub fn into_stream(mut self) -> impl Stream> + Unpin { + Box::pin(try_stream! { + loop { + r#yield!(self.recv().await?); + } + }) + } +} + +impl Drop for PgListener { + fn drop(&mut self) { + if let Some(mut conn) = self.connection.take() { + let fut = async move { + let _ = conn.execute("UNLISTEN *").await; + + // inline the drop handler from `PoolConnection` so it doesn't try to spawn another task + // otherwise, it may trigger a panic if this task is dropped because the runtime is going away: + // https://github.com/launchbadge/sqlx/issues/1389 + conn.return_to_pool().await; + }; + + // Unregister any listeners before returning the connection to the pool. + crate::rt::spawn(fut.in_current_span()); + } + } +} + +impl<'c> Acquire<'c> for &'c mut PgListener { + type Database = Postgres; + type Connection = &'c mut PgConnection; + + fn acquire(self) -> BoxFuture<'c, Result> { + self.connection().boxed() + } + + fn begin(self) -> BoxFuture<'c, Result, Error>> { + self.connection().and_then(|c| c.begin()).boxed() + } +} + +impl<'c> Executor<'c> for &'c mut PgListener { + type Database = Postgres; + + fn fetch_many<'e, 'q, E>( + self, + query: E, + ) -> BoxStream<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + futures_util::stream::once(async move { + // need some basic type annotation to help the compiler a bit + let res: Result<_, Error> = Ok(self.connection().await?.fetch_many(query)); + res + }) + .try_flatten() + .boxed() + } + + fn fetch_optional<'e, 'q, E>(self, query: E) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + E: Execute<'q, Self::Database>, + 'q: 'e, + E: 'q, + { + async move { self.connection().await?.fetch_optional(query).await }.boxed() + } + + fn prepare_with<'e, 'q: 'e>( + self, + query: &'q str, + parameters: &'e [PgTypeInfo], + ) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + async move { + self.connection() + .await? + .prepare_with(query, parameters) + .await + } + .boxed() + } + + #[doc(hidden)] + fn describe<'e, 'q: 'e>( + self, + query: &'q str, + ) -> BoxFuture<'e, Result, Error>> + where + 'c: 'e, + { + async move { self.connection().await?.describe(query).await }.boxed() + } +} + +impl PgNotification { + /// The process ID of the notifying backend process. + #[inline] + pub fn process_id(&self) -> u32 { + self.0.process_id + } + + /// The channel that the notify has been raised on. This can be thought + /// of as the message topic. + #[inline] + pub fn channel(&self) -> &str { + from_utf8(&self.0.channel).unwrap() + } + + /// The payload of the notification. An empty payload is received as an + /// empty string. + #[inline] + pub fn payload(&self) -> &str { + from_utf8(&self.0.payload).unwrap() + } +} + +impl Debug for PgListener { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("PgListener").finish() + } +} + +impl Debug for PgNotification { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("PgNotification") + .field("process_id", &self.process_id()) + .field("channel", &self.channel()) + .field("payload", &self.payload()) + .finish() + } +} + +fn ident(mut name: &str) -> String { + // If the input string contains a NUL byte, we should truncate the + // identifier. + if let Some(index) = name.find('\0') { + name = &name[..index]; + } + + // Any double quotes must be escaped + name.replace('"', "\"\"") +} + +fn build_listen_all_query(channels: impl IntoIterator>) -> String { + channels.into_iter().fold(String::new(), |mut acc, chan| { + acc.push_str(r#"LISTEN ""#); + acc.push_str(&ident(chan.as_ref())); + acc.push_str(r#"";"#); + acc + }) +} + +#[test] +fn test_build_listen_all_query_with_single_channel() { + let output = build_listen_all_query(&["test"]); + assert_eq!(output.as_str(), r#"LISTEN "test";"#); +} + +#[test] +fn test_build_listen_all_query_with_multiple_channels() { + let output = build_listen_all_query(&["channel.0", "channel.1"]); + assert_eq!(output.as_str(), r#"LISTEN "channel.0";LISTEN "channel.1";"#); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/authentication.rs b/src-tauri/vendor/sqlx-postgres/src/message/authentication.rs new file mode 100644 index 00000000..3a3cf7ff --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/authentication.rs @@ -0,0 +1,193 @@ +use std::str::from_utf8; + +use memchr::memchr; +use sqlx_core::bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::ProtocolDecode; + +use crate::message::{BackendMessage, BackendMessageFormat}; +use base64::prelude::{Engine as _, BASE64_STANDARD}; +// On startup, the server sends an appropriate authentication request message, +// to which the frontend must reply with an appropriate authentication +// response message (such as a password). + +// For all authentication methods except GSSAPI, SSPI and SASL, there is at +// most one request and one response. In some methods, no response at all is +// needed from the frontend, and so no authentication request occurs. + +// For GSSAPI, SSPI and SASL, multiple exchanges of packets may +// be needed to complete the authentication. + +// +// + +#[derive(Debug)] +pub enum Authentication { + /// The authentication exchange is successfully completed. + Ok, + + /// The frontend must now send a [PasswordMessage] containing the + /// password in clear-text form. + CleartextPassword, + + /// The frontend must now send a [PasswordMessage] containing the + /// password (with user name) encrypted via MD5, then encrypted + /// again using the 4-byte random salt. + Md5Password(AuthenticationMd5Password), + + /// The frontend must now initiate a SASL negotiation, + /// using one of the SASL mechanisms listed in the message. + /// + /// The frontend will send a [SaslInitialResponse] with the name + /// of the selected mechanism, and the first part of the SASL + /// data stream in response to this. + /// + /// If further messages are needed, the server will + /// respond with [Authentication::SaslContinue]. + Sasl(AuthenticationSasl), + + /// This message contains challenge data from the previous step of SASL negotiation. + /// + /// The frontend must respond with a [SaslResponse] message. + SaslContinue(AuthenticationSaslContinue), + + /// SASL authentication has completed with additional mechanism-specific + /// data for the client. + /// + /// The server will next send [Authentication::Ok] to + /// indicate successful authentication. + SaslFinal(AuthenticationSaslFinal), +} + +impl BackendMessage for Authentication { + const FORMAT: BackendMessageFormat = BackendMessageFormat::Authentication; + + fn decode_body(mut buf: Bytes) -> Result { + Ok(match buf.get_u32() { + 0 => Authentication::Ok, + + 3 => Authentication::CleartextPassword, + + 5 => { + let mut salt = [0; 4]; + buf.copy_to_slice(&mut salt); + + Authentication::Md5Password(AuthenticationMd5Password { salt }) + } + + 10 => Authentication::Sasl(AuthenticationSasl(buf)), + 11 => Authentication::SaslContinue(AuthenticationSaslContinue::decode(buf)?), + 12 => Authentication::SaslFinal(AuthenticationSaslFinal::decode(buf)?), + + ty => { + return Err(err_protocol!("unknown authentication method: {}", ty)); + } + }) + } +} + +/// Body of [Authentication::Md5Password]. +#[derive(Debug)] +pub struct AuthenticationMd5Password { + pub salt: [u8; 4], +} + +/// Body of [Authentication::Sasl]. +#[derive(Debug)] +pub struct AuthenticationSasl(Bytes); + +impl AuthenticationSasl { + #[inline] + pub fn mechanisms(&self) -> SaslMechanisms<'_> { + SaslMechanisms(&self.0) + } +} + +/// An iterator over the SASL authentication mechanisms provided by the server. +pub struct SaslMechanisms<'a>(&'a [u8]); + +impl<'a> Iterator for SaslMechanisms<'a> { + type Item = &'a str; + + fn next(&mut self) -> Option { + if !self.0.is_empty() && self.0[0] == b'\0' { + return None; + } + + let mechanism = memchr(b'\0', self.0).and_then(|nul| from_utf8(&self.0[..nul]).ok())?; + + self.0 = &self.0[(mechanism.len() + 1)..]; + + Some(mechanism) + } +} + +#[derive(Debug)] +pub struct AuthenticationSaslContinue { + pub salt: Vec, + pub iterations: u32, + pub nonce: String, + pub message: String, +} + +impl ProtocolDecode<'_> for AuthenticationSaslContinue { + fn decode_with(buf: Bytes, _: ()) -> Result { + let mut iterations: u32 = 4096; + let mut salt = Vec::new(); + let mut nonce = Bytes::new(); + + // [Example] + // r=/z+giZiTxAH7r8sNAeHr7cvpqV3uo7G/bJBIJO3pjVM7t3ng,s=4UV68bIkC8f9/X8xH7aPhg==,i=4096 + + for item in buf.split(|b| *b == b',') { + let key = item[0]; + let value = &item[2..]; + + match key { + b'r' => { + nonce = buf.slice_ref(value); + } + + b'i' => { + iterations = atoi::atoi(value).unwrap_or(4096); + } + + b's' => { + salt = BASE64_STANDARD.decode(value).map_err(Error::protocol)?; + } + + _ => {} + } + } + + Ok(Self { + iterations, + salt, + nonce: from_utf8(&nonce).map_err(Error::protocol)?.to_owned(), + message: from_utf8(&buf).map_err(Error::protocol)?.to_owned(), + }) + } +} + +#[derive(Debug)] +pub struct AuthenticationSaslFinal { + pub verifier: Vec, +} + +impl ProtocolDecode<'_> for AuthenticationSaslFinal { + fn decode_with(buf: Bytes, _: ()) -> Result { + let mut verifier = Vec::new(); + + for item in buf.split(|b| *b == b',') { + let key = item[0]; + let value = &item[2..]; + + if let b'v' = key { + verifier = BASE64_STANDARD.decode(value).map_err(Error::protocol)?; + } + } + + Ok(Self { verifier }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/backend_key_data.rs b/src-tauri/vendor/sqlx-postgres/src/message/backend_key_data.rs new file mode 100644 index 00000000..f2dc2f23 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/backend_key_data.rs @@ -0,0 +1,50 @@ +use byteorder::{BigEndian, ByteOrder}; +use sqlx_core::bytes::Bytes; + +use crate::error::Error; +use crate::message::{BackendMessage, BackendMessageFormat}; + +/// Contains cancellation key data. The frontend must save these values if it +/// wishes to be able to issue `CancelRequest` messages later. +#[derive(Debug)] +pub struct BackendKeyData { + /// The process ID of this database. + pub process_id: u32, + + /// The secret key of this database. + pub secret_key: u32, +} + +impl BackendMessage for BackendKeyData { + const FORMAT: BackendMessageFormat = BackendMessageFormat::BackendKeyData; + + fn decode_body(buf: Bytes) -> Result { + let process_id = BigEndian::read_u32(&buf); + let secret_key = BigEndian::read_u32(&buf[4..]); + + Ok(Self { + process_id, + secret_key, + }) + } +} + +#[test] +fn test_decode_backend_key_data() { + const DATA: &[u8] = b"\0\0'\xc6\x89R\xc5+"; + + let m = BackendKeyData::decode_body(DATA.into()).unwrap(); + + assert_eq!(m.process_id, 10182); + assert_eq!(m.secret_key, 2303903019); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_backend_key_data(b: &mut test::Bencher) { + const DATA: &[u8] = b"\0\0'\xc6\x89R\xc5+"; + + b.iter(|| { + BackendKeyData::decode_body(test::black_box(Bytes::from_static(DATA))).unwrap(); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/bind.rs b/src-tauri/vendor/sqlx-postgres/src/message/bind.rs new file mode 100644 index 00000000..4f58bdf5 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/bind.rs @@ -0,0 +1,106 @@ +use crate::io::{PgBufMutExt, PortalId, StatementId}; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use crate::PgValueFormat; +use std::num::Saturating; + +/// +/// +/// ## Note: +/// +/// The integer values for number of bind parameters, number of parameter format codes, +/// and number of result format codes all are interpreted as *unsigned*! +#[derive(Debug)] +pub struct Bind<'a> { + /// The ID of the destination portal (`PortalId::UNNAMED` selects the unnamed portal). + pub portal: PortalId, + + /// The id of the source prepared statement. + pub statement: StatementId, + + /// The parameter format codes. Each must presently be zero (text) or one (binary). + /// + /// There can be zero to indicate that there are no parameters or that the parameters all use the + /// default format (text); or one, in which case the specified format code is applied to all + /// parameters; or it can equal the actual number of parameters. + pub formats: &'a [PgValueFormat], + + // Note: interpreted as unsigned, as is `formats.len()` and `result_formats.len()` + /// The number of parameters. + /// + /// May be different from `formats.len()` + pub num_params: u16, + + /// The value of each parameter, in the indicated format. + pub params: &'a [u8], + + /// The result-column format codes. Each must presently be zero (text) or one (binary). + /// + /// There can be zero to indicate that there are no result columns or that the + /// result columns should all use the default format (text); or one, in which + /// case the specified format code is applied to all result columns (if any); + /// or it can equal the actual number of result columns of the query. + pub result_formats: &'a [PgValueFormat], +} + +impl FrontendMessage for Bind<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Bind; + + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + size += self.portal.name_len(); + size += self.statement.name_len(); + + // Parameter formats and length prefix + size += 2; + size += self.formats.len(); + + // `num_params` + size += 2; + + size += self.params.len(); + + // Result formats and length prefix + size += 2; + size += self.result_formats.len(); + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), crate::Error> { + buf.put_portal_name(self.portal); + + buf.put_statement_name(self.statement); + + // NOTE: the integer values for the number of parameters and format codes in this message + // are all interpreted as *unsigned*! + // + // https://github.com/launchbadge/sqlx/issues/3464 + let formats_len = u16::try_from(self.formats.len()).map_err(|_| { + err_protocol!("too many parameter format codes ({})", self.formats.len()) + })?; + + buf.extend(formats_len.to_be_bytes()); + + for &format in self.formats { + buf.extend((format as i16).to_be_bytes()); + } + + buf.extend(self.num_params.to_be_bytes()); + + buf.extend(self.params); + + let result_formats_len = u16::try_from(self.formats.len()) + .map_err(|_| err_protocol!("too many result format codes ({})", self.formats.len()))?; + + buf.extend(result_formats_len.to_be_bytes()); + + for &format in self.result_formats { + buf.extend((format as i16).to_be_bytes()); + } + + Ok(()) + } +} + +// TODO: Unit Test Bind +// TODO: Benchmark Bind diff --git a/src-tauri/vendor/sqlx-postgres/src/message/close.rs b/src-tauri/vendor/sqlx-postgres/src/message/close.rs new file mode 100644 index 00000000..172f244c --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/close.rs @@ -0,0 +1,45 @@ +use crate::io::{PgBufMutExt, PortalId, StatementId}; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use std::num::Saturating; + +const CLOSE_PORTAL: u8 = b'P'; +const CLOSE_STATEMENT: u8 = b'S'; + +#[derive(Debug)] +#[allow(dead_code)] +pub enum Close { + Statement(StatementId), + Portal(PortalId), +} + +impl FrontendMessage for Close { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Close; + + fn body_size_hint(&self) -> Saturating { + // Either `CLOSE_PORTAL` or `CLOSE_STATEMENT` + let mut size = Saturating(1); + + match self { + Close::Statement(id) => size += id.name_len(), + Close::Portal(id) => size += id.name_len(), + } + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), crate::Error> { + match self { + Close::Statement(id) => { + buf.push(CLOSE_STATEMENT); + buf.put_statement_name(*id); + } + + Close::Portal(id) => { + buf.push(CLOSE_PORTAL); + buf.put_portal_name(*id); + } + } + + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/command_complete.rs b/src-tauri/vendor/sqlx-postgres/src/message/command_complete.rs new file mode 100644 index 00000000..eb33c512 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/command_complete.rs @@ -0,0 +1,82 @@ +use atoi::atoi; +use memchr::memrchr; +use sqlx_core::bytes::Bytes; + +use crate::error::Error; +use crate::message::{BackendMessage, BackendMessageFormat}; + +#[derive(Debug)] +pub struct CommandComplete { + /// The command tag. This is usually a single word that identifies which SQL command + /// was completed. + tag: Bytes, +} + +impl BackendMessage for CommandComplete { + const FORMAT: BackendMessageFormat = BackendMessageFormat::CommandComplete; + + fn decode_body(bytes: Bytes) -> Result { + Ok(CommandComplete { tag: bytes }) + } +} + +impl CommandComplete { + /// Returns the number of rows affected. + /// If the command does not return rows (e.g., "CREATE TABLE"), returns 0. + pub fn rows_affected(&self) -> u64 { + // Look backwards for the first SPACE + memrchr(b' ', &self.tag) + // This is either a word or the number of rows affected + .and_then(|i| atoi(&self.tag[(i + 1)..])) + .unwrap_or(0) + } +} + +#[test] +fn test_decode_command_complete_for_insert() { + const DATA: &[u8] = b"INSERT 0 1214\0"; + + let cc = CommandComplete::decode_body(Bytes::from_static(DATA)).unwrap(); + + assert_eq!(cc.rows_affected(), 1214); +} + +#[test] +fn test_decode_command_complete_for_begin() { + const DATA: &[u8] = b"BEGIN\0"; + + let cc = CommandComplete::decode_body(Bytes::from_static(DATA)).unwrap(); + + assert_eq!(cc.rows_affected(), 0); +} + +#[test] +fn test_decode_command_complete_for_update() { + const DATA: &[u8] = b"UPDATE 5\0"; + + let cc = CommandComplete::decode_body(Bytes::from_static(DATA)).unwrap(); + + assert_eq!(cc.rows_affected(), 5); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_command_complete(b: &mut test::Bencher) { + const DATA: &[u8] = b"INSERT 0 1214\0"; + + b.iter(|| { + let _ = CommandComplete::decode_body(test::black_box(Bytes::from_static(DATA))); + }); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_command_complete_rows_affected(b: &mut test::Bencher) { + const DATA: &[u8] = b"INSERT 0 1214\0"; + + let data = CommandComplete::decode_body(Bytes::from_static(DATA)).unwrap(); + + b.iter(|| { + let _rows = test::black_box(&data).rows_affected(); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/copy.rs b/src-tauri/vendor/sqlx-postgres/src/message/copy.rs new file mode 100644 index 00000000..837d849a --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/copy.rs @@ -0,0 +1,141 @@ +use crate::error::Result; +use crate::io::BufMutExt; +use crate::message::{ + BackendMessage, BackendMessageFormat, FrontendMessage, FrontendMessageFormat, +}; +use sqlx_core::bytes::{Buf, Bytes}; +use sqlx_core::Error; +use std::num::Saturating; +use std::ops::Deref; + +/// The same structure is sent for both `CopyInResponse` and `CopyOutResponse` +pub struct CopyResponseData { + pub format: i8, + pub num_columns: i16, + pub format_codes: Vec, +} + +pub struct CopyInResponse(pub CopyResponseData); + +#[allow(dead_code)] +pub struct CopyOutResponse(pub CopyResponseData); + +pub struct CopyData(pub B); + +pub struct CopyFail { + pub message: String, +} + +pub struct CopyDone; + +impl CopyResponseData { + #[inline] + fn decode(mut buf: Bytes) -> Result { + let format = buf.get_i8(); + let num_columns = buf.get_i16(); + + let format_codes = (0..num_columns).map(|_| buf.get_i16()).collect(); + + Ok(CopyResponseData { + format, + num_columns, + format_codes, + }) + } +} + +impl BackendMessage for CopyInResponse { + const FORMAT: BackendMessageFormat = BackendMessageFormat::CopyInResponse; + + #[inline(always)] + fn decode_body(buf: Bytes) -> std::result::Result { + Ok(Self(CopyResponseData::decode(buf)?)) + } +} + +impl BackendMessage for CopyOutResponse { + const FORMAT: BackendMessageFormat = BackendMessageFormat::CopyOutResponse; + + #[inline(always)] + fn decode_body(buf: Bytes) -> std::result::Result { + Ok(Self(CopyResponseData::decode(buf)?)) + } +} + +impl BackendMessage for CopyData { + const FORMAT: BackendMessageFormat = BackendMessageFormat::CopyData; + + #[inline(always)] + fn decode_body(buf: Bytes) -> std::result::Result { + Ok(Self(buf)) + } +} + +impl> FrontendMessage for CopyData { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::CopyData; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(self.0.len()) + } + + #[inline(always)] + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + buf.extend_from_slice(&self.0); + Ok(()) + } +} + +impl FrontendMessage for CopyFail { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::CopyFail; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(self.message.len()) + } + + #[inline(always)] + fn encode_body(&self, buf: &mut Vec) -> std::result::Result<(), Error> { + buf.put_str_nul(&self.message); + Ok(()) + } +} + +impl CopyFail { + #[inline(always)] + pub fn new(msg: impl Into) -> CopyFail { + CopyFail { + message: msg.into(), + } + } +} + +impl FrontendMessage for CopyDone { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::CopyDone; + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(0) + } + + #[inline(always)] + fn encode_body(&self, _buf: &mut Vec) -> std::result::Result<(), Error> { + Ok(()) + } +} + +impl BackendMessage for CopyDone { + const FORMAT: BackendMessageFormat = BackendMessageFormat::CopyDone; + + #[inline(always)] + fn decode_body(bytes: Bytes) -> std::result::Result { + if !bytes.is_empty() { + // Not fatal but may indicate a protocol change + tracing::debug!( + "Postgres backend returned non-empty message for CopyDone: \"{}\"", + bytes.escape_ascii() + ) + } + + Ok(CopyDone) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/data_row.rs b/src-tauri/vendor/sqlx-postgres/src/message/data_row.rs new file mode 100644 index 00000000..ae9d0d9b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/data_row.rs @@ -0,0 +1,138 @@ +use byteorder::{BigEndian, ByteOrder}; +use sqlx_core::bytes::Bytes; +use std::ops::Range; + +use crate::error::Error; +use crate::message::{BackendMessage, BackendMessageFormat}; + +/// A row of data from the database. +#[derive(Debug)] +pub struct DataRow { + pub(crate) storage: Bytes, + + /// Ranges into the stored row data. + /// This uses `u32` instead of usize to reduce the size of this type. Values cannot be larger + /// than `i32` in postgres. + pub(crate) values: Vec>>, +} + +impl DataRow { + #[inline] + pub(crate) fn get(&self, index: usize) -> Option<&'_ [u8]> { + self.values[index] + .as_ref() + .map(|col| &self.storage[(col.start as usize)..(col.end as usize)]) + } +} + +impl BackendMessage for DataRow { + const FORMAT: BackendMessageFormat = BackendMessageFormat::DataRow; + + fn decode_body(buf: Bytes) -> Result { + if buf.len() < 2 { + return Err(err_protocol!( + "expected at least 2 bytes, got {}", + buf.len() + )); + } + + let cnt = BigEndian::read_u16(&buf) as usize; + + let mut values = Vec::with_capacity(cnt); + let mut offset: u32 = 2; + + for _ in 0..cnt { + let value_start = offset + .checked_add(4) + .ok_or_else(|| err_protocol!("next value start out of range (offset: {offset})"))?; + + // widen both to a larger type for a safe comparison + if (buf.len() as u64) < (value_start as u64) { + return Err(err_protocol!( + "expected 4 bytes at offset {offset}, got {}", + (value_start as u64) - (buf.len() as u64) + )); + } + + // Length of the column value, in bytes (this count does not include itself). + // Can be zero. As a special case, -1 indicates a NULL column value. + // No value bytes follow in the NULL case. + // + // we know `offset` is within range of `buf.len()` from the above check + #[allow(clippy::cast_possible_truncation)] + let length = BigEndian::read_i32(&buf[(offset as usize)..]); + + if let Ok(length) = u32::try_from(length) { + let value_end = value_start.checked_add(length).ok_or_else(|| { + err_protocol!("value_start + length out of range ({offset} + {length})") + })?; + + values.push(Some(value_start..value_end)); + offset = value_end; + } else { + // Negative values signify NULL + values.push(None); + // `value_start` is actually the next value now. + offset = value_start; + } + } + + Ok(Self { + storage: buf, + values, + }) + } +} + +#[test] +fn test_decode_data_row() { + const DATA: &[u8] = b"\ + \x00\x08\ + \xff\xff\xff\xff\ + \x00\x00\x00\x04\ + \x00\x00\x00\n\ + \xff\xff\xff\xff\ + \x00\x00\x00\x04\ + \x00\x00\x00\x14\ + \xff\xff\xff\xff\ + \x00\x00\x00\x04\ + \x00\x00\x00(\ + \xff\xff\xff\xff\ + \x00\x00\x00\x04\ + \x00\x00\x00P"; + + let row = DataRow::decode_body(DATA.into()).unwrap(); + + assert_eq!(row.values.len(), 8); + + assert!(row.get(0).is_none()); + assert_eq!(row.get(1).unwrap(), &[0_u8, 0, 0, 10][..]); + assert!(row.get(2).is_none()); + assert_eq!(row.get(3).unwrap(), &[0_u8, 0, 0, 20][..]); + assert!(row.get(4).is_none()); + assert_eq!(row.get(5).unwrap(), &[0_u8, 0, 0, 40][..]); + assert!(row.get(6).is_none()); + assert_eq!(row.get(7).unwrap(), &[0_u8, 0, 0, 80][..]); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_data_row_get(b: &mut test::Bencher) { + const DATA: &[u8] = b"\x00\x08\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00\n\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00\x14\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00(\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00P"; + + let row = DataRow::decode_body(test::black_box(Bytes::from_static(DATA))).unwrap(); + + b.iter(|| { + let _value = test::black_box(&row).get(3); + }); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_data_row(b: &mut test::Bencher) { + const DATA: &[u8] = b"\x00\x08\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00\n\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00\x14\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00(\xff\xff\xff\xff\x00\x00\x00\x04\x00\x00\x00P"; + + b.iter(|| { + let _ = DataRow::decode_body(test::black_box(Bytes::from_static(DATA))); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/describe.rs b/src-tauri/vendor/sqlx-postgres/src/message/describe.rs new file mode 100644 index 00000000..d6ea7e89 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/describe.rs @@ -0,0 +1,103 @@ +use crate::io::{PgBufMutExt, PortalId, StatementId}; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +const DESCRIBE_PORTAL: u8 = b'P'; +const DESCRIBE_STATEMENT: u8 = b'S'; + +/// Note: will emit both a RowDescription and a ParameterDescription message +#[derive(Debug)] +#[allow(dead_code)] +pub enum Describe { + Statement(StatementId), + Portal(PortalId), +} + +impl FrontendMessage for Describe { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Describe; + + fn body_size_hint(&self) -> Saturating { + // Either `DESCRIBE_PORTAL` or `DESCRIBE_STATEMENT` + let mut size = Saturating(1); + + match self { + Describe::Statement(id) => size += id.name_len(), + Describe::Portal(id) => size += id.name_len(), + } + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + match self { + // #[likely] + Describe::Statement(id) => { + buf.push(DESCRIBE_STATEMENT); + buf.put_statement_name(*id); + } + + Describe::Portal(id) => { + buf.push(DESCRIBE_PORTAL); + buf.put_portal_name(*id); + } + } + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::message::FrontendMessage; + + use super::{Describe, PortalId, StatementId}; + + #[test] + fn test_encode_describe_portal() { + const EXPECTED: &[u8] = b"D\0\0\0\x17Psqlx_p_1234567890\0"; + + let mut buf = Vec::new(); + let m = Describe::Portal(PortalId::TEST_VAL); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[test] + fn test_encode_describe_unnamed_portal() { + const EXPECTED: &[u8] = b"D\0\0\0\x06P\0"; + + let mut buf = Vec::new(); + let m = Describe::Portal(PortalId::UNNAMED); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[test] + fn test_encode_describe_statement() { + const EXPECTED: &[u8] = b"D\0\0\0\x17Ssqlx_s_1234567890\0"; + + let mut buf = Vec::new(); + let m = Describe::Statement(StatementId::TEST_VAL); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[test] + fn test_encode_describe_unnamed_statement() { + const EXPECTED: &[u8] = b"D\0\0\0\x06S\0"; + + let mut buf = Vec::new(); + let m = Describe::Statement(StatementId::UNNAMED); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/execute.rs b/src-tauri/vendor/sqlx-postgres/src/message/execute.rs new file mode 100644 index 00000000..f82b7884 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/execute.rs @@ -0,0 +1,73 @@ +use std::num::Saturating; + +use sqlx_core::Error; + +use crate::io::{PgBufMutExt, PortalId}; +use crate::message::{FrontendMessage, FrontendMessageFormat}; + +pub struct Execute { + /// The id of the portal to execute. + pub portal: PortalId, + + /// Maximum number of rows to return, if portal contains a query + /// that returns rows (ignored otherwise). Zero denotes “no limit”. + pub limit: u32, +} + +impl FrontendMessage for Execute { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Execute; + + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + + size += self.portal.name_len(); + size += 2; // limit + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + buf.put_portal_name(self.portal); + buf.extend(&self.limit.to_be_bytes()); + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::io::PortalId; + use crate::message::FrontendMessage; + + use super::Execute; + + #[test] + fn test_encode_execute_named_portal() { + const EXPECTED: &[u8] = b"E\0\0\0\x1Asqlx_p_1234567890\0\0\0\0\x02"; + + let mut buf = Vec::new(); + let m = Execute { + portal: PortalId::TEST_VAL, + limit: 2, + }; + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[test] + fn test_encode_execute_unnamed_portal() { + const EXPECTED: &[u8] = b"E\0\0\0\x09\0\x49\x96\x02\xD2"; + + let mut buf = Vec::new(); + let m = Execute { + portal: PortalId::UNNAMED, + limit: 1234567890, + }; + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/flush.rs b/src-tauri/vendor/sqlx-postgres/src/message/flush.rs new file mode 100644 index 00000000..d1dfabbf --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/flush.rs @@ -0,0 +1,25 @@ +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +/// The Flush message does not cause any specific output to be generated, +/// but forces the backend to deliver any data pending in its output buffers. +/// +/// A Flush must be sent after any extended-query command except Sync, if the +/// frontend wishes to examine the results of that command before issuing more commands. +#[derive(Debug)] +pub struct Flush; + +impl FrontendMessage for Flush { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Flush; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(0) + } + + #[inline(always)] + fn encode_body(&self, _buf: &mut Vec) -> Result<(), Error> { + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/mod.rs b/src-tauri/vendor/sqlx-postgres/src/message/mod.rs new file mode 100644 index 00000000..e62f9beb --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/mod.rs @@ -0,0 +1,229 @@ +use sqlx_core::bytes::Bytes; +use std::num::Saturating; + +use crate::error::Error; +use crate::io::PgBufMutExt; + +mod authentication; +mod backend_key_data; +mod bind; +mod close; +mod command_complete; +mod copy; +mod data_row; +mod describe; +mod execute; +mod flush; +mod notification; +mod parameter_description; +mod parameter_status; +mod parse; +mod parse_complete; +mod password; +mod query; +mod ready_for_query; +mod response; +mod row_description; +mod sasl; +mod ssl_request; +mod startup; +mod sync; +mod terminate; + +pub use authentication::{Authentication, AuthenticationSasl}; +pub use backend_key_data::BackendKeyData; +pub use bind::Bind; +pub use close::Close; +pub use command_complete::CommandComplete; +pub use copy::{CopyData, CopyDone, CopyFail, CopyInResponse, CopyOutResponse, CopyResponseData}; +pub use data_row::DataRow; +pub use describe::Describe; +pub use execute::Execute; +#[allow(unused_imports)] +pub use flush::Flush; +pub use notification::Notification; +pub use parameter_description::ParameterDescription; +pub use parameter_status::ParameterStatus; +pub use parse::Parse; +pub use parse_complete::ParseComplete; +pub use password::Password; +pub use query::Query; +pub use ready_for_query::{ReadyForQuery, TransactionStatus}; +pub use response::{Notice, PgSeverity}; +pub use row_description::RowDescription; +pub use sasl::{SaslInitialResponse, SaslResponse}; +use sqlx_core::io::ProtocolEncode; +pub use ssl_request::SslRequest; +pub use startup::Startup; +pub use sync::Sync; +pub use terminate::Terminate; + +// Note: we can't use the same enum for both frontend and backend message formats +// because there are duplicated format codes between them. +// +// For example, `Close` (frontend) and `CommandComplete` (backend) both use format code `C`. +// +#[derive(Debug, PartialOrd, PartialEq)] +#[repr(u8)] +pub enum FrontendMessageFormat { + Bind = b'B', + Close = b'C', + CopyData = b'd', + CopyDone = b'c', + CopyFail = b'f', + Describe = b'D', + Execute = b'E', + Flush = b'H', + Parse = b'P', + /// This message format is polymorphic. It's used for: + /// + /// * Plain password responses + /// * MD5 password responses + /// * SASL responses + /// * GSSAPI/SSPI responses + PasswordPolymorphic = b'p', + Query = b'Q', + Sync = b'S', + Terminate = b'X', +} + +#[derive(Debug, PartialOrd, PartialEq)] +#[repr(u8)] +pub enum BackendMessageFormat { + Authentication, + BackendKeyData, + BindComplete, + CloseComplete, + CommandComplete, + CopyData, + CopyDone, + CopyInResponse, + CopyOutResponse, + DataRow, + EmptyQueryResponse, + ErrorResponse, + NoData, + NoticeResponse, + NotificationResponse, + ParameterDescription, + ParameterStatus, + ParseComplete, + PortalSuspended, + ReadyForQuery, + RowDescription, +} + +#[derive(Debug)] +pub struct ReceivedMessage { + pub format: BackendMessageFormat, + pub contents: Bytes, +} + +impl ReceivedMessage { + #[inline] + pub fn decode(self) -> Result + where + T: BackendMessage, + { + if T::FORMAT != self.format { + return Err(err_protocol!( + "Postgres protocol error: expected {:?}, got {:?}", + T::FORMAT, + self.format + )); + } + + T::decode_body(self.contents).map_err(|e| match e { + Error::Protocol(s) => { + err_protocol!("Postgres protocol error (reading {:?}): {s}", self.format) + } + other => other, + }) + } +} + +impl BackendMessageFormat { + pub fn try_from_u8(v: u8) -> Result { + // https://www.postgresql.org/docs/current/protocol-message-formats.html + + Ok(match v { + b'1' => BackendMessageFormat::ParseComplete, + b'2' => BackendMessageFormat::BindComplete, + b'3' => BackendMessageFormat::CloseComplete, + b'C' => BackendMessageFormat::CommandComplete, + b'd' => BackendMessageFormat::CopyData, + b'c' => BackendMessageFormat::CopyDone, + b'G' => BackendMessageFormat::CopyInResponse, + b'H' => BackendMessageFormat::CopyOutResponse, + b'D' => BackendMessageFormat::DataRow, + b'E' => BackendMessageFormat::ErrorResponse, + b'I' => BackendMessageFormat::EmptyQueryResponse, + b'A' => BackendMessageFormat::NotificationResponse, + b'K' => BackendMessageFormat::BackendKeyData, + b'N' => BackendMessageFormat::NoticeResponse, + b'R' => BackendMessageFormat::Authentication, + b'S' => BackendMessageFormat::ParameterStatus, + b'T' => BackendMessageFormat::RowDescription, + b'Z' => BackendMessageFormat::ReadyForQuery, + b'n' => BackendMessageFormat::NoData, + b's' => BackendMessageFormat::PortalSuspended, + b't' => BackendMessageFormat::ParameterDescription, + + _ => return Err(err_protocol!("unknown message type: {:?}", v as char)), + }) + } +} + +pub(crate) trait FrontendMessage: Sized { + /// The format prefix of this message. + const FORMAT: FrontendMessageFormat; + + /// Return the amount of space, in bytes, to reserve in the buffer passed to [`Self::encode_body()`]. + fn body_size_hint(&self) -> Saturating; + + /// Encode this type as a Frontend message in the Postgres protocol. + /// + /// The implementation should *not* include `Self::FORMAT` or the length prefix. + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error>; + + #[inline(always)] + #[cfg_attr(not(test), allow(dead_code))] + fn encode_msg(self, buf: &mut Vec) -> Result<(), Error> { + EncodeMessage(self).encode(buf) + } +} + +pub(crate) trait BackendMessage: Sized { + /// The expected message format. + /// + /// + const FORMAT: BackendMessageFormat; + + /// Decode this type from a Backend message in the Postgres protocol. + /// + /// The format code and length prefix have already been read and are not at the start of `bytes`. + fn decode_body(buf: Bytes) -> Result; +} + +pub struct EncodeMessage(pub F); + +impl ProtocolEncode<'_, ()> for EncodeMessage { + fn encode_with(&self, buf: &mut Vec, _context: ()) -> Result<(), Error> { + let mut size_hint = self.0.body_size_hint(); + // plus format code and length prefix + size_hint += 5; + + // don't panic if `size_hint` is ridiculous + buf.try_reserve(size_hint.0).map_err(|e| { + err_protocol!( + "Postgres protocol: error allocating {} bytes for encoding message {:?}: {e}", + size_hint.0, + F::FORMAT, + ) + })?; + + buf.push(F::FORMAT as u8); + + buf.put_length_prefixed(|buf| self.0.encode_body(buf)) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/notification.rs b/src-tauri/vendor/sqlx-postgres/src/message/notification.rs new file mode 100644 index 00000000..7bf02983 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/notification.rs @@ -0,0 +1,39 @@ +use sqlx_core::bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::BufExt; +use crate::message::{BackendMessage, BackendMessageFormat}; + +#[derive(Debug)] +pub struct Notification { + pub(crate) process_id: u32, + pub(crate) channel: Bytes, + pub(crate) payload: Bytes, +} + +impl BackendMessage for Notification { + const FORMAT: BackendMessageFormat = BackendMessageFormat::NotificationResponse; + + fn decode_body(mut buf: Bytes) -> Result { + let process_id = buf.get_u32(); + let channel = buf.get_bytes_nul()?; + let payload = buf.get_bytes_nul()?; + + Ok(Self { + process_id, + channel, + payload, + }) + } +} + +#[test] +fn test_decode_notification_response() { + const NOTIFICATION_RESPONSE: &[u8] = b"\x34\x20\x10\x02TEST-CHANNEL\0THIS IS A TEST\0"; + + let message = Notification::decode_body(Bytes::from(NOTIFICATION_RESPONSE)).unwrap(); + + assert_eq!(message.process_id, 0x34201002); + assert_eq!(&*message.channel, &b"TEST-CHANNEL"[..]); + assert_eq!(&*message.payload, &b"THIS IS A TEST"[..]); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/parameter_description.rs b/src-tauri/vendor/sqlx-postgres/src/message/parameter_description.rs new file mode 100644 index 00000000..f0b25ccc --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/parameter_description.rs @@ -0,0 +1,58 @@ +use smallvec::SmallVec; +use sqlx_core::bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::message::{BackendMessage, BackendMessageFormat}; +use crate::types::Oid; + +#[derive(Debug)] +pub struct ParameterDescription { + pub types: SmallVec<[Oid; 6]>, +} + +impl BackendMessage for ParameterDescription { + const FORMAT: BackendMessageFormat = BackendMessageFormat::ParameterDescription; + + fn decode_body(mut buf: Bytes) -> Result { + // Note: this is correct, max parameters is 65535, not 32767 + // https://github.com/launchbadge/sqlx/issues/3464 + let cnt = buf.get_u16(); + let mut types = SmallVec::with_capacity(cnt as usize); + + for _ in 0..cnt { + types.push(Oid(buf.get_u32())); + } + + Ok(Self { types }) + } +} + +#[test] +fn test_decode_parameter_description() { + const DATA: &[u8] = b"\x00\x02\x00\x00\x00\x00\x00\x00\x05\x00"; + + let m = ParameterDescription::decode_body(DATA.into()).unwrap(); + + assert_eq!(m.types.len(), 2); + assert_eq!(m.types[0], Oid(0x0000_0000)); + assert_eq!(m.types[1], Oid(0x0000_0500)); +} + +#[test] +fn test_decode_empty_parameter_description() { + const DATA: &[u8] = b"\x00\x00"; + + let m = ParameterDescription::decode_body(DATA.into()).unwrap(); + + assert!(m.types.is_empty()); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_parameter_description(b: &mut test::Bencher) { + const DATA: &[u8] = b"\x00\x02\x00\x00\x00\x00\x00\x00\x05\x00"; + + b.iter(|| { + ParameterDescription::decode_body(test::black_box(Bytes::from_static(DATA))).unwrap(); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/parameter_status.rs b/src-tauri/vendor/sqlx-postgres/src/message/parameter_status.rs new file mode 100644 index 00000000..d979d189 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/parameter_status.rs @@ -0,0 +1,65 @@ +use sqlx_core::bytes::Bytes; + +use crate::error::Error; +use crate::io::BufExt; +use crate::message::{BackendMessage, BackendMessageFormat}; + +#[derive(Debug)] +pub struct ParameterStatus { + pub name: String, + pub value: String, +} + +impl BackendMessage for ParameterStatus { + const FORMAT: BackendMessageFormat = BackendMessageFormat::ParameterStatus; + + fn decode_body(mut buf: Bytes) -> Result { + let name = buf.get_str_nul()?; + let value = buf.get_str_nul()?; + + Ok(Self { name, value }) + } +} + +#[test] +fn test_decode_parameter_status() { + const DATA: &[u8] = b"client_encoding\x00UTF8\x00"; + + let m = ParameterStatus::decode_body(DATA.into()).unwrap(); + + assert_eq!(&m.name, "client_encoding"); + assert_eq!(&m.value, "UTF8") +} + +#[test] +fn test_decode_empty_parameter_status() { + const DATA: &[u8] = b"\x00\x00"; + + let m = ParameterStatus::decode_body(DATA.into()).unwrap(); + + assert!(m.name.is_empty()); + assert!(m.value.is_empty()); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_parameter_status(b: &mut test::Bencher) { + const DATA: &[u8] = b"client_encoding\x00UTF8\x00"; + + b.iter(|| { + ParameterStatus::decode_body(test::black_box(Bytes::from_static(DATA))).unwrap(); + }); +} + +#[test] +fn test_decode_parameter_status_response() { + const PARAMETER_STATUS_RESPONSE: &[u8] = b"crdb_version\0CockroachDB CCL v21.1.0 (x86_64-unknown-linux-gnu, built 2021/05/17 13:49:40, go1.15.11)\0"; + + let message = ParameterStatus::decode_body(Bytes::from(PARAMETER_STATUS_RESPONSE)).unwrap(); + + assert_eq!(message.name, "crdb_version"); + assert_eq!( + message.value, + "CockroachDB CCL v21.1.0 (x86_64-unknown-linux-gnu, built 2021/05/17 13:49:40, go1.15.11)" + ); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/parse.rs b/src-tauri/vendor/sqlx-postgres/src/message/parse.rs new file mode 100644 index 00000000..62f57a1c --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/parse.rs @@ -0,0 +1,95 @@ +use crate::io::BufMutExt; +use crate::io::{PgBufMutExt, StatementId}; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use crate::types::Oid; +use sqlx_core::Error; +use std::num::Saturating; + +#[derive(Debug)] +pub struct Parse<'a> { + /// The ID of the destination prepared statement. + pub statement: StatementId, + + /// The query string to be parsed. + pub query: &'a str, + + /// The parameter data types specified (could be zero). Note that this is not an + /// indication of the number of parameters that might appear in the query string, + /// only the number that the frontend wants to pre-specify types for. + pub param_types: &'a [Oid], +} + +impl FrontendMessage for Parse<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Parse; + + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + + size += self.statement.name_len(); + + size += self.query.len(); + size += 1; // NUL terminator + + size += 2; // param_types_len + + // `param_types` + size += self.param_types.len().saturating_mul(4); + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + buf.put_statement_name(self.statement); + + buf.put_str_nul(self.query); + + // Note: actually interpreted as unsigned + // https://github.com/launchbadge/sqlx/issues/3464 + let param_types_len = u16::try_from(self.param_types.len()).map_err(|_| { + err_protocol!( + "param_types.len() too large for binary protocol: {}", + self.param_types.len() + ) + })?; + + buf.extend(param_types_len.to_be_bytes()); + + for &oid in self.param_types { + buf.extend(oid.0.to_be_bytes()); + } + + Ok(()) + } +} + +#[test] +fn test_encode_parse() { + const EXPECTED: &[u8] = b"P\0\0\0\x26sqlx_s_1234567890\0SELECT $1\0\0\x01\0\0\0\x19"; + + let mut buf = Vec::new(); + let m = Parse { + statement: StatementId::TEST_VAL, + query: "SELECT $1", + param_types: &[Oid(25)], + }; + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); +} + +#[test] +fn test_encode_parse_unnamed_statement() { + const EXPECTED: &[u8] = b"P\0\0\0\x15\0SELECT $1\0\0\x01\0\0\0\x19"; + + let mut buf = Vec::new(); + let m = Parse { + statement: StatementId::UNNAMED, + query: "SELECT $1", + param_types: &[Oid(25)], + }; + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/parse_complete.rs b/src-tauri/vendor/sqlx-postgres/src/message/parse_complete.rs new file mode 100644 index 00000000..3051f5ff --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/parse_complete.rs @@ -0,0 +1,13 @@ +use crate::message::{BackendMessage, BackendMessageFormat}; +use sqlx_core::bytes::Bytes; +use sqlx_core::Error; + +pub struct ParseComplete; + +impl BackendMessage for ParseComplete { + const FORMAT: BackendMessageFormat = BackendMessageFormat::ParseComplete; + + fn decode_body(_bytes: Bytes) -> Result { + Ok(ParseComplete) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/password.rs b/src-tauri/vendor/sqlx-postgres/src/message/password.rs new file mode 100644 index 00000000..4eaaeb15 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/password.rs @@ -0,0 +1,153 @@ +use crate::io::BufMutExt; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use md5::{Digest, Md5}; +use sqlx_core::Error; +use std::fmt::Write; +use std::num::Saturating; + +#[derive(Debug)] +pub enum Password<'a> { + Cleartext(&'a str), + + Md5 { + password: &'a str, + username: &'a str, + salt: [u8; 4], + }, +} + +impl FrontendMessage for Password<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::PasswordPolymorphic; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + + match self { + Password::Cleartext(password) => { + // To avoid reporting the exact password length anywhere, + // we deliberately give a bad estimate. + // + // This shouldn't affect performance in the long run. + size += password + .len() + .saturating_add(1) // NUL terminator + .checked_next_power_of_two() + .unwrap_or(usize::MAX); + } + Password::Md5 { .. } => { + // "md5<32 hex chars>\0" + size += 36; + } + } + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + match self { + Password::Cleartext(password) => { + buf.put_str_nul(password); + } + + Password::Md5 { + username, + password, + salt, + } => { + // The actual `PasswordMessage` can be computed in SQL as + // `concat('md5', md5(concat(md5(concat(password, username)), random-salt)))`. + + // Keep in mind the md5() function returns its result as a hex string. + + let mut hasher = Md5::new(); + + hasher.update(password); + hasher.update(username); + + let mut output = String::with_capacity(35); + + let _ = write!(output, "{:x}", hasher.finalize_reset()); + + hasher.update(&output); + hasher.update(salt); + + output.clear(); + + let _ = write!(output, "md5{:x}", hasher.finalize()); + + buf.put_str_nul(&output); + } + } + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use crate::message::FrontendMessage; + + use super::Password; + + #[test] + fn test_encode_clear_password() { + const EXPECTED: &[u8] = b"p\0\0\0\rpassword\0"; + + let mut buf = Vec::new(); + let m = Password::Cleartext("password"); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[test] + fn test_encode_md5_password() { + const EXPECTED: &[u8] = b"p\0\0\0(md53e2c9d99d49b201ef867a36f3f9ed62c\0"; + + let mut buf = Vec::new(); + let m = Password::Md5 { + password: "password", + username: "root", + salt: [147, 24, 57, 152], + }; + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); + } + + #[cfg(all(test, not(debug_assertions)))] + #[bench] + fn bench_encode_clear_password(b: &mut test::Bencher) { + use test::black_box; + + let mut buf = Vec::with_capacity(128); + + b.iter(|| { + buf.clear(); + + black_box(Password::Cleartext("password")).encode_msg(&mut buf); + }); + } + + #[cfg(all(test, not(debug_assertions)))] + #[bench] + fn bench_encode_md5_password(b: &mut test::Bencher) { + use test::black_box; + + let mut buf = Vec::with_capacity(128); + + b.iter(|| { + buf.clear(); + + black_box(Password::Md5 { + password: "password", + username: "root", + salt: [147, 24, 57, 152], + }) + .encode_msg(&mut buf); + }); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/query.rs b/src-tauri/vendor/sqlx-postgres/src/message/query.rs new file mode 100644 index 00000000..788d7808 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/query.rs @@ -0,0 +1,37 @@ +use crate::io::BufMutExt; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +#[derive(Debug)] +pub struct Query<'a>(pub &'a str); + +impl FrontendMessage for Query<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Query; + + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + + size += self.0.len(); + size += 1; // NUL terminator + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + buf.put_str_nul(self.0); + Ok(()) + } +} + +#[test] +fn test_encode_query() { + const EXPECTED: &[u8] = b"Q\0\0\0\x0DSELECT 1\0"; + + let mut buf = Vec::new(); + let m = Query("SELECT 1"); + + m.encode_msg(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/ready_for_query.rs b/src-tauri/vendor/sqlx-postgres/src/message/ready_for_query.rs new file mode 100644 index 00000000..a1f6761b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/ready_for_query.rs @@ -0,0 +1,56 @@ +use sqlx_core::bytes::Bytes; + +use crate::error::Error; +use crate::message::{BackendMessage, BackendMessageFormat}; + +#[derive(Debug)] +#[repr(u8)] +pub enum TransactionStatus { + /// Not in a transaction block. + Idle = b'I', + + /// In a transaction block. + Transaction = b'T', + + /// In a _failed_ transaction block. Queries will be rejected until block is ended. + Error = b'E', +} + +#[derive(Debug)] +pub struct ReadyForQuery { + pub transaction_status: TransactionStatus, +} + +impl BackendMessage for ReadyForQuery { + const FORMAT: BackendMessageFormat = BackendMessageFormat::ReadyForQuery; + + fn decode_body(buf: Bytes) -> Result { + let status = match buf[0] { + b'I' => TransactionStatus::Idle, + b'T' => TransactionStatus::Transaction, + b'E' => TransactionStatus::Error, + + status => { + return Err(err_protocol!( + "unknown transaction status: {:?}", + status as char + )); + } + }; + + Ok(Self { + transaction_status: status, + }) + } +} + +#[test] +fn test_decode_ready_for_query() -> Result<(), Error> { + const DATA: &[u8] = b"E"; + + let m = ReadyForQuery::decode_body(Bytes::from_static(DATA))?; + + assert!(matches!(m.transaction_status, TransactionStatus::Error)); + + Ok(()) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/response.rs b/src-tauri/vendor/sqlx-postgres/src/message/response.rs new file mode 100644 index 00000000..d6e43e08 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/response.rs @@ -0,0 +1,272 @@ +use std::ops::Range; +use std::str::from_utf8; + +use memchr::memchr; + +use sqlx_core::bytes::Bytes; + +use crate::error::Error; +use crate::io::ProtocolDecode; +use crate::message::{BackendMessage, BackendMessageFormat}; + +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +#[repr(u8)] +pub enum PgSeverity { + Panic, + Fatal, + Error, + Warning, + Notice, + Debug, + Info, + Log, +} + +impl PgSeverity { + #[inline] + pub fn is_error(self) -> bool { + matches!(self, Self::Panic | Self::Fatal | Self::Error) + } +} + +impl TryFrom<&str> for PgSeverity { + type Error = Error; + + fn try_from(s: &str) -> Result { + let result = match s { + "PANIC" => PgSeverity::Panic, + "FATAL" => PgSeverity::Fatal, + "ERROR" => PgSeverity::Error, + "WARNING" => PgSeverity::Warning, + "NOTICE" => PgSeverity::Notice, + "DEBUG" => PgSeverity::Debug, + "INFO" => PgSeverity::Info, + "LOG" => PgSeverity::Log, + + severity => { + return Err(err_protocol!("unknown severity: {:?}", severity)); + } + }; + + Ok(result) + } +} + +#[derive(Debug)] +pub struct Notice { + storage: Bytes, + severity: PgSeverity, + message: Range, + code: Range, +} + +impl Notice { + #[inline] + pub fn severity(&self) -> PgSeverity { + self.severity + } + + #[inline] + pub fn code(&self) -> &str { + self.get_cached_str(self.code.clone()) + } + + #[inline] + pub fn message(&self) -> &str { + self.get_cached_str(self.message.clone()) + } + + // Field descriptions available here: + // https://www.postgresql.org/docs/current/protocol-error-fields.html + + #[inline] + pub fn get(&self, ty: u8) -> Option<&str> { + self.get_raw(ty).and_then(|v| from_utf8(v).ok()) + } + + pub fn get_raw(&self, ty: u8) -> Option<&[u8]> { + self.fields() + .filter(|(field, _)| *field == ty) + .map(|(_, range)| &self.storage[range]) + .next() + } +} + +impl Notice { + #[inline] + fn fields(&self) -> Fields<'_> { + Fields { + storage: &self.storage, + offset: 0, + } + } + + #[inline] + fn get_cached_str(&self, cache: Range) -> &str { + // unwrap: this cannot fail at this stage + from_utf8(&self.storage[cache]).unwrap() + } +} + +impl ProtocolDecode<'_> for Notice { + fn decode_with(buf: Bytes, _: ()) -> Result { + // In order to support PostgreSQL 9.5 and older we need to parse the localized S field. + // Newer versions additionally come with the V field that is guaranteed to be in English. + // We thus read both versions and prefer the unlocalized one if available. + const DEFAULT_SEVERITY: PgSeverity = PgSeverity::Log; + let mut severity_v = None; + let mut severity_s = None; + let mut message = 0..0; + let mut code = 0..0; + + // we cache the three always present fields + // this enables to keep the access time down for the fields most likely accessed + + let fields = Fields { + storage: &buf, + offset: 0, + }; + + for (field, v) in fields { + if !(message.is_empty() || code.is_empty()) { + // stop iterating when we have the 3 fields we were looking for + // we assume V (severity) was the first field as it should be + break; + } + + match field { + b'S' => { + severity_s = from_utf8(&buf[v.clone()]) + // If the error string is not UTF-8, we have no hope of interpreting it, + // localized or not. The `V` field would likely fail to parse as well. + .map_err(|_| notice_protocol_err())? + .try_into() + // If we couldn't parse the severity here, it might just be localized. + .ok(); + } + + b'V' => { + // Propagate errors here, because V is not localized and + // thus we are missing a possible variant. + severity_v = Some( + from_utf8(&buf[v.clone()]) + .map_err(|_| notice_protocol_err())? + .try_into()?, + ); + } + + b'M' => { + _ = from_utf8(&buf[v.clone()]).map_err(|_| notice_protocol_err())?; + message = v; + } + + b'C' => { + _ = from_utf8(&buf[v.clone()]).map_err(|_| notice_protocol_err())?; + code = v; + } + + // If more fields are added, make sure to check that they are valid UTF-8, + // otherwise the get_cached_str method will panic. + _ => {} + } + } + + Ok(Self { + severity: severity_v.or(severity_s).unwrap_or(DEFAULT_SEVERITY), + message, + code, + storage: buf, + }) + } +} + +impl BackendMessage for Notice { + const FORMAT: BackendMessageFormat = BackendMessageFormat::NoticeResponse; + + fn decode_body(buf: Bytes) -> Result { + // Keeping both impls for now + Self::decode_with(buf, ()) + } +} + +/// An iterator over each field in the Error (or Notice) response. +struct Fields<'a> { + storage: &'a [u8], + offset: usize, +} + +impl<'a> Iterator for Fields<'a> { + type Item = (u8, Range); + + fn next(&mut self) -> Option { + // The fields in the response body are sequentially stored as [tag][string], + // ending in a final, additional [nul] + + let ty = *self.storage.get(self.offset)?; + + if ty == 0 { + return None; + } + + // Consume the type byte + self.offset = self.offset.checked_add(1)?; + + let start = self.offset; + + let len = memchr(b'\0', self.storage.get(start..)?)?; + + // Neither can overflow as they will always be `<= self.storage.len()`. + let end = self.offset + len; + self.offset = end + 1; + + Some((ty, start..end)) + } +} + +fn notice_protocol_err() -> Error { + // https://github.com/launchbadge/sqlx/issues/1144 + Error::Protocol( + "Postgres returned a non-UTF-8 string for its error message. \ + This is most likely due to an error that occurred during authentication and \ + the default lc_messages locale is not binary-compatible with UTF-8. \ + See the server logs for the error details." + .into(), + ) +} + +#[test] +fn test_decode_error_response() { + const DATA: &[u8] = b"SNOTICE\0VNOTICE\0C42710\0Mextension \"uuid-ossp\" already exists, skipping\0Fextension.c\0L1656\0RCreateExtension\0\0"; + + let m = Notice::decode(Bytes::from_static(DATA)).unwrap(); + + assert_eq!( + m.message(), + "extension \"uuid-ossp\" already exists, skipping" + ); + + assert!(matches!(m.severity(), PgSeverity::Notice)); + assert_eq!(m.code(), "42710"); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_error_response_get_message(b: &mut test::Bencher) { + const DATA: &[u8] = b"SNOTICE\0VNOTICE\0C42710\0Mextension \"uuid-ossp\" already exists, skipping\0Fextension.c\0L1656\0RCreateExtension\0\0"; + + let res = Notice::decode(test::black_box(Bytes::from_static(DATA))).unwrap(); + + b.iter(|| { + let _ = test::black_box(&res).message(); + }); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_decode_error_response(b: &mut test::Bencher) { + const DATA: &[u8] = b"SNOTICE\0VNOTICE\0C42710\0Mextension \"uuid-ossp\" already exists, skipping\0Fextension.c\0L1656\0RCreateExtension\0\0"; + + b.iter(|| { + let _ = Notice::decode(test::black_box(Bytes::from_static(DATA))); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/row_description.rs b/src-tauri/vendor/sqlx-postgres/src/message/row_description.rs new file mode 100644 index 00000000..668f31ed --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/row_description.rs @@ -0,0 +1,99 @@ +use sqlx_core::bytes::{Buf, Bytes}; + +use crate::error::Error; +use crate::io::BufExt; +use crate::message::{BackendMessage, BackendMessageFormat}; +use crate::types::Oid; + +#[derive(Debug)] +pub struct RowDescription { + pub fields: Vec, +} + +#[derive(Debug)] +pub struct Field { + /// The name of the field. + pub name: String, + + /// If the field can be identified as a column of a specific table, the + /// object ID of the table; otherwise zero. + pub relation_id: Option, + + /// If the field can be identified as a column of a specific table, the attribute number of + /// the column; otherwise zero. + pub relation_attribute_no: Option, + + /// The object ID of the field's data type. + pub data_type_id: Oid, + + /// The data type size (see pg_type.typlen). Note that negative values denote + /// variable-width types. + #[allow(dead_code)] + pub data_type_size: i16, + + /// The type modifier (see pg_attribute.atttypmod). The meaning of the + /// modifier is type-specific. + #[allow(dead_code)] + pub type_modifier: i32, + + /// The format code being used for the field. + #[allow(dead_code)] + pub format: i16, +} + +impl BackendMessage for RowDescription { + const FORMAT: BackendMessageFormat = BackendMessageFormat::RowDescription; + + fn decode_body(mut buf: Bytes) -> Result { + if buf.len() < 2 { + return Err(err_protocol!( + "expected at least 2 bytes, got {}", + buf.len() + )); + } + + let cnt = buf.get_u16(); + let mut fields = Vec::with_capacity(cnt as usize); + + for _ in 0..cnt { + let name = buf.get_str_nul()?.to_owned(); + + if buf.len() < 18 { + return Err(err_protocol!( + "expected at least 18 bytes after field name {name:?}, got {}", + buf.len() + )); + } + + let relation_id = buf.get_u32(); + let relation_attribute_no = buf.get_i16(); + let data_type_id = Oid(buf.get_u32()); + let data_type_size = buf.get_i16(); + let type_modifier = buf.get_i32(); + let format = buf.get_i16(); + + fields.push(Field { + name, + relation_id: if relation_id == 0 { + None + } else { + Some(Oid(relation_id)) + }, + relation_attribute_no: if relation_attribute_no == 0 { + None + } else { + Some(relation_attribute_no) + }, + data_type_id, + data_type_size, + type_modifier, + format, + }) + } + + Ok(Self { fields }) + } +} + +// TODO: Unit Test RowDescription +// TODO: Benchmark RowDescription diff --git a/src-tauri/vendor/sqlx-postgres/src/message/sasl.rs b/src-tauri/vendor/sqlx-postgres/src/message/sasl.rs new file mode 100644 index 00000000..9d393189 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/sasl.rs @@ -0,0 +1,69 @@ +use crate::io::BufMutExt; +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +pub struct SaslInitialResponse<'a> { + pub response: &'a str, + pub plus: bool, +} + +impl SaslInitialResponse<'_> { + #[inline(always)] + fn selected_mechanism(&self) -> &'static str { + if self.plus { + "SCRAM-SHA-256-PLUS" + } else { + "SCRAM-SHA-256" + } + } +} + +impl FrontendMessage for SaslInitialResponse<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::PasswordPolymorphic; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + let mut size = Saturating(0); + + size += self.selected_mechanism().len(); + size += 1; // NUL terminator + + size += 4; // response_len + size += self.response.len(); + + size + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + // name of the SASL authentication mechanism that the client selected + buf.put_str_nul(self.selected_mechanism()); + + let response_len = i32::try_from(self.response.len()).map_err(|_| { + err_protocol!( + "SASL Initial Response length too long for protcol: {}", + self.response.len() + ) + })?; + + buf.extend_from_slice(&response_len.to_be_bytes()); + buf.extend_from_slice(self.response.as_bytes()); + + Ok(()) + } +} + +pub struct SaslResponse<'a>(pub &'a str); + +impl FrontendMessage for SaslResponse<'_> { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::PasswordPolymorphic; + + fn body_size_hint(&self) -> Saturating { + Saturating(self.0.len()) + } + + fn encode_body(&self, buf: &mut Vec) -> Result<(), Error> { + buf.extend(self.0.as_bytes()); + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/ssl_request.rs b/src-tauri/vendor/sqlx-postgres/src/message/ssl_request.rs new file mode 100644 index 00000000..09c88622 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/ssl_request.rs @@ -0,0 +1,38 @@ +use crate::io::ProtocolEncode; + +pub struct SslRequest; + +impl SslRequest { + // https://www.postgresql.org/docs/current/protocol-message-formats.html#PROTOCOL-MESSAGE-FORMATS-SSLREQUEST + pub const BYTES: &'static [u8] = b"\x00\x00\x00\x08\x04\xd2\x16\x2f"; +} + +// Cannot impl FrontendMessage because it does not have a format code +impl ProtocolEncode<'_> for SslRequest { + #[inline(always)] + fn encode_with(&self, buf: &mut Vec, _context: ()) -> Result<(), crate::Error> { + buf.extend_from_slice(Self::BYTES); + Ok(()) + } +} + +#[test] +fn test_encode_ssl_request() { + let mut buf = Vec::new(); + + // Int32(8) + // Length of message contents in bytes, including self. + buf.extend_from_slice(&8_u32.to_be_bytes()); + + // Int32(80877103) + // The SSL request code. The value is chosen to contain 1234 in the most significant 16 bits, + // and 5679 in the least significant 16 bits. + // (To avoid confusion, this code must not be the same as any protocol version number.) + buf.extend_from_slice(&(((1234 << 16) | 5679) as u32).to_be_bytes()); + + let mut encoded = Vec::new(); + SslRequest.encode(&mut encoded).unwrap(); + + assert_eq!(buf, SslRequest::BYTES); + assert_eq!(buf, encoded); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/startup.rs b/src-tauri/vendor/sqlx-postgres/src/message/startup.rs new file mode 100644 index 00000000..1c6d735a --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/startup.rs @@ -0,0 +1,96 @@ +use crate::io::PgBufMutExt; +use crate::io::{BufMutExt, ProtocolEncode}; + +// To begin a session, a frontend opens a connection to the server and sends a startup message. +// This message includes the names of the user and of the database the user wants to connect to; +// it also identifies the particular protocol version to be used. + +// Optionally, the startup message can include additional settings for run-time parameters. + +pub struct Startup<'a> { + /// The database user name to connect as. Required; there is no default. + pub username: Option<&'a str>, + + /// The database to connect to. Defaults to the user name. + pub database: Option<&'a str>, + + /// Additional start-up params. + /// + pub params: &'a [(&'a str, &'a str)], +} + +// Startup cannot impl FrontendMessage because it doesn't have a format code. +impl ProtocolEncode<'_> for Startup<'_> { + fn encode_with(&self, buf: &mut Vec, _context: ()) -> Result<(), crate::Error> { + buf.reserve(120); + + buf.put_length_prefixed(|buf| { + // The protocol version number. The most significant 16 bits are the + // major version number (3 for the protocol described here). The least + // significant 16 bits are the minor version number (0 + // for the protocol described here) + buf.extend(&196_608_i32.to_be_bytes()); + + if let Some(username) = self.username { + // The database user name to connect as. + encode_startup_param(buf, "user", username); + } + + if let Some(database) = self.database { + // The database to connect to. Defaults to the user name. + encode_startup_param(buf, "database", database); + } + + for (name, value) in self.params { + encode_startup_param(buf, name, value); + } + + // A zero byte is required as a terminator + // after the last name/value pair. + buf.push(0); + + Ok(()) + }) + } +} + +#[inline] +fn encode_startup_param(buf: &mut Vec, name: &str, value: &str) { + buf.put_str_nul(name); + buf.put_str_nul(value); +} + +#[test] +fn test_encode_startup() { + const EXPECTED: &[u8] = b"\0\0\0)\0\x03\0\0user\0postgres\0database\0postgres\0\0"; + + let mut buf = Vec::new(); + let m = Startup { + username: Some("postgres"), + database: Some("postgres"), + params: &[], + }; + + m.encode(&mut buf).unwrap(); + + assert_eq!(buf, EXPECTED); +} + +#[cfg(all(test, not(debug_assertions)))] +#[bench] +fn bench_encode_startup(b: &mut test::Bencher) { + use test::black_box; + + let mut buf = Vec::with_capacity(128); + + b.iter(|| { + buf.clear(); + + black_box(Startup { + username: Some("postgres"), + database: Some("postgres"), + params: &[], + }) + .encode(&mut buf); + }); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/sync.rs b/src-tauri/vendor/sqlx-postgres/src/message/sync.rs new file mode 100644 index 00000000..56f44987 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/sync.rs @@ -0,0 +1,20 @@ +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +#[derive(Debug)] +pub struct Sync; + +impl FrontendMessage for Sync { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Sync; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(0) + } + + #[inline(always)] + fn encode_body(&self, _buf: &mut Vec) -> Result<(), Error> { + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/message/terminate.rs b/src-tauri/vendor/sqlx-postgres/src/message/terminate.rs new file mode 100644 index 00000000..39f8ff6e --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/message/terminate.rs @@ -0,0 +1,19 @@ +use crate::message::{FrontendMessage, FrontendMessageFormat}; +use sqlx_core::Error; +use std::num::Saturating; + +pub struct Terminate; + +impl FrontendMessage for Terminate { + const FORMAT: FrontendMessageFormat = FrontendMessageFormat::Terminate; + + #[inline(always)] + fn body_size_hint(&self) -> Saturating { + Saturating(0) + } + + #[inline(always)] + fn encode_body(&self, _buf: &mut Vec) -> Result<(), Error> { + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/migrate.rs b/src-tauri/vendor/sqlx-postgres/src/migrate.rs new file mode 100644 index 00000000..c37e92f4 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/migrate.rs @@ -0,0 +1,330 @@ +use std::str::FromStr; +use std::time::Duration; +use std::time::Instant; + +use futures_core::future::BoxFuture; + +pub(crate) use sqlx_core::migrate::MigrateError; +pub(crate) use sqlx_core::migrate::{AppliedMigration, Migration}; +pub(crate) use sqlx_core::migrate::{Migrate, MigrateDatabase}; + +use crate::connection::{ConnectOptions, Connection}; +use crate::error::Error; +use crate::executor::Executor; +use crate::query::query; +use crate::query_as::query_as; +use crate::query_scalar::query_scalar; +use crate::{PgConnectOptions, PgConnection, Postgres}; + +fn parse_for_maintenance(url: &str) -> Result<(PgConnectOptions, String), Error> { + let mut options = PgConnectOptions::from_str(url)?; + + // pull out the name of the database to create + let database = options + .database + .as_deref() + .unwrap_or(&options.username) + .to_owned(); + + // switch us to the maintenance database + // use `postgres` _unless_ the database is postgres, in which case, use `template1` + // this matches the behavior of the `createdb` util + options.database = if database == "postgres" { + Some("template1".into()) + } else { + Some("postgres".into()) + }; + + Ok((options, database)) +} + +impl MigrateDatabase for Postgres { + fn create_database(url: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let _ = conn + .execute(&*format!( + "CREATE DATABASE \"{}\"", + database.replace('"', "\"\"") + )) + .await?; + + Ok(()) + }) + } + + fn database_exists(url: &str) -> BoxFuture<'_, Result> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let exists: bool = + query_scalar("select exists(SELECT 1 from pg_database WHERE datname = $1)") + .bind(database) + .fetch_one(&mut conn) + .await?; + + Ok(exists) + }) + } + + fn drop_database(url: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let _ = conn + .execute(&*format!( + "DROP DATABASE IF EXISTS \"{}\"", + database.replace('"', "\"\"") + )) + .await?; + + Ok(()) + }) + } + + fn force_drop_database(url: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let (options, database) = parse_for_maintenance(url)?; + let mut conn = options.connect().await?; + + let row: (String,) = query_as("SELECT current_setting('server_version_num')") + .fetch_one(&mut conn) + .await?; + + let version = row.0.parse::().unwrap(); + + let pid_type = if version >= 90200 { "pid" } else { "procpid" }; + + conn.execute(&*format!( + "SELECT pg_terminate_backend(pg_stat_activity.{pid_type}) FROM pg_stat_activity \ + WHERE pg_stat_activity.datname = '{database}' AND {pid_type} <> pg_backend_pid()" + )) + .await?; + + Self::drop_database(url).await + }) + } +} + +impl Migrate for PgConnection { + fn ensure_migrations_table(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + // language=SQL + self.execute( + r#" +CREATE TABLE IF NOT EXISTS _sqlx_migrations ( + version BIGINT PRIMARY KEY, + description TEXT NOT NULL, + installed_on TIMESTAMPTZ NOT NULL DEFAULT now(), + success BOOLEAN NOT NULL, + checksum BYTEA NOT NULL, + execution_time BIGINT NOT NULL +); + "#, + ) + .await?; + + Ok(()) + }) + } + + fn dirty_version(&mut self) -> BoxFuture<'_, Result, MigrateError>> { + Box::pin(async move { + // language=SQL + let row: Option<(i64,)> = query_as( + "SELECT version FROM _sqlx_migrations WHERE success = false ORDER BY version LIMIT 1", + ) + .fetch_optional(self) + .await?; + + Ok(row.map(|r| r.0)) + }) + } + + fn list_applied_migrations( + &mut self, + ) -> BoxFuture<'_, Result, MigrateError>> { + Box::pin(async move { + // language=SQL + let rows: Vec<(i64, Vec)> = + query_as("SELECT version, checksum FROM _sqlx_migrations ORDER BY version") + .fetch_all(self) + .await?; + + let migrations = rows + .into_iter() + .map(|(version, checksum)| AppliedMigration { + version, + checksum: checksum.into(), + }) + .collect(); + + Ok(migrations) + }) + } + + fn lock(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + let database_name = current_database(self).await?; + let lock_id = generate_lock_id(&database_name); + + // create an application lock over the database + // this function will not return until the lock is acquired + + // https://www.postgresql.org/docs/current/explicit-locking.html#ADVISORY-LOCKS + // https://www.postgresql.org/docs/current/functions-admin.html#FUNCTIONS-ADVISORY-LOCKS-TABLE + + // language=SQL + let _ = query("SELECT pg_advisory_lock($1)") + .bind(lock_id) + .execute(self) + .await?; + + Ok(()) + }) + } + + fn unlock(&mut self) -> BoxFuture<'_, Result<(), MigrateError>> { + Box::pin(async move { + let database_name = current_database(self).await?; + let lock_id = generate_lock_id(&database_name); + + // language=SQL + let _ = query("SELECT pg_advisory_unlock($1)") + .bind(lock_id) + .execute(self) + .await?; + + Ok(()) + }) + } + + fn apply<'e: 'm, 'm>( + &'e mut self, + migration: &'m Migration, + ) -> BoxFuture<'m, Result> { + Box::pin(async move { + let start = Instant::now(); + + // execute migration queries + if migration.no_tx { + execute_migration(self, migration).await?; + } else { + // Use a single transaction for the actual migration script and the essential bookeeping so we never + // execute migrations twice. See https://github.com/launchbadge/sqlx/issues/1966. + // The `execution_time` however can only be measured for the whole transaction. This value _only_ exists for + // data lineage and debugging reasons, so it is not super important if it is lost. So we initialize it to -1 + // and update it once the actual transaction completed. + let mut tx = self.begin().await?; + execute_migration(&mut tx, migration).await?; + tx.commit().await?; + } + + // Update `elapsed_time`. + // NOTE: The process may disconnect/die at this point, so the elapsed time value might be lost. We accept + // this small risk since this value is not super important. + let elapsed = start.elapsed(); + + // language=SQL + #[allow(clippy::cast_possible_truncation)] + let _ = query( + r#" + UPDATE _sqlx_migrations + SET execution_time = $1 + WHERE version = $2 + "#, + ) + .bind(elapsed.as_nanos() as i64) + .bind(migration.version) + .execute(self) + .await?; + + Ok(elapsed) + }) + } + + fn revert<'e: 'm, 'm>( + &'e mut self, + migration: &'m Migration, + ) -> BoxFuture<'m, Result> { + Box::pin(async move { + let start = Instant::now(); + + // execute migration queries + if migration.no_tx { + revert_migration(self, migration).await?; + } else { + // Use a single transaction for the actual migration script and the essential bookeeping so we never + // execute migrations twice. See https://github.com/launchbadge/sqlx/issues/1966. + let mut tx = self.begin().await?; + revert_migration(&mut tx, migration).await?; + tx.commit().await?; + } + + let elapsed = start.elapsed(); + + Ok(elapsed) + }) + } +} + +async fn execute_migration( + conn: &mut PgConnection, + migration: &Migration, +) -> Result<(), MigrateError> { + let _ = conn + .execute(&*migration.sql) + .await + .map_err(|e| MigrateError::ExecuteMigration(e, migration.version))?; + + // language=SQL + let _ = query( + r#" + INSERT INTO _sqlx_migrations ( version, description, success, checksum, execution_time ) + VALUES ( $1, $2, TRUE, $3, -1 ) + "#, + ) + .bind(migration.version) + .bind(&*migration.description) + .bind(&*migration.checksum) + .execute(conn) + .await?; + + Ok(()) +} + +async fn revert_migration( + conn: &mut PgConnection, + migration: &Migration, +) -> Result<(), MigrateError> { + let _ = conn + .execute(&*migration.sql) + .await + .map_err(|e| MigrateError::ExecuteMigration(e, migration.version))?; + + // language=SQL + let _ = query(r#"DELETE FROM _sqlx_migrations WHERE version = $1"#) + .bind(migration.version) + .execute(conn) + .await?; + + Ok(()) +} + +async fn current_database(conn: &mut PgConnection) -> Result { + // language=SQL + Ok(query_scalar("SELECT current_database()") + .fetch_one(conn) + .await?) +} + +// inspired from rails: https://github.com/rails/rails/blob/6e49cc77ab3d16c06e12f93158eaf3e507d4120e/activerecord/lib/active_record/migration.rb#L1308 +fn generate_lock_id(database_name: &str) -> i64 { + const CRC_IEEE: crc::Crc = crc::Crc::::new(&crc::CRC_32_ISO_HDLC); + // 0x3d32ad9e chosen by fair dice roll + 0x3d32ad9e * (CRC_IEEE.checksum(database_name.as_bytes()) as i64) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/options/connect.rs b/src-tauri/vendor/sqlx-postgres/src/options/connect.rs new file mode 100644 index 00000000..bc6e4adc --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/connect.rs @@ -0,0 +1,36 @@ +use crate::connection::ConnectOptions; +use crate::error::Error; +use crate::{PgConnectOptions, PgConnection}; +use futures_core::future::BoxFuture; +use log::LevelFilter; +use sqlx_core::Url; +use std::time::Duration; + +impl ConnectOptions for PgConnectOptions { + type Connection = PgConnection; + + fn from_url(url: &Url) -> Result { + Self::parse_from_url(url) + } + + fn to_url_lossy(&self) -> Url { + self.build_url() + } + + fn connect(&self) -> BoxFuture<'_, Result> + where + Self::Connection: Sized, + { + Box::pin(PgConnection::establish(self)) + } + + fn log_statements(mut self, level: LevelFilter) -> Self { + self.log_settings.log_statements(level); + self + } + + fn log_slow_statements(mut self, level: LevelFilter, duration: Duration) -> Self { + self.log_settings.log_slow_statements(level, duration); + self + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/options/doc.md b/src-tauri/vendor/sqlx-postgres/src/options/doc.md new file mode 100644 index 00000000..15c2459c --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/doc.md @@ -0,0 +1,185 @@ +Options and flags which can be used to configure a PostgreSQL connection. + +A value of `PgConnectOptions` can be parsed from a connection URL, +as described by [libpq][libpq-connstring]. + +The general form for a connection URL is: + +```text +postgresql://[user[:password]@][host][:port][/dbname][?param1=value1&...] +``` + +The URL scheme designator can be either `postgresql://` or `postgres://`. +Each of the URL parts is optional. For defaults, see the next section. + +This type also implements [`FromStr`][std::str::FromStr] so you can parse it from a string +containing a connection URL and then further adjust options if necessary (see example below). + +Note that characters not allowed in URLs must be [percent-encoded]. + +# Parameters + +This API accepts many of the same parameters as [libpq][libpq-params]; +if a parameter is not passed in via URL, it is populated by reading +[environment variables][libpq-envars] or choosing customary defaults. + +| Parameter | Environment Variable | Default / Remarks | +|--------------------|----------------------|-------------------------------------------------------------| +| `user` | `PGUSER` | The `whoami` of the currently running process. | +| `password` | `PGPASSWORD` | Read from [`passfile`], if it exists. | +| [`passfile`] | `PGPASSFILE` | `~/.pgpass` or `%APPDATA%\postgresql\pgpass.conf` (Windows) | +| `host` | `PGHOST` | See [Note: Default Host](#note-default-host). | +| `hostaddr` | `PGHOSTADDR` | See [Note: Default Host](#note-default-host). | +| `port` | `PGPORT` | `5432` | +| `dbname` | `PGDATABASE` | Unset; defaults to the username server-side. | +| `sslmode` | `PGSSLMODE` | `prefer`. See [`PgSslMode`] for details. | +| `sslrootcert` | `PGSSLROOTCERT` | Unset. See [Note: SSL](#note-ssl). | +| `sslcert` | `PGSSLCERT` | Unset. See [Note: SSL](#note-ssl). | +| `sslkey` | `PGSSLKEY` | Unset. See [Note: SSL](#note-ssl). | +| `options` | `PGOPTIONS` | Unset. | +| `application_name` | `PGAPPNAME` | Unset. | + +[`passfile`] handling may be bypassed using [`PgConnectOptions::new_without_pgpass()`]. + +## SQLx-Specific +SQLx also parses some bespoke parameters. These are _not_ configurable by environment variable. +Instead, the name is linked to the method to set the value. + +| Parameter | Default | +|--------------------------------------------------------------|-------------------------------| +| [`statement-cache-capacity`][Self::statement_cache_capacity] | `100` | + +# Example URLs +```text +postgresql:// +postgresql://:5433 +postgresql://localhost +postgresql://localhost:5433 +postgresql://localhost/mydb +postgresql://user@localhost +postgresql://user:secret@localhost +postgresql://user:correct%20horse%20battery%20staple@localhost +postgresql://localhost?dbname=mydb&user=postgres&password=postgres +``` + +See also [Note: Unix Domain Sockets](#note-unix-domain-sockets) below. + +# Note: Default Host +If the connection URL does not contain a hostname and `PGHOST` is not set, +this constructor looks for an open Unix domain socket in one of a few standard locations +(configured when Postgres is built): + +* `/var/run/postgresql/.s.PGSQL.{port}` (Debian) +* `/private/tmp/.s.PGSQL.{port}` (macOS when installed through Homebrew) +* `/tmp/.s.PGSQL.{port}` (default otherwise) + +This depends on the value of `port` being correct. +If Postgres is using a port other than the default (`5432`), `port` must be set. + +If no Unix domain socket is found, `localhost` is assumed. + +Note: this description is updated on a best-effort basis. +See `default_host()` in the same source file as this method for the current behavior. + +# Note: SSL +## Root Certs +If `sslrootcert` is not set, the default root certificates used depends on Cargo features: + +* If `tls-native-tls` is enabled, the system root certificates are used. +* If `tls-rustls-native-roots` is enabled, the system root certificates are used. +* Otherwise, TLS roots are populated using the [`webpki-roots`] crate. + +## Environment Variables +Unlike with `libpq`, the following environment variables may be _either_ +a path to a file _or_ a string value containing a [PEM-encoded value][rfc7468]: + +* `PGSSLROOTCERT` +* `PGSSLCERT` +* `PGSSLKEY` + +If the string begins with the standard `-----BEGIN -----` header +and ends with the standard `-----END -----` footer, +it is parsed directly. + +This behavior is _only_ implemented for the environment variables, not the URL parameters. + +Note: passing the SSL private key via environment variable may be a security risk. + +# Note: Unix Domain Sockets +If you want to connect to Postgres over a Unix domain socket, you can pass the path +to the _directory_ containing the socket as the `host` parameter. + +The final path to the socket will be `{host}/.s.PGSQL.{port}` as is standard for Postgres. + +If you're passing the domain socket path as the host segment of the URL, forward slashes +in the path must be [percent-encoded] (replacing `/` with `%2F`), e.g.: + +```text +postgres://%2Fvar%2Frun%2Fpostgresql/dbname + +Different port: +postgres://%2Fvar%2Frun%2Fpostgresql:5433/dbname + +With username and password: +postgres://user:password@%2Fvar%2Frun%2Fpostgresql/dbname + +With username and password, and different port: +postgres://user:password@%2Fvar%2Frun%2Fpostgresql:5432/dbname +``` + +Instead, the hostname can be passed in the query segment of the URL, +which does not require forward-slashes to be percent-encoded +(however, [other characters are][percent-encoded]): + +```text +postgres:dbname?host=/var/run/postgresql + +Different port: +postgres://:5433/dbname?host=/var/run/postgresql + +With username and password: +postgres://user:password@/dbname?host=/var/run/postgresql + +With username and password, and different port: +postgres://user:password@:5433/dbname?host=/var/run/postgresql +``` + +# Example + +```rust,no_run +use sqlx::{Connection, ConnectOptions}; +use sqlx::postgres::{PgConnectOptions, PgConnection, PgPool, PgSslMode}; + +# async fn example() -> sqlx::Result<()> { +// URL connection string +let conn = PgConnection::connect("postgres://localhost/mydb").await?; + +// Manually-constructed options +let conn = PgConnectOptions::new() + .host("secret-host") + .port(2525) + .username("secret-user") + .password("secret-password") + .ssl_mode(PgSslMode::Require) + .connect() + .await?; + +// Modifying options parsed from a string +let mut opts: PgConnectOptions = "postgres://localhost/mydb".parse()?; + +// Change the log verbosity level for queries. +// Information about SQL queries is logged at `DEBUG` level by default. +opts = opts.log_statements(log::LevelFilter::Trace); + +let pool = PgPool::connect_with(opts).await?; +# Ok(()) +# } +``` + +[percent-encoded]: https://developer.mozilla.org/en-US/docs/Glossary/Percent-encoding +[`passfile`]: https://www.postgresql.org/docs/current/libpq-pgpass.html +[libpq-connstring]: https://www.postgresql.org/docs/current/libpq-connect.html#LIBPQ-CONNSTRING +[libpq-params]: https://www.postgresql.org/docs/current/libpq-connect.html#LIBPQ-PARAMKEYWORDS +[libpq-envars]: https://www.postgresql.org/docs/current/libpq-envars.html +[rfc7468]: https://datatracker.ietf.org/doc/html/rfc7468 +[`webpki-roots`]: https://docs.rs/webpki-roots \ No newline at end of file diff --git a/src-tauri/vendor/sqlx-postgres/src/options/mod.rs b/src-tauri/vendor/sqlx-postgres/src/options/mod.rs new file mode 100644 index 00000000..723721a9 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/mod.rs @@ -0,0 +1,609 @@ +use std::borrow::Cow; +use std::env::var; +use std::fmt::{Display, Write}; +use std::path::{Path, PathBuf}; + +pub use ssl_mode::PgSslMode; + +use crate::{connection::LogSettings, net::tls::CertificateInput}; + +mod connect; +mod parse; +mod pgpass; +mod ssl_mode; + +#[doc = include_str!("doc.md")] +#[derive(Debug, Clone)] +pub struct PgConnectOptions { + pub(crate) host: String, + pub(crate) port: u16, + pub(crate) socket: Option, + pub(crate) username: String, + pub(crate) password: Option, + pub(crate) database: Option, + pub(crate) ssl_mode: PgSslMode, + pub(crate) ssl_root_cert: Option, + pub(crate) ssl_client_cert: Option, + pub(crate) ssl_client_key: Option, + pub(crate) statement_cache_capacity: usize, + pub(crate) application_name: Option, + pub(crate) log_settings: LogSettings, + pub(crate) extra_float_digits: Option>, + pub(crate) options: Option, +} + +impl Default for PgConnectOptions { + fn default() -> Self { + Self::new_without_pgpass().apply_pgpass() + } +} + +impl PgConnectOptions { + /// Create a default set of connection options populated from the current environment. + /// + /// This behaves as if parsed from the connection string `postgres://` + /// + /// See the type-level documentation for details. + pub fn new() -> Self { + Self::new_without_pgpass().apply_pgpass() + } + + /// Create a default set of connection options _without_ reading from `passfile`. + /// + /// Equivalent to [`PgConnectOptions::new()`] but `passfile` is ignored. + /// + /// See the type-level documentation for details. + pub fn new_without_pgpass() -> Self { + let port = var("PGPORT") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(5432); + + let host = var("PGHOSTADDR") + .ok() + .or_else(|| var("PGHOST").ok()) + .unwrap_or_else(|| default_host(port)); + + let username = var("PGUSER").ok().unwrap_or_else(whoami::username); + + let database = var("PGDATABASE").ok(); + + PgConnectOptions { + port, + host, + socket: None, + username, + password: var("PGPASSWORD").ok(), + database, + ssl_root_cert: var("PGSSLROOTCERT").ok().map(CertificateInput::from), + ssl_client_cert: var("PGSSLCERT").ok().map(CertificateInput::from), + // As of writing, the implementation of `From` only looks for + // `-----BEGIN CERTIFICATE-----` and so will not attempt to parse + // a PEM-encoded private key. + ssl_client_key: var("PGSSLKEY").ok().map(CertificateInput::from), + ssl_mode: var("PGSSLMODE") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or_default(), + statement_cache_capacity: 100, + application_name: var("PGAPPNAME").ok(), + extra_float_digits: Some("2".into()), + log_settings: Default::default(), + options: var("PGOPTIONS").ok(), + } + } + + pub(crate) fn apply_pgpass(mut self) -> Self { + if self.password.is_none() { + self.password = pgpass::load_password( + &self.host, + self.port, + &self.username, + self.database.as_deref(), + ); + } + + self + } + + /// Sets the name of the host to connect to. + /// + /// If a host name begins with a slash, it specifies + /// Unix-domain communication rather than TCP/IP communication; the value is the name of + /// the directory in which the socket file is stored. + /// + /// The default behavior when host is not specified, or is empty, + /// is to connect to a Unix-domain socket + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .host("localhost"); + /// ``` + pub fn host(mut self, host: &str) -> Self { + host.clone_into(&mut self.host); + self + } + + /// Sets the port to connect to at the server host. + /// + /// The default port for PostgreSQL is `5432`. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .port(5432); + /// ``` + pub fn port(mut self, port: u16) -> Self { + self.port = port; + self + } + + /// Sets a custom path to a directory containing a unix domain socket, + /// switching the connection method from TCP to the corresponding socket. + /// + /// By default set to `None`. + pub fn socket(mut self, path: impl AsRef) -> Self { + self.socket = Some(path.as_ref().to_path_buf()); + self + } + + /// Sets the username to connect as. + /// + /// Defaults to be the same as the operating system name of + /// the user running the application. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .username("postgres"); + /// ``` + pub fn username(mut self, username: &str) -> Self { + username.clone_into(&mut self.username); + self + } + + /// Sets the password to use if the server demands password authentication. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .username("root") + /// .password("safe-and-secure"); + /// ``` + pub fn password(mut self, password: &str) -> Self { + self.password = Some(password.to_owned()); + self + } + + /// Sets the database name. Defaults to be the same as the user name. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .database("postgres"); + /// ``` + pub fn database(mut self, database: &str) -> Self { + self.database = Some(database.to_owned()); + self + } + + /// Sets whether or with what priority a secure SSL TCP/IP connection will be negotiated + /// with the server. + /// + /// By default, the SSL mode is [`Prefer`](PgSslMode::Prefer), and the client will + /// first attempt an SSL connection but fallback to a non-SSL connection on failure. + /// + /// Ignored for Unix domain socket communication. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// let options = PgConnectOptions::new() + /// .ssl_mode(PgSslMode::Require); + /// ``` + pub fn ssl_mode(mut self, mode: PgSslMode) -> Self { + self.ssl_mode = mode; + self + } + + /// Sets the name of a file containing SSL certificate authority (CA) certificate(s). + /// If the file exists, the server's certificate will be verified to be signed by + /// one of these authorities. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_root_cert("./ca-certificate.crt"); + /// ``` + pub fn ssl_root_cert(mut self, cert: impl AsRef) -> Self { + self.ssl_root_cert = Some(CertificateInput::File(cert.as_ref().to_path_buf())); + self + } + + /// Sets the name of a file containing SSL client certificate. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_client_cert("./client.crt"); + /// ``` + pub fn ssl_client_cert(mut self, cert: impl AsRef) -> Self { + self.ssl_client_cert = Some(CertificateInput::File(cert.as_ref().to_path_buf())); + self + } + + /// Sets the SSL client certificate as a PEM-encoded byte slice. + /// + /// This should be an ASCII-encoded blob that starts with `-----BEGIN CERTIFICATE-----`. + /// + /// # Example + /// Note: embedding SSL certificates and keys in the binary is not advised. + /// This is for illustration purposes only. + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// + /// const CERT: &[u8] = b"\ + /// -----BEGIN CERTIFICATE----- + /// + /// -----END CERTIFICATE-----"; + /// + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_client_cert_from_pem(CERT); + /// ``` + pub fn ssl_client_cert_from_pem(mut self, cert: impl AsRef<[u8]>) -> Self { + self.ssl_client_cert = Some(CertificateInput::Inline(cert.as_ref().to_vec())); + self + } + + /// Sets the name of a file containing SSL client key. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_client_key("./client.key"); + /// ``` + pub fn ssl_client_key(mut self, key: impl AsRef) -> Self { + self.ssl_client_key = Some(CertificateInput::File(key.as_ref().to_path_buf())); + self + } + + /// Sets the SSL client key as a PEM-encoded byte slice. + /// + /// This should be an ASCII-encoded blob that starts with `-----BEGIN PRIVATE KEY-----`. + /// + /// # Example + /// Note: embedding SSL certificates and keys in the binary is not advised. + /// This is for illustration purposes only. + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// + /// const KEY: &[u8] = b"\ + /// -----BEGIN PRIVATE KEY----- + /// + /// -----END PRIVATE KEY-----"; + /// + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_client_key_from_pem(KEY); + /// ``` + pub fn ssl_client_key_from_pem(mut self, key: impl AsRef<[u8]>) -> Self { + self.ssl_client_key = Some(CertificateInput::Inline(key.as_ref().to_vec())); + self + } + + /// Sets PEM encoded trusted SSL Certificate Authorities (CA). + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgSslMode, PgConnectOptions}; + /// let options = PgConnectOptions::new() + /// // Providing a CA certificate with less than VerifyCa is pointless + /// .ssl_mode(PgSslMode::VerifyCa) + /// .ssl_root_cert_from_pem(vec![]); + /// ``` + pub fn ssl_root_cert_from_pem(mut self, pem_certificate: Vec) -> Self { + self.ssl_root_cert = Some(CertificateInput::Inline(pem_certificate)); + self + } + + /// Sets the capacity of the connection's statement cache in a number of stored + /// distinct statements. Caching is handled using LRU, meaning when the + /// amount of queries hits the defined limit, the oldest statement will get + /// dropped. + /// + /// The default cache capacity is 100 statements. + pub fn statement_cache_capacity(mut self, capacity: usize) -> Self { + self.statement_cache_capacity = capacity; + self + } + + /// Sets the application name. Defaults to None + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .application_name("my-app"); + /// ``` + pub fn application_name(mut self, application_name: &str) -> Self { + self.application_name = Some(application_name.to_owned()); + self + } + + /// Sets or removes the `extra_float_digits` connection option. + /// + /// This changes the default precision of floating-point values returned in text mode (when + /// not using prepared statements such as calling methods of [`Executor`] directly). + /// + /// Historically, Postgres would by default round floating-point values to 6 and 15 digits + /// for `float4`/`REAL` (`f32`) and `float8`/`DOUBLE` (`f64`), respectively, which would mean + /// that the returned value may not be exactly the same as its representation in Postgres. + /// + /// The nominal range for this value is `-15` to `3`, where negative values for this option + /// cause floating-points to be rounded to that many fewer digits than normal (`-1` causes + /// `float4` to be rounded to 5 digits instead of six, or 14 instead of 15 for `float8`), + /// positive values cause Postgres to emit that many extra digits of precision over default + /// (or simply use maximum precision in Postgres 12 and later), + /// and 0 means keep the default behavior (or the "old" behavior described above + /// as of Postgres 12). + /// + /// SQLx sets this value to 3 by default, which tells Postgres to return floating-point values + /// at their maximum precision in the hope that the parsed value will be identical to its + /// counterpart in Postgres. This is also the default in Postgres 12 and later anyway. + /// + /// However, older versions of Postgres and alternative implementations that talk the Postgres + /// protocol may not support this option, or the full range of values. + /// + /// If you get an error like "unknown option `extra_float_digits`" when connecting, try + /// setting this to `None` or consult the manual of your database for the allowed range + /// of values. + /// + /// For more information, see: + /// * [Postgres manual, 20.11.2: Client Connection Defaults; Locale and Formatting][20.11.2] + /// * [Postgres manual, 8.1.3: Numeric Types; Floating-point Types][8.1.3] + /// + /// [`Executor`]: crate::executor::Executor + /// [20.11.2]: https://www.postgresql.org/docs/current/runtime-config-client.html#RUNTIME-CONFIG-CLIENT-FORMAT + /// [8.1.3]: https://www.postgresql.org/docs/current/datatype-numeric.html#DATATYPE-FLOAT + /// + /// ### Examples + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// + /// let mut options = PgConnectOptions::new() + /// // for Redshift and Postgres 10 + /// .extra_float_digits(2); + /// + /// let mut options = PgConnectOptions::new() + /// // don't send the option at all (Postgres 9 and older) + /// .extra_float_digits(None); + /// ``` + pub fn extra_float_digits(mut self, extra_float_digits: impl Into>) -> Self { + self.extra_float_digits = extra_float_digits.into().map(|it| it.to_string().into()); + self + } + + /// Set additional startup options for the connection as a list of key-value pairs. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .options([("geqo", "off"), ("statement_timeout", "5min")]); + /// ``` + pub fn options(mut self, options: I) -> Self + where + K: Display, + V: Display, + I: IntoIterator, + { + // Do this in here so `options_str` is only set if we have an option to insert + let options_str = self.options.get_or_insert_with(String::new); + for (k, v) in options { + if !options_str.is_empty() { + options_str.push(' '); + } + + write!(options_str, "-c {k}={v}").expect("failed to write an option to the string"); + } + self + } + + /// We try using a socket if hostname starts with `/` or if socket parameter + /// is specified. + pub(crate) fn fetch_socket(&self) -> Option { + match self.socket { + Some(ref socket) => { + let full_path = format!("{}/.s.PGSQL.{}", socket.display(), self.port); + Some(full_path) + } + None if self.host.starts_with('/') => { + let full_path = format!("{}/.s.PGSQL.{}", self.host, self.port); + Some(full_path) + } + _ => None, + } + } +} + +impl PgConnectOptions { + /// Get the current host. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .host("127.0.0.1"); + /// assert_eq!(options.get_host(), "127.0.0.1"); + /// ``` + pub fn get_host(&self) -> &str { + &self.host + } + + /// Get the server's port. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .port(6543); + /// assert_eq!(options.get_port(), 6543); + /// ``` + pub fn get_port(&self) -> u16 { + self.port + } + + /// Get the socket path. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .socket("/tmp"); + /// assert!(options.get_socket().is_some()); + /// ``` + pub fn get_socket(&self) -> Option<&PathBuf> { + self.socket.as_ref() + } + + /// Get the server's port. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .username("foo"); + /// assert_eq!(options.get_username(), "foo"); + /// ``` + pub fn get_username(&self) -> &str { + &self.username + } + + /// Get the current database name. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .database("postgres"); + /// assert!(options.get_database().is_some()); + /// ``` + pub fn get_database(&self) -> Option<&str> { + self.database.as_deref() + } + + /// Get the SSL mode. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::{PgConnectOptions, PgSslMode}; + /// let options = PgConnectOptions::new(); + /// assert!(matches!(options.get_ssl_mode(), PgSslMode::Prefer)); + /// ``` + pub fn get_ssl_mode(&self) -> PgSslMode { + self.ssl_mode + } + + /// Get the application name. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .application_name("service"); + /// assert!(options.get_application_name().is_some()); + /// ``` + pub fn get_application_name(&self) -> Option<&str> { + self.application_name.as_deref() + } + + /// Get the options. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_postgres::PgConnectOptions; + /// let options = PgConnectOptions::new() + /// .options([("foo", "bar")]); + /// assert!(options.get_options().is_some()); + /// ``` + pub fn get_options(&self) -> Option<&str> { + self.options.as_deref() + } +} + +fn default_host(port: u16) -> String { + // try to check for the existence of a unix socket and uses that + let socket = format!(".s.PGSQL.{port}"); + let candidates = [ + "/var/run/postgresql", // Debian + "/private/tmp", // OSX (homebrew) + "/tmp", // Default + ]; + + for candidate in &candidates { + if Path::new(candidate).join(&socket).exists() { + return candidate.to_string(); + } + } + + // fallback to localhost if no socket was found + "localhost".to_owned() +} + +#[test] +fn test_options_formatting() { + let options = PgConnectOptions::new().options([("geqo", "off")]); + assert_eq!(options.options, Some("-c geqo=off".to_string())); + let options = options.options([("search_path", "sqlx")]); + assert_eq!( + options.options, + Some("-c geqo=off -c search_path=sqlx".to_string()) + ); + let options = PgConnectOptions::new().options([("geqo", "off"), ("statement_timeout", "5min")]); + assert_eq!( + options.options, + Some("-c geqo=off -c statement_timeout=5min".to_string()) + ); + let options = PgConnectOptions::new(); + assert_eq!(options.options, None); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/options/parse.rs b/src-tauri/vendor/sqlx-postgres/src/options/parse.rs new file mode 100644 index 00000000..efbf85d8 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/parse.rs @@ -0,0 +1,342 @@ +use crate::error::Error; +use crate::{PgConnectOptions, PgSslMode}; +use sqlx_core::percent_encoding::{percent_decode_str, utf8_percent_encode, NON_ALPHANUMERIC}; +use sqlx_core::Url; +use std::net::IpAddr; +use std::str::FromStr; + +impl PgConnectOptions { + pub(crate) fn parse_from_url(url: &Url) -> Result { + let mut options = Self::new_without_pgpass(); + + if let Some(host) = url.host_str() { + let host_decoded = percent_decode_str(host); + options = match host_decoded.clone().next() { + Some(b'/') => options.socket(&*host_decoded.decode_utf8().map_err(Error::config)?), + _ => options.host(host), + } + } + + if let Some(port) = url.port() { + options = options.port(port); + } + + let username = url.username(); + if !username.is_empty() { + options = options.username( + &percent_decode_str(username) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + if let Some(password) = url.password() { + options = options.password( + &percent_decode_str(password) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + let path = url.path().trim_start_matches('/'); + if !path.is_empty() { + options = options.database( + &percent_decode_str(path) + .decode_utf8() + .map_err(Error::config)?, + ); + } + + for (key, value) in url.query_pairs().into_iter() { + match &*key { + "sslmode" | "ssl-mode" => { + options = options.ssl_mode(value.parse().map_err(Error::config)?); + } + + "sslrootcert" | "ssl-root-cert" | "ssl-ca" => { + options = options.ssl_root_cert(&*value); + } + + "sslcert" | "ssl-cert" => options = options.ssl_client_cert(&*value), + + "sslkey" | "ssl-key" => options = options.ssl_client_key(&*value), + + "statement-cache-capacity" => { + options = + options.statement_cache_capacity(value.parse().map_err(Error::config)?); + } + + "host" => { + if value.starts_with('/') { + options = options.socket(&*value); + } else { + options = options.host(&value); + } + } + + "hostaddr" => { + value.parse::().map_err(Error::config)?; + options = options.host(&value) + } + + "port" => options = options.port(value.parse().map_err(Error::config)?), + + "dbname" => options = options.database(&value), + + "user" => options = options.username(&value), + + "password" => options = options.password(&value), + + "application_name" => options = options.application_name(&value), + + "options" => { + if let Some(options) = options.options.as_mut() { + options.push(' '); + options.push_str(&value); + } else { + options.options = Some(value.to_string()); + } + } + + k if k.starts_with("options[") => { + if let Some(key) = k.strip_prefix("options[").unwrap().strip_suffix(']') { + options = options.options([(key, &*value)]); + } + } + + _ => tracing::warn!(%key, %value, "ignoring unrecognized connect parameter"), + } + } + + let options = options.apply_pgpass(); + + Ok(options) + } + + pub(crate) fn build_url(&self) -> Url { + let host = match &self.socket { + Some(socket) => { + utf8_percent_encode(&socket.to_string_lossy(), NON_ALPHANUMERIC).to_string() + } + None => self.host.to_owned(), + }; + + let mut url = Url::parse(&format!( + "postgres://{}@{}:{}", + self.username, host, self.port + )) + .expect("BUG: generated un-parseable URL"); + + if let Some(password) = &self.password { + let password = utf8_percent_encode(password, NON_ALPHANUMERIC).to_string(); + let _ = url.set_password(Some(&password)); + } + + if let Some(database) = &self.database { + url.set_path(database); + } + + let ssl_mode = match self.ssl_mode { + PgSslMode::Allow => "allow", + PgSslMode::Disable => "disable", + PgSslMode::Prefer => "prefer", + PgSslMode::Require => "require", + PgSslMode::VerifyCa => "verify-ca", + PgSslMode::VerifyFull => "verify-full", + }; + url.query_pairs_mut().append_pair("sslmode", ssl_mode); + + if let Some(ssl_root_cert) = &self.ssl_root_cert { + url.query_pairs_mut() + .append_pair("sslrootcert", &ssl_root_cert.to_string()); + } + + if let Some(ssl_client_cert) = &self.ssl_client_cert { + url.query_pairs_mut() + .append_pair("sslcert", &ssl_client_cert.to_string()); + } + + if let Some(ssl_client_key) = &self.ssl_client_key { + url.query_pairs_mut() + .append_pair("sslkey", &ssl_client_key.to_string()); + } + + url.query_pairs_mut().append_pair( + "statement-cache-capacity", + &self.statement_cache_capacity.to_string(), + ); + + url + } +} + +impl FromStr for PgConnectOptions { + type Err = Error; + + fn from_str(s: &str) -> Result { + let url: Url = s.parse().map_err(Error::config)?; + + Self::parse_from_url(&url) + } +} + +#[test] +fn it_parses_socket_correctly_from_parameter() { + let url = "postgres:///?host=/var/run/postgres/"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(Some("/var/run/postgres/".into()), opts.socket); +} + +#[test] +fn it_parses_host_correctly_from_parameter() { + let url = "postgres:///?host=google.database.com"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!("google.database.com", &opts.host); +} + +#[test] +fn it_parses_hostaddr_correctly_from_parameter() { + let url = "postgres:///?hostaddr=8.8.8.8"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!("8.8.8.8", &opts.host); +} + +#[test] +fn it_parses_port_correctly_from_parameter() { + let url = "postgres:///?port=1234"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!(1234, opts.port); +} + +#[test] +fn it_parses_dbname_correctly_from_parameter() { + let url = "postgres:///?dbname=some_db"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!(Some("some_db"), opts.database.as_deref()); +} + +#[test] +fn it_parses_user_correctly_from_parameter() { + let url = "postgres:///?user=some_user"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!("some_user", opts.username); +} + +#[test] +fn it_parses_password_correctly_from_parameter() { + let url = "postgres:///?password=some_pass"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(None, opts.socket); + assert_eq!(Some("some_pass"), opts.password.as_deref()); +} + +#[test] +fn it_parses_application_name_correctly_from_parameter() { + let url = "postgres:///?application_name=some_name"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(Some("some_name"), opts.application_name.as_deref()); +} + +#[test] +fn it_parses_username_with_at_sign_correctly() { + let url = "postgres://user@hostname:password@hostname:5432/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!("user@hostname", &opts.username); +} + +#[test] +fn it_parses_password_with_non_ascii_chars_correctly() { + let url = "postgres://username:p@ssw0rd@hostname:5432/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(Some("p@ssw0rd".into()), opts.password); +} + +#[test] +fn it_parses_socket_correctly_percent_encoded() { + let url = "postgres://%2Fvar%2Flib%2Fpostgres/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!(Some("/var/lib/postgres/".into()), opts.socket); +} +#[test] +fn it_parses_socket_correctly_with_username_percent_encoded() { + let url = "postgres://some_user@%2Fvar%2Flib%2Fpostgres/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!("some_user", opts.username); + assert_eq!(Some("/var/lib/postgres/".into()), opts.socket); + assert_eq!(Some("database"), opts.database.as_deref()); +} +#[test] +fn it_parses_libpq_options_correctly() { + let url = "postgres:///?options=-c%20synchronous_commit%3Doff%20--search_path%3Dpostgres"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!( + Some("-c synchronous_commit=off --search_path=postgres".into()), + opts.options + ); +} +#[test] +fn it_parses_sqlx_options_correctly() { + let url = "postgres:///?options[synchronous_commit]=off&options[search_path]=postgres"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + assert_eq!( + Some("-c synchronous_commit=off -c search_path=postgres".into()), + opts.options + ); +} + +#[test] +fn it_returns_the_parsed_url_when_socket() { + let url = "postgres://username@%2Fvar%2Flib%2Fpostgres/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + let mut expected_url = Url::parse(url).unwrap(); + // PgConnectOptions defaults + let query_string = "sslmode=prefer&statement-cache-capacity=100"; + let port = 5432; + expected_url.set_query(Some(query_string)); + let _ = expected_url.set_port(Some(port)); + + assert_eq!(expected_url, opts.build_url()); +} + +#[test] +fn it_returns_the_parsed_url_when_host() { + let url = "postgres://username:p@ssw0rd@hostname:5432/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + let mut expected_url = Url::parse(url).unwrap(); + // PgConnectOptions defaults + let query_string = "sslmode=prefer&statement-cache-capacity=100"; + expected_url.set_query(Some(query_string)); + + assert_eq!(expected_url, opts.build_url()); +} + +#[test] +fn built_url_can_be_parsed() { + let url = "postgres://username:p@ssw0rd@hostname:5432/database"; + let opts = PgConnectOptions::from_str(url).unwrap(); + + let parsed = PgConnectOptions::from_str(&opts.build_url().to_string()); + + assert!(parsed.is_ok()); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/options/pgpass.rs b/src-tauri/vendor/sqlx-postgres/src/options/pgpass.rs new file mode 100644 index 00000000..bf165595 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/pgpass.rs @@ -0,0 +1,351 @@ +use std::borrow::Cow; +use std::env::var_os; +use std::fs::File; +use std::io::{BufRead, BufReader}; +use std::path::PathBuf; + +/// try to load a password from the various pgpass file locations +pub fn load_password( + host: &str, + port: u16, + username: &str, + database: Option<&str>, +) -> Option { + let custom_file = var_os("PGPASSFILE"); + if let Some(file) = custom_file { + if let Some(password) = + load_password_from_file(PathBuf::from(file), host, port, username, database) + { + return Some(password); + } + } + + #[cfg(not(target_os = "windows"))] + let default_file = home::home_dir().map(|path| path.join(".pgpass")); + #[cfg(target_os = "windows")] + let default_file = { + use etcetera::BaseStrategy; + + etcetera::base_strategy::Windows::new() + .ok() + .map(|basedirs| basedirs.data_dir().join("postgres").join("pgpass.conf")) + }; + load_password_from_file(default_file?, host, port, username, database) +} + +/// try to extract a password from a pgpass file +fn load_password_from_file( + path: PathBuf, + host: &str, + port: u16, + username: &str, + database: Option<&str>, +) -> Option { + let file = File::open(&path) + .map_err(|e| { + match e.kind() { + std::io::ErrorKind::NotFound => { + tracing::debug!( + path = %path.display(), + "`.pgpass` file not found", + ); + } + _ => { + tracing::warn!( + path = %path.display(), + "Failed to open `.pgpass` file: {e:?}", + ); + } + }; + }) + .ok()?; + + #[cfg(target_os = "linux")] + { + use std::os::unix::fs::PermissionsExt; + + // check file permissions on linux + + let metadata = file.metadata().ok()?; + let permissions = metadata.permissions(); + let mode = permissions.mode(); + if mode & 0o77 != 0 { + tracing::warn!( + path = %path.display(), + permissions = format!("{mode:o}"), + "Ignoring path. Permissions are not strict enough", + ); + return None; + } + } + + let reader = BufReader::new(file); + load_password_from_reader(reader, host, port, username, database) +} + +fn load_password_from_reader( + mut reader: impl BufRead, + host: &str, + port: u16, + username: &str, + database: Option<&str>, +) -> Option { + let mut line = String::new(); + + // https://stackoverflow.com/a/55041833 + fn trim_newline(s: &mut String) { + if s.ends_with('\n') { + s.pop(); + if s.ends_with('\r') { + s.pop(); + } + } + } + + while let Ok(n) = reader.read_line(&mut line) { + if n == 0 { + break; + } + + if line.starts_with('#') { + // comment, do nothing + } else { + // try to load password from line + trim_newline(&mut line); + if let Some(password) = load_password_from_line(&line, host, port, username, database) { + return Some(password); + } + } + + line.clear(); + } + + None +} + +/// try to check all fields & extract the password +fn load_password_from_line( + mut line: &str, + host: &str, + port: u16, + username: &str, + database: Option<&str>, +) -> Option { + let whole_line = line; + + // Pgpass line ordering: hostname, port, database, username, password + // See: https://www.postgresql.org/docs/9.3/libpq-pgpass.html + match line.trim_start().chars().next() { + None | Some('#') => None, + _ => { + matches_next_field(whole_line, &mut line, host)?; + matches_next_field(whole_line, &mut line, &port.to_string())?; + matches_next_field(whole_line, &mut line, database.unwrap_or_default())?; + matches_next_field(whole_line, &mut line, username)?; + Some(line.to_owned()) + } + } +} + +/// check if the next field matches the provided value +fn matches_next_field(whole_line: &str, line: &mut &str, value: &str) -> Option<()> { + let field = find_next_field(line); + match field { + Some(field) => { + if field == "*" || field == value { + Some(()) + } else { + None + } + } + None => { + tracing::warn!(line = whole_line, "Malformed line in pgpass file"); + None + } + } +} + +/// extract the next value from a line in a pgpass file +/// +/// `line` will get updated to point behind the field and delimiter +fn find_next_field<'a>(line: &mut &'a str) -> Option> { + let mut escaping = false; + let mut escaped_string = None; + let mut last_added = 0; + + let char_indicies = line.char_indices(); + for (idx, c) in char_indicies { + if c == ':' && !escaping { + let (field, rest) = line.split_at(idx); + *line = &rest[1..]; + + if let Some(mut escaped_string) = escaped_string { + escaped_string += &field[last_added..]; + return Some(Cow::Owned(escaped_string)); + } else { + return Some(Cow::Borrowed(field)); + } + } else if c == '\\' { + let s = escaped_string.get_or_insert_with(String::new); + + if escaping { + s.push('\\'); + } else { + *s += &line[last_added..idx]; + } + + escaping = !escaping; + last_added = idx + 1; + } else { + escaping = false; + } + } + + None +} + +#[cfg(test)] +mod tests { + use super::{find_next_field, load_password_from_line, load_password_from_reader}; + use std::borrow::Cow; + + #[test] + fn test_find_next_field() { + fn test_case<'a>(mut input: &'a str, result: Option>, rest: &str) { + assert_eq!(find_next_field(&mut input), result); + assert_eq!(input, rest); + } + + // normal field + test_case("foo:bar:baz", Some(Cow::Borrowed("foo")), "bar:baz"); + // \ escaped + test_case( + "foo\\\\:bar:baz", + Some(Cow::Owned("foo\\".to_owned())), + "bar:baz", + ); + // : escaped + test_case( + "foo\\::bar:baz", + Some(Cow::Owned("foo:".to_owned())), + "bar:baz", + ); + // unnecessary escape + test_case( + "foo\\a:bar:baz", + Some(Cow::Owned("fooa".to_owned())), + "bar:baz", + ); + // other text after escape + test_case( + "foo\\\\a:bar:baz", + Some(Cow::Owned("foo\\a".to_owned())), + "bar:baz", + ); + // double escape + test_case( + "foo\\\\\\\\a:bar:baz", + Some(Cow::Owned("foo\\\\a".to_owned())), + "bar:baz", + ); + // utf8 support + test_case("🦀:bar:baz", Some(Cow::Borrowed("🦀")), "bar:baz"); + + // missing delimiter (eof) + test_case("foo", None, "foo"); + // missing delimiter after escape + test_case("foo\\:", None, "foo\\:"); + // missing delimiter after unused trailing escape + test_case("foo\\", None, "foo\\"); + } + + #[test] + fn test_load_password_from_line() { + // normal + assert_eq!( + load_password_from_line( + "localhost:5432:bar:foo:baz", + "localhost", + 5432, + "foo", + Some("bar") + ), + Some("baz".to_owned()) + ); + // wildcard + assert_eq!( + load_password_from_line("*:5432:bar:foo:baz", "localhost", 5432, "foo", Some("bar")), + Some("baz".to_owned()) + ); + // accept wildcard with missing db + assert_eq!( + load_password_from_line("localhost:5432:*:foo:baz", "localhost", 5432, "foo", None), + Some("baz".to_owned()) + ); + + // doesn't match + assert_eq!( + load_password_from_line( + "thishost:5432:bar:foo:baz", + "thathost", + 5432, + "foo", + Some("bar") + ), + None + ); + // malformed entry + assert_eq!( + load_password_from_line( + "localhost:5432:bar:foo", + "localhost", + 5432, + "foo", + Some("bar") + ), + None + ); + } + + #[test] + fn test_load_password_from_reader() { + let file = b"\ + localhost:5432:bar:foo:baz\n\ + # mixed line endings (also a comment!)\n\ + *:5432:bar:foo:baz\r\n\ + # trailing space, comment with CRLF! \r\n\ + thishost:5432:bar:foo:baz \n\ + # malformed line \n\ + thathost:5432:foobar:foo\n\ + # missing trailing newline\n\ + localhost:5432:*:foo:baz + "; + + // normal + assert_eq!( + load_password_from_reader(&mut &file[..], "localhost", 5432, "foo", Some("bar")), + Some("baz".to_owned()) + ); + // wildcard + assert_eq!( + load_password_from_reader(&mut &file[..], "localhost", 5432, "foo", Some("foobar")), + Some("baz".to_owned()) + ); + // accept wildcard with missing db + assert_eq!( + load_password_from_reader(&mut &file[..], "localhost", 5432, "foo", None), + Some("baz".to_owned()) + ); + + // doesn't match + assert_eq!( + load_password_from_reader(&mut &file[..], "thathost", 5432, "foo", Some("foobar")), + None + ); + // malformed entry + assert_eq!( + load_password_from_reader(&mut &file[..], "thathost", 5432, "foo", Some("foobar")), + None + ); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/options/ssl_mode.rs b/src-tauri/vendor/sqlx-postgres/src/options/ssl_mode.rs new file mode 100644 index 00000000..657728ab --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/options/ssl_mode.rs @@ -0,0 +1,53 @@ +use crate::error::Error; +use std::str::FromStr; + +/// Options for controlling the level of protection provided for PostgreSQL SSL connections. +/// +/// It is used by the [`ssl_mode`](super::PgConnectOptions::ssl_mode) method. +#[derive(Debug, Clone, Copy, Default)] +pub enum PgSslMode { + /// Only try a non-SSL connection. + Disable, + + /// First try a non-SSL connection; if that fails, try an SSL connection. + Allow, + + /// First try an SSL connection; if that fails, try a non-SSL connection. + /// + /// This is the default if no other mode is specified. + #[default] + Prefer, + + /// Only try an SSL connection. If a root CA file is present, verify the connection + /// in the same way as if `VerifyCa` was specified. + Require, + + /// Only try an SSL connection, and verify that the server certificate is issued by a + /// trusted certificate authority (CA). + VerifyCa, + + /// Only try an SSL connection; verify that the server certificate is issued by a trusted + /// CA and that the requested server host name matches that in the certificate. + VerifyFull, +} + +impl FromStr for PgSslMode { + type Err = Error; + + fn from_str(s: &str) -> Result { + Ok(match &*s.to_ascii_lowercase() { + "disable" => PgSslMode::Disable, + "allow" => PgSslMode::Allow, + "prefer" => PgSslMode::Prefer, + "require" => PgSslMode::Require, + "verify-ca" => PgSslMode::VerifyCa, + "verify-full" => PgSslMode::VerifyFull, + + _ => { + return Err(Error::Configuration( + format!("unknown value {s:?} for `ssl_mode`").into(), + )); + } + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/query_result.rs b/src-tauri/vendor/sqlx-postgres/src/query_result.rs new file mode 100644 index 00000000..3a243f3e --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/query_result.rs @@ -0,0 +1,30 @@ +use std::iter::{Extend, IntoIterator}; + +#[derive(Debug, Default)] +pub struct PgQueryResult { + pub(super) rows_affected: u64, +} + +impl PgQueryResult { + pub fn rows_affected(&self) -> u64 { + self.rows_affected + } +} + +impl Extend for PgQueryResult { + fn extend>(&mut self, iter: T) { + for elem in iter { + self.rows_affected += elem.rows_affected; + } + } +} + +#[cfg(feature = "any")] +impl From for sqlx_core::any::AnyQueryResult { + fn from(done: PgQueryResult) -> Self { + sqlx_core::any::AnyQueryResult { + rows_affected: done.rows_affected, + last_insert_id: None, + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/row.rs b/src-tauri/vendor/sqlx-postgres/src/row.rs new file mode 100644 index 00000000..f9e43bb9 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/row.rs @@ -0,0 +1,75 @@ +use crate::column::ColumnIndex; +use crate::error::Error; +use crate::message::DataRow; +use crate::statement::PgStatementMetadata; +use crate::value::PgValueFormat; +use crate::{PgColumn, PgValueRef, Postgres}; +pub(crate) use sqlx_core::row::Row; +use sqlx_core::type_checking::TypeChecking; +use sqlx_core::value::ValueRef; +use std::fmt::Debug; +use std::sync::Arc; + +/// Implementation of [`Row`] for PostgreSQL. +pub struct PgRow { + pub(crate) data: DataRow, + pub(crate) format: PgValueFormat, + pub(crate) metadata: Arc, +} + +impl Row for PgRow { + type Database = Postgres; + + fn columns(&self) -> &[PgColumn] { + &self.metadata.columns + } + + fn try_get_raw(&self, index: I) -> Result, Error> + where + I: ColumnIndex, + { + let index = index.index(self)?; + let column = &self.metadata.columns[index]; + let value = self.data.get(index); + + Ok(PgValueRef { + format: self.format, + row: Some(&self.data.storage), + type_info: column.type_info.clone(), + value, + }) + } +} + +impl ColumnIndex for &'_ str { + fn index(&self, row: &PgRow) -> Result { + row.metadata + .column_names + .get(*self) + .ok_or_else(|| Error::ColumnNotFound((*self).into())) + .copied() + } +} + +impl Debug for PgRow { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "PgRow ")?; + + let mut debug_map = f.debug_map(); + for (index, column) in self.columns().iter().enumerate() { + match self.try_get_raw(index) { + Ok(value) => { + debug_map.entry( + &column.name, + &Postgres::fmt_value_debug(&::to_owned(&value)), + ); + } + Err(error) => { + debug_map.entry(&column.name, &format!("decode error: {error:?}")); + } + } + } + + debug_map.finish() + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/statement.rs b/src-tauri/vendor/sqlx-postgres/src/statement.rs new file mode 100644 index 00000000..abd553af --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/statement.rs @@ -0,0 +1,86 @@ +use super::{PgColumn, PgTypeInfo}; +use crate::column::ColumnIndex; +use crate::error::Error; +use crate::ext::ustr::UStr; +use crate::{PgArguments, Postgres}; +use std::borrow::Cow; +use std::sync::Arc; + +pub(crate) use sqlx_core::statement::Statement; +use sqlx_core::{Either, HashMap}; + +#[derive(Debug, Clone)] +pub struct PgStatement<'q> { + pub(crate) sql: Cow<'q, str>, + pub(crate) metadata: Arc, +} + +#[derive(Debug, Default)] +pub(crate) struct PgStatementMetadata { + pub(crate) columns: Vec, + // This `Arc` is not redundant; it's used to avoid deep-copying this map for the `Any` backend. + // See `sqlx-postgres/src/any.rs` + pub(crate) column_names: Arc>, + pub(crate) parameters: Vec, +} + +impl<'q> Statement<'q> for PgStatement<'q> { + type Database = Postgres; + + fn to_owned(&self) -> PgStatement<'static> { + PgStatement::<'static> { + sql: Cow::Owned(self.sql.clone().into_owned()), + metadata: self.metadata.clone(), + } + } + + fn sql(&self) -> &str { + &self.sql + } + + fn parameters(&self) -> Option> { + Some(Either::Left(&self.metadata.parameters)) + } + + fn columns(&self) -> &[PgColumn] { + &self.metadata.columns + } + + impl_statement_query!(PgArguments); +} + +impl ColumnIndex> for &'_ str { + fn index(&self, statement: &PgStatement<'_>) -> Result { + statement + .metadata + .column_names + .get(*self) + .ok_or_else(|| Error::ColumnNotFound((*self).into())) + .copied() + } +} + +// #[cfg(feature = "any")] +// impl<'q> From> for crate::any::AnyStatement<'q> { +// #[inline] +// fn from(statement: PgStatement<'q>) -> Self { +// crate::any::AnyStatement::<'q> { +// columns: statement +// .metadata +// .columns +// .iter() +// .map(|col| col.clone().into()) +// .collect(), +// column_names: statement.metadata.column_names.clone(), +// parameters: Some(Either::Left( +// statement +// .metadata +// .parameters +// .iter() +// .map(|ty| ty.clone().into()) +// .collect(), +// )), +// sql: statement.sql, +// } +// } +// } diff --git a/src-tauri/vendor/sqlx-postgres/src/testing/mod.rs b/src-tauri/vendor/sqlx-postgres/src/testing/mod.rs new file mode 100644 index 00000000..af20fe87 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/testing/mod.rs @@ -0,0 +1,201 @@ +use std::fmt::Write; +use std::ops::Deref; +use std::str::FromStr; +use std::time::Duration; + +use futures_core::future::BoxFuture; + +use once_cell::sync::OnceCell; +use sqlx_core::connection::Connection; +use sqlx_core::query_scalar::query_scalar; + +use crate::error::Error; +use crate::executor::Executor; +use crate::pool::{Pool, PoolOptions}; +use crate::query::query; +use crate::{PgConnectOptions, PgConnection, Postgres}; + +pub(crate) use sqlx_core::testing::*; + +// Using a blocking `OnceCell` here because the critical sections are short. +static MASTER_POOL: OnceCell> = OnceCell::new(); +// Automatically delete any databases created before the start of the test binary. + +impl TestSupport for Postgres { + fn test_context(args: &TestArgs) -> BoxFuture<'_, Result, Error>> { + Box::pin(async move { test_context(args).await }) + } + + fn cleanup_test(db_name: &str) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + let mut conn = MASTER_POOL + .get() + .expect("cleanup_test() invoked outside `#[sqlx::test]`") + .acquire() + .await?; + + do_cleanup(&mut conn, db_name).await + }) + } + + fn cleanup_test_dbs() -> BoxFuture<'static, Result, Error>> { + Box::pin(async move { + let url = dotenvy::var("DATABASE_URL").expect("DATABASE_URL must be set"); + + let mut conn = PgConnection::connect(&url).await?; + + let delete_db_names: Vec = + query_scalar("select db_name from _sqlx_test.databases") + .fetch_all(&mut conn) + .await?; + + if delete_db_names.is_empty() { + return Ok(None); + } + + let mut deleted_db_names = Vec::with_capacity(delete_db_names.len()); + + let mut command = String::new(); + + for db_name in &delete_db_names { + command.clear(); + writeln!(command, "drop database if exists {db_name:?};").ok(); + match conn.execute(&*command).await { + Ok(_deleted) => { + deleted_db_names.push(db_name); + } + // Assume a database error just means the DB is still in use. + Err(Error::Database(dbe)) => { + eprintln!("could not clean test database {db_name:?}: {dbe}") + } + // Bubble up other errors + Err(e) => return Err(e), + } + } + + query("delete from _sqlx_test.databases where db_name = any($1::text[])") + .bind(&deleted_db_names) + .execute(&mut conn) + .await?; + + let _ = conn.close().await; + Ok(Some(delete_db_names.len())) + }) + } + + fn snapshot( + _conn: &mut Self::Connection, + ) -> BoxFuture<'_, Result, Error>> { + // TODO: I want to get the testing feature out the door so this will have to wait, + // but I'm keeping the code around for now because I plan to come back to it. + todo!() + } +} + +async fn test_context(args: &TestArgs) -> Result, Error> { + let url = dotenvy::var("DATABASE_URL").expect("DATABASE_URL must be set"); + + let master_opts = PgConnectOptions::from_str(&url).expect("failed to parse DATABASE_URL"); + + let pool = PoolOptions::new() + // Postgres' normal connection limit is 100 plus 3 superuser connections + // We don't want to use the whole cap and there may be fuzziness here due to + // concurrently running tests anyway. + .max_connections(20) + // Immediately close master connections. Tokio's I/O streams don't like hopping runtimes. + .after_release(|_conn, _| Box::pin(async move { Ok(false) })) + .connect_lazy_with(master_opts); + + let master_pool = match MASTER_POOL.try_insert(pool) { + Ok(inserted) => inserted, + Err((existing, pool)) => { + // Sanity checks. + assert_eq!( + existing.connect_options().host, + pool.connect_options().host, + "DATABASE_URL changed at runtime, host differs" + ); + + assert_eq!( + existing.connect_options().database, + pool.connect_options().database, + "DATABASE_URL changed at runtime, database differs" + ); + + existing + } + }; + + let mut conn = master_pool.acquire().await?; + + // language=PostgreSQL + conn.execute( + // Explicit lock avoids this latent bug: https://stackoverflow.com/a/29908840 + // I couldn't find a bug on the mailing list for `CREATE SCHEMA` specifically, + // but a clearly related bug with `CREATE TABLE` has been known since 2007: + // https://www.postgresql.org/message-id/200710222037.l9MKbCJZ098744%40wwwmaster.postgresql.org + // magic constant 8318549251334697844 is just 8 ascii bytes 'sqlxtest'. + r#" + select pg_advisory_xact_lock(8318549251334697844); + + create schema if not exists _sqlx_test; + + create table if not exists _sqlx_test.databases ( + db_name text primary key, + test_path text not null, + created_at timestamptz not null default now() + ); + + create index if not exists databases_created_at + on _sqlx_test.databases(created_at); + + create sequence if not exists _sqlx_test.database_ids; + "#, + ) + .await?; + + let db_name = Postgres::db_name(args); + do_cleanup(&mut conn, &db_name).await?; + + query( + r#" + insert into _sqlx_test.databases(db_name, test_path) values ($1, $2) + "#, + ) + .bind(&db_name) + .bind(args.test_path) + .execute(&mut *conn) + .await?; + + let create_command = format!("create database {db_name:?}"); + debug_assert!(create_command.starts_with("create database \"")); + conn.execute(&(create_command)[..]).await?; + + Ok(TestContext { + pool_opts: PoolOptions::new() + // Don't allow a single test to take all the connections. + // Most tests shouldn't require more than 5 connections concurrently, + // or else they're likely doing too much in one test. + .max_connections(5) + // Close connections ASAP if left in the idle queue. + .idle_timeout(Some(Duration::from_secs(1))) + .parent(master_pool.clone()), + connect_opts: master_pool + .connect_options() + .deref() + .clone() + .database(&db_name), + db_name, + }) +} + +async fn do_cleanup(conn: &mut PgConnection, db_name: &str) -> Result<(), Error> { + let delete_db_command = format!("drop database if exists {db_name:?};"); + conn.execute(&*delete_db_command).await?; + query("delete from _sqlx_test.databases where db_name = $1::text") + .bind(db_name) + .execute(&mut *conn) + .await?; + + Ok(()) +} diff --git a/src-tauri/vendor/sqlx-postgres/src/transaction.rs b/src-tauri/vendor/sqlx-postgres/src/transaction.rs new file mode 100644 index 00000000..23352a8d --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/transaction.rs @@ -0,0 +1,110 @@ +use futures_core::future::BoxFuture; +use sqlx_core::database::Database; +use std::borrow::Cow; + +use crate::error::Error; +use crate::executor::Executor; + +use crate::{PgConnection, Postgres}; + +pub(crate) use sqlx_core::transaction::*; + +/// Implementation of [`TransactionManager`] for PostgreSQL. +pub struct PgTransactionManager; + +impl TransactionManager for PgTransactionManager { + type Database = Postgres; + + fn begin<'conn>( + conn: &'conn mut PgConnection, + statement: Option>, + ) -> BoxFuture<'conn, Result<(), Error>> { + Box::pin(async move { + let depth = conn.inner.transaction_depth; + let statement = match statement { + // custom `BEGIN` statements are not allowed if we're already in + // a transaction (we need to issue a `SAVEPOINT` instead) + Some(_) if depth > 0 => return Err(Error::InvalidSavePointStatement), + Some(statement) => statement, + None => begin_ansi_transaction_sql(depth), + }; + + let rollback = Rollback::new(conn); + rollback.conn.queue_simple_query(&statement)?; + rollback.conn.wait_until_ready().await?; + if !rollback.conn.in_transaction() { + return Err(Error::BeginFailed); + } + rollback.conn.inner.transaction_depth += 1; + rollback.defuse(); + + Ok(()) + }) + } + + fn commit(conn: &mut PgConnection) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + if conn.inner.transaction_depth > 0 { + conn.execute(&*commit_ansi_transaction_sql(conn.inner.transaction_depth)) + .await?; + + conn.inner.transaction_depth -= 1; + } + + Ok(()) + }) + } + + fn rollback(conn: &mut PgConnection) -> BoxFuture<'_, Result<(), Error>> { + Box::pin(async move { + if conn.inner.transaction_depth > 0 { + conn.execute(&*rollback_ansi_transaction_sql( + conn.inner.transaction_depth, + )) + .await?; + + conn.inner.transaction_depth -= 1; + } + + Ok(()) + }) + } + + fn start_rollback(conn: &mut PgConnection) { + if conn.inner.transaction_depth > 0 { + conn.queue_simple_query(&rollback_ansi_transaction_sql(conn.inner.transaction_depth)) + .expect("BUG: Rollback query somehow too large for protocol"); + + conn.inner.transaction_depth -= 1; + } + } + + fn get_transaction_depth(conn: &::Connection) -> usize { + conn.inner.transaction_depth + } +} + +struct Rollback<'c> { + conn: &'c mut PgConnection, + defuse: bool, +} + +impl Drop for Rollback<'_> { + fn drop(&mut self) { + if !self.defuse { + PgTransactionManager::start_rollback(self.conn) + } + } +} + +impl<'c> Rollback<'c> { + fn new(conn: &'c mut PgConnection) -> Self { + Self { + conn, + defuse: false, + } + } + fn defuse(mut self) { + self.defuse = true; + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/type_checking.rs b/src-tauri/vendor/sqlx-postgres/src/type_checking.rs new file mode 100644 index 00000000..672d9f73 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/type_checking.rs @@ -0,0 +1,239 @@ +use crate::Postgres; + +// The paths used below will also be emitted by the macros so they have to match the final facade. +#[allow(unused_imports, dead_code)] +mod sqlx { + pub use crate as postgres; + pub use sqlx_core::*; +} + +impl_type_checking!( + Postgres { + (), + bool, + String | &str, + i8, + i16, + i32, + i64, + f32, + f64, + Vec | &[u8], + + sqlx::postgres::types::Oid, + + sqlx::postgres::types::PgInterval, + + sqlx::postgres::types::PgMoney, + + sqlx::postgres::types::PgLTree, + + sqlx::postgres::types::PgLQuery, + + sqlx::postgres::types::PgCube, + + sqlx::postgres::types::PgPoint, + + sqlx::postgres::types::PgLine, + + sqlx::postgres::types::PgLSeg, + + sqlx::postgres::types::PgBox, + + sqlx::postgres::types::PgPath, + + sqlx::postgres::types::PgPolygon, + + sqlx::postgres::types::PgCircle, + + #[cfg(feature = "uuid")] + sqlx::types::Uuid, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveTime, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveDate, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::NaiveDateTime, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::types::chrono::DateTime | sqlx::types::chrono::DateTime<_>, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::postgres::types::PgTimeTz, + + #[cfg(feature = "time")] + sqlx::types::time::Time, + + #[cfg(feature = "time")] + sqlx::types::time::Date, + + #[cfg(feature = "time")] + sqlx::types::time::PrimitiveDateTime, + + #[cfg(feature = "time")] + sqlx::types::time::OffsetDateTime, + + #[cfg(feature = "time")] + sqlx::postgres::types::PgTimeTz, + + #[cfg(feature = "bigdecimal")] + sqlx::types::BigDecimal, + + #[cfg(feature = "rust_decimal")] + sqlx::types::Decimal, + + #[cfg(feature = "ipnetwork")] + sqlx::types::ipnetwork::IpNetwork, + + #[cfg(feature = "ipnet")] + sqlx::types::ipnet::IpNet, + + #[cfg(feature = "mac_address")] + sqlx::types::mac_address::MacAddress, + + #[cfg(feature = "json")] + sqlx::types::JsonValue, + + #[cfg(feature = "bit-vec")] + sqlx::types::BitVec, + + sqlx::postgres::types::PgHstore, + // Arrays + + Vec | &[bool], + Vec | &[String], + Vec> | &[Vec], + Vec | &[i8], + Vec | &[i16], + Vec | &[i32], + Vec | &[i64], + Vec | &[f32], + Vec | &[f64], + Vec | &[sqlx::postgres::types::Oid], + Vec | &[sqlx::postgres::types::PgMoney], + Vec | &[sqlx::postgres::types::PgInterval], + + #[cfg(feature = "uuid")] + Vec | &[sqlx::types::Uuid], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec | &[sqlx::types::chrono::NaiveTime], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec | &[sqlx::types::chrono::NaiveDate], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec | &[sqlx::types::chrono::NaiveDateTime], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec> | &[sqlx::types::chrono::DateTime<_>], + + #[cfg(feature = "time")] + Vec | &[sqlx::types::time::Time], + + #[cfg(feature = "time")] + Vec | &[sqlx::types::time::Date], + + #[cfg(feature = "time")] + Vec | &[sqlx::types::time::PrimitiveDateTime], + + #[cfg(feature = "time")] + Vec | &[sqlx::types::time::OffsetDateTime], + + #[cfg(feature = "bigdecimal")] + Vec | &[sqlx::types::BigDecimal], + + #[cfg(feature = "rust_decimal")] + Vec | &[sqlx::types::Decimal], + + #[cfg(feature = "ipnetwork")] + Vec | &[sqlx::types::ipnetwork::IpNetwork], + + #[cfg(feature = "ipnet")] + Vec | &[sqlx::types::ipnet::IpNet], + + #[cfg(feature = "mac_address")] + Vec | &[sqlx::types::mac_address::MacAddress], + + #[cfg(feature = "json")] + Vec | &[sqlx::types::JsonValue], + + Vec | &[sqlx::postgres::types::PgHstore], + + // Ranges + + sqlx::postgres::types::PgRange, + sqlx::postgres::types::PgRange, + + #[cfg(feature = "bigdecimal")] + sqlx::postgres::types::PgRange, + + #[cfg(feature = "rust_decimal")] + sqlx::postgres::types::PgRange, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::postgres::types::PgRange, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::postgres::types::PgRange, + + #[cfg(all(feature = "chrono", not(feature = "time")))] + sqlx::postgres::types::PgRange> | + sqlx::postgres::types::PgRange>, + + #[cfg(feature = "time")] + sqlx::postgres::types::PgRange, + + #[cfg(feature = "time")] + sqlx::postgres::types::PgRange, + + #[cfg(feature = "time")] + sqlx::postgres::types::PgRange, + + // Range arrays + + Vec> | &[sqlx::postgres::types::PgRange], + Vec> | &[sqlx::postgres::types::PgRange], + + #[cfg(feature = "bigdecimal")] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(feature = "rust_decimal")] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec>> | + &[sqlx::postgres::types::PgRange>], + + #[cfg(all(feature = "chrono", not(feature = "time")))] + Vec>> | + &[sqlx::postgres::types::PgRange>], + + #[cfg(feature = "time")] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(feature = "time")] + Vec> | + &[sqlx::postgres::types::PgRange], + + #[cfg(feature = "time")] + Vec> | + &[sqlx::postgres::types::PgRange], + }, + ParamChecking::Strong, + feature-types: info => info.__type_feature_gate(), +); diff --git a/src-tauri/vendor/sqlx-postgres/src/type_info.rs b/src-tauri/vendor/sqlx-postgres/src/type_info.rs new file mode 100644 index 00000000..28c56758 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/type_info.rs @@ -0,0 +1,1390 @@ +#![allow(dead_code)] + +use std::borrow::Cow; +use std::fmt::{self, Display, Formatter}; +use std::ops::Deref; +use std::sync::Arc; + +use crate::ext::ustr::UStr; +use crate::types::Oid; + +pub(crate) use sqlx_core::type_info::TypeInfo; + +/// Type information for a PostgreSQL type. +/// +/// ### Note: Implementation of `==` ([`PartialEq::eq()`]) +/// Because `==` on [`TypeInfo`]s has been used throughout the SQLx API as a synonym for type compatibility, +/// e.g. in the default impl of [`Type::compatible()`][sqlx_core::types::Type::compatible], +/// some concessions have been made in the implementation. +/// +/// When comparing two `PgTypeInfo`s using the `==` operator ([`PartialEq::eq()`]), +/// if one was constructed with [`Self::with_oid()`] and the other with [`Self::with_name()`] or +/// [`Self::array_of()`], `==` will return `true`: +/// +/// ``` +/// # use sqlx::postgres::{types::Oid, PgTypeInfo}; +/// // Potentially surprising result, this assert will pass: +/// assert_eq!(PgTypeInfo::with_oid(Oid(1)), PgTypeInfo::with_name("definitely_not_real")); +/// ``` +/// +/// Since it is not possible in this case to prove the types are _not_ compatible (because +/// both `PgTypeInfo`s need to be resolved by an active connection to know for sure) +/// and type compatibility is mainly done as a sanity check anyway, +/// it was deemed acceptable to fudge equality in this very specific case. +/// +/// This also applies when querying with the text protocol (not using prepared statements, +/// e.g. [`sqlx::raw_sql()`][sqlx_core::raw_sql::raw_sql]), as the connection will be unable +/// to look up the type info like it normally does when preparing a statement: it won't know +/// what the OIDs of the output columns will be until it's in the middle of reading the result, +/// and by that time it's too late. +/// +/// To compare types for exact equality, use [`Self::type_eq()`] instead. +#[derive(Debug, Clone, PartialEq)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct PgTypeInfo(pub(crate) PgType); + +impl Deref for PgTypeInfo { + type Target = PgType; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +#[repr(u32)] +pub enum PgType { + Bool, + Bytea, + Char, + Name, + Int8, + Int2, + Int4, + Text, + Oid, + Json, + JsonArray, + Point, + Lseg, + Path, + Box, + Polygon, + Line, + LineArray, + Cidr, + CidrArray, + Float4, + Float8, + Unknown, + Circle, + CircleArray, + Macaddr8, + Macaddr8Array, + Macaddr, + Inet, + BoolArray, + ByteaArray, + CharArray, + NameArray, + Int2Array, + Int4Array, + TextArray, + BpcharArray, + VarcharArray, + Int8Array, + PointArray, + LsegArray, + PathArray, + BoxArray, + Float4Array, + Float8Array, + PolygonArray, + OidArray, + MacaddrArray, + InetArray, + Bpchar, + Varchar, + Date, + Time, + Timestamp, + TimestampArray, + DateArray, + TimeArray, + Timestamptz, + TimestamptzArray, + Interval, + IntervalArray, + NumericArray, + Timetz, + TimetzArray, + Bit, + BitArray, + Varbit, + VarbitArray, + Numeric, + Record, + RecordArray, + Uuid, + UuidArray, + Jsonb, + JsonbArray, + Int4Range, + Int4RangeArray, + NumRange, + NumRangeArray, + TsRange, + TsRangeArray, + TstzRange, + TstzRangeArray, + DateRange, + DateRangeArray, + Int8Range, + Int8RangeArray, + Jsonpath, + JsonpathArray, + Money, + MoneyArray, + + // https://www.postgresql.org/docs/9.3/datatype-pseudo.html + Void, + + // A realized user-defined type. When a connection sees a DeclareXX variant it resolves + // into this one before passing it along to `accepts` or inside of `Value` objects. + Custom(Arc), + + // From [`PgTypeInfo::with_name`] + DeclareWithName(UStr), + + // NOTE: Do we want to bring back type declaration by ID? It's notoriously fragile but + // someone may have a user for it + DeclareWithOid(Oid), + + DeclareArrayOf(Arc), +} + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct PgCustomType { + #[cfg_attr(feature = "offline", serde(skip))] + pub(crate) oid: Oid, + pub(crate) name: UStr, + pub(crate) kind: PgTypeKind, +} + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub enum PgTypeKind { + Simple, + Pseudo, + Domain(PgTypeInfo), + Composite(Arc<[(String, PgTypeInfo)]>), + Array(PgTypeInfo), + Enum(Arc<[String]>), + Range(PgTypeInfo), +} + +#[derive(Debug, Clone)] +#[cfg_attr(feature = "offline", derive(serde::Serialize, serde::Deserialize))] +pub struct PgArrayOf { + pub(crate) elem_name: UStr, + pub(crate) name: Box, +} + +impl PgTypeInfo { + /// Returns the corresponding `PgTypeInfo` if the OID is a built-in type and recognized by SQLx. + pub(crate) fn try_from_oid(oid: Oid) -> Option { + PgType::try_from_oid(oid).map(Self) + } + + /// Returns the _kind_ (simple, array, enum, etc.) for this type. + pub fn kind(&self) -> &PgTypeKind { + self.0.kind() + } + + /// Returns the OID for this type, if available. + /// + /// The OID may not be available if SQLx only knows the type by name. + /// It will have to be resolved by a `PgConnection` at runtime which + /// will yield a new and semantically distinct `TypeInfo` instance. + /// + /// This method does not perform any such lookup. + /// + /// ### Note + /// With the exception of [the default `pg_type` catalog][pg_type], type OIDs are *not* stable in PostgreSQL. + /// If a type is added by an extension, its OID will be assigned when the `CREATE EXTENSION` statement is executed, + /// and so can change depending on what extensions are installed and in what order, as well as the exact + /// version of PostgreSQL. + /// + /// [pg_type]: https://github.com/postgres/postgres/blob/master/src/include/catalog/pg_type.dat + pub fn oid(&self) -> Option { + self.0.try_oid() + } + + #[doc(hidden)] + pub fn __type_feature_gate(&self) -> Option<&'static str> { + if [ + PgTypeInfo::DATE, + PgTypeInfo::TIME, + PgTypeInfo::TIMESTAMP, + PgTypeInfo::TIMESTAMPTZ, + PgTypeInfo::DATE_ARRAY, + PgTypeInfo::TIME_ARRAY, + PgTypeInfo::TIMESTAMP_ARRAY, + PgTypeInfo::TIMESTAMPTZ_ARRAY, + ] + .contains(self) + { + Some("time") + } else if [PgTypeInfo::UUID, PgTypeInfo::UUID_ARRAY].contains(self) { + Some("uuid") + } else if [ + PgTypeInfo::JSON, + PgTypeInfo::JSONB, + PgTypeInfo::JSON_ARRAY, + PgTypeInfo::JSONB_ARRAY, + ] + .contains(self) + { + Some("json") + } else if [ + PgTypeInfo::CIDR, + PgTypeInfo::INET, + PgTypeInfo::CIDR_ARRAY, + PgTypeInfo::INET_ARRAY, + ] + .contains(self) + { + Some("ipnetwork") + } else if [PgTypeInfo::MACADDR].contains(self) { + Some("mac_address") + } else if [PgTypeInfo::NUMERIC, PgTypeInfo::NUMERIC_ARRAY].contains(self) { + Some("bigdecimal") + } else { + None + } + } + + /// Create a `PgTypeInfo` from a type name. + /// + /// The OID for the type will be fetched from Postgres on use of + /// a value of this type. The fetched OID will be cached per-connection. + /// + /// ### Note: Type Names Prefixed with `_` + /// In `pg_catalog.pg_type`, Postgres prefixes a type name with `_` to denote an array of that + /// type, e.g. `int4[]` actually exists in `pg_type` as `_int4`. + /// + /// Previously, it was necessary in manual [`PgHasArrayType`][crate::PgHasArrayType] impls + /// to return [`PgTypeInfo::with_name()`] with the type name prefixed with `_` to denote + /// an array type, but this would not work with schema-qualified names. + /// + /// As of 0.8, [`PgTypeInfo::array_of()`] is used to declare an array type, + /// and the Postgres driver is now able to properly resolve arrays of custom types, + /// even in other schemas, which was not previously supported. + /// + /// It is highly recommended to migrate existing usages to [`PgTypeInfo::array_of()`] where + /// applicable. + /// + /// However, to maintain compatibility, the driver now infers any type name prefixed with `_` + /// to be an array of that type. This may introduce some breakages for types which use + /// a `_` prefix but which are not arrays. + /// + /// As a workaround, type names with `_` as a prefix but which are not arrays should be wrapped + /// in quotes, e.g.: + /// ``` + /// use sqlx::postgres::PgTypeInfo; + /// use sqlx::{Type, TypeInfo}; + /// + /// /// `CREATE TYPE "_foo" AS ENUM ('Bar', 'Baz');` + /// #[derive(sqlx::Type)] + /// // Will prevent SQLx from inferring `_foo` as an array type. + /// #[sqlx(type_name = r#""_foo""#)] + /// enum Foo { + /// Bar, + /// Baz + /// } + /// + /// assert_eq!(Foo::type_info().name(), r#""_foo""#); + /// ``` + pub const fn with_name(name: &'static str) -> Self { + Self(PgType::DeclareWithName(UStr::Static(name))) + } + + /// Create a `PgTypeInfo` of an array from the name of its element type. + /// + /// The array type OID will be fetched from Postgres on use of a value of this type. + /// The fetched OID will be cached per-connection. + pub fn array_of(elem_name: &'static str) -> Self { + // to satisfy `name()` and `display_name()`, we need to construct strings to return + Self(PgType::DeclareArrayOf(Arc::new(PgArrayOf { + elem_name: elem_name.into(), + name: format!("{elem_name}[]").into(), + }))) + } + + /// Create a `PgTypeInfo` from an OID. + /// + /// Note that the OID for a type is very dependent on the environment. If you only ever use + /// one database or if this is an unhandled built-in type, you should be fine. Otherwise, + /// you will be better served using [`Self::with_name()`]. + /// + /// ### Note: Interaction with `==` + /// This constructor may give surprising results with `==`. + /// + /// See [the type-level docs][Self] for details. + pub const fn with_oid(oid: Oid) -> Self { + Self(PgType::DeclareWithOid(oid)) + } + + /// Returns `true` if `self` can be compared exactly to `other`. + /// + /// Unlike `==`, this will return false if + pub fn type_eq(&self, other: &Self) -> bool { + self.eq_impl(other, false) + } +} + +// DEVELOPER PRO TIP: find builtin type OIDs easily by grepping this file +// https://github.com/postgres/postgres/blob/master/src/include/catalog/pg_type.dat +// +// If you have Postgres running locally you can also try +// SELECT oid, typarray FROM pg_type where typname = '' + +impl PgType { + /// Returns the corresponding `PgType` if the OID is a built-in type and recognized by SQLx. + pub(crate) fn try_from_oid(oid: Oid) -> Option { + Some(match oid.0 { + 16 => PgType::Bool, + 17 => PgType::Bytea, + 18 => PgType::Char, + 19 => PgType::Name, + 20 => PgType::Int8, + 21 => PgType::Int2, + 23 => PgType::Int4, + 25 => PgType::Text, + 26 => PgType::Oid, + 114 => PgType::Json, + 199 => PgType::JsonArray, + 600 => PgType::Point, + 601 => PgType::Lseg, + 602 => PgType::Path, + 603 => PgType::Box, + 604 => PgType::Polygon, + 628 => PgType::Line, + 629 => PgType::LineArray, + 650 => PgType::Cidr, + 651 => PgType::CidrArray, + 700 => PgType::Float4, + 701 => PgType::Float8, + 705 => PgType::Unknown, + 718 => PgType::Circle, + 719 => PgType::CircleArray, + 774 => PgType::Macaddr8, + 775 => PgType::Macaddr8Array, + 790 => PgType::Money, + 791 => PgType::MoneyArray, + 829 => PgType::Macaddr, + 869 => PgType::Inet, + 1000 => PgType::BoolArray, + 1001 => PgType::ByteaArray, + 1002 => PgType::CharArray, + 1003 => PgType::NameArray, + 1005 => PgType::Int2Array, + 1007 => PgType::Int4Array, + 1009 => PgType::TextArray, + 1014 => PgType::BpcharArray, + 1015 => PgType::VarcharArray, + 1016 => PgType::Int8Array, + 1017 => PgType::PointArray, + 1018 => PgType::LsegArray, + 1019 => PgType::PathArray, + 1020 => PgType::BoxArray, + 1021 => PgType::Float4Array, + 1022 => PgType::Float8Array, + 1027 => PgType::PolygonArray, + 1028 => PgType::OidArray, + 1040 => PgType::MacaddrArray, + 1041 => PgType::InetArray, + 1042 => PgType::Bpchar, + 1043 => PgType::Varchar, + 1082 => PgType::Date, + 1083 => PgType::Time, + 1114 => PgType::Timestamp, + 1115 => PgType::TimestampArray, + 1182 => PgType::DateArray, + 1183 => PgType::TimeArray, + 1184 => PgType::Timestamptz, + 1185 => PgType::TimestamptzArray, + 1186 => PgType::Interval, + 1187 => PgType::IntervalArray, + 1231 => PgType::NumericArray, + 1266 => PgType::Timetz, + 1270 => PgType::TimetzArray, + 1560 => PgType::Bit, + 1561 => PgType::BitArray, + 1562 => PgType::Varbit, + 1563 => PgType::VarbitArray, + 1700 => PgType::Numeric, + 2278 => PgType::Void, + 2249 => PgType::Record, + 2287 => PgType::RecordArray, + 2950 => PgType::Uuid, + 2951 => PgType::UuidArray, + 3802 => PgType::Jsonb, + 3807 => PgType::JsonbArray, + 3904 => PgType::Int4Range, + 3905 => PgType::Int4RangeArray, + 3906 => PgType::NumRange, + 3907 => PgType::NumRangeArray, + 3908 => PgType::TsRange, + 3909 => PgType::TsRangeArray, + 3910 => PgType::TstzRange, + 3911 => PgType::TstzRangeArray, + 3912 => PgType::DateRange, + 3913 => PgType::DateRangeArray, + 3926 => PgType::Int8Range, + 3927 => PgType::Int8RangeArray, + 4072 => PgType::Jsonpath, + 4073 => PgType::JsonpathArray, + + _ => { + return None; + } + }) + } + + pub(crate) fn oid(&self) -> Oid { + match self.try_oid() { + Some(oid) => oid, + None => unreachable!("(bug) use of unresolved type declaration [oid]"), + } + } + + pub(crate) fn try_oid(&self) -> Option { + Some(match self { + PgType::Bool => Oid(16), + PgType::Bytea => Oid(17), + PgType::Char => Oid(18), + PgType::Name => Oid(19), + PgType::Int8 => Oid(20), + PgType::Int2 => Oid(21), + PgType::Int4 => Oid(23), + PgType::Text => Oid(25), + PgType::Oid => Oid(26), + PgType::Json => Oid(114), + PgType::JsonArray => Oid(199), + PgType::Point => Oid(600), + PgType::Lseg => Oid(601), + PgType::Path => Oid(602), + PgType::Box => Oid(603), + PgType::Polygon => Oid(604), + PgType::Line => Oid(628), + PgType::LineArray => Oid(629), + PgType::Cidr => Oid(650), + PgType::CidrArray => Oid(651), + PgType::Float4 => Oid(700), + PgType::Float8 => Oid(701), + PgType::Unknown => Oid(705), + PgType::Circle => Oid(718), + PgType::CircleArray => Oid(719), + PgType::Macaddr8 => Oid(774), + PgType::Macaddr8Array => Oid(775), + PgType::Money => Oid(790), + PgType::MoneyArray => Oid(791), + PgType::Macaddr => Oid(829), + PgType::Inet => Oid(869), + PgType::BoolArray => Oid(1000), + PgType::ByteaArray => Oid(1001), + PgType::CharArray => Oid(1002), + PgType::NameArray => Oid(1003), + PgType::Int2Array => Oid(1005), + PgType::Int4Array => Oid(1007), + PgType::TextArray => Oid(1009), + PgType::BpcharArray => Oid(1014), + PgType::VarcharArray => Oid(1015), + PgType::Int8Array => Oid(1016), + PgType::PointArray => Oid(1017), + PgType::LsegArray => Oid(1018), + PgType::PathArray => Oid(1019), + PgType::BoxArray => Oid(1020), + PgType::Float4Array => Oid(1021), + PgType::Float8Array => Oid(1022), + PgType::PolygonArray => Oid(1027), + PgType::OidArray => Oid(1028), + PgType::MacaddrArray => Oid(1040), + PgType::InetArray => Oid(1041), + PgType::Bpchar => Oid(1042), + PgType::Varchar => Oid(1043), + PgType::Date => Oid(1082), + PgType::Time => Oid(1083), + PgType::Timestamp => Oid(1114), + PgType::TimestampArray => Oid(1115), + PgType::DateArray => Oid(1182), + PgType::TimeArray => Oid(1183), + PgType::Timestamptz => Oid(1184), + PgType::TimestamptzArray => Oid(1185), + PgType::Interval => Oid(1186), + PgType::IntervalArray => Oid(1187), + PgType::NumericArray => Oid(1231), + PgType::Timetz => Oid(1266), + PgType::TimetzArray => Oid(1270), + PgType::Bit => Oid(1560), + PgType::BitArray => Oid(1561), + PgType::Varbit => Oid(1562), + PgType::VarbitArray => Oid(1563), + PgType::Numeric => Oid(1700), + PgType::Void => Oid(2278), + PgType::Record => Oid(2249), + PgType::RecordArray => Oid(2287), + PgType::Uuid => Oid(2950), + PgType::UuidArray => Oid(2951), + PgType::Jsonb => Oid(3802), + PgType::JsonbArray => Oid(3807), + PgType::Int4Range => Oid(3904), + PgType::Int4RangeArray => Oid(3905), + PgType::NumRange => Oid(3906), + PgType::NumRangeArray => Oid(3907), + PgType::TsRange => Oid(3908), + PgType::TsRangeArray => Oid(3909), + PgType::TstzRange => Oid(3910), + PgType::TstzRangeArray => Oid(3911), + PgType::DateRange => Oid(3912), + PgType::DateRangeArray => Oid(3913), + PgType::Int8Range => Oid(3926), + PgType::Int8RangeArray => Oid(3927), + PgType::Jsonpath => Oid(4072), + PgType::JsonpathArray => Oid(4073), + + PgType::Custom(ty) => ty.oid, + + PgType::DeclareWithOid(oid) => *oid, + PgType::DeclareWithName(_) => { + return None; + } + PgType::DeclareArrayOf(_) => { + return None; + } + }) + } + + pub(crate) fn display_name(&self) -> &str { + match self { + PgType::Bool => "BOOL", + PgType::Bytea => "BYTEA", + PgType::Char => "\"CHAR\"", + PgType::Name => "NAME", + PgType::Int8 => "INT8", + PgType::Int2 => "INT2", + PgType::Int4 => "INT4", + PgType::Text => "TEXT", + PgType::Oid => "OID", + PgType::Json => "JSON", + PgType::JsonArray => "JSON[]", + PgType::Point => "POINT", + PgType::Lseg => "LSEG", + PgType::Path => "PATH", + PgType::Box => "BOX", + PgType::Polygon => "POLYGON", + PgType::Line => "LINE", + PgType::LineArray => "LINE[]", + PgType::Cidr => "CIDR", + PgType::CidrArray => "CIDR[]", + PgType::Float4 => "FLOAT4", + PgType::Float8 => "FLOAT8", + PgType::Unknown => "UNKNOWN", + PgType::Circle => "CIRCLE", + PgType::CircleArray => "CIRCLE[]", + PgType::Macaddr8 => "MACADDR8", + PgType::Macaddr8Array => "MACADDR8[]", + PgType::Macaddr => "MACADDR", + PgType::Inet => "INET", + PgType::BoolArray => "BOOL[]", + PgType::ByteaArray => "BYTEA[]", + PgType::CharArray => "\"CHAR\"[]", + PgType::NameArray => "NAME[]", + PgType::Int2Array => "INT2[]", + PgType::Int4Array => "INT4[]", + PgType::TextArray => "TEXT[]", + PgType::BpcharArray => "CHAR[]", + PgType::VarcharArray => "VARCHAR[]", + PgType::Int8Array => "INT8[]", + PgType::PointArray => "POINT[]", + PgType::LsegArray => "LSEG[]", + PgType::PathArray => "PATH[]", + PgType::BoxArray => "BOX[]", + PgType::Float4Array => "FLOAT4[]", + PgType::Float8Array => "FLOAT8[]", + PgType::PolygonArray => "POLYGON[]", + PgType::OidArray => "OID[]", + PgType::MacaddrArray => "MACADDR[]", + PgType::InetArray => "INET[]", + PgType::Bpchar => "CHAR", + PgType::Varchar => "VARCHAR", + PgType::Date => "DATE", + PgType::Time => "TIME", + PgType::Timestamp => "TIMESTAMP", + PgType::TimestampArray => "TIMESTAMP[]", + PgType::DateArray => "DATE[]", + PgType::TimeArray => "TIME[]", + PgType::Timestamptz => "TIMESTAMPTZ", + PgType::TimestamptzArray => "TIMESTAMPTZ[]", + PgType::Interval => "INTERVAL", + PgType::IntervalArray => "INTERVAL[]", + PgType::NumericArray => "NUMERIC[]", + PgType::Timetz => "TIMETZ", + PgType::TimetzArray => "TIMETZ[]", + PgType::Bit => "BIT", + PgType::BitArray => "BIT[]", + PgType::Varbit => "VARBIT", + PgType::VarbitArray => "VARBIT[]", + PgType::Numeric => "NUMERIC", + PgType::Record => "RECORD", + PgType::RecordArray => "RECORD[]", + PgType::Uuid => "UUID", + PgType::UuidArray => "UUID[]", + PgType::Jsonb => "JSONB", + PgType::JsonbArray => "JSONB[]", + PgType::Int4Range => "INT4RANGE", + PgType::Int4RangeArray => "INT4RANGE[]", + PgType::NumRange => "NUMRANGE", + PgType::NumRangeArray => "NUMRANGE[]", + PgType::TsRange => "TSRANGE", + PgType::TsRangeArray => "TSRANGE[]", + PgType::TstzRange => "TSTZRANGE", + PgType::TstzRangeArray => "TSTZRANGE[]", + PgType::DateRange => "DATERANGE", + PgType::DateRangeArray => "DATERANGE[]", + PgType::Int8Range => "INT8RANGE", + PgType::Int8RangeArray => "INT8RANGE[]", + PgType::Jsonpath => "JSONPATH", + PgType::JsonpathArray => "JSONPATH[]", + PgType::Money => "MONEY", + PgType::MoneyArray => "MONEY[]", + PgType::Void => "VOID", + PgType::Custom(ty) => &ty.name, + PgType::DeclareWithOid(_) => "?", + PgType::DeclareWithName(name) => name, + PgType::DeclareArrayOf(array) => &array.name, + } + } + + pub(crate) fn name(&self) -> &str { + match self { + PgType::Bool => "bool", + PgType::Bytea => "bytea", + PgType::Char => "char", + PgType::Name => "name", + PgType::Int8 => "int8", + PgType::Int2 => "int2", + PgType::Int4 => "int4", + PgType::Text => "text", + PgType::Oid => "oid", + PgType::Json => "json", + PgType::JsonArray => "_json", + PgType::Point => "point", + PgType::Lseg => "lseg", + PgType::Path => "path", + PgType::Box => "box", + PgType::Polygon => "polygon", + PgType::Line => "line", + PgType::LineArray => "_line", + PgType::Cidr => "cidr", + PgType::CidrArray => "_cidr", + PgType::Float4 => "float4", + PgType::Float8 => "float8", + PgType::Unknown => "unknown", + PgType::Circle => "circle", + PgType::CircleArray => "_circle", + PgType::Macaddr8 => "macaddr8", + PgType::Macaddr8Array => "_macaddr8", + PgType::Macaddr => "macaddr", + PgType::Inet => "inet", + PgType::BoolArray => "_bool", + PgType::ByteaArray => "_bytea", + PgType::CharArray => "_char", + PgType::NameArray => "_name", + PgType::Int2Array => "_int2", + PgType::Int4Array => "_int4", + PgType::TextArray => "_text", + PgType::BpcharArray => "_bpchar", + PgType::VarcharArray => "_varchar", + PgType::Int8Array => "_int8", + PgType::PointArray => "_point", + PgType::LsegArray => "_lseg", + PgType::PathArray => "_path", + PgType::BoxArray => "_box", + PgType::Float4Array => "_float4", + PgType::Float8Array => "_float8", + PgType::PolygonArray => "_polygon", + PgType::OidArray => "_oid", + PgType::MacaddrArray => "_macaddr", + PgType::InetArray => "_inet", + PgType::Bpchar => "bpchar", + PgType::Varchar => "varchar", + PgType::Date => "date", + PgType::Time => "time", + PgType::Timestamp => "timestamp", + PgType::TimestampArray => "_timestamp", + PgType::DateArray => "_date", + PgType::TimeArray => "_time", + PgType::Timestamptz => "timestamptz", + PgType::TimestamptzArray => "_timestamptz", + PgType::Interval => "interval", + PgType::IntervalArray => "_interval", + PgType::NumericArray => "_numeric", + PgType::Timetz => "timetz", + PgType::TimetzArray => "_timetz", + PgType::Bit => "bit", + PgType::BitArray => "_bit", + PgType::Varbit => "varbit", + PgType::VarbitArray => "_varbit", + PgType::Numeric => "numeric", + PgType::Record => "record", + PgType::RecordArray => "_record", + PgType::Uuid => "uuid", + PgType::UuidArray => "_uuid", + PgType::Jsonb => "jsonb", + PgType::JsonbArray => "_jsonb", + PgType::Int4Range => "int4range", + PgType::Int4RangeArray => "_int4range", + PgType::NumRange => "numrange", + PgType::NumRangeArray => "_numrange", + PgType::TsRange => "tsrange", + PgType::TsRangeArray => "_tsrange", + PgType::TstzRange => "tstzrange", + PgType::TstzRangeArray => "_tstzrange", + PgType::DateRange => "daterange", + PgType::DateRangeArray => "_daterange", + PgType::Int8Range => "int8range", + PgType::Int8RangeArray => "_int8range", + PgType::Jsonpath => "jsonpath", + PgType::JsonpathArray => "_jsonpath", + PgType::Money => "money", + PgType::MoneyArray => "_money", + PgType::Void => "void", + PgType::Custom(ty) => &ty.name, + PgType::DeclareWithOid(_) => "?", + PgType::DeclareWithName(name) => name, + PgType::DeclareArrayOf(array) => &array.name, + } + } + + pub(crate) fn kind(&self) -> &PgTypeKind { + match self { + PgType::Bool => &PgTypeKind::Simple, + PgType::Bytea => &PgTypeKind::Simple, + PgType::Char => &PgTypeKind::Simple, + PgType::Name => &PgTypeKind::Simple, + PgType::Int8 => &PgTypeKind::Simple, + PgType::Int2 => &PgTypeKind::Simple, + PgType::Int4 => &PgTypeKind::Simple, + PgType::Text => &PgTypeKind::Simple, + PgType::Oid => &PgTypeKind::Simple, + PgType::Json => &PgTypeKind::Simple, + PgType::JsonArray => &PgTypeKind::Array(PgTypeInfo(PgType::Json)), + PgType::Point => &PgTypeKind::Simple, + PgType::Lseg => &PgTypeKind::Simple, + PgType::Path => &PgTypeKind::Simple, + PgType::Box => &PgTypeKind::Simple, + PgType::Polygon => &PgTypeKind::Simple, + PgType::Line => &PgTypeKind::Simple, + PgType::LineArray => &PgTypeKind::Array(PgTypeInfo(PgType::Line)), + PgType::Cidr => &PgTypeKind::Simple, + PgType::CidrArray => &PgTypeKind::Array(PgTypeInfo(PgType::Cidr)), + PgType::Float4 => &PgTypeKind::Simple, + PgType::Float8 => &PgTypeKind::Simple, + PgType::Unknown => &PgTypeKind::Simple, + PgType::Circle => &PgTypeKind::Simple, + PgType::CircleArray => &PgTypeKind::Array(PgTypeInfo(PgType::Circle)), + PgType::Macaddr8 => &PgTypeKind::Simple, + PgType::Macaddr8Array => &PgTypeKind::Array(PgTypeInfo(PgType::Macaddr8)), + PgType::Macaddr => &PgTypeKind::Simple, + PgType::Inet => &PgTypeKind::Simple, + PgType::BoolArray => &PgTypeKind::Array(PgTypeInfo(PgType::Bool)), + PgType::ByteaArray => &PgTypeKind::Array(PgTypeInfo(PgType::Bytea)), + PgType::CharArray => &PgTypeKind::Array(PgTypeInfo(PgType::Char)), + PgType::NameArray => &PgTypeKind::Array(PgTypeInfo(PgType::Name)), + PgType::Int2Array => &PgTypeKind::Array(PgTypeInfo(PgType::Int2)), + PgType::Int4Array => &PgTypeKind::Array(PgTypeInfo(PgType::Int4)), + PgType::TextArray => &PgTypeKind::Array(PgTypeInfo(PgType::Text)), + PgType::BpcharArray => &PgTypeKind::Array(PgTypeInfo(PgType::Bpchar)), + PgType::VarcharArray => &PgTypeKind::Array(PgTypeInfo(PgType::Varchar)), + PgType::Int8Array => &PgTypeKind::Array(PgTypeInfo(PgType::Int8)), + PgType::PointArray => &PgTypeKind::Array(PgTypeInfo(PgType::Point)), + PgType::LsegArray => &PgTypeKind::Array(PgTypeInfo(PgType::Lseg)), + PgType::PathArray => &PgTypeKind::Array(PgTypeInfo(PgType::Path)), + PgType::BoxArray => &PgTypeKind::Array(PgTypeInfo(PgType::Box)), + PgType::Float4Array => &PgTypeKind::Array(PgTypeInfo(PgType::Float4)), + PgType::Float8Array => &PgTypeKind::Array(PgTypeInfo(PgType::Float8)), + PgType::PolygonArray => &PgTypeKind::Array(PgTypeInfo(PgType::Polygon)), + PgType::OidArray => &PgTypeKind::Array(PgTypeInfo(PgType::Oid)), + PgType::MacaddrArray => &PgTypeKind::Array(PgTypeInfo(PgType::Macaddr)), + PgType::InetArray => &PgTypeKind::Array(PgTypeInfo(PgType::Inet)), + PgType::Bpchar => &PgTypeKind::Simple, + PgType::Varchar => &PgTypeKind::Simple, + PgType::Date => &PgTypeKind::Simple, + PgType::Time => &PgTypeKind::Simple, + PgType::Timestamp => &PgTypeKind::Simple, + PgType::TimestampArray => &PgTypeKind::Array(PgTypeInfo(PgType::Timestamp)), + PgType::DateArray => &PgTypeKind::Array(PgTypeInfo(PgType::Date)), + PgType::TimeArray => &PgTypeKind::Array(PgTypeInfo(PgType::Time)), + PgType::Timestamptz => &PgTypeKind::Simple, + PgType::TimestamptzArray => &PgTypeKind::Array(PgTypeInfo(PgType::Timestamptz)), + PgType::Interval => &PgTypeKind::Simple, + PgType::IntervalArray => &PgTypeKind::Array(PgTypeInfo(PgType::Interval)), + PgType::NumericArray => &PgTypeKind::Array(PgTypeInfo(PgType::Numeric)), + PgType::Timetz => &PgTypeKind::Simple, + PgType::TimetzArray => &PgTypeKind::Array(PgTypeInfo(PgType::Timetz)), + PgType::Bit => &PgTypeKind::Simple, + PgType::BitArray => &PgTypeKind::Array(PgTypeInfo(PgType::Bit)), + PgType::Varbit => &PgTypeKind::Simple, + PgType::VarbitArray => &PgTypeKind::Array(PgTypeInfo(PgType::Varbit)), + PgType::Numeric => &PgTypeKind::Simple, + PgType::Record => &PgTypeKind::Simple, + PgType::RecordArray => &PgTypeKind::Array(PgTypeInfo(PgType::Record)), + PgType::Uuid => &PgTypeKind::Simple, + PgType::UuidArray => &PgTypeKind::Array(PgTypeInfo(PgType::Uuid)), + PgType::Jsonb => &PgTypeKind::Simple, + PgType::JsonbArray => &PgTypeKind::Array(PgTypeInfo(PgType::Jsonb)), + PgType::Int4Range => &PgTypeKind::Range(PgTypeInfo::INT4), + PgType::Int4RangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::Int4Range)), + PgType::NumRange => &PgTypeKind::Range(PgTypeInfo::NUMERIC), + PgType::NumRangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::NumRange)), + PgType::TsRange => &PgTypeKind::Range(PgTypeInfo::TIMESTAMP), + PgType::TsRangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::TsRange)), + PgType::TstzRange => &PgTypeKind::Range(PgTypeInfo::TIMESTAMPTZ), + PgType::TstzRangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::TstzRange)), + PgType::DateRange => &PgTypeKind::Range(PgTypeInfo::DATE), + PgType::DateRangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::DateRange)), + PgType::Int8Range => &PgTypeKind::Range(PgTypeInfo::INT8), + PgType::Int8RangeArray => &PgTypeKind::Array(PgTypeInfo(PgType::Int8Range)), + PgType::Jsonpath => &PgTypeKind::Simple, + PgType::JsonpathArray => &PgTypeKind::Array(PgTypeInfo(PgType::Jsonpath)), + PgType::Money => &PgTypeKind::Simple, + PgType::MoneyArray => &PgTypeKind::Array(PgTypeInfo(PgType::Money)), + + PgType::Void => &PgTypeKind::Pseudo, + + PgType::Custom(ty) => &ty.kind, + + PgType::DeclareWithOid(oid) => { + unreachable!("(bug) use of unresolved type declaration [oid={}]", oid.0); + } + PgType::DeclareWithName(name) => { + unreachable!("(bug) use of unresolved type declaration [name={name}]"); + } + PgType::DeclareArrayOf(array) => { + unreachable!( + "(bug) use of unresolved type declaration [array of={}]", + array.elem_name + ); + } + } + } + + /// If `self` is an array type, return the type info for its element. + pub(crate) fn try_array_element(&self) -> Option> { + // We explicitly match on all the `None` cases to ensure an exhaustive match. + match self { + PgType::Bool => None, + PgType::BoolArray => Some(Cow::Owned(PgTypeInfo(PgType::Bool))), + PgType::Bytea => None, + PgType::ByteaArray => Some(Cow::Owned(PgTypeInfo(PgType::Bytea))), + PgType::Char => None, + PgType::CharArray => Some(Cow::Owned(PgTypeInfo(PgType::Char))), + PgType::Name => None, + PgType::NameArray => Some(Cow::Owned(PgTypeInfo(PgType::Name))), + PgType::Int8 => None, + PgType::Int8Array => Some(Cow::Owned(PgTypeInfo(PgType::Int8))), + PgType::Int2 => None, + PgType::Int2Array => Some(Cow::Owned(PgTypeInfo(PgType::Int2))), + PgType::Int4 => None, + PgType::Int4Array => Some(Cow::Owned(PgTypeInfo(PgType::Int4))), + PgType::Text => None, + PgType::TextArray => Some(Cow::Owned(PgTypeInfo(PgType::Text))), + PgType::Oid => None, + PgType::OidArray => Some(Cow::Owned(PgTypeInfo(PgType::Oid))), + PgType::Json => None, + PgType::JsonArray => Some(Cow::Owned(PgTypeInfo(PgType::Json))), + PgType::Point => None, + PgType::PointArray => Some(Cow::Owned(PgTypeInfo(PgType::Point))), + PgType::Lseg => None, + PgType::LsegArray => Some(Cow::Owned(PgTypeInfo(PgType::Lseg))), + PgType::Path => None, + PgType::PathArray => Some(Cow::Owned(PgTypeInfo(PgType::Path))), + PgType::Box => None, + PgType::BoxArray => Some(Cow::Owned(PgTypeInfo(PgType::Box))), + PgType::Polygon => None, + PgType::PolygonArray => Some(Cow::Owned(PgTypeInfo(PgType::Polygon))), + PgType::Line => None, + PgType::LineArray => Some(Cow::Owned(PgTypeInfo(PgType::Line))), + PgType::Cidr => None, + PgType::CidrArray => Some(Cow::Owned(PgTypeInfo(PgType::Cidr))), + PgType::Float4 => None, + PgType::Float4Array => Some(Cow::Owned(PgTypeInfo(PgType::Float4))), + PgType::Float8 => None, + PgType::Float8Array => Some(Cow::Owned(PgTypeInfo(PgType::Float8))), + PgType::Circle => None, + PgType::CircleArray => Some(Cow::Owned(PgTypeInfo(PgType::Circle))), + PgType::Macaddr8 => None, + PgType::Macaddr8Array => Some(Cow::Owned(PgTypeInfo(PgType::Macaddr8))), + PgType::Money => None, + PgType::MoneyArray => Some(Cow::Owned(PgTypeInfo(PgType::Money))), + PgType::Macaddr => None, + PgType::MacaddrArray => Some(Cow::Owned(PgTypeInfo(PgType::Macaddr))), + PgType::Inet => None, + PgType::InetArray => Some(Cow::Owned(PgTypeInfo(PgType::Inet))), + PgType::Bpchar => None, + PgType::BpcharArray => Some(Cow::Owned(PgTypeInfo(PgType::Bpchar))), + PgType::Varchar => None, + PgType::VarcharArray => Some(Cow::Owned(PgTypeInfo(PgType::Varchar))), + PgType::Date => None, + PgType::DateArray => Some(Cow::Owned(PgTypeInfo(PgType::Date))), + PgType::Time => None, + PgType::TimeArray => Some(Cow::Owned(PgTypeInfo(PgType::Time))), + PgType::Timestamp => None, + PgType::TimestampArray => Some(Cow::Owned(PgTypeInfo(PgType::Timestamp))), + PgType::Timestamptz => None, + PgType::TimestamptzArray => Some(Cow::Owned(PgTypeInfo(PgType::Timestamptz))), + PgType::Interval => None, + PgType::IntervalArray => Some(Cow::Owned(PgTypeInfo(PgType::Interval))), + PgType::Timetz => None, + PgType::TimetzArray => Some(Cow::Owned(PgTypeInfo(PgType::Timetz))), + PgType::Bit => None, + PgType::BitArray => Some(Cow::Owned(PgTypeInfo(PgType::Bit))), + PgType::Varbit => None, + PgType::VarbitArray => Some(Cow::Owned(PgTypeInfo(PgType::Varbit))), + PgType::Numeric => None, + PgType::NumericArray => Some(Cow::Owned(PgTypeInfo(PgType::Numeric))), + PgType::Record => None, + PgType::RecordArray => Some(Cow::Owned(PgTypeInfo(PgType::Record))), + PgType::Uuid => None, + PgType::UuidArray => Some(Cow::Owned(PgTypeInfo(PgType::Uuid))), + PgType::Jsonb => None, + PgType::JsonbArray => Some(Cow::Owned(PgTypeInfo(PgType::Jsonb))), + PgType::Int4Range => None, + PgType::Int4RangeArray => Some(Cow::Owned(PgTypeInfo(PgType::Int4Range))), + PgType::NumRange => None, + PgType::NumRangeArray => Some(Cow::Owned(PgTypeInfo(PgType::NumRange))), + PgType::TsRange => None, + PgType::TsRangeArray => Some(Cow::Owned(PgTypeInfo(PgType::TsRange))), + PgType::TstzRange => None, + PgType::TstzRangeArray => Some(Cow::Owned(PgTypeInfo(PgType::TstzRange))), + PgType::DateRange => None, + PgType::DateRangeArray => Some(Cow::Owned(PgTypeInfo(PgType::DateRange))), + PgType::Int8Range => None, + PgType::Int8RangeArray => Some(Cow::Owned(PgTypeInfo(PgType::Int8Range))), + PgType::Jsonpath => None, + PgType::JsonpathArray => Some(Cow::Owned(PgTypeInfo(PgType::Jsonpath))), + // There is no `UnknownArray` + PgType::Unknown => None, + // There is no `VoidArray` + PgType::Void => None, + + PgType::Custom(ty) => match &ty.kind { + PgTypeKind::Simple => None, + PgTypeKind::Pseudo => None, + PgTypeKind::Domain(_) => None, + PgTypeKind::Composite(_) => None, + PgTypeKind::Array(ref elem_type_info) => Some(Cow::Borrowed(elem_type_info)), + PgTypeKind::Enum(_) => None, + PgTypeKind::Range(_) => None, + }, + PgType::DeclareWithOid(_) => None, + PgType::DeclareWithName(name) => { + // LEGACY: infer the array element name from a `_` prefix + UStr::strip_prefix(name, "_") + .map(|elem| Cow::Owned(PgTypeInfo(PgType::DeclareWithName(elem)))) + } + PgType::DeclareArrayOf(array) => Some(Cow::Owned(PgTypeInfo(PgType::DeclareWithName( + array.elem_name.clone(), + )))), + } + } + + /// Returns `true` if this type cannot be matched by name. + fn is_declare_with_oid(&self) -> bool { + matches!(self, Self::DeclareWithOid(_)) + } + + /// Compare two `PgType`s, first by OID, then by array element, then by name. + /// + /// If `soft_eq` is true and `self` or `other` is `DeclareWithOid` but not both, return `true` + /// before checking names. + fn eq_impl(&self, other: &Self, soft_eq: bool) -> bool { + if let (Some(a), Some(b)) = (self.try_oid(), other.try_oid()) { + // If there are OIDs available, use OIDs to perform a direct match + return a == b; + } + + if soft_eq && (self.is_declare_with_oid() || other.is_declare_with_oid()) { + // If we get to this point, one instance is `DeclareWithOid()` and the other is + // `DeclareArrayOf()` or `DeclareWithName()`, which means we can't compare the two. + // + // Since this is only likely to occur when using the text protocol where we can't + // resolve type names before executing a query, we can just opt out of typechecking. + return true; + } + + if let (Some(elem_a), Some(elem_b)) = (self.try_array_element(), other.try_array_element()) + { + return elem_a == elem_b; + } + + // Otherwise, perform a match on the name + name_eq(self.name(), other.name()) + } +} + +impl TypeInfo for PgTypeInfo { + fn name(&self) -> &str { + self.0.display_name() + } + + fn is_null(&self) -> bool { + false + } + + fn is_void(&self) -> bool { + matches!(self.0, PgType::Void) + } + + fn type_compatible(&self, other: &Self) -> bool + where + Self: Sized, + { + self == other + } +} + +impl PartialEq for PgCustomType { + fn eq(&self, other: &PgCustomType) -> bool { + other.oid == self.oid + } +} + +impl PgTypeInfo { + // boolean, state of true or false + pub(crate) const BOOL: Self = Self(PgType::Bool); + pub(crate) const BOOL_ARRAY: Self = Self(PgType::BoolArray); + + // binary data types, variable-length binary string + pub(crate) const BYTEA: Self = Self(PgType::Bytea); + pub(crate) const BYTEA_ARRAY: Self = Self(PgType::ByteaArray); + + // uuid + pub(crate) const UUID: Self = Self(PgType::Uuid); + pub(crate) const UUID_ARRAY: Self = Self(PgType::UuidArray); + + // record + pub(crate) const RECORD: Self = Self(PgType::Record); + pub(crate) const RECORD_ARRAY: Self = Self(PgType::RecordArray); + + // + // JSON types + // https://www.postgresql.org/docs/current/datatype-json.html + // + + pub(crate) const JSON: Self = Self(PgType::Json); + pub(crate) const JSON_ARRAY: Self = Self(PgType::JsonArray); + + pub(crate) const JSONB: Self = Self(PgType::Jsonb); + pub(crate) const JSONB_ARRAY: Self = Self(PgType::JsonbArray); + + pub(crate) const JSONPATH: Self = Self(PgType::Jsonpath); + pub(crate) const JSONPATH_ARRAY: Self = Self(PgType::JsonpathArray); + + // + // network address types + // https://www.postgresql.org/docs/current/datatype-net-types.html + // + + pub(crate) const CIDR: Self = Self(PgType::Cidr); + pub(crate) const CIDR_ARRAY: Self = Self(PgType::CidrArray); + + pub(crate) const INET: Self = Self(PgType::Inet); + pub(crate) const INET_ARRAY: Self = Self(PgType::InetArray); + + pub(crate) const MACADDR: Self = Self(PgType::Macaddr); + pub(crate) const MACADDR_ARRAY: Self = Self(PgType::MacaddrArray); + + pub(crate) const MACADDR8: Self = Self(PgType::Macaddr8); + pub(crate) const MACADDR8_ARRAY: Self = Self(PgType::Macaddr8Array); + + // + // character types + // https://www.postgresql.org/docs/current/datatype-character.html + // + + // internal type for object names + pub(crate) const NAME: Self = Self(PgType::Name); + pub(crate) const NAME_ARRAY: Self = Self(PgType::NameArray); + + // character type, fixed-length, blank-padded + pub(crate) const BPCHAR: Self = Self(PgType::Bpchar); + pub(crate) const BPCHAR_ARRAY: Self = Self(PgType::BpcharArray); + + // character type, variable-length with limit + pub(crate) const VARCHAR: Self = Self(PgType::Varchar); + pub(crate) const VARCHAR_ARRAY: Self = Self(PgType::VarcharArray); + + // character type, variable-length + pub(crate) const TEXT: Self = Self(PgType::Text); + pub(crate) const TEXT_ARRAY: Self = Self(PgType::TextArray); + + // unknown type, transmitted as text + pub(crate) const UNKNOWN: Self = Self(PgType::Unknown); + + // + // numeric types + // https://www.postgresql.org/docs/current/datatype-numeric.html + // + + // single-byte internal type + pub(crate) const CHAR: Self = Self(PgType::Char); + pub(crate) const CHAR_ARRAY: Self = Self(PgType::CharArray); + + // internal type for type ids + pub(crate) const OID: Self = Self(PgType::Oid); + pub(crate) const OID_ARRAY: Self = Self(PgType::OidArray); + + // small-range integer; -32768 to +32767 + pub(crate) const INT2: Self = Self(PgType::Int2); + pub(crate) const INT2_ARRAY: Self = Self(PgType::Int2Array); + + // typical choice for integer; -2147483648 to +2147483647 + pub(crate) const INT4: Self = Self(PgType::Int4); + pub(crate) const INT4_ARRAY: Self = Self(PgType::Int4Array); + + // large-range integer; -9223372036854775808 to +9223372036854775807 + pub(crate) const INT8: Self = Self(PgType::Int8); + pub(crate) const INT8_ARRAY: Self = Self(PgType::Int8Array); + + // variable-precision, inexact, 6 decimal digits precision + pub(crate) const FLOAT4: Self = Self(PgType::Float4); + pub(crate) const FLOAT4_ARRAY: Self = Self(PgType::Float4Array); + + // variable-precision, inexact, 15 decimal digits precision + pub(crate) const FLOAT8: Self = Self(PgType::Float8); + pub(crate) const FLOAT8_ARRAY: Self = Self(PgType::Float8Array); + + // user-specified precision, exact + pub(crate) const NUMERIC: Self = Self(PgType::Numeric); + pub(crate) const NUMERIC_ARRAY: Self = Self(PgType::NumericArray); + + // user-specified precision, exact + pub(crate) const MONEY: Self = Self(PgType::Money); + pub(crate) const MONEY_ARRAY: Self = Self(PgType::MoneyArray); + + // + // date/time types + // https://www.postgresql.org/docs/current/datatype-datetime.html + // + + // both date and time (no time zone) + pub(crate) const TIMESTAMP: Self = Self(PgType::Timestamp); + pub(crate) const TIMESTAMP_ARRAY: Self = Self(PgType::TimestampArray); + + // both date and time (with time zone) + pub(crate) const TIMESTAMPTZ: Self = Self(PgType::Timestamptz); + pub(crate) const TIMESTAMPTZ_ARRAY: Self = Self(PgType::TimestamptzArray); + + // date (no time of day) + pub(crate) const DATE: Self = Self(PgType::Date); + pub(crate) const DATE_ARRAY: Self = Self(PgType::DateArray); + + // time of day (no date) + pub(crate) const TIME: Self = Self(PgType::Time); + pub(crate) const TIME_ARRAY: Self = Self(PgType::TimeArray); + + // time of day (no date), with time zone + pub(crate) const TIMETZ: Self = Self(PgType::Timetz); + pub(crate) const TIMETZ_ARRAY: Self = Self(PgType::TimetzArray); + + // time interval + pub(crate) const INTERVAL: Self = Self(PgType::Interval); + pub(crate) const INTERVAL_ARRAY: Self = Self(PgType::IntervalArray); + + // + // geometric types + // https://www.postgresql.org/docs/current/datatype-geometric.html + // + + // point on a plane + pub(crate) const POINT: Self = Self(PgType::Point); + pub(crate) const POINT_ARRAY: Self = Self(PgType::PointArray); + + // infinite line + pub(crate) const LINE: Self = Self(PgType::Line); + pub(crate) const LINE_ARRAY: Self = Self(PgType::LineArray); + + // finite line segment + pub(crate) const LSEG: Self = Self(PgType::Lseg); + pub(crate) const LSEG_ARRAY: Self = Self(PgType::LsegArray); + + // rectangular box + pub(crate) const BOX: Self = Self(PgType::Box); + pub(crate) const BOX_ARRAY: Self = Self(PgType::BoxArray); + + // open or closed path + pub(crate) const PATH: Self = Self(PgType::Path); + pub(crate) const PATH_ARRAY: Self = Self(PgType::PathArray); + + // polygon + pub(crate) const POLYGON: Self = Self(PgType::Polygon); + pub(crate) const POLYGON_ARRAY: Self = Self(PgType::PolygonArray); + + // circle + pub(crate) const CIRCLE: Self = Self(PgType::Circle); + pub(crate) const CIRCLE_ARRAY: Self = Self(PgType::CircleArray); + + // + // bit string types + // https://www.postgresql.org/docs/current/datatype-bit.html + // + + pub(crate) const BIT: Self = Self(PgType::Bit); + pub(crate) const BIT_ARRAY: Self = Self(PgType::BitArray); + + pub(crate) const VARBIT: Self = Self(PgType::Varbit); + pub(crate) const VARBIT_ARRAY: Self = Self(PgType::VarbitArray); + + // + // range types + // https://www.postgresql.org/docs/current/rangetypes.html + // + + pub(crate) const INT4_RANGE: Self = Self(PgType::Int4Range); + pub(crate) const INT4_RANGE_ARRAY: Self = Self(PgType::Int4RangeArray); + + pub(crate) const NUM_RANGE: Self = Self(PgType::NumRange); + pub(crate) const NUM_RANGE_ARRAY: Self = Self(PgType::NumRangeArray); + + pub(crate) const TS_RANGE: Self = Self(PgType::TsRange); + pub(crate) const TS_RANGE_ARRAY: Self = Self(PgType::TsRangeArray); + + pub(crate) const TSTZ_RANGE: Self = Self(PgType::TstzRange); + pub(crate) const TSTZ_RANGE_ARRAY: Self = Self(PgType::TstzRangeArray); + + pub(crate) const DATE_RANGE: Self = Self(PgType::DateRange); + pub(crate) const DATE_RANGE_ARRAY: Self = Self(PgType::DateRangeArray); + + pub(crate) const INT8_RANGE: Self = Self(PgType::Int8Range); + pub(crate) const INT8_RANGE_ARRAY: Self = Self(PgType::Int8RangeArray); + + // + // pseudo types + // https://www.postgresql.org/docs/9.3/datatype-pseudo.html + // + + pub(crate) const VOID: Self = Self(PgType::Void); +} + +impl Display for PgTypeInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.pad(self.name()) + } +} + +impl PartialEq for PgType { + fn eq(&self, other: &PgType) -> bool { + self.eq_impl(other, true) + } +} + +/// Check type names for equality, respecting Postgres' case sensitivity rules for identifiers. +/// +/// https://www.postgresql.org/docs/current/sql-syntax-lexical.html#SQL-SYNTAX-IDENTIFIERS +fn name_eq(name1: &str, name2: &str) -> bool { + // Cop-out of processing Unicode escapes by just using string equality. + if name1.starts_with("U&") { + // If `name2` doesn't start with `U&` this will automatically be `false`. + return name1 == name2; + } + + let mut chars1 = identifier_chars(name1); + let mut chars2 = identifier_chars(name2); + + while let (Some(a), Some(b)) = (chars1.next(), chars2.next()) { + if !a.eq(&b) { + return false; + } + } + + chars1.next().is_none() && chars2.next().is_none() +} + +struct IdentifierChar { + ch: char, + case_sensitive: bool, +} + +impl IdentifierChar { + fn eq(&self, other: &Self) -> bool { + if self.case_sensitive || other.case_sensitive { + self.ch == other.ch + } else { + self.ch.eq_ignore_ascii_case(&other.ch) + } + } +} + +/// Return an iterator over all significant characters of an identifier. +/// +/// Ignores non-escaped quotation marks. +fn identifier_chars(ident: &str) -> impl Iterator + '_ { + let mut case_sensitive = false; + let mut last_char_quote = false; + + ident.chars().filter_map(move |ch| { + if ch == '"' { + if last_char_quote { + last_char_quote = false; + } else { + last_char_quote = true; + return None; + } + } else if last_char_quote { + last_char_quote = false; + case_sensitive = !case_sensitive; + } + + Some(IdentifierChar { ch, case_sensitive }) + }) +} + +#[test] +fn test_name_eq() { + let test_values = [ + ("foo", "foo", true), + ("foo", "Foo", true), + ("foo", "FOO", true), + ("foo", r#""foo""#, true), + ("foo", r#""Foo""#, false), + ("foo", "foo.foo", false), + ("foo.foo", "foo.foo", true), + ("foo.foo", "foo.Foo", true), + ("foo.foo", "foo.FOO", true), + ("foo.foo", "Foo.foo", true), + ("foo.foo", "Foo.Foo", true), + ("foo.foo", "FOO.FOO", true), + ("foo.foo", "foo", false), + ("foo.foo", r#"foo."foo""#, true), + ("foo.foo", r#"foo."Foo""#, false), + ("foo.foo", r#"foo."FOO""#, false), + ]; + + for (left, right, eq) in test_values { + assert_eq!( + name_eq(left, right), + eq, + "failed check for name_eq({left:?}, {right:?})" + ); + assert_eq!( + name_eq(right, left), + eq, + "failed check for name_eq({right:?}, {left:?})" + ); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/array.rs b/src-tauri/vendor/sqlx-postgres/src/types/array.rs new file mode 100644 index 00000000..9b8be634 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/array.rs @@ -0,0 +1,356 @@ +use sqlx_core::bytes::Buf; +use sqlx_core::types::Text; +use std::borrow::Cow; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::type_info::PgType; +use crate::types::Oid; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +/// Provides information necessary to encode and decode Postgres arrays as compatible Rust types. +/// +/// Implementing this trait for some type `T` enables relevant `Type`,`Encode` and `Decode` impls +/// for `Vec`, `&[T]` (slices), `[T; N]` (arrays), etc. +/// +/// ### Note: `#[derive(sqlx::Type)]` +/// If you have the `postgres` feature enabled, `#[derive(sqlx::Type)]` will also generate +/// an impl of this trait for your type if your wrapper is marked `#[sqlx(transparent)]`: +/// +/// ```rust,ignore +/// #[derive(sqlx::Type)] +/// #[sqlx(transparent)] +/// struct UserId(i64); +/// +/// let user_ids: Vec = sqlx::query_scalar("select '{ 123, 456 }'::int8[]") +/// .fetch(&mut pg_connection) +/// .await?; +/// ``` +/// +/// However, this may cause an error if the type being wrapped does not implement `PgHasArrayType`, +/// e.g. `Vec` itself, because we don't currently support multidimensional arrays: +/// +/// ```rust,ignore +/// #[derive(sqlx::Type)] // ERROR: `Vec` does not implement `PgHasArrayType` +/// #[sqlx(transparent)] +/// struct UserIds(Vec); +/// ``` +/// +/// To remedy this, add `#[sqlx(no_pg_array)]`, which disables the generation +/// of the `PgHasArrayType` impl: +/// +/// ```rust,ignore +/// #[derive(sqlx::Type)] +/// #[sqlx(transparent, no_pg_array)] +/// struct UserIds(Vec); +/// ``` +/// +/// See [the documentation of `Type`][Type] for more details. +pub trait PgHasArrayType { + fn array_type_info() -> PgTypeInfo; + fn array_compatible(ty: &PgTypeInfo) -> bool { + *ty == Self::array_type_info() + } +} + +impl PgHasArrayType for &T +where + T: PgHasArrayType, +{ + fn array_type_info() -> PgTypeInfo { + T::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + T::array_compatible(ty) + } +} + +impl PgHasArrayType for Option +where + T: PgHasArrayType, +{ + fn array_type_info() -> PgTypeInfo { + T::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + T::array_compatible(ty) + } +} + +impl PgHasArrayType for Text { + fn array_type_info() -> PgTypeInfo { + String::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + String::array_compatible(ty) + } +} + +impl Type for [T] +where + T: PgHasArrayType, +{ + fn type_info() -> PgTypeInfo { + T::array_type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + T::array_compatible(ty) + } +} + +impl Type for Vec +where + T: PgHasArrayType, +{ + fn type_info() -> PgTypeInfo { + T::array_type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + T::array_compatible(ty) + } +} + +impl Type for [T; N] +where + T: PgHasArrayType, +{ + fn type_info() -> PgTypeInfo { + T::array_type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + T::array_compatible(ty) + } +} + +impl<'q, T> Encode<'q, Postgres> for Vec +where + for<'a> &'a [T]: Encode<'q, Postgres>, + T: Encode<'q, Postgres>, +{ + #[inline] + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.as_slice().encode_by_ref(buf) + } +} + +impl<'q, T, const N: usize> Encode<'q, Postgres> for [T; N] +where + for<'a> &'a [T]: Encode<'q, Postgres>, + T: Encode<'q, Postgres>, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.as_slice().encode_by_ref(buf) + } +} + +impl<'q, T> Encode<'q, Postgres> for &'_ [T] +where + T: Encode<'q, Postgres> + Type, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + let type_info = self + .first() + .and_then(Encode::produces) + .unwrap_or_else(T::type_info); + + buf.extend(&1_i32.to_be_bytes()); // number of dimensions + buf.extend(&0_i32.to_be_bytes()); // flags + + // element type + match type_info.0 { + PgType::DeclareWithName(name) => buf.patch_type_by_name(&name), + PgType::DeclareArrayOf(array) => buf.patch_array_type(array), + + ty => { + buf.extend(&ty.oid().0.to_be_bytes()); + } + } + + let array_len = i32::try_from(self.len()).map_err(|_| { + format!( + "encoded array length is too large for Postgres: {}", + self.len() + ) + })?; + + buf.extend(array_len.to_be_bytes()); // len + buf.extend(&1_i32.to_be_bytes()); // lower bound + + for element in self.iter() { + buf.encode(element)?; + } + + Ok(IsNull::No) + } +} + +impl<'r, T, const N: usize> Decode<'r, Postgres> for [T; N] +where + T: for<'a> Decode<'a, Postgres> + Type, +{ + fn decode(value: PgValueRef<'r>) -> Result { + // This could be done more efficiently by refactoring the Vec decoding below so that it can + // be used for arrays and Vec. + let vec: Vec = Decode::decode(value)?; + let array: [T; N] = vec.try_into().map_err(|_| "wrong number of elements")?; + Ok(array) + } +} + +impl<'r, T> Decode<'r, Postgres> for Vec +where + T: for<'a> Decode<'a, Postgres> + Type, +{ + fn decode(value: PgValueRef<'r>) -> Result { + let format = value.format(); + + match format { + PgValueFormat::Binary => { + // https://github.com/postgres/postgres/blob/a995b371ae29de2d38c4b7881cf414b1560e9746/src/backend/utils/adt/arrayfuncs.c#L1548 + + let mut buf = value.as_bytes()?; + + // number of dimensions in the array + let ndim = buf.get_i32(); + + if ndim == 0 { + // zero dimensions is an empty array + return Ok(Vec::new()); + } + + if ndim != 1 { + return Err(format!("encountered an array of {ndim} dimensions; only one-dimensional arrays are supported").into()); + } + + // appears to have been used in the past to communicate potential NULLS + // but reading source code back through our supported postgres versions (9.5+) + // this is never used for anything + let _flags = buf.get_i32(); + + // the OID of the element + let element_type_oid = Oid(buf.get_u32()); + let element_type_info: PgTypeInfo = PgTypeInfo::try_from_oid(element_type_oid) + .or_else(|| value.type_info.try_array_element().map(Cow::into_owned)) + .ok_or_else(|| { + BoxDynError::from(format!( + "failed to resolve array element type for oid {}", + element_type_oid.0 + )) + })?; + + // length of the array axis + let len = buf.get_i32(); + + let len = usize::try_from(len) + .map_err(|_| format!("overflow converting array len ({len}) to usize"))?; + + // the lower bound, we only support arrays starting from "1" + let lower = buf.get_i32(); + + if lower != 1 { + return Err(format!("encountered an array with a lower bound of {lower} in the first dimension; only arrays starting at one are supported").into()); + } + + let mut elements = Vec::with_capacity(len); + + for _ in 0..len { + let value_ref = PgValueRef::get(&mut buf, format, element_type_info.clone())?; + + elements.push(T::decode(value_ref)?); + } + + Ok(elements) + } + + PgValueFormat::Text => { + // no type is provided from the database for the element + let element_type_info = T::type_info(); + + let s = value.as_str()?; + + // https://github.com/postgres/postgres/blob/a995b371ae29de2d38c4b7881cf414b1560e9746/src/backend/utils/adt/arrayfuncs.c#L718 + + // trim the wrapping braces + let s = &s[1..(s.len() - 1)]; + + if s.is_empty() { + // short-circuit empty arrays up here + return Ok(Vec::new()); + } + + // NOTE: Nearly *all* types use ',' as the sequence delimiter. Yes, there is one + // that does not. The BOX (not PostGIS) type uses ';' as a delimiter. + + // TODO: When we add support for BOX we need to figure out some way to make the + // delimiter selection + + let delimiter = ','; + let mut done = false; + let mut in_quotes = false; + let mut in_escape = false; + let mut value = String::with_capacity(10); + let mut chars = s.chars(); + let mut elements = Vec::with_capacity(4); + + while !done { + loop { + match chars.next() { + Some(ch) => match ch { + _ if in_escape => { + value.push(ch); + in_escape = false; + } + + '"' => { + in_quotes = !in_quotes; + } + + '\\' => { + in_escape = true; + } + + _ if ch == delimiter && !in_quotes => { + break; + } + + _ => { + value.push(ch); + } + }, + + None => { + done = true; + break; + } + } + } + + let value_opt = if value == "NULL" { + None + } else { + Some(value.as_bytes()) + }; + + elements.push(T::decode(PgValueRef { + value: value_opt, + row: None, + type_info: element_type_info.clone(), + format, + })?); + + value.clear(); + } + + Ok(elements) + } + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal-range.md b/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal-range.md new file mode 100644 index 00000000..5d4ee502 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal-range.md @@ -0,0 +1,20 @@ +#### Note: `BigDecimal` Has a Larger Range than `NUMERIC` +`BigDecimal` can represent values with a far, far greater range than the `NUMERIC` type in Postgres can. + +`NUMERIC` is limited to 131,072 digits before the decimal point, and 16,384 digits after it. +See [Section 8.1, Numeric Types] of the Postgres manual for details. + +Meanwhile, `BigDecimal` can theoretically represent a value with an arbitrary number of decimal digits, albeit +with a maximum of 263 significant figures. + +Because encoding in the current API design _must_ be infallible, +when attempting to encode a `BigDecimal` that cannot fit in the wire representation of `NUMERIC`, +SQLx may instead encode a sentinel value that falls outside the allowed range but is still representable. + +This will cause the query to return a `DatabaseError` with code `22P03` (`invalid_binary_representation`) +and the error message `invalid scale in external "numeric" value` (though this may be subject to change). + +However, `BigDecimal` should be able to decode any `NUMERIC` value except `NaN`, +for which it has no representation. + +[Section 8.1, Numeric Types]: https://www.postgresql.org/docs/current/datatype-numeric.html diff --git a/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal.rs b/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal.rs new file mode 100644 index 00000000..869f8507 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/bigdecimal.rs @@ -0,0 +1,477 @@ +use bigdecimal::BigDecimal; +use num_bigint::{BigInt, Sign}; +use std::cmp; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::numeric::{PgNumeric, PgNumericSign}; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl Type for BigDecimal { + fn type_info() -> PgTypeInfo { + PgTypeInfo::NUMERIC + } +} + +impl PgHasArrayType for BigDecimal { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::NUMERIC_ARRAY + } +} + +impl TryFrom for BigDecimal { + type Error = BoxDynError; + + fn try_from(numeric: PgNumeric) -> Result { + Self::try_from(&numeric) + } +} + +impl TryFrom<&'_ PgNumeric> for BigDecimal { + type Error = BoxDynError; + + fn try_from(numeric: &'_ PgNumeric) -> Result { + let (digits, sign, weight) = match *numeric { + PgNumeric::Number { + ref digits, + sign, + weight, + .. + } => (digits, sign, weight), + + PgNumeric::NotANumber => { + return Err("BigDecimal does not support NaN values".into()); + } + }; + + if digits.is_empty() { + // Postgres returns an empty digit array for 0 but BigInt expects at least one zero + return Ok(0u64.into()); + } + + let sign = match sign { + PgNumericSign::Positive => Sign::Plus, + PgNumericSign::Negative => Sign::Minus, + }; + + // weight is 0 if the decimal point falls after the first base-10000 digit + // + // `Vec` capacity cannot exceed `isize::MAX` bytes, so this cast can't wrap in practice. + #[allow(clippy::cast_possible_wrap)] + let scale = (digits.len() as i64 - weight as i64 - 1) * 4; + + // no optimized algorithm for base-10 so use base-100 for faster processing + let mut cents = Vec::with_capacity(digits.len() * 2); + + #[allow( + clippy::cast_possible_truncation, + clippy::cast_possible_wrap, + clippy::cast_sign_loss + )] + for (i, &digit) in digits.iter().enumerate() { + if !PgNumeric::is_valid_digit(digit) { + return Err(format!( + "PgNumeric to BigDecimal: {i}th digit is out of range {digit}" + ) + .into()); + } + + cents.push((digit / 100) as u8); + cents.push((digit % 100) as u8); + } + + let bigint = BigInt::from_radix_be(sign, ¢s, 100) + .ok_or("PgNumeric contained an out-of-range digit")?; + + Ok(BigDecimal::new(bigint, scale)) + } +} + +impl TryFrom<&'_ BigDecimal> for PgNumeric { + type Error = BoxDynError; + + fn try_from(decimal: &BigDecimal) -> Result { + let base_10_to_10000 = |chunk: &[u8]| chunk.iter().fold(0i16, |a, &d| a * 10 + d as i16); + + // NOTE: this unfortunately copies the BigInt internally + let (integer, exp) = decimal.as_bigint_and_exponent(); + + // this routine is specifically optimized for base-10 + // FIXME: is there a way to iterate over the digits to avoid the Vec allocation + let (sign, base_10) = integer.to_radix_be(10); + + let base_10_len = i64::try_from(base_10.len()).map_err(|_| { + format!( + "BigDecimal base-10 length out of range for PgNumeric: {}", + base_10.len() + ) + })?; + + // weight is positive power of 10000 + // exp is the negative power of 10 + let weight_10 = base_10_len - exp; + + // scale is only nonzero when we have fractional digits + // since `exp` is the _negative_ decimal exponent, it tells us + // exactly what our scale should be + let scale: i16 = cmp::max(0, exp).try_into()?; + + // there's an implicit +1 offset in the interpretation + let weight: i16 = if weight_10 <= 0 { + weight_10 / 4 - 1 + } else { + // the `-1` is a fix for an off by 1 error (4 digits should still be 0 weight) + (weight_10 - 1) / 4 + } + .try_into()?; + + let digits_len = if base_10.len() % 4 != 0 { + base_10.len() / 4 + 1 + } else { + base_10.len() / 4 + }; + + // For efficiency, we want to process the base-10 digits in chunks of 4, + // but that means we need to deal with the non-divisible remainder first. + let offset = weight_10.rem_euclid(4); + + // Do a checked conversion to the smallest integer, + // so we can widen arbitrarily without triggering lints. + let offset = u8::try_from(offset).unwrap_or_else(|_| { + panic!("BUG: `offset` should be in the range [0, 4) but is {offset}") + }); + + let mut digits = Vec::with_capacity(digits_len); + + if let Some(first) = base_10.get(..offset as usize) { + if !first.is_empty() { + digits.push(base_10_to_10000(first)); + } + } else if offset != 0 { + // If we didn't hit the `if let Some` branch, + // then `base_10.len()` must strictly be smaller + #[allow(clippy::cast_possible_truncation)] + let power = (offset as usize - base_10.len()) as u32; + + digits.push(base_10_to_10000(&base_10) * 10i16.pow(power)); + } + + if let Some(rest) = base_10.get(offset as usize..) { + // `chunk.len()` is always between 1 and 4 + #[allow(clippy::cast_possible_truncation)] + digits.extend( + rest.chunks(4) + .map(|chunk| base_10_to_10000(chunk) * 10i16.pow(4 - chunk.len() as u32)), + ); + } + + while let Some(&0) = digits.last() { + digits.pop(); + } + + Ok(PgNumeric::Number { + sign: sign_to_pg(sign), + scale, + weight, + digits, + }) + } +} + +#[doc=include_str!("bigdecimal-range.md")] +impl Encode<'_, Postgres> for BigDecimal { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + PgNumeric::try_from(self)?.encode(buf)?; + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + PgNumeric::size_hint(self.digits()) + } +} + +/// ### Note: `NaN` +/// `BigDecimal` has a greater range than `NUMERIC` (see the corresponding `Encode` impl for details) +/// but cannot represent `NaN`, so decoding may return an error. +impl Decode<'_, Postgres> for BigDecimal { + fn decode(value: PgValueRef<'_>) -> Result { + match value.format() { + PgValueFormat::Binary => PgNumeric::decode(value.as_bytes()?)?.try_into(), + PgValueFormat::Text => Ok(value.as_str()?.parse::()?), + } + } +} + +fn sign_to_pg(sign: Sign) -> PgNumericSign { + match sign { + Sign::Plus | Sign::NoSign => PgNumericSign::Positive, + Sign::Minus => PgNumericSign::Negative, + } +} + +#[cfg(test)] +mod bigdecimal_to_pgnumeric { + use super::{BigDecimal, PgNumeric, PgNumericSign}; + use std::convert::TryFrom; + + #[test] + fn zero() { + let zero: BigDecimal = "0".parse().unwrap(); + + assert_eq!( + PgNumeric::try_from(&zero).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![] + } + ); + } + + #[test] + fn one() { + let one: BigDecimal = "1".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![1] + } + ); + } + + #[test] + fn ten() { + let ten: BigDecimal = "10".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&ten).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![10] + } + ); + } + + #[test] + fn one_hundred() { + let one_hundred: BigDecimal = "100".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_hundred).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![100] + } + ); + } + + #[test] + fn ten_thousand() { + // BigDecimal doesn't normalize here + let ten_thousand: BigDecimal = "10000".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&ten_thousand).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1] + } + ); + } + + #[test] + fn two_digits() { + let two_digits: BigDecimal = "12345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&two_digits).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1, 2345] + } + ); + } + + #[test] + fn one_tenth() { + let one_tenth: BigDecimal = "0.1".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_tenth).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 1, + weight: -1, + digits: vec![1000] + } + ); + } + + #[test] + fn one_hundredth() { + let one_hundredth: BigDecimal = "0.01".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_hundredth).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 2, + weight: -1, + digits: vec![100] + } + ); + } + + #[test] + fn twelve_thousandths() { + let twelve_thousandths: BigDecimal = "0.012".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&twelve_thousandths).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 3, + weight: -1, + digits: vec![120] + } + ); + } + + #[test] + fn decimal_1() { + let decimal: BigDecimal = "1.2345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 4, + weight: 0, + digits: vec![1, 2345] + } + ); + } + + #[test] + fn decimal_2() { + let decimal: BigDecimal = "0.12345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: -1, + digits: vec![1234, 5000] + } + ); + } + + #[test] + fn decimal_3() { + let decimal: BigDecimal = "0.01234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: -1, + digits: vec![0123, 4000] + } + ); + } + + #[test] + fn decimal_4() { + let decimal: BigDecimal = "12345.67890".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: 1, + digits: vec![1, 2345, 6789] + } + ); + } + + #[test] + fn one_digit_decimal() { + let one_digit_decimal: BigDecimal = "0.00001234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_digit_decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 8, + weight: -2, + digits: vec![1234] + } + ); + } + + #[test] + fn issue_423_four_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let four_digit: BigDecimal = "1234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&four_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![1234] + } + ); + } + + #[test] + fn issue_423_negative_four_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let negative_four_digit: BigDecimal = "-1234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&negative_four_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Negative, + scale: 0, + weight: 0, + digits: vec![1234] + } + ); + } + + #[test] + fn issue_423_eight_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let eight_digit: BigDecimal = "12345678".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&eight_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1234, 5678] + } + ); + } + + #[test] + fn issue_423_negative_eight_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let negative_eight_digit: BigDecimal = "-12345678".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&negative_eight_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Negative, + scale: 0, + weight: 1, + digits: vec![1234, 5678] + } + ); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/bit_vec.rs b/src-tauri/vendor/sqlx-postgres/src/types/bit_vec.rs new file mode 100644 index 00000000..b519a5f2 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/bit_vec.rs @@ -0,0 +1,99 @@ +use crate::arguments::value_size_int4_checked; +use crate::{ + decode::Decode, + encode::{Encode, IsNull}, + error::BoxDynError, + types::Type, + PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres, +}; +use bit_vec::BitVec; +use sqlx_core::bytes::Buf; +use std::{io, mem}; + +impl Type for BitVec { + fn type_info() -> PgTypeInfo { + PgTypeInfo::VARBIT + } + + fn compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::BIT || *ty == PgTypeInfo::VARBIT + } +} + +impl PgHasArrayType for BitVec { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::VARBIT_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::BIT_ARRAY || *ty == PgTypeInfo::VARBIT_ARRAY + } +} + +impl Encode<'_, Postgres> for BitVec { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + let len = value_size_int4_checked(self.len())?; + + buf.extend(len.to_be_bytes()); + buf.extend(self.to_bytes()); + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + self.len() + } +} + +impl Decode<'_, Postgres> for BitVec { + fn decode(value: PgValueRef<'_>) -> Result { + match value.format() { + PgValueFormat::Binary => { + let mut bytes = value.as_bytes()?; + let len = bytes.get_i32(); + + let len = usize::try_from(len).map_err(|_| format!("invalid VARBIT len: {len}"))?; + + // The smallest amount of data we can read is one byte + let bytes_len = (len + 7) / 8; + + if bytes.remaining() != bytes_len { + Err(io::Error::new( + io::ErrorKind::InvalidData, + "VARBIT length mismatch.", + ))?; + } + + let mut bitvec = BitVec::from_bytes(bytes); + + // Chop off zeroes from the back. We get bits in bytes, so if + // our bitvec is not in full bytes, extra zeroes are added to + // the end. + while bitvec.len() > len { + bitvec.pop(); + } + + Ok(bitvec) + } + PgValueFormat::Text => { + let s = value.as_str()?; + let mut bit_vec = BitVec::with_capacity(s.len()); + + for c in s.chars() { + match c { + '0' => bit_vec.push(false), + '1' => bit_vec.push(true), + _ => { + Err(io::Error::new( + io::ErrorKind::InvalidData, + "VARBIT data contains other characters than 1 or 0.", + ))?; + } + } + } + + Ok(bit_vec) + } + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/bool.rs b/src-tauri/vendor/sqlx-postgres/src/types/bool.rs new file mode 100644 index 00000000..8c3e140d --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/bool.rs @@ -0,0 +1,42 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl Type for bool { + fn type_info() -> PgTypeInfo { + PgTypeInfo::BOOL + } +} + +impl PgHasArrayType for bool { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::BOOL_ARRAY + } +} + +impl Encode<'_, Postgres> for bool { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.push(*self as u8); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for bool { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => value.as_bytes()?[0] != 0, + + PgValueFormat::Text => match value.as_str()? { + "t" => true, + "f" => false, + + s => { + return Err(format!("unexpected value {s:?} for boolean").into()); + } + }, + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/bytes.rs b/src-tauri/vendor/sqlx-postgres/src/types/bytes.rs new file mode 100644 index 00000000..45968837 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/bytes.rs @@ -0,0 +1,112 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl PgHasArrayType for u8 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::BYTEA + } +} + +impl PgHasArrayType for &'_ [u8] { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::BYTEA_ARRAY + } +} + +impl PgHasArrayType for Box<[u8]> { + fn array_type_info() -> PgTypeInfo { + <[&[u8]] as Type>::type_info() + } +} + +impl PgHasArrayType for Vec { + fn array_type_info() -> PgTypeInfo { + <[&[u8]] as Type>::type_info() + } +} + +impl PgHasArrayType for [u8; N] { + fn array_type_info() -> PgTypeInfo { + <[&[u8]] as Type>::type_info() + } +} + +impl Encode<'_, Postgres> for &'_ [u8] { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend_from_slice(self); + + Ok(IsNull::No) + } +} + +impl Encode<'_, Postgres> for Box<[u8]> { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&[u8] as Encode>::encode(self.as_ref(), buf) + } +} + +impl Encode<'_, Postgres> for Vec { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&[u8] as Encode>::encode(self, buf) + } +} + +impl Encode<'_, Postgres> for [u8; N] { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&[u8] as Encode>::encode(self.as_slice(), buf) + } +} + +impl<'r> Decode<'r, Postgres> for &'r [u8] { + fn decode(value: PgValueRef<'r>) -> Result { + match value.format() { + PgValueFormat::Binary => value.as_bytes(), + PgValueFormat::Text => { + Err("unsupported decode to `&[u8]` of BYTEA in a simple query; use a prepared query or decode to `Vec`".into()) + } + } + } +} + +fn text_hex_decode_input(value: PgValueRef<'_>) -> Result<&[u8], BoxDynError> { + // BYTEA is formatted as \x followed by hex characters + value + .as_bytes()? + .strip_prefix(b"\\x") + .ok_or("text does not start with \\x") + .map_err(Into::into) +} + +impl Decode<'_, Postgres> for Box<[u8]> { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => Box::from(value.as_bytes()?), + PgValueFormat::Text => Box::from(hex::decode(text_hex_decode_input(value)?)?), + }) + } +} + +impl Decode<'_, Postgres> for Vec { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => value.as_bytes()?.to_owned(), + PgValueFormat::Text => hex::decode(text_hex_decode_input(value)?)?, + }) + } +} + +impl Decode<'_, Postgres> for [u8; N] { + fn decode(value: PgValueRef<'_>) -> Result { + let mut bytes = [0u8; N]; + match value.format() { + PgValueFormat::Binary => { + bytes = value.as_bytes()?.try_into()?; + } + PgValueFormat::Text => hex::decode_to_slice(text_hex_decode_input(value)?, &mut bytes)?, + }; + Ok(bytes) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/chrono/date.rs b/src-tauri/vendor/sqlx-postgres/src/types/chrono/date.rs new file mode 100644 index 00000000..0327d5c4 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/chrono/date.rs @@ -0,0 +1,64 @@ +use std::mem; + +use chrono::{NaiveDate, TimeDelta}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl Type for NaiveDate { + fn type_info() -> PgTypeInfo { + PgTypeInfo::DATE + } +} + +impl PgHasArrayType for NaiveDate { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::DATE_ARRAY + } +} + +impl Encode<'_, Postgres> for NaiveDate { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // DATE is encoded as the days since epoch + let days: i32 = (*self - postgres_epoch_date()) + .num_days() + .try_into() + .map_err(|_| { + format!("value {self:?} would overflow binary encoding for Postgres DATE") + })?; + + Encode::::encode(days, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for NaiveDate { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // DATE is encoded as the days since epoch + let days: i32 = Decode::::decode(value)?; + + let days = TimeDelta::try_days(days.into()) + .unwrap_or_else(|| { + unreachable!("BUG: days ({days}) as `i32` multiplied into seconds should not overflow `i64`") + }); + + postgres_epoch_date() + days + } + + PgValueFormat::Text => NaiveDate::parse_from_str(value.as_str()?, "%Y-%m-%d")?, + }) + } +} + +#[inline] +fn postgres_epoch_date() -> NaiveDate { + NaiveDate::from_ymd_opt(2000, 1, 1).expect("expected 2000-01-01 to be a valid NaiveDate") +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/chrono/datetime.rs b/src-tauri/vendor/sqlx-postgres/src/types/chrono/datetime.rs new file mode 100644 index 00000000..77f900d4 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/chrono/datetime.rs @@ -0,0 +1,133 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use chrono::{ + DateTime, Duration, FixedOffset, Local, NaiveDate, NaiveDateTime, Offset, TimeZone, Utc, +}; +use std::mem; + +impl Type for NaiveDateTime { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMP + } +} + +impl Type for DateTime { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMPTZ + } +} + +impl PgHasArrayType for NaiveDateTime { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMP_ARRAY + } +} + +impl PgHasArrayType for DateTime { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMPTZ_ARRAY + } +} + +impl Encode<'_, Postgres> for NaiveDateTime { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // TIMESTAMP is encoded as the microseconds since the epoch + let micros = (*self - postgres_epoch_datetime()) + .num_microseconds() + .ok_or_else(|| format!("NaiveDateTime out of range for Postgres: {self:?}"))?; + + Encode::::encode(micros, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for NaiveDateTime { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // TIMESTAMP is encoded as the microseconds since the epoch + let us = Decode::::decode(value)?; + postgres_epoch_datetime() + Duration::microseconds(us) + } + + PgValueFormat::Text => { + let s = value.as_str()?; + NaiveDateTime::parse_from_str( + s, + if s.contains('+') { + // Contains a time-zone specifier + // This is given for timestamptz for some reason + // Postgres already guarantees this to always be UTC + "%Y-%m-%d %H:%M:%S%.f%#z" + } else { + "%Y-%m-%d %H:%M:%S%.f" + }, + )? + } + }) + } +} + +impl Encode<'_, Postgres> for DateTime { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + Encode::::encode(self.naive_utc(), buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for DateTime { + fn decode(value: PgValueRef<'r>) -> Result { + let fixed = as Decode>::decode(value)?; + Ok(Local.from_utc_datetime(&fixed.naive_utc())) + } +} + +impl<'r> Decode<'r, Postgres> for DateTime { + fn decode(value: PgValueRef<'r>) -> Result { + let fixed = as Decode>::decode(value)?; + Ok(Utc.from_utc_datetime(&fixed.naive_utc())) + } +} + +impl<'r> Decode<'r, Postgres> for DateTime { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + let naive = >::decode(value)?; + Utc.fix().from_utc_datetime(&naive) + } + + PgValueFormat::Text => { + let s = value.as_str()?; + DateTime::parse_from_str( + s, + if s.contains('+') || s.contains('-') { + // Contains a time-zone specifier + // This is given for timestamptz for some reason + // Postgres already guarantees this to always be UTC + "%Y-%m-%d %H:%M:%S%.f%#z" + } else { + "%Y-%m-%d %H:%M:%S%.f" + }, + )? + } + }) + } +} + +#[inline] +fn postgres_epoch_datetime() -> NaiveDateTime { + NaiveDate::from_ymd_opt(2000, 1, 1) + .expect("expected 2000-01-01 to be a valid NaiveDate") + .and_hms_opt(0, 0, 0) + .expect("expected 2000-01-01T00:00:00 to be a valid NaiveDateTime") +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/chrono/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/chrono/mod.rs new file mode 100644 index 00000000..bd27c4d2 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/chrono/mod.rs @@ -0,0 +1,3 @@ +mod date; +mod datetime; +mod time; diff --git a/src-tauri/vendor/sqlx-postgres/src/types/chrono/time.rs b/src-tauri/vendor/sqlx-postgres/src/types/chrono/time.rs new file mode 100644 index 00000000..ca66f389 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/chrono/time.rs @@ -0,0 +1,58 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use chrono::{Duration, NaiveTime}; +use std::mem; + +impl Type for NaiveTime { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIME + } +} + +impl PgHasArrayType for NaiveTime { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIME_ARRAY + } +} + +impl Encode<'_, Postgres> for NaiveTime { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // TIME is encoded as the microseconds since midnight + let micros = (*self - NaiveTime::default()) + .num_microseconds() + .ok_or_else(|| format!("Time out of range for PostgreSQL: {self}"))?; + + Encode::::encode(micros, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for NaiveTime { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // TIME is encoded as the microseconds since midnight + let us: i64 = Decode::::decode(value)?; + NaiveTime::default() + Duration::microseconds(us) + } + + PgValueFormat::Text => NaiveTime::parse_from_str(value.as_str()?, "%H:%M:%S%.f")?, + }) + } +} + +#[test] +fn check_naive_time_default_is_midnight() { + // Just a canary in case this changes. + assert_eq!( + NaiveTime::from_hms_opt(0, 0, 0), + Some(NaiveTime::default()), + "implementation assumes `NaiveTime::default()` equals midnight" + ); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/citext.rs b/src-tauri/vendor/sqlx-postgres/src/types/citext.rs new file mode 100644 index 00000000..c0316ac8 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/citext.rs @@ -0,0 +1,106 @@ +use crate::types::array_compatible; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueRef, Postgres}; +use sqlx_core::decode::Decode; +use sqlx_core::encode::{Encode, IsNull}; +use sqlx_core::error::BoxDynError; +use sqlx_core::types::Type; +use std::fmt; +use std::fmt::{Debug, Display, Formatter}; +use std::ops::Deref; +use std::str::FromStr; + +/// Case-insensitive text (`citext`) support for Postgres. +/// +/// Note that SQLx considers the `citext` type to be compatible with `String` +/// and its various derivatives, so direct usage of this type is generally unnecessary. +/// +/// However, it may be needed, for example, when binding a `citext[]` array, +/// as Postgres will generally not accept a `text[]` array (mapped from `Vec`) in its place. +/// +/// See [the Postgres manual, Appendix F, Section 10][PG.F.10] for details on using `citext`. +/// +/// [PG.F.10]: https://www.postgresql.org/docs/current/citext.html +/// +/// ### Note: Extension Required +/// The `citext` extension is not enabled by default in Postgres. You will need to do so explicitly: +/// +/// ```ignore +/// CREATE EXTENSION IF NOT EXISTS "citext"; +/// ``` +/// +/// ### Note: `PartialEq` is Case-Sensitive +/// This type derives `PartialEq` which forwards to the implementation on `String`, which +/// is case-sensitive. This impl exists mainly for testing. +/// +/// To properly emulate the case-insensitivity of `citext` would require use of locale-aware +/// functions in `libc`, and even then would require querying the locale of the database server +/// and setting it locally, which is unsafe. +#[derive(Clone, Debug, Default, PartialEq)] +pub struct PgCiText(pub String); + +impl Type for PgCiText { + fn type_info() -> PgTypeInfo { + // Since `citext` is enabled by an extension, it does not have a stable OID. + PgTypeInfo::with_name("citext") + } + + fn compatible(ty: &PgTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Deref for PgCiText { + type Target = str; + + fn deref(&self) -> &Self::Target { + self.0.as_str() + } +} + +impl From for PgCiText { + fn from(value: String) -> Self { + Self(value) + } +} + +impl From for String { + fn from(value: PgCiText) -> Self { + value.0 + } +} + +impl FromStr for PgCiText { + type Err = core::convert::Infallible; + + fn from_str(s: &str) -> Result { + Ok(PgCiText(s.parse()?)) + } +} + +impl Display for PgCiText { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + +impl PgHasArrayType for PgCiText { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_citext") + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + array_compatible::<&str>(ty) + } +} + +impl Encode<'_, Postgres> for PgCiText { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&str as Encode>::encode(&**self, buf) + } +} + +impl Decode<'_, Postgres> for PgCiText { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(PgCiText(value.as_str()?.to_owned())) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/cube.rs b/src-tauri/vendor/sqlx-postgres/src/types/cube.rs new file mode 100644 index 00000000..cc2a0160 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/cube.rs @@ -0,0 +1,537 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use sqlx_core::Error; +use std::mem; +use std::str::FromStr; + +const BYTE_WIDTH: usize = 8; + +/// +const MAX_DIMENSIONS: usize = 100; + +const IS_POINT_FLAG: u32 = 1 << 31; + +// FIXME(breaking): these variants are confusingly named and structured +// consider changing them or making this an opaque wrapper around `Vec` +#[derive(Debug, Clone, PartialEq)] +pub enum PgCube { + /// A one-dimensional point. + // FIXME: `Point1D(f64)` + Point(f64), + /// An N-dimensional point ("represented internally as a zero-volume cube"). + // FIXME: `PointND(f64)` + ZeroVolume(Vec), + + /// A one-dimensional interval with starting and ending points. + // FIXME: `Interval1D { start: f64, end: f64 }` + OneDimensionInterval(f64, f64), + + // FIXME: add `Cube3D { lower_left: [f64; 3], upper_right: [f64; 3] }`? + /// An N-dimensional cube with points representing lower-left and upper-right corners, respectively. + // FIXME: `CubeND { lower_left: Vec, upper_right: Vec }` + MultiDimension(Vec>), +} + +#[derive(Copy, Clone, Debug, PartialEq, Eq)] +struct Header { + dimensions: usize, + is_point: bool, +} + +#[derive(Debug, thiserror::Error)] +#[error("error decoding CUBE (is_point: {is_point}, dimensions: {dimensions})")] +struct DecodeError { + is_point: bool, + dimensions: usize, + message: String, +} + +impl Type for PgCube { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("cube") + } +} + +impl PgHasArrayType for PgCube { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_cube") + } +} + +impl<'r> Decode<'r, Postgres> for PgCube { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgCube::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgCube::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgCube { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("cube")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + self.header().encoded_size() + } +} + +impl FromStr for PgCube { + type Err = Error; + + fn from_str(s: &str) -> Result { + let content = s + .trim_start_matches('(') + .trim_start_matches('[') + .trim_end_matches(')') + .trim_end_matches(']') + .replace(' ', ""); + + if !content.contains('(') && !content.contains(',') { + return parse_point(&content); + } + + if !content.contains("),(") { + return parse_zero_volume(&content); + } + + let point_vecs = content.split("),(").collect::>(); + if point_vecs.len() == 2 && !point_vecs.iter().any(|pv| pv.contains(',')) { + return parse_one_dimensional_interval(point_vecs); + } + + parse_multidimensional_interval(point_vecs) + } +} + +impl PgCube { + fn header(&self) -> Header { + match self { + PgCube::Point(..) => Header { + is_point: true, + dimensions: 1, + }, + PgCube::ZeroVolume(values) => Header { + is_point: true, + dimensions: values.len(), + }, + PgCube::OneDimensionInterval(..) => Header { + is_point: false, + dimensions: 1, + }, + PgCube::MultiDimension(multi_values) => Header { + is_point: false, + dimensions: multi_values.first().map(|arr| arr.len()).unwrap_or(0), + }, + } + } + + fn from_bytes(mut bytes: &[u8]) -> Result { + let header = Header::try_read(&mut bytes)?; + + if bytes.len() != header.data_size() { + return Err(DecodeError::new( + &header, + format!( + "expected {} bytes after header, got {}", + header.data_size(), + bytes.len() + ), + ) + .into()); + } + + match (header.is_point, header.dimensions) { + (true, 1) => Ok(PgCube::Point(bytes.get_f64())), + (true, _) => Ok(PgCube::ZeroVolume( + read_vec(&mut bytes).map_err(|e| DecodeError::new(&header, e))?, + )), + (false, 1) => Ok(PgCube::OneDimensionInterval( + bytes.get_f64(), + bytes.get_f64(), + )), + (false, _) => Ok(PgCube::MultiDimension(read_cube(&header, bytes)?)), + } + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), String> { + let header = self.header(); + + buff.reserve(header.data_size()); + + header.try_write(buff)?; + + match self { + PgCube::Point(value) => { + buff.extend_from_slice(&value.to_be_bytes()); + } + PgCube::ZeroVolume(values) => { + buff.extend(values.iter().flat_map(|v| v.to_be_bytes())); + } + PgCube::OneDimensionInterval(x, y) => { + buff.extend_from_slice(&x.to_be_bytes()); + buff.extend_from_slice(&y.to_be_bytes()); + } + PgCube::MultiDimension(multi_values) => { + if multi_values.len() != 2 { + return Err(format!("invalid CUBE value: {self:?}")); + } + + buff.extend( + multi_values + .iter() + .flat_map(|point| point.iter().flat_map(|scalar| scalar.to_be_bytes())), + ); + } + }; + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +fn read_vec(bytes: &mut &[u8]) -> Result, String> { + if bytes.len() % BYTE_WIDTH != 0 { + return Err(format!( + "data length not divisible by {BYTE_WIDTH}: {}", + bytes.len() + )); + } + + let mut out = Vec::with_capacity(bytes.len() / BYTE_WIDTH); + + while bytes.has_remaining() { + out.push(bytes.get_f64()); + } + + Ok(out) +} + +fn read_cube(header: &Header, mut bytes: &[u8]) -> Result>, String> { + if bytes.len() != header.data_size() { + return Err(format!( + "expected {} bytes, got {}", + header.data_size(), + bytes.len() + )); + } + + let mut out = Vec::with_capacity(2); + + // Expecting exactly 2 N-dimensional points + for _ in 0..2 { + let mut point = Vec::new(); + + for _ in 0..header.dimensions { + point.push(bytes.get_f64()); + } + + out.push(point); + } + + Ok(out) +} + +fn parse_float_from_str(s: &str, error_msg: &str) -> Result { + s.parse().map_err(|_| Error::Decode(error_msg.into())) +} + +fn parse_point(str: &str) -> Result { + Ok(PgCube::Point(parse_float_from_str( + str, + "Failed to parse point", + )?)) +} + +fn parse_zero_volume(content: &str) -> Result { + content + .split(',') + .map(|p| parse_float_from_str(p, "Failed to parse into zero-volume cube")) + .collect::, _>>() + .map(PgCube::ZeroVolume) +} + +fn parse_one_dimensional_interval(point_vecs: Vec<&str>) -> Result { + let x = parse_float_from_str( + &remove_parentheses(point_vecs.first().ok_or(Error::Decode( + format!("Could not decode cube interval x: {:?}", point_vecs).into(), + ))?), + "Failed to parse X in one-dimensional interval", + )?; + let y = parse_float_from_str( + &remove_parentheses(point_vecs.get(1).ok_or(Error::Decode( + format!("Could not decode cube interval y: {:?}", point_vecs).into(), + ))?), + "Failed to parse Y in one-dimensional interval", + )?; + Ok(PgCube::OneDimensionInterval(x, y)) +} + +fn parse_multidimensional_interval(point_vecs: Vec<&str>) -> Result { + point_vecs + .iter() + .map(|&point_vec| { + point_vec + .split(',') + .map(|point| { + parse_float_from_str( + &remove_parentheses(point), + "Failed to parse into multi-dimension cube", + ) + }) + .collect::, _>>() + }) + .collect::, _>>() + .map(PgCube::MultiDimension) +} + +fn remove_parentheses(s: &str) -> String { + s.trim_matches(|c| c == '(' || c == ')').to_string() +} + +impl Header { + const PACKED_WIDTH: usize = mem::size_of::(); + + fn encoded_size(&self) -> usize { + Self::PACKED_WIDTH + self.data_size() + } + + fn data_size(&self) -> usize { + if self.is_point { + self.dimensions * BYTE_WIDTH + } else { + self.dimensions * BYTE_WIDTH * 2 + } + } + + fn try_write(&self, buff: &mut PgArgumentBuffer) -> Result<(), String> { + if self.dimensions > MAX_DIMENSIONS { + return Err(format!( + "CUBE dimensionality exceeds allowed maximum ({} > {MAX_DIMENSIONS})", + self.dimensions + )); + } + + // Cannot overflow thanks to the above check. + #[allow(clippy::cast_possible_truncation)] + let mut packed = self.dimensions as u32; + + // https://github.com/postgres/postgres/blob/e3ec9dc1bf4983fcedb6f43c71ea12ee26aefc7a/contrib/cube/cubedata.h#L18-L24 + if self.is_point { + packed |= IS_POINT_FLAG; + } + + buff.extend(packed.to_be_bytes()); + + Ok(()) + } + + fn try_read(buf: &mut &[u8]) -> Result { + if buf.len() < Self::PACKED_WIDTH { + return Err(format!( + "expected CUBE data to contain at least {} bytes, got {}", + Self::PACKED_WIDTH, + buf.len() + )); + } + + let packed = buf.get_u32(); + + let is_point = packed & IS_POINT_FLAG != 0; + let dimensions = packed & !IS_POINT_FLAG; + + // can only overflow on 16-bit platforms + let dimensions = usize::try_from(dimensions) + .ok() + .filter(|&it| it <= MAX_DIMENSIONS) + .ok_or_else(|| format!("received CUBE data with higher than expected dimensionality: {dimensions} (is_point: {is_point})"))?; + + Ok(Self { + is_point, + dimensions, + }) + } +} + +impl DecodeError { + fn new(header: &Header, message: String) -> Self { + DecodeError { + is_point: header.is_point, + dimensions: header.dimensions, + message, + } + } +} + +#[cfg(test)] +mod cube_tests { + + use std::str::FromStr; + + use super::PgCube; + + const POINT_BYTES: &[u8] = &[128, 0, 0, 1, 64, 0, 0, 0, 0, 0, 0, 0]; + const ZERO_VOLUME_BYTES: &[u8] = &[ + 128, 0, 0, 2, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, + ]; + const ONE_DIMENSIONAL_INTERVAL_BYTES: &[u8] = &[ + 0, 0, 0, 1, 64, 28, 0, 0, 0, 0, 0, 0, 64, 32, 0, 0, 0, 0, 0, 0, + ]; + const MULTI_DIMENSION_2_DIM_BYTES: &[u8] = &[ + 0, 0, 0, 2, 63, 240, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, + 64, 16, 0, 0, 0, 0, 0, 0, + ]; + const MULTI_DIMENSION_3_DIM_BYTES: &[u8] = &[ + 0, 0, 0, 3, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, 64, 16, 0, 0, 0, 0, 0, 0, 64, + 20, 0, 0, 0, 0, 0, 0, 64, 24, 0, 0, 0, 0, 0, 0, 64, 28, 0, 0, 0, 0, 0, 0, + ]; + + #[test] + fn can_deserialise_point_type_byes() { + let cube = PgCube::from_bytes(POINT_BYTES).unwrap(); + assert_eq!(cube, PgCube::Point(2.)) + } + + #[test] + fn can_deserialise_point_type_str() { + let cube_1 = PgCube::from_str("(2)").unwrap(); + assert_eq!(cube_1, PgCube::Point(2.)); + let cube_2 = PgCube::from_str("2").unwrap(); + assert_eq!(cube_2, PgCube::Point(2.)); + } + + #[test] + fn can_serialise_point_type() { + assert_eq!(PgCube::Point(2.).serialize_to_vec(), POINT_BYTES,) + } + #[test] + fn can_deserialise_zero_volume_bytes() { + let cube = PgCube::from_bytes(ZERO_VOLUME_BYTES).unwrap(); + assert_eq!(cube, PgCube::ZeroVolume(vec![2., 3.])); + } + + #[test] + fn can_deserialise_zero_volume_string() { + let cube_1 = PgCube::from_str("(2,3,4)").unwrap(); + assert_eq!(cube_1, PgCube::ZeroVolume(vec![2., 3., 4.])); + let cube_2 = PgCube::from_str("2,3,4").unwrap(); + assert_eq!(cube_2, PgCube::ZeroVolume(vec![2., 3., 4.])); + } + + #[test] + fn can_serialise_zero_volume() { + assert_eq!( + PgCube::ZeroVolume(vec![2., 3.]).serialize_to_vec(), + ZERO_VOLUME_BYTES + ); + } + + #[test] + fn can_deserialise_one_dimension_interval_bytes() { + let cube = PgCube::from_bytes(ONE_DIMENSIONAL_INTERVAL_BYTES).unwrap(); + assert_eq!(cube, PgCube::OneDimensionInterval(7., 8.)) + } + + #[test] + fn can_deserialise_one_dimension_interval_string() { + let cube_1 = PgCube::from_str("((7),(8))").unwrap(); + assert_eq!(cube_1, PgCube::OneDimensionInterval(7., 8.)); + let cube_2 = PgCube::from_str("(7),(8)").unwrap(); + assert_eq!(cube_2, PgCube::OneDimensionInterval(7., 8.)); + } + + #[test] + fn can_serialise_one_dimension_interval() { + assert_eq!( + PgCube::OneDimensionInterval(7., 8.).serialize_to_vec(), + ONE_DIMENSIONAL_INTERVAL_BYTES + ) + } + + #[test] + fn can_deserialise_multi_dimension_2_dimension_byte() { + let cube = PgCube::from_bytes(MULTI_DIMENSION_2_DIM_BYTES).unwrap(); + assert_eq!( + cube, + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]) + ) + } + + #[test] + fn can_deserialise_multi_dimension_2_dimension_string() { + let cube_1 = PgCube::from_str("((1,2),(3,4))").unwrap(); + assert_eq!( + cube_1, + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]) + ); + let cube_2 = PgCube::from_str("((1, 2), (3, 4))").unwrap(); + assert_eq!( + cube_2, + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]) + ); + let cube_3 = PgCube::from_str("(1,2),(3,4)").unwrap(); + assert_eq!( + cube_3, + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]) + ); + let cube_4 = PgCube::from_str("(1, 2), (3, 4)").unwrap(); + assert_eq!( + cube_4, + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]) + ) + } + + #[test] + fn can_serialise_multi_dimension_2_dimension() { + assert_eq!( + PgCube::MultiDimension(vec![vec![1., 2.], vec![3., 4.]]).serialize_to_vec(), + MULTI_DIMENSION_2_DIM_BYTES + ) + } + + #[test] + fn can_deserialise_multi_dimension_3_dimension_bytes() { + let cube = PgCube::from_bytes(MULTI_DIMENSION_3_DIM_BYTES).unwrap(); + assert_eq!( + cube, + PgCube::MultiDimension(vec![vec![2., 3., 4.], vec![5., 6., 7.]]) + ) + } + + #[test] + fn can_deserialise_multi_dimension_3_dimension_string() { + let cube = PgCube::from_str("((2,3,4),(5,6,7))").unwrap(); + assert_eq!( + cube, + PgCube::MultiDimension(vec![vec![2., 3., 4.], vec![5., 6., 7.]]) + ); + let cube_2 = PgCube::from_str("(2,3,4),(5,6,7)").unwrap(); + assert_eq!( + cube_2, + PgCube::MultiDimension(vec![vec![2., 3., 4.], vec![5., 6., 7.]]) + ); + } + + #[test] + fn can_serialise_multi_dimension_3_dimension() { + assert_eq!( + PgCube::MultiDimension(vec![vec![2., 3., 4.], vec![5., 6., 7.]]).serialize_to_vec(), + MULTI_DIMENSION_3_DIM_BYTES + ) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/float.rs b/src-tauri/vendor/sqlx-postgres/src/types/float.rs new file mode 100644 index 00000000..116a28c2 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/float.rs @@ -0,0 +1,65 @@ +use byteorder::{BigEndian, ByteOrder}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl Type for f32 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::FLOAT4 + } +} + +impl PgHasArrayType for f32 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::FLOAT4_ARRAY + } +} + +impl Encode<'_, Postgres> for f32 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for f32 { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => BigEndian::read_f32(value.as_bytes()?), + PgValueFormat::Text => value.as_str()?.parse()?, + }) + } +} + +impl Type for f64 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::FLOAT8 + } +} + +impl PgHasArrayType for f64 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::FLOAT8_ARRAY + } +} + +impl Encode<'_, Postgres> for f64 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for f64 { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => BigEndian::read_f64(value.as_bytes()?), + PgValueFormat::Text => value.as_str()?.parse()?, + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/box.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/box.rs new file mode 100644 index 00000000..28016b27 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/box.rs @@ -0,0 +1,324 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use std::str::FromStr; + +const ERROR: &str = "error decoding BOX"; + +/// ## Postgres Geometric Box type +/// +/// Description: Rectangular box +/// Representation: `((upper_right_x,upper_right_y),(lower_left_x,lower_left_y))` +/// +/// Boxes are represented by pairs of points that are opposite corners of the box. Values of type box are specified using any of the following syntaxes: +/// +/// ```text +/// ( ( upper_right_x , upper_right_y ) , ( lower_left_x , lower_left_y ) ) +/// ( upper_right_x , upper_right_y ) , ( lower_left_x , lower_left_y ) +/// upper_right_x , upper_right_y , lower_left_x , lower_left_y +/// ``` +/// where `(upper_right_x,upper_right_y) and (lower_left_x,lower_left_y)` are any two opposite corners of the box. +/// Any two opposite corners can be supplied on input, but the values will be reordered as needed to store the upper right and lower left corners, in that order. +/// +/// See [Postgres Manual, Section 8.8.4: Geometric Types - Boxes][PG.S.8.8.4] for details. +/// +/// [PG.S.8.8.4]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-GEOMETRIC-BOXES +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgBox { + pub upper_right_x: f64, + pub upper_right_y: f64, + pub lower_left_x: f64, + pub lower_left_y: f64, +} + +impl Type for PgBox { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("box") + } +} + +impl PgHasArrayType for PgBox { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_box") + } +} + +impl<'r> Decode<'r, Postgres> for PgBox { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgBox::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgBox::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgBox { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("box")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgBox { + type Err = BoxDynError; + + fn from_str(s: &str) -> Result { + let sanitised = s.replace(['(', ')', '[', ']', ' '], ""); + let mut parts = sanitised.split(','); + + let upper_right_x = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get upper_right_x from {}", ERROR, s))?; + + let upper_right_y = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get upper_right_y from {}", ERROR, s))?; + + let lower_left_x = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get lower_left_x from {}", ERROR, s))?; + + let lower_left_y = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get lower_left_y from {}", ERROR, s))?; + + if parts.next().is_some() { + return Err(format!("{}: too many numbers inputted in {}", ERROR, s).into()); + } + + Ok(PgBox { + upper_right_x, + upper_right_y, + lower_left_x, + lower_left_y, + }) + } +} + +impl PgBox { + fn from_bytes(mut bytes: &[u8]) -> Result { + let upper_right_x = bytes.get_f64(); + let upper_right_y = bytes.get_f64(); + let lower_left_x = bytes.get_f64(); + let lower_left_y = bytes.get_f64(); + + Ok(PgBox { + upper_right_x, + upper_right_y, + lower_left_x, + lower_left_y, + }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), String> { + let min_x = &self.upper_right_x.min(self.lower_left_x); + let min_y = &self.upper_right_y.min(self.lower_left_y); + let max_x = &self.upper_right_x.max(self.lower_left_x); + let max_y = &self.upper_right_y.max(self.lower_left_y); + + buff.extend_from_slice(&max_x.to_be_bytes()); + buff.extend_from_slice(&max_y.to_be_bytes()); + buff.extend_from_slice(&min_x.to_be_bytes()); + buff.extend_from_slice(&min_y.to_be_bytes()); + + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +#[cfg(test)] +mod box_tests { + + use std::str::FromStr; + + use super::PgBox; + + const BOX_BYTES: &[u8] = &[ + 64, 0, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 192, 0, 0, 0, 0, 0, 0, 0, 192, 0, 0, 0, + 0, 0, 0, 0, + ]; + + #[test] + fn can_deserialise_box_type_bytes_in_order() { + let pg_box = PgBox::from_bytes(BOX_BYTES).unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 2., + upper_right_y: 2., + lower_left_x: -2., + lower_left_y: -2. + } + ) + } + + #[test] + fn can_deserialise_box_type_str_first_syntax() { + let pg_box = PgBox::from_str("[( 1, 2), (3, 4 )]").unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 1., + upper_right_y: 2., + lower_left_x: 3., + lower_left_y: 4. + } + ); + } + #[test] + fn can_deserialise_box_type_str_second_syntax() { + let pg_box = PgBox::from_str("(( 1, 2), (3, 4 ))").unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 1., + upper_right_y: 2., + lower_left_x: 3., + lower_left_y: 4. + } + ); + } + + #[test] + fn can_deserialise_box_type_str_third_syntax() { + let pg_box = PgBox::from_str("(1, 2), (3, 4 )").unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 1., + upper_right_y: 2., + lower_left_x: 3., + lower_left_y: 4. + } + ); + } + + #[test] + fn can_deserialise_box_type_str_fourth_syntax() { + let pg_box = PgBox::from_str("1, 2, 3, 4").unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 1., + upper_right_y: 2., + lower_left_x: 3., + lower_left_y: 4. + } + ); + } + + #[test] + fn cannot_deserialise_too_many_numbers() { + let input_str = "1, 2, 3, 4, 5"; + let pg_box = PgBox::from_str(input_str); + assert!(pg_box.is_err()); + if let Err(err) = pg_box { + assert_eq!( + err.to_string(), + format!("error decoding BOX: too many numbers inputted in {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_too_few_numbers() { + let input_str = "1, 2, 3 "; + let pg_box = PgBox::from_str(input_str); + assert!(pg_box.is_err()); + if let Err(err) = pg_box { + assert_eq!( + err.to_string(), + format!("error decoding BOX: could not get lower_left_y from {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_invalid_numbers() { + let input_str = "1, 2, 3, FOUR"; + let pg_box = PgBox::from_str(input_str); + assert!(pg_box.is_err()); + if let Err(err) = pg_box { + assert_eq!( + err.to_string(), + format!("error decoding BOX: could not get lower_left_y from {input_str}") + ) + } + } + + #[test] + fn can_deserialise_box_type_str_float() { + let pg_box = PgBox::from_str("(1.1, 2.2), (3.3, 4.4)").unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 1.1, + upper_right_y: 2.2, + lower_left_x: 3.3, + lower_left_y: 4.4 + } + ); + } + + #[test] + fn can_serialise_box_type_in_order() { + let pg_box = PgBox { + upper_right_x: 2., + lower_left_x: -2., + upper_right_y: -2., + lower_left_y: 2., + }; + assert_eq!(pg_box.serialize_to_vec(), BOX_BYTES,) + } + + #[test] + fn can_serialise_box_type_out_of_order() { + let pg_box = PgBox { + upper_right_x: -2., + lower_left_x: 2., + upper_right_y: 2., + lower_left_y: -2., + }; + assert_eq!(pg_box.serialize_to_vec(), BOX_BYTES,) + } + + #[test] + fn can_order_box() { + let pg_box = PgBox { + upper_right_x: -2., + lower_left_x: 2., + upper_right_y: 2., + lower_left_y: -2., + }; + let bytes = pg_box.serialize_to_vec(); + + let pg_box = PgBox::from_bytes(&bytes).unwrap(); + assert_eq!( + pg_box, + PgBox { + upper_right_x: 2., + upper_right_y: 2., + lower_left_x: -2., + lower_left_y: -2. + } + ) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/circle.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/circle.rs new file mode 100644 index 00000000..dde54dd2 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/circle.rs @@ -0,0 +1,250 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use sqlx_core::Error; +use std::str::FromStr; + +const ERROR: &str = "error decoding CIRCLE"; + +/// ## Postgres Geometric Circle type +/// +/// Description: Circle +/// Representation: `< (x, y), radius >` (center point and radius) +/// +/// ```text +/// < ( x , y ) , radius > +/// ( ( x , y ) , radius ) +/// ( x , y ) , radius +/// x , y , radius +/// ``` +/// where `(x,y)` is the center point. +/// +/// See [Postgres Manual, Section 8.8.7, Geometric Types - Circles][PG.S.8.8.7] for details. +/// +/// [PG.S.8.8.7]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-CIRCLE +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgCircle { + pub x: f64, + pub y: f64, + pub radius: f64, +} + +impl Type for PgCircle { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("circle") + } +} + +impl PgHasArrayType for PgCircle { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_circle") + } +} + +impl<'r> Decode<'r, Postgres> for PgCircle { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgCircle::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgCircle::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgCircle { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("circle")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgCircle { + type Err = BoxDynError; + + fn from_str(s: &str) -> Result { + let sanitised = s.replace(['<', '>', '(', ')', ' '], ""); + let mut parts = sanitised.split(','); + + let x = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get x from {}", ERROR, s))?; + + let y = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get y from {}", ERROR, s))?; + + let radius = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get radius from {}", ERROR, s))?; + + if parts.next().is_some() { + return Err(format!("{}: too many numbers inputted in {}", ERROR, s).into()); + } + + if radius < 0. { + return Err(format!("{}: cannot have negative radius: {}", ERROR, s).into()); + } + + Ok(PgCircle { x, y, radius }) + } +} + +impl PgCircle { + fn from_bytes(mut bytes: &[u8]) -> Result { + let x = bytes.get_f64(); + let y = bytes.get_f64(); + let r = bytes.get_f64(); + Ok(PgCircle { x, y, radius: r }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), Error> { + buff.extend_from_slice(&self.x.to_be_bytes()); + buff.extend_from_slice(&self.y.to_be_bytes()); + buff.extend_from_slice(&self.radius.to_be_bytes()); + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +#[cfg(test)] +mod circle_tests { + + use std::str::FromStr; + + use super::PgCircle; + + const CIRCLE_BYTES: &[u8] = &[ + 63, 241, 153, 153, 153, 153, 153, 154, 64, 1, 153, 153, 153, 153, 153, 154, 64, 10, 102, + 102, 102, 102, 102, 102, + ]; + + #[test] + fn can_deserialise_circle_type_bytes() { + let circle = PgCircle::from_bytes(CIRCLE_BYTES).unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.1, + y: 2.2, + radius: 3.3 + } + ) + } + + #[test] + fn can_deserialise_circle_type_str() { + let circle = PgCircle::from_str("<(1, 2), 3 >").unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.0, + y: 2.0, + radius: 3.0 + } + ); + } + + #[test] + fn can_deserialise_circle_type_str_second_syntax() { + let circle = PgCircle::from_str("((1, 2), 3 )").unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.0, + y: 2.0, + radius: 3.0 + } + ); + } + + #[test] + fn can_deserialise_circle_type_str_third_syntax() { + let circle = PgCircle::from_str("(1, 2), 3 ").unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.0, + y: 2.0, + radius: 3.0 + } + ); + } + + #[test] + fn can_deserialise_circle_type_str_fourth_syntax() { + let circle = PgCircle::from_str("1, 2, 3 ").unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.0, + y: 2.0, + radius: 3.0 + } + ); + } + + #[test] + fn cannot_deserialise_circle_invalid_numbers() { + let input_str = "1, 2, Three"; + let circle = PgCircle::from_str(input_str); + assert!(circle.is_err()); + if let Err(err) = circle { + assert_eq!( + err.to_string(), + format!("error decoding CIRCLE: could not get radius from {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_circle_negative_radius() { + let input_str = "1, 2, -3"; + let circle = PgCircle::from_str(input_str); + assert!(circle.is_err()); + if let Err(err) = circle { + assert_eq!( + err.to_string(), + format!("error decoding CIRCLE: cannot have negative radius: {input_str}") + ) + } + } + + #[test] + fn can_deserialise_circle_type_str_float() { + let circle = PgCircle::from_str("<(1.1, 2.2), 3.3>").unwrap(); + assert_eq!( + circle, + PgCircle { + x: 1.1, + y: 2.2, + radius: 3.3 + } + ); + } + + #[test] + fn can_serialise_circle_type() { + let circle = PgCircle { + x: 1.1, + y: 2.2, + radius: 3.3, + }; + assert_eq!(circle.serialize_to_vec(), CIRCLE_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/line.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/line.rs new file mode 100644 index 00000000..8f08c949 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/line.rs @@ -0,0 +1,214 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use std::str::FromStr; + +const ERROR: &str = "error decoding LINE"; + +/// ## Postgres Geometric Line type +/// +/// Description: Infinite line +/// Representation: `{A, B, C}` +/// +/// Lines are represented by the linear equation Ax + By + C = 0, where A and B are not both zero. +/// +/// See [Postgres Manual, Section 8.8.2, Geometric Types - Lines][PG.S.8.8.2] for details. +/// +/// [PG.S.8.8.2]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-LINE +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgLine { + pub a: f64, + pub b: f64, + pub c: f64, +} + +impl Type for PgLine { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("line") + } +} + +impl PgHasArrayType for PgLine { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_line") + } +} + +impl<'r> Decode<'r, Postgres> for PgLine { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgLine::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgLine::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgLine { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("line")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgLine { + type Err = BoxDynError; + + fn from_str(s: &str) -> Result { + let mut parts = s + .trim_matches(|c| c == '{' || c == '}' || c == ' ') + .split(','); + + let a = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get a from {}", ERROR, s))?; + + let b = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get b from {}", ERROR, s))?; + + let c = parts + .next() + .and_then(|s| s.trim().parse::().ok()) + .ok_or_else(|| format!("{}: could not get c from {}", ERROR, s))?; + + if parts.next().is_some() { + return Err(format!("{}: too many numbers inputted in {}", ERROR, s).into()); + } + + Ok(PgLine { a, b, c }) + } +} + +impl PgLine { + fn from_bytes(mut bytes: &[u8]) -> Result { + let a = bytes.get_f64(); + let b = bytes.get_f64(); + let c = bytes.get_f64(); + Ok(PgLine { a, b, c }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), BoxDynError> { + buff.extend_from_slice(&self.a.to_be_bytes()); + buff.extend_from_slice(&self.b.to_be_bytes()); + buff.extend_from_slice(&self.c.to_be_bytes()); + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +#[cfg(test)] +mod line_tests { + + use std::str::FromStr; + + use super::PgLine; + + const LINE_BYTES: &[u8] = &[ + 63, 241, 153, 153, 153, 153, 153, 154, 64, 1, 153, 153, 153, 153, 153, 154, 64, 10, 102, + 102, 102, 102, 102, 102, + ]; + + #[test] + fn can_deserialise_line_type_bytes() { + let line = PgLine::from_bytes(LINE_BYTES).unwrap(); + assert_eq!( + line, + PgLine { + a: 1.1, + b: 2.2, + c: 3.3 + } + ) + } + + #[test] + fn can_deserialise_line_type_str() { + let line = PgLine::from_str("{ 1, 2, 3 }").unwrap(); + assert_eq!( + line, + PgLine { + a: 1.0, + b: 2.0, + c: 3.0 + } + ); + } + + #[test] + fn cannot_deserialise_line_too_few_numbers() { + let input_str = "{ 1, 2 }"; + let line = PgLine::from_str(input_str); + assert!(line.is_err()); + if let Err(err) = line { + assert_eq!( + err.to_string(), + format!("error decoding LINE: could not get c from {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_line_too_many_numbers() { + let input_str = "{ 1, 2, 3, 4 }"; + let line = PgLine::from_str(input_str); + assert!(line.is_err()); + if let Err(err) = line { + assert_eq!( + err.to_string(), + format!("error decoding LINE: too many numbers inputted in {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_line_invalid_numbers() { + let input_str = "{ 1, 2, three }"; + let line = PgLine::from_str(input_str); + assert!(line.is_err()); + if let Err(err) = line { + assert_eq!( + err.to_string(), + format!("error decoding LINE: could not get c from {input_str}") + ) + } + } + + #[test] + fn can_deserialise_line_type_str_float() { + let line = PgLine::from_str("{1.1, 2.2, 3.3}").unwrap(); + assert_eq!( + line, + PgLine { + a: 1.1, + b: 2.2, + c: 3.3 + } + ); + } + + #[test] + fn can_serialise_line_type() { + let line = PgLine { + a: 1.1, + b: 2.2, + c: 3.3, + }; + assert_eq!(line.serialize_to_vec(), LINE_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/line_segment.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/line_segment.rs new file mode 100644 index 00000000..cd08e4da --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/line_segment.rs @@ -0,0 +1,286 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use std::str::FromStr; + +const ERROR: &str = "error decoding LSEG"; + +/// ## Postgres Geometric Line Segment type +/// +/// Description: Finite line segment +/// Representation: `((start_x,start_y),(end_x,end_y))` +/// +/// +/// Line segments are represented by pairs of points that are the endpoints of the segment. Values of type lseg are specified using any of the following syntaxes: +/// ```text +/// [ ( start_x , start_y ) , ( end_x , end_y ) ] +/// ( ( start_x , start_y ) , ( end_x , end_y ) ) +/// ( start_x , start_y ) , ( end_x , end_y ) +/// start_x , start_y , end_x , end_y +/// ``` +/// where `(start_x,start_y) and (end_x,end_y)` are the end points of the line segment. +/// +/// See [Postgres Manual, Section 8.8.3, Geometric Types - Line Segments][PG.S.8.8.3] for details. +/// +/// [PG.S.8.8.3]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-LSEG +/// +#[doc(alias = "line segment")] +#[derive(Debug, Clone, PartialEq)] +pub struct PgLSeg { + pub start_x: f64, + pub start_y: f64, + pub end_x: f64, + pub end_y: f64, +} + +impl Type for PgLSeg { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("lseg") + } +} + +impl PgHasArrayType for PgLSeg { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_lseg") + } +} + +impl<'r> Decode<'r, Postgres> for PgLSeg { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgLSeg::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgLSeg::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgLSeg { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("lseg")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgLSeg { + type Err = BoxDynError; + + fn from_str(s: &str) -> Result { + let sanitised = s.replace(['(', ')', '[', ']', ' '], ""); + let mut parts = sanitised.split(','); + + let start_x = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get start_x from {}", ERROR, s))?; + + let start_y = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get start_y from {}", ERROR, s))?; + + let end_x = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get end_x from {}", ERROR, s))?; + + let end_y = parts + .next() + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| format!("{}: could not get end_y from {}", ERROR, s))?; + + if parts.next().is_some() { + return Err(format!("{}: too many numbers inputted in {}", ERROR, s).into()); + } + + Ok(PgLSeg { + start_x, + start_y, + end_x, + end_y, + }) + } +} + +impl PgLSeg { + fn from_bytes(mut bytes: &[u8]) -> Result { + let start_x = bytes.get_f64(); + let start_y = bytes.get_f64(); + let end_x = bytes.get_f64(); + let end_y = bytes.get_f64(); + + Ok(PgLSeg { + start_x, + start_y, + end_x, + end_y, + }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), BoxDynError> { + buff.extend_from_slice(&self.start_x.to_be_bytes()); + buff.extend_from_slice(&self.start_y.to_be_bytes()); + buff.extend_from_slice(&self.end_x.to_be_bytes()); + buff.extend_from_slice(&self.end_y.to_be_bytes()); + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +#[cfg(test)] +mod lseg_tests { + + use std::str::FromStr; + + use super::PgLSeg; + + const LINE_SEGMENT_BYTES: &[u8] = &[ + 63, 241, 153, 153, 153, 153, 153, 154, 64, 1, 153, 153, 153, 153, 153, 154, 64, 10, 102, + 102, 102, 102, 102, 102, 64, 17, 153, 153, 153, 153, 153, 154, + ]; + + #[test] + fn can_deserialise_lseg_type_bytes() { + let lseg = PgLSeg::from_bytes(LINE_SEGMENT_BYTES).unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1.1, + start_y: 2.2, + end_x: 3.3, + end_y: 4.4 + } + ) + } + + #[test] + fn can_deserialise_lseg_type_str_first_syntax() { + let lseg = PgLSeg::from_str("[( 1, 2), (3, 4 )]").unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1., + start_y: 2., + end_x: 3., + end_y: 4. + } + ); + } + #[test] + fn can_deserialise_lseg_type_str_second_syntax() { + let lseg = PgLSeg::from_str("(( 1, 2), (3, 4 ))").unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1., + start_y: 2., + end_x: 3., + end_y: 4. + } + ); + } + + #[test] + fn can_deserialise_lseg_type_str_third_syntax() { + let lseg = PgLSeg::from_str("(1, 2), (3, 4 )").unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1., + start_y: 2., + end_x: 3., + end_y: 4. + } + ); + } + + #[test] + fn can_deserialise_lseg_type_str_fourth_syntax() { + let lseg = PgLSeg::from_str("1, 2, 3, 4").unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1., + start_y: 2., + end_x: 3., + end_y: 4. + } + ); + } + + #[test] + fn can_deserialise_too_many_numbers() { + let input_str = "1, 2, 3, 4, 5"; + let lseg = PgLSeg::from_str(input_str); + assert!(lseg.is_err()); + if let Err(err) = lseg { + assert_eq!( + err.to_string(), + format!("error decoding LSEG: too many numbers inputted in {input_str}") + ) + } + } + + #[test] + fn can_deserialise_too_few_numbers() { + let input_str = "1, 2, 3"; + let lseg = PgLSeg::from_str(input_str); + assert!(lseg.is_err()); + if let Err(err) = lseg { + assert_eq!( + err.to_string(), + format!("error decoding LSEG: could not get end_y from {input_str}") + ) + } + } + + #[test] + fn can_deserialise_invalid_numbers() { + let input_str = "1, 2, 3, FOUR"; + let lseg = PgLSeg::from_str(input_str); + assert!(lseg.is_err()); + if let Err(err) = lseg { + assert_eq!( + err.to_string(), + format!("error decoding LSEG: could not get end_y from {input_str}") + ) + } + } + + #[test] + fn can_deserialise_lseg_type_str_float() { + let lseg = PgLSeg::from_str("(1.1, 2.2), (3.3, 4.4)").unwrap(); + assert_eq!( + lseg, + PgLSeg { + start_x: 1.1, + start_y: 2.2, + end_x: 3.3, + end_y: 4.4 + } + ); + } + + #[test] + fn can_serialise_lseg_type() { + let lseg = PgLSeg { + start_x: 1.1, + start_y: 2.2, + end_x: 3.3, + end_y: 4.4, + }; + assert_eq!(lseg.serialize_to_vec(), LINE_SEGMENT_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/mod.rs new file mode 100644 index 00000000..c3142145 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/mod.rs @@ -0,0 +1,7 @@ +pub mod r#box; +pub mod circle; +pub mod line; +pub mod line_segment; +pub mod path; +pub mod point; +pub mod polygon; diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/path.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/path.rs new file mode 100644 index 00000000..6799289f --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/path.rs @@ -0,0 +1,375 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::{PgPoint, Type}; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use sqlx_core::Error; +use std::mem; +use std::str::FromStr; + +const BYTE_WIDTH: usize = mem::size_of::(); + +/// ## Postgres Geometric Path type +/// +/// Description: Open path or Closed path (similar to polygon) +/// Representation: Open `[(x1,y1),...]`, Closed `((x1,y1),...)` +/// +/// Paths are represented by lists of connected points. Paths can be open, where the first and last points in the list are considered not connected, or closed, where the first and last points are considered connected. +/// Values of type path are specified using any of the following syntaxes: +/// ```text +/// [ ( x1 , y1 ) , ... , ( xn , yn ) ] +/// ( ( x1 , y1 ) , ... , ( xn , yn ) ) +/// ( x1 , y1 ) , ... , ( xn , yn ) +/// ( x1 , y1 , ... , xn , yn ) +/// x1 , y1 , ... , xn , yn +/// ``` +/// where the points are the end points of the line segments comprising the path. Square brackets `([])` indicate an open path, while parentheses `(())` indicate a closed path. +/// When the outermost parentheses are omitted, as in the third through fifth syntaxes, a closed path is assumed. +/// +/// See [Postgres Manual, Section 8.8.5, Geometric Types - Paths][PG.S.8.8.5] for details. +/// +/// [PG.S.8.8.5]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-GEOMETRIC-PATHS +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgPath { + pub closed: bool, + pub points: Vec, +} + +#[derive(Copy, Clone, Debug, PartialEq, Eq)] +struct Header { + is_closed: bool, + length: usize, +} + +impl Type for PgPath { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("path") + } +} + +impl PgHasArrayType for PgPath { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_path") + } +} + +impl<'r> Decode<'r, Postgres> for PgPath { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgPath::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgPath::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgPath { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("path")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgPath { + type Err = Error; + + fn from_str(s: &str) -> Result { + let closed = !s.contains('['); + let sanitised = s.replace(['(', ')', '[', ']', ' '], ""); + let parts = sanitised.split(',').collect::>(); + + let mut points = vec![]; + + if parts.len() % 2 != 0 { + return Err(Error::Decode( + format!("Unmatched pair in PATH: {}", s).into(), + )); + } + + for chunk in parts.chunks_exact(2) { + if let [x_str, y_str] = chunk { + let x = parse_float_from_str(x_str, "could not get x")?; + let y = parse_float_from_str(y_str, "could not get y")?; + + let point = PgPoint { x, y }; + points.push(point); + } + } + + if !points.is_empty() { + return Ok(PgPath { points, closed }); + } + + Err(Error::Decode( + format!("could not get path from {}", s).into(), + )) + } +} + +impl PgPath { + fn header(&self) -> Header { + Header { + is_closed: self.closed, + length: self.points.len(), + } + } + + fn from_bytes(mut bytes: &[u8]) -> Result { + let header = Header::try_read(&mut bytes)?; + + if bytes.len() != header.data_size() { + return Err(format!( + "expected {} bytes after header, got {}", + header.data_size(), + bytes.len() + ) + .into()); + } + + if bytes.len() % BYTE_WIDTH * 2 != 0 { + return Err(format!( + "data length not divisible by pairs of {BYTE_WIDTH}: {}", + bytes.len() + ) + .into()); + } + + let mut out_points = Vec::with_capacity(bytes.len() / (BYTE_WIDTH * 2)); + + while bytes.has_remaining() { + let point = PgPoint { + x: bytes.get_f64(), + y: bytes.get_f64(), + }; + out_points.push(point) + } + Ok(PgPath { + closed: header.is_closed, + points: out_points, + }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), BoxDynError> { + let header = self.header(); + buff.reserve(header.data_size()); + header.try_write(buff)?; + + for point in &self.points { + buff.extend_from_slice(&point.x.to_be_bytes()); + buff.extend_from_slice(&point.y.to_be_bytes()); + } + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +impl Header { + const HEADER_WIDTH: usize = mem::size_of::() + mem::size_of::(); + + fn data_size(&self) -> usize { + self.length * BYTE_WIDTH * 2 + } + + fn try_read(buf: &mut &[u8]) -> Result { + if buf.len() < Self::HEADER_WIDTH { + return Err(format!( + "expected PATH data to contain at least {} bytes, got {}", + Self::HEADER_WIDTH, + buf.len() + )); + } + + let is_closed = buf.get_i8(); + let length = buf.get_i32(); + + let length = usize::try_from(length).ok().ok_or_else(|| { + format!( + "received PATH data length: {length}. Expected length between 0 and {}", + usize::MAX + ) + })?; + + Ok(Self { + is_closed: is_closed != 0, + length, + }) + } + + fn try_write(&self, buff: &mut PgArgumentBuffer) -> Result<(), String> { + let is_closed = self.is_closed as i8; + + let length = i32::try_from(self.length).map_err(|_| { + format!( + "PATH length exceeds allowed maximum ({} > {})", + self.length, + i32::MAX + ) + })?; + + buff.extend(is_closed.to_be_bytes()); + buff.extend(length.to_be_bytes()); + + Ok(()) + } +} + +fn parse_float_from_str(s: &str, error_msg: &str) -> Result { + s.parse().map_err(|_| Error::Decode(error_msg.into())) +} + +#[cfg(test)] +mod path_tests { + + use std::str::FromStr; + + use crate::types::PgPoint; + + use super::PgPath; + + const PATH_CLOSED_BYTES: &[u8] = &[ + 1, 0, 0, 0, 2, 63, 240, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, + 64, 16, 0, 0, 0, 0, 0, 0, + ]; + + const PATH_OPEN_BYTES: &[u8] = &[ + 0, 0, 0, 0, 2, 63, 240, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, + 64, 16, 0, 0, 0, 0, 0, 0, + ]; + + const PATH_UNEVEN_POINTS: &[u8] = &[ + 0, 0, 0, 0, 2, 63, 240, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, + 64, 16, 0, 0, + ]; + + #[test] + fn can_deserialise_path_type_bytes_closed() { + let path = PgPath::from_bytes(PATH_CLOSED_BYTES).unwrap(); + assert_eq!( + path, + PgPath { + closed: true, + points: vec![PgPoint { x: 1.0, y: 2.0 }, PgPoint { x: 3.0, y: 4.0 }] + } + ) + } + + #[test] + fn cannot_deserialise_path_type_uneven_point_bytes() { + let path = PgPath::from_bytes(PATH_UNEVEN_POINTS); + assert!(path.is_err()); + + if let Err(err) = path { + assert_eq!( + err.to_string(), + format!("expected 32 bytes after header, got 28") + ) + } + } + + #[test] + fn can_deserialise_path_type_bytes_open() { + let path = PgPath::from_bytes(PATH_OPEN_BYTES).unwrap(); + assert_eq!( + path, + PgPath { + closed: false, + points: vec![PgPoint { x: 1.0, y: 2.0 }, PgPoint { x: 3.0, y: 4.0 }] + } + ) + } + + #[test] + fn can_deserialise_path_type_str_first_syntax() { + let path = PgPath::from_str("[( 1, 2), (3, 4 )]").unwrap(); + assert_eq!( + path, + PgPath { + closed: false, + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn cannot_deserialise_path_type_str_uneven_points_first_syntax() { + let input_str = "[( 1, 2), (3)]"; + let path = PgPath::from_str(input_str); + + assert!(path.is_err()); + + if let Err(err) = path { + assert_eq!( + err.to_string(), + format!("error occurred while decoding: Unmatched pair in PATH: {input_str}") + ) + } + } + + #[test] + fn can_deserialise_path_type_str_second_syntax() { + let path = PgPath::from_str("(( 1, 2), (3, 4 ))").unwrap(); + assert_eq!( + path, + PgPath { + closed: true, + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_path_type_str_third_syntax() { + let path = PgPath::from_str("(1, 2), (3, 4 )").unwrap(); + assert_eq!( + path, + PgPath { + closed: true, + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_path_type_str_fourth_syntax() { + let path = PgPath::from_str("1, 2, 3, 4").unwrap(); + assert_eq!( + path, + PgPath { + closed: true, + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_path_type_str_float() { + let path = PgPath::from_str("(1.1, 2.2), (3.3, 4.4)").unwrap(); + assert_eq!( + path, + PgPath { + closed: true, + points: vec![PgPoint { x: 1.1, y: 2.2 }, PgPoint { x: 3.3, y: 4.4 }] + } + ); + } + + #[test] + fn can_serialise_path_type() { + let path = PgPath { + closed: true, + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }], + }; + assert_eq!(path.serialize_to_vec(), PATH_CLOSED_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/point.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/point.rs new file mode 100644 index 00000000..5078ce1e --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/point.rs @@ -0,0 +1,141 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use sqlx_core::Error; +use std::str::FromStr; + +/// ## Postgres Geometric Point type +/// +/// Description: Point on a plane +/// Representation: `(x, y)` +/// +/// Points are the fundamental two-dimensional building block for geometric types. Values of type point are specified using either of the following syntaxes: +/// ```text +/// ( x , y ) +/// x , y +/// ```` +/// where x and y are the respective coordinates, as floating-point numbers. +/// +/// See [Postgres Manual, Section 8.8.1, Geometric Types - Points][PG.S.8.8.1] for details. +/// +/// [PG.S.8.8.1]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-GEOMETRIC-POINTS +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgPoint { + pub x: f64, + pub y: f64, +} + +impl Type for PgPoint { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("point") + } +} + +impl PgHasArrayType for PgPoint { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_point") + } +} + +impl<'r> Decode<'r, Postgres> for PgPoint { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgPoint::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgPoint::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgPoint { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("point")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +fn parse_float_from_str(s: &str, error_msg: &str) -> Result { + s.trim() + .parse() + .map_err(|_| Error::Decode(error_msg.into())) +} + +impl FromStr for PgPoint { + type Err = BoxDynError; + + fn from_str(s: &str) -> Result { + let (x_str, y_str) = s + .trim_matches(|c| c == '(' || c == ')' || c == ' ') + .split_once(',') + .ok_or_else(|| format!("error decoding POINT: could not get x and y from {}", s))?; + + let x = parse_float_from_str(x_str, "error decoding POINT: could not get x")?; + let y = parse_float_from_str(y_str, "error decoding POINT: could not get y")?; + + Ok(PgPoint { x, y }) + } +} + +impl PgPoint { + fn from_bytes(mut bytes: &[u8]) -> Result { + let x = bytes.get_f64(); + let y = bytes.get_f64(); + Ok(PgPoint { x, y }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), BoxDynError> { + buff.extend_from_slice(&self.x.to_be_bytes()); + buff.extend_from_slice(&self.y.to_be_bytes()); + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +#[cfg(test)] +mod point_tests { + + use std::str::FromStr; + + use super::PgPoint; + + const POINT_BYTES: &[u8] = &[ + 64, 0, 204, 204, 204, 204, 204, 205, 64, 20, 204, 204, 204, 204, 204, 205, + ]; + + #[test] + fn can_deserialise_point_type_bytes() { + let point = PgPoint::from_bytes(POINT_BYTES).unwrap(); + assert_eq!(point, PgPoint { x: 2.1, y: 5.2 }) + } + + #[test] + fn can_deserialise_point_type_str() { + let point = PgPoint::from_str("(2, 3)").unwrap(); + assert_eq!(point, PgPoint { x: 2., y: 3. }); + } + + #[test] + fn can_deserialise_point_type_str_float() { + let point = PgPoint::from_str("(2.5, 3.4)").unwrap(); + assert_eq!(point, PgPoint { x: 2.5, y: 3.4 }); + } + + #[test] + fn can_serialise_point_type() { + let point = PgPoint { x: 2.1, y: 5.2 }; + assert_eq!(point.serialize_to_vec(), POINT_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/geometry/polygon.rs b/src-tauri/vendor/sqlx-postgres/src/types/geometry/polygon.rs new file mode 100644 index 00000000..a5a203c6 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/geometry/polygon.rs @@ -0,0 +1,366 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::{PgPoint, Type}; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use sqlx_core::bytes::Buf; +use sqlx_core::Error; +use std::mem; +use std::str::FromStr; + +const BYTE_WIDTH: usize = mem::size_of::(); + +/// ## Postgres Geometric Polygon type +/// +/// Description: Polygon (similar to closed polygon) +/// Representation: `((x1,y1),...)` +/// +/// Polygons are represented by lists of points (the vertexes of the polygon). Polygons are very similar to closed paths; the essential semantic difference is that a polygon is considered to include the area within it, while a path is not. +/// An important implementation difference between polygons and paths is that the stored representation of a polygon includes its smallest bounding box. This speeds up certain search operations, although computing the bounding box adds overhead while constructing new polygons. +/// Values of type polygon are specified using any of the following syntaxes: +/// +/// ```text +/// ( ( x1 , y1 ) , ... , ( xn , yn ) ) +/// ( x1 , y1 ) , ... , ( xn , yn ) +/// ( x1 , y1 , ... , xn , yn ) +/// x1 , y1 , ... , xn , yn +/// ``` +/// +/// where the points are the end points of the line segments comprising the boundary of the polygon. +/// +/// See [Postgres Manual, Section 8.8.6, Geometric Types - Polygons][PG.S.8.8.6] for details. +/// +/// [PG.S.8.8.6]: https://www.postgresql.org/docs/current/datatype-geometric.html#DATATYPE-POLYGON +/// +#[derive(Debug, Clone, PartialEq)] +pub struct PgPolygon { + pub points: Vec, +} + +#[derive(Copy, Clone, Debug, PartialEq, Eq)] +struct Header { + length: usize, +} + +impl Type for PgPolygon { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("polygon") + } +} + +impl PgHasArrayType for PgPolygon { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_polygon") + } +} + +impl<'r> Decode<'r, Postgres> for PgPolygon { + fn decode(value: PgValueRef<'r>) -> Result> { + match value.format() { + PgValueFormat::Text => Ok(PgPolygon::from_str(value.as_str()?)?), + PgValueFormat::Binary => Ok(PgPolygon::from_bytes(value.as_bytes()?)?), + } + } +} + +impl<'q> Encode<'q, Postgres> for PgPolygon { + fn produces(&self) -> Option { + Some(PgTypeInfo::with_name("polygon")) + } + + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + self.serialize(buf)?; + Ok(IsNull::No) + } +} + +impl FromStr for PgPolygon { + type Err = Error; + + fn from_str(s: &str) -> Result { + let sanitised = s.replace(['(', ')', '[', ']', ' '], ""); + let parts = sanitised.split(',').collect::>(); + + let mut points = vec![]; + + if parts.len() % 2 != 0 { + return Err(Error::Decode( + format!("Unmatched pair in POLYGON: {}", s).into(), + )); + } + + for chunk in parts.chunks_exact(2) { + if let [x_str, y_str] = chunk { + let x = parse_float_from_str(x_str, "could not get x")?; + let y = parse_float_from_str(y_str, "could not get y")?; + + let point = PgPoint { x, y }; + points.push(point); + } + } + + if !points.is_empty() { + return Ok(PgPolygon { points }); + } + + Err(Error::Decode( + format!("could not get polygon from {}", s).into(), + )) + } +} + +impl PgPolygon { + fn header(&self) -> Header { + Header { + length: self.points.len(), + } + } + + fn from_bytes(mut bytes: &[u8]) -> Result { + let header = Header::try_read(&mut bytes)?; + + if bytes.len() != header.data_size() { + return Err(format!( + "expected {} bytes after header, got {}", + header.data_size(), + bytes.len() + ) + .into()); + } + + if bytes.len() % BYTE_WIDTH * 2 != 0 { + return Err(format!( + "data length not divisible by pairs of {BYTE_WIDTH}: {}", + bytes.len() + ) + .into()); + } + + let mut out_points = Vec::with_capacity(bytes.len() / (BYTE_WIDTH * 2)); + while bytes.has_remaining() { + let point = PgPoint { + x: bytes.get_f64(), + y: bytes.get_f64(), + }; + out_points.push(point) + } + Ok(PgPolygon { points: out_points }) + } + + fn serialize(&self, buff: &mut PgArgumentBuffer) -> Result<(), BoxDynError> { + let header = self.header(); + buff.reserve(header.data_size()); + header.try_write(buff)?; + + for point in &self.points { + buff.extend_from_slice(&point.x.to_be_bytes()); + buff.extend_from_slice(&point.y.to_be_bytes()); + } + Ok(()) + } + + #[cfg(test)] + fn serialize_to_vec(&self) -> Vec { + let mut buff = PgArgumentBuffer::default(); + self.serialize(&mut buff).unwrap(); + buff.to_vec() + } +} + +impl Header { + const HEADER_WIDTH: usize = mem::size_of::() + mem::size_of::(); + + fn data_size(&self) -> usize { + self.length * BYTE_WIDTH * 2 + } + + fn try_read(buf: &mut &[u8]) -> Result { + if buf.len() < Self::HEADER_WIDTH { + return Err(format!( + "expected polygon data to contain at least {} bytes, got {}", + Self::HEADER_WIDTH, + buf.len() + )); + } + + let length = buf.get_i32(); + + let length = usize::try_from(length).ok().ok_or_else(|| { + format!( + "received polygon with length: {length}. Expected length between 0 and {}", + usize::MAX + ) + })?; + + Ok(Self { length }) + } + + fn try_write(&self, buff: &mut PgArgumentBuffer) -> Result<(), String> { + let length = i32::try_from(self.length).map_err(|_| { + format!( + "polygon length exceeds allowed maximum ({} > {})", + self.length, + i32::MAX + ) + })?; + + buff.extend(length.to_be_bytes()); + + Ok(()) + } +} + +fn parse_float_from_str(s: &str, error_msg: &str) -> Result { + s.parse().map_err(|_| Error::Decode(error_msg.into())) +} + +#[cfg(test)] +mod polygon_tests { + + use std::str::FromStr; + + use crate::types::PgPoint; + + use super::PgPolygon; + + const POLYGON_BYTES: &[u8] = &[ + 0, 0, 0, 12, 192, 0, 0, 0, 0, 0, 0, 0, 192, 8, 0, 0, 0, 0, 0, 0, 191, 240, 0, 0, 0, 0, 0, + 0, 192, 8, 0, 0, 0, 0, 0, 0, 191, 240, 0, 0, 0, 0, 0, 0, 191, 240, 0, 0, 0, 0, 0, 0, 63, + 240, 0, 0, 0, 0, 0, 0, 63, 240, 0, 0, 0, 0, 0, 0, 63, 240, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, + 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 64, 8, 0, 0, 0, 0, 0, 0, 64, 0, 0, 0, 0, 0, 0, 0, 192, + 8, 0, 0, 0, 0, 0, 0, 63, 240, 0, 0, 0, 0, 0, 0, 192, 8, 0, 0, 0, 0, 0, 0, 63, 240, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 191, 240, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 191, + 240, 0, 0, 0, 0, 0, 0, 192, 0, 0, 0, 0, 0, 0, 0, 192, 0, 0, 0, 0, 0, 0, 0, 192, 0, 0, 0, 0, + 0, 0, 0, + ]; + + #[test] + fn can_deserialise_polygon_type_bytes() { + let polygon = PgPolygon::from_bytes(POLYGON_BYTES).unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![ + PgPoint { x: -2., y: -3. }, + PgPoint { x: -1., y: -3. }, + PgPoint { x: -1., y: -1. }, + PgPoint { x: 1., y: 1. }, + PgPoint { x: 1., y: 3. }, + PgPoint { x: 2., y: 3. }, + PgPoint { x: 2., y: -3. }, + PgPoint { x: 1., y: -3. }, + PgPoint { x: 1., y: 0. }, + PgPoint { x: -1., y: 0. }, + PgPoint { x: -1., y: -2. }, + PgPoint { x: -2., y: -2. } + ] + } + ) + } + + #[test] + fn can_deserialise_polygon_type_str_first_syntax() { + let polygon = PgPolygon::from_str("[( 1, 2), (3, 4 )]").unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_polygon_type_str_second_syntax() { + let polygon = PgPolygon::from_str("(( 1, 2), (3, 4 ))").unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn cannot_deserialise_polygon_type_str_uneven_points_first_syntax() { + let input_str = "[( 1, 2), (3)]"; + let polygon = PgPolygon::from_str(input_str); + + assert!(polygon.is_err()); + + if let Err(err) = polygon { + assert_eq!( + err.to_string(), + format!("error occurred while decoding: Unmatched pair in POLYGON: {input_str}") + ) + } + } + + #[test] + fn cannot_deserialise_polygon_type_str_invalid_numbers() { + let input_str = "[( 1, 2), (2, three)]"; + let polygon = PgPolygon::from_str(input_str); + + assert!(polygon.is_err()); + + if let Err(err) = polygon { + assert_eq!( + err.to_string(), + format!("error occurred while decoding: could not get y") + ) + } + } + + #[test] + fn can_deserialise_polygon_type_str_third_syntax() { + let polygon = PgPolygon::from_str("(1, 2), (3, 4 )").unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_polygon_type_str_fourth_syntax() { + let polygon = PgPolygon::from_str("1, 2, 3, 4").unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![PgPoint { x: 1., y: 2. }, PgPoint { x: 3., y: 4. }] + } + ); + } + + #[test] + fn can_deserialise_polygon_type_str_float() { + let polygon = PgPolygon::from_str("(1.1, 2.2), (3.3, 4.4)").unwrap(); + assert_eq!( + polygon, + PgPolygon { + points: vec![PgPoint { x: 1.1, y: 2.2 }, PgPoint { x: 3.3, y: 4.4 }] + } + ); + } + + #[test] + fn can_serialise_polygon_type() { + let polygon = PgPolygon { + points: vec![ + PgPoint { x: -2., y: -3. }, + PgPoint { x: -1., y: -3. }, + PgPoint { x: -1., y: -1. }, + PgPoint { x: 1., y: 1. }, + PgPoint { x: 1., y: 3. }, + PgPoint { x: 2., y: 3. }, + PgPoint { x: 2., y: -3. }, + PgPoint { x: 1., y: -3. }, + PgPoint { x: 1., y: 0. }, + PgPoint { x: -1., y: 0. }, + PgPoint { x: -1., y: -2. }, + PgPoint { x: -2., y: -2. }, + ], + }; + assert_eq!(polygon.serialize_to_vec(), POLYGON_BYTES,) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/hstore.rs b/src-tauri/vendor/sqlx-postgres/src/types/hstore.rs new file mode 100644 index 00000000..a03970fb --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/hstore.rs @@ -0,0 +1,329 @@ +use std::{ + collections::{btree_map, BTreeMap}, + mem, + ops::{Deref, DerefMut}, + str, +}; + +use crate::{ + decode::Decode, + encode::{Encode, IsNull}, + error::BoxDynError, + types::Type, + PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueRef, Postgres, +}; +use serde::{Deserialize, Serialize}; +use sqlx_core::bytes::Buf; + +/// Key-value support (`hstore`) for Postgres. +/// +/// SQLx currently maps `hstore` to a `BTreeMap>` but this may be expanded in +/// future to allow for user defined types. +/// +/// See [the Postgres manual, Appendix F, Section 18][PG.F.18] +/// +/// [PG.F.18]: https://www.postgresql.org/docs/current/hstore.html +/// +/// ### Note: Requires Postgres 8.3+ +/// Introduced as a method for storing unstructured data, the `hstore` extension was first added in +/// Postgres 8.3. +/// +/// +/// ### Note: Extension Required +/// The `hstore` extension is not enabled by default in Postgres. You will need to do so explicitly: +/// +/// ```ignore +/// CREATE EXTENSION IF NOT EXISTS hstore; +/// ``` +/// +/// # Examples +/// +/// ``` +/// # use sqlx_postgres::types::PgHstore; +/// // Shows basic usage of the PgHstore type. +/// // +/// #[derive(Clone, Debug, Default, Eq, PartialEq)] +/// struct UserCreate<'a> { +/// username: &'a str, +/// password: &'a str, +/// additional_data: PgHstore +/// } +/// +/// let mut new_user = UserCreate { +/// username: "name.surname@email.com", +/// password: "@super_secret_1", +/// ..Default::default() +/// }; +/// +/// new_user.additional_data.insert("department".to_string(), Some("IT".to_string())); +/// new_user.additional_data.insert("equipment_issued".to_string(), None); +/// ``` +/// ```ignore +/// query_scalar::<_, i64>( +/// "insert into user(username, password, additional_data) values($1, $2, $3) returning id" +/// ) +/// .bind(new_user.username) +/// .bind(new_user.password) +/// .bind(new_user.additional_data) +/// .fetch_one(pg_conn) +/// .await?; +/// ``` +/// +/// ``` +/// # use sqlx_postgres::types::PgHstore; +/// // PgHstore implements FromIterator to simplify construction. +/// // +/// let additional_data = PgHstore::from_iter([ +/// ("department".to_string(), Some("IT".to_string())), +/// ("equipment_issued".to_string(), None), +/// ]); +/// +/// assert_eq!(additional_data["department"], Some("IT".to_string())); +/// assert_eq!(additional_data["equipment_issued"], None); +/// +/// // Also IntoIterator for ease of iteration. +/// // +/// for (key, value) in additional_data { +/// println!("{key}: {value:?}"); +/// } +/// ``` +/// +#[derive(Clone, Debug, Default, Eq, PartialEq, Deserialize, Serialize)] +pub struct PgHstore(pub BTreeMap>); + +impl Deref for PgHstore { + type Target = BTreeMap>; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl DerefMut for PgHstore { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.0 + } +} + +impl FromIterator<(String, String)> for PgHstore { + fn from_iter>(iter: T) -> Self { + iter.into_iter().map(|(k, v)| (k, Some(v))).collect() + } +} + +impl FromIterator<(String, Option)> for PgHstore { + fn from_iter)>>(iter: T) -> Self { + let mut result = Self::default(); + + for (key, value) in iter { + result.0.insert(key, value); + } + + result + } +} + +impl IntoIterator for PgHstore { + type Item = (String, Option); + type IntoIter = btree_map::IntoIter>; + + fn into_iter(self) -> Self::IntoIter { + self.0.into_iter() + } +} + +impl Type for PgHstore { + fn type_info() -> PgTypeInfo { + PgTypeInfo::with_name("hstore") + } +} + +impl PgHasArrayType for PgHstore { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::array_of("hstore") + } +} + +impl<'r> Decode<'r, Postgres> for PgHstore { + fn decode(value: PgValueRef<'r>) -> Result { + let mut buf = <&[u8] as Decode>::decode(value)?; + let len = read_length(&mut buf)?; + + let len = + usize::try_from(len).map_err(|_| format!("PgHstore: length out of range: {len}"))?; + + let mut result = Self::default(); + + for i in 0..len { + let key = read_string(&mut buf) + .map_err(|e| format!("PgHstore: error reading {i}th key: {e}"))? + .ok_or_else(|| format!("PgHstore: expected {i}th key, got nothing"))?; + + let value = read_string(&mut buf) + .map_err(|e| format!("PgHstore: error reading value for key {key:?}: {e}"))?; + + result.insert(key, value); + } + + if !buf.is_empty() { + tracing::warn!("{} unread bytes at the end of HSTORE value", buf.len()); + } + + Ok(result) + } +} + +impl Encode<'_, Postgres> for PgHstore { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend_from_slice(&i32::to_be_bytes( + self.0 + .len() + .try_into() + .map_err(|_| format!("PgHstore length out of range: {}", self.0.len()))?, + )); + + for (i, (key, val)) in self.0.iter().enumerate() { + let key_bytes = key.as_bytes(); + + let key_len = i32::try_from(key_bytes.len()).map_err(|_| { + // Doesn't make sense to print the key itself: it's more than 2 GiB long! + format!( + "PgHstore: length of {i}th key out of range: {} bytes", + key_bytes.len() + ) + })?; + + buf.extend_from_slice(&i32::to_be_bytes(key_len)); + buf.extend_from_slice(key_bytes); + + match val { + Some(val) => { + let val_bytes = val.as_bytes(); + + let val_len = i32::try_from(val_bytes.len()).map_err(|_| { + format!( + "PgHstore: value length for key {key:?} out of range: {} bytes", + val_bytes.len() + ) + })?; + buf.extend_from_slice(&i32::to_be_bytes(val_len)); + buf.extend_from_slice(val_bytes); + } + None => { + buf.extend_from_slice(&i32::to_be_bytes(-1)); + } + } + } + + Ok(IsNull::No) + } +} + +fn read_length(buf: &mut &[u8]) -> Result { + if buf.len() < mem::size_of::() { + return Err(format!( + "expected {} bytes, got {}", + mem::size_of::(), + buf.len() + )); + } + + Ok(buf.get_i32()) +} + +fn read_string(buf: &mut &[u8]) -> Result, String> { + let len = read_length(buf)?; + + match len { + -1 => Ok(None), + len => { + let len = + usize::try_from(len).map_err(|_| format!("string length out of range: {len}"))?; + + if buf.len() < len { + return Err(format!("expected {len} bytes, got {}", buf.len())); + } + + let (val, rest) = buf.split_at(len); + *buf = rest; + + Ok(Some( + str::from_utf8(val).map_err(|e| e.to_string())?.to_string(), + )) + } + } +} + +#[cfg(test)] +mod test { + use super::*; + use crate::PgValueFormat; + + const EMPTY: &str = "00000000"; + + const NAME_SURNAME_AGE: &str = + "0000000300000003616765ffffffff000000046e616d65000000044a6f686e000000077375726e616d6500000003446f65"; + + #[test] + fn hstore_deserialize_ok() { + let empty = hex::decode(EMPTY).unwrap(); + let name_surname_age = hex::decode(NAME_SURNAME_AGE).unwrap(); + + let empty = PgValueRef { + value: Some(empty.as_slice()), + row: None, + type_info: PgTypeInfo::with_name("hstore"), + format: PgValueFormat::Binary, + }; + + let name_surname = PgValueRef { + value: Some(name_surname_age.as_slice()), + row: None, + type_info: PgTypeInfo::with_name("hstore"), + format: PgValueFormat::Binary, + }; + + let res_empty = PgHstore::decode(empty).unwrap(); + let res_name_surname = PgHstore::decode(name_surname).unwrap(); + + assert!(res_empty.is_empty()); + assert_eq!(res_name_surname["name"], Some("John".to_string())); + assert_eq!(res_name_surname["surname"], Some("Doe".to_string())); + assert_eq!(res_name_surname["age"], None); + } + + #[test] + #[should_panic(expected = "PgHstore: length out of range: -5")] + fn hstore_deserialize_buffer_length_error() { + let buf = PgValueRef { + value: Some(&[255, 255, 255, 251]), + row: None, + type_info: PgTypeInfo::with_name("hstore"), + format: PgValueFormat::Binary, + }; + + PgHstore::decode(buf).unwrap(); + } + + #[test] + fn hstore_serialize_ok() { + let mut buff = PgArgumentBuffer::default(); + let _ = PgHstore::from_iter::<[(String, String); 0]>([]) + .encode_by_ref(&mut buff) + .unwrap(); + + assert_eq!(hex::encode(buff.as_slice()), EMPTY); + + buff.clear(); + + let _ = PgHstore::from_iter([ + ("name".to_string(), Some("John".to_string())), + ("surname".to_string(), Some("Doe".to_string())), + ("age".to_string(), None), + ]) + .encode_by_ref(&mut buff) + .unwrap(); + + assert_eq!(hex::encode(buff.as_slice()), NAME_SURNAME_AGE); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/int.rs b/src-tauri/vendor/sqlx-postgres/src/types/int.rs new file mode 100644 index 00000000..b8255f1b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/int.rs @@ -0,0 +1,176 @@ +use byteorder::{BigEndian, ByteOrder}; +use std::num::{NonZeroI16, NonZeroI32, NonZeroI64}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +fn int_decode(value: PgValueRef<'_>) -> Result { + Ok(match value.format() { + PgValueFormat::Text => value.as_str()?.parse()?, + PgValueFormat::Binary => { + let buf = value.as_bytes()?; + + // Return error if buf is empty or is more than 8 bytes + match buf.len() { + 0 => { + return Err("Value Buffer found empty while decoding to integer type".into()); + } + buf_len @ 9.. => { + return Err(format!( + "Value Buffer exceeds 8 bytes while decoding to integer type. Buffer size = {} bytes ", buf_len + ) + .into()); + } + _ => {} + } + + BigEndian::read_int(buf, buf.len()) + } + }) +} + +impl Type for i8 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::CHAR + } +} + +impl PgHasArrayType for i8 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::CHAR_ARRAY + } +} + +impl Encode<'_, Postgres> for i8 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for i8 { + fn decode(value: PgValueRef<'_>) -> Result { + // note: decoding here is for the `"char"` type as Postgres does not have a native 1-byte integer type. + // https://github.com/postgres/postgres/blob/master/src/backend/utils/adt/char.c#L58-L60 + match value.format() { + PgValueFormat::Binary => int_decode(value)?.try_into().map_err(Into::into), + PgValueFormat::Text => { + let text = value.as_str()?; + + // A value of 0 is represented with the empty string. + if text.is_empty() { + return Ok(0); + } + + if text.starts_with('\\') { + // For values between 0x80 and 0xFF, it's encoded in octal. + return Ok(i8::from_str_radix(text.trim_start_matches('\\'), 8)?); + } + + // Wrapping is the whole idea. + #[allow(clippy::cast_possible_wrap)] + Ok(text.as_bytes()[0] as i8) + } + } + } +} + +impl Type for i16 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INT2 + } +} + +impl PgHasArrayType for i16 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT2_ARRAY + } +} + +impl Encode<'_, Postgres> for i16 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for i16 { + fn decode(value: PgValueRef<'_>) -> Result { + int_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Type for i32 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INT4 + } +} + +impl PgHasArrayType for i32 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT4_ARRAY + } +} + +impl Encode<'_, Postgres> for i32 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for i32 { + fn decode(value: PgValueRef<'_>) -> Result { + int_decode(value)?.try_into().map_err(Into::into) + } +} + +impl Type for i64 { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INT8 + } +} + +impl PgHasArrayType for i64 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT8_ARRAY + } +} + +impl Encode<'_, Postgres> for i64 { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for i64 { + fn decode(value: PgValueRef<'_>) -> Result { + int_decode(value) + } +} + +impl PgHasArrayType for NonZeroI16 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT2_ARRAY + } +} + +impl PgHasArrayType for NonZeroI32 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT4_ARRAY + } +} + +impl PgHasArrayType for NonZeroI64 { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT8_ARRAY + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/interval.rs b/src-tauri/vendor/sqlx-postgres/src/types/interval.rs new file mode 100644 index 00000000..02b1faa6 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/interval.rs @@ -0,0 +1,399 @@ +use std::mem; + +use byteorder::{NetworkEndian, ReadBytesExt}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +// `PgInterval` is available for direct access to the INTERVAL type + +#[derive(Debug, Eq, PartialEq, Clone, Copy, Hash, Default)] +pub struct PgInterval { + pub months: i32, + pub days: i32, + pub microseconds: i64, +} + +impl Type for PgInterval { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL + } +} + +impl PgHasArrayType for PgInterval { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL_ARRAY + } +} + +impl<'de> Decode<'de, Postgres> for PgInterval { + fn decode(value: PgValueRef<'de>) -> Result { + match value.format() { + PgValueFormat::Binary => { + let mut buf = value.as_bytes()?; + let microseconds = buf.read_i64::()?; + let days = buf.read_i32::()?; + let months = buf.read_i32::()?; + + Ok(PgInterval { + months, + days, + microseconds, + }) + } + + // TODO: Implement parsing of text mode + PgValueFormat::Text => { + Err("not implemented: decode `INTERVAL` in text mode (unprepared queries)".into()) + } + } + } +} + +impl Encode<'_, Postgres> for PgInterval { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.microseconds.to_be_bytes()); + buf.extend(&self.days.to_be_bytes()); + buf.extend(&self.months.to_be_bytes()); + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + 2 * mem::size_of::() + } +} + +// We then implement Encode + Type for std Duration, chrono Duration, and time Duration +// This is to enable ease-of-use for encoding when its simple + +impl Type for std::time::Duration { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL + } +} + +impl PgHasArrayType for std::time::Duration { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL_ARRAY + } +} + +impl Encode<'_, Postgres> for std::time::Duration { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + PgInterval::try_from(*self)?.encode_by_ref(buf) + } + + fn size_hint(&self) -> usize { + 2 * mem::size_of::() + } +} + +impl TryFrom for PgInterval { + type Error = BoxDynError; + + /// Convert a `std::time::Duration` to a `PgInterval` + /// + /// This returns an error if there is a loss of precision using nanoseconds or if there is a + /// microsecond overflow. + fn try_from(value: std::time::Duration) -> Result { + if value.as_nanos() % 1000 != 0 { + return Err("PostgreSQL `INTERVAL` does not support nanoseconds precision".into()); + } + + Ok(Self { + months: 0, + days: 0, + microseconds: value.as_micros().try_into()?, + }) + } +} + +#[cfg(feature = "chrono")] +impl Type for chrono::Duration { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL + } +} + +#[cfg(feature = "chrono")] +impl PgHasArrayType for chrono::Duration { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL_ARRAY + } +} + +#[cfg(feature = "chrono")] +impl Encode<'_, Postgres> for chrono::Duration { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + let pg_interval = PgInterval::try_from(*self)?; + pg_interval.encode_by_ref(buf) + } + + fn size_hint(&self) -> usize { + 2 * mem::size_of::() + } +} + +#[cfg(feature = "chrono")] +impl TryFrom for PgInterval { + type Error = BoxDynError; + + /// Convert a `chrono::Duration` to a `PgInterval`. + /// + /// This returns an error if there is a loss of precision using nanoseconds or if there is a + /// nanosecond overflow. + fn try_from(value: chrono::Duration) -> Result { + value + .num_nanoseconds() + .map_or::, _>( + Err("Overflow has occurred for PostgreSQL `INTERVAL`".into()), + |nanoseconds| { + if nanoseconds % 1000 != 0 { + return Err( + "PostgreSQL `INTERVAL` does not support nanoseconds precision".into(), + ); + } + Ok(()) + }, + )?; + + value.num_microseconds().map_or( + Err("Overflow has occurred for PostgreSQL `INTERVAL`".into()), + |microseconds| { + Ok(Self { + months: 0, + days: 0, + microseconds, + }) + }, + ) + } +} + +#[cfg(feature = "time")] +impl Type for time::Duration { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL + } +} + +#[cfg(feature = "time")] +impl PgHasArrayType for time::Duration { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INTERVAL_ARRAY + } +} + +#[cfg(feature = "time")] +impl Encode<'_, Postgres> for time::Duration { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + let pg_interval = PgInterval::try_from(*self)?; + pg_interval.encode_by_ref(buf) + } + + fn size_hint(&self) -> usize { + 2 * mem::size_of::() + } +} + +#[cfg(feature = "time")] +impl TryFrom for PgInterval { + type Error = BoxDynError; + + /// Convert a `time::Duration` to a `PgInterval`. + /// + /// This returns an error if there is a loss of precision using nanoseconds or if there is a + /// microsecond overflow. + fn try_from(value: time::Duration) -> Result { + if value.whole_nanoseconds() % 1000 != 0 { + return Err("PostgreSQL `INTERVAL` does not support nanoseconds precision".into()); + } + + Ok(Self { + months: 0, + days: 0, + microseconds: value.whole_microseconds().try_into()?, + }) + } +} + +#[test] +fn test_encode_interval() { + let mut buf = PgArgumentBuffer::default(); + + let interval = PgInterval { + months: 0, + days: 0, + microseconds: 0, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!(&**buf, [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]); + buf.clear(); + + let interval = PgInterval { + months: 0, + days: 0, + microseconds: 1_000, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!(&**buf, [0, 0, 0, 0, 0, 0, 3, 232, 0, 0, 0, 0, 0, 0, 0, 0]); + buf.clear(); + + let interval = PgInterval { + months: 0, + days: 0, + microseconds: 1_000_000, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!(&**buf, [0, 0, 0, 0, 0, 15, 66, 64, 0, 0, 0, 0, 0, 0, 0, 0]); + buf.clear(); + + let interval = PgInterval { + months: 0, + days: 0, + microseconds: 3_600_000_000, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!( + &**buf, + [0, 0, 0, 0, 214, 147, 164, 0, 0, 0, 0, 0, 0, 0, 0, 0] + ); + buf.clear(); + + let interval = PgInterval { + months: 0, + days: 1, + microseconds: 0, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!(&**buf, [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0]); + buf.clear(); + + let interval = PgInterval { + months: 1, + days: 0, + microseconds: 0, + }; + assert!(matches!( + Encode::::encode(&interval, &mut buf), + Ok(IsNull::No) + )); + assert_eq!(&**buf, [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1]); + buf.clear(); + + assert_eq!( + PgInterval::default(), + PgInterval { + months: 0, + days: 0, + microseconds: 0, + } + ); +} + +#[test] +fn test_pginterval_std() { + // Case for positive duration + let interval = PgInterval { + days: 0, + months: 0, + microseconds: 27_000, + }; + assert_eq!( + &PgInterval::try_from(std::time::Duration::from_micros(27_000)).unwrap(), + &interval + ); + + // Case when precision loss occurs + assert!(PgInterval::try_from(std::time::Duration::from_nanos(27_000_001)).is_err()); + + // Case when microsecond overflow occurs + assert!(PgInterval::try_from(std::time::Duration::from_secs(20_000_000_000_000)).is_err()); +} + +#[test] +#[cfg(feature = "chrono")] +fn test_pginterval_chrono() { + // Case for positive duration + let interval = PgInterval { + days: 0, + months: 0, + microseconds: 27_000, + }; + assert_eq!( + &PgInterval::try_from(chrono::Duration::microseconds(27_000)).unwrap(), + &interval + ); + + // Case for negative duration + let interval = PgInterval { + days: 0, + months: 0, + microseconds: -27_000, + }; + assert_eq!( + &PgInterval::try_from(chrono::Duration::microseconds(-27_000)).unwrap(), + &interval + ); + + // Case when precision loss occurs + assert!(PgInterval::try_from(chrono::Duration::nanoseconds(27_000_001)).is_err()); + assert!(PgInterval::try_from(chrono::Duration::nanoseconds(-27_000_001)).is_err()); + + // Case when nanosecond overflow occurs + assert!(PgInterval::try_from(chrono::Duration::seconds(10_000_000_000)).is_err()); + assert!(PgInterval::try_from(chrono::Duration::seconds(-10_000_000_000)).is_err()); +} + +#[test] +#[cfg(feature = "time")] +fn test_pginterval_time() { + // Case for positive duration + let interval = PgInterval { + days: 0, + months: 0, + microseconds: 27_000, + }; + assert_eq!( + &PgInterval::try_from(time::Duration::microseconds(27_000)).unwrap(), + &interval + ); + + // Case for negative duration + let interval = PgInterval { + days: 0, + months: 0, + microseconds: -27_000, + }; + assert_eq!( + &PgInterval::try_from(time::Duration::microseconds(-27_000)).unwrap(), + &interval + ); + + // Case when precision loss occurs + assert!(PgInterval::try_from(time::Duration::nanoseconds(27_000_001)).is_err()); + assert!(PgInterval::try_from(time::Duration::nanoseconds(-27_000_001)).is_err()); + + // Case when microsecond overflow occurs + assert!(PgInterval::try_from(time::Duration::seconds(10_000_000_000_000)).is_err()); + assert!(PgInterval::try_from(time::Duration::seconds(-10_000_000_000_000)).is_err()); +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipaddr.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipaddr.rs new file mode 100644 index 00000000..b157eff3 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipaddr.rs @@ -0,0 +1,62 @@ +use std::net::IpAddr; + +use ipnet::IpNet; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueRef, Postgres}; + +impl Type for IpAddr +where + IpNet: Type, +{ + fn type_info() -> PgTypeInfo { + IpNet::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + IpNet::compatible(ty) + } +} + +impl PgHasArrayType for IpAddr { + fn array_type_info() -> PgTypeInfo { + ::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + ::array_compatible(ty) + } +} + +impl<'db> Encode<'db, Postgres> for IpAddr +where + IpNet: Encode<'db, Postgres>, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + IpNet::from(*self).encode_by_ref(buf) + } + + fn size_hint(&self) -> usize { + IpNet::from(*self).size_hint() + } +} + +impl<'db> Decode<'db, Postgres> for IpAddr +where + IpNet: Decode<'db, Postgres>, +{ + fn decode(value: PgValueRef<'db>) -> Result { + let ipnet = IpNet::decode(value)?; + + if matches!(ipnet, IpNet::V4(net) if net.prefix_len() != 32) + || matches!(ipnet, IpNet::V6(net) if net.prefix_len() != 128) + { + Err("lossy decode from inet/cidr")? + } + + Ok(ipnet.addr()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipnet.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipnet.rs new file mode 100644 index 00000000..1f986174 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/ipnet.rs @@ -0,0 +1,130 @@ +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; + +#[cfg(feature = "ipnet")] +use ipnet::{IpNet, Ipv4Net, Ipv6Net}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +// https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/include/utils/inet.h#L39 + +// Technically this is a magic number here but it doesn't make sense to drag in the whole of `libc` +// just for one constant. +const PGSQL_AF_INET: u8 = 2; // AF_INET +const PGSQL_AF_INET6: u8 = PGSQL_AF_INET + 1; + +impl Type for IpNet { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INET + } + + fn compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::CIDR || *ty == PgTypeInfo::INET + } +} + +impl PgHasArrayType for IpNet { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INET_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::CIDR_ARRAY || *ty == PgTypeInfo::INET_ARRAY + } +} + +impl Encode<'_, Postgres> for IpNet { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/backend/utils/adt/network.c#L293 + // https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/backend/utils/adt/network.c#L271 + + match self { + IpNet::V4(net) => { + buf.push(PGSQL_AF_INET); // ip_family + buf.push(net.prefix_len()); // ip_bits + buf.push(0); // is_cidr + buf.push(4); // nb (number of bytes) + buf.extend_from_slice(&net.addr().octets()) // address + } + + IpNet::V6(net) => { + buf.push(PGSQL_AF_INET6); // ip_family + buf.push(net.prefix_len()); // ip_bits + buf.push(0); // is_cidr + buf.push(16); // nb (number of bytes) + buf.extend_from_slice(&net.addr().octets()); // address + } + } + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + match self { + IpNet::V4(_) => 8, + IpNet::V6(_) => 20, + } + } +} + +impl Decode<'_, Postgres> for IpNet { + fn decode(value: PgValueRef<'_>) -> Result { + let bytes = match value.format() { + PgValueFormat::Binary => value.as_bytes()?, + PgValueFormat::Text => { + let s = value.as_str()?; + println!("{s}"); + if s.contains('/') { + return Ok(s.parse()?); + } + // IpNet::from_str doesn't handle conversion from IpAddr to IpNet + let addr: IpAddr = s.parse()?; + return Ok(addr.into()); + } + }; + + if bytes.len() >= 8 { + let family = bytes[0]; + let prefix = bytes[1]; + let _is_cidr = bytes[2] != 0; + let len = bytes[3]; + + match family { + PGSQL_AF_INET => { + if bytes.len() == 8 && len == 4 { + let inet = Ipv4Net::new( + Ipv4Addr::new(bytes[4], bytes[5], bytes[6], bytes[7]), + prefix, + )?; + + return Ok(IpNet::V4(inet)); + } + } + + PGSQL_AF_INET6 => { + if bytes.len() == 20 && len == 16 { + let inet = Ipv6Net::new( + Ipv6Addr::from([ + bytes[4], bytes[5], bytes[6], bytes[7], bytes[8], bytes[9], + bytes[10], bytes[11], bytes[12], bytes[13], bytes[14], bytes[15], + bytes[16], bytes[17], bytes[18], bytes[19], + ]), + prefix, + )?; + + return Ok(IpNet::V6(inet)); + } + } + + _ => { + return Err(format!("unknown ip family {family}").into()); + } + } + } + + Err("invalid data received when expecting an INET".into()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnet/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/mod.rs new file mode 100644 index 00000000..cd40cf30 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnet/mod.rs @@ -0,0 +1,7 @@ +// Prefer `ipnetwork` over `ipnet` because it was implemented first (want to avoid breaking change). +#[cfg(not(feature = "ipnetwork"))] +mod ipaddr; + +// Parent module is named after the `ipnet` crate, this is named after the `IpNet` type. +#[allow(clippy::module_inception)] +mod ipnet; diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipaddr.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipaddr.rs new file mode 100644 index 00000000..ee587eda --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipaddr.rs @@ -0,0 +1,62 @@ +use std::net::IpAddr; + +use ipnetwork::IpNetwork; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueRef, Postgres}; + +impl Type for IpAddr +where + IpNetwork: Type, +{ + fn type_info() -> PgTypeInfo { + IpNetwork::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + IpNetwork::compatible(ty) + } +} + +impl PgHasArrayType for IpAddr { + fn array_type_info() -> PgTypeInfo { + ::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + ::array_compatible(ty) + } +} + +impl<'db> Encode<'db, Postgres> for IpAddr +where + IpNetwork: Encode<'db, Postgres>, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + IpNetwork::from(*self).encode_by_ref(buf) + } + + fn size_hint(&self) -> usize { + IpNetwork::from(*self).size_hint() + } +} + +impl<'db> Decode<'db, Postgres> for IpAddr +where + IpNetwork: Decode<'db, Postgres>, +{ + fn decode(value: PgValueRef<'db>) -> Result { + let ipnetwork = IpNetwork::decode(value)?; + + if ipnetwork.is_ipv4() && ipnetwork.prefix() != 32 + || ipnetwork.is_ipv6() && ipnetwork.prefix() != 128 + { + Err("lossy decode from inet/cidr")? + } + + Ok(ipnetwork.ip()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipnetwork.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipnetwork.rs new file mode 100644 index 00000000..4f619ba9 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/ipnetwork.rs @@ -0,0 +1,122 @@ +use std::net::{Ipv4Addr, Ipv6Addr}; + +use ipnetwork::{IpNetwork, Ipv4Network, Ipv6Network}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +// https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/include/utils/inet.h#L39 + +// Technically this is a magic number here but it doesn't make sense to drag in the whole of `libc` +// just for one constant. +const PGSQL_AF_INET: u8 = 2; // AF_INET +const PGSQL_AF_INET6: u8 = PGSQL_AF_INET + 1; + +impl Type for IpNetwork { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INET + } + + fn compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::CIDR || *ty == PgTypeInfo::INET + } +} + +impl PgHasArrayType for IpNetwork { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INET_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::CIDR_ARRAY || *ty == PgTypeInfo::INET_ARRAY + } +} + +impl Encode<'_, Postgres> for IpNetwork { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/backend/utils/adt/network.c#L293 + // https://github.com/postgres/postgres/blob/574925bfd0a8175f6e161936ea11d9695677ba09/src/backend/utils/adt/network.c#L271 + + match self { + IpNetwork::V4(net) => { + buf.push(PGSQL_AF_INET); // ip_family + buf.push(net.prefix()); // ip_bits + buf.push(0); // is_cidr + buf.push(4); // nb (number of bytes) + buf.extend_from_slice(&net.ip().octets()) // address + } + + IpNetwork::V6(net) => { + buf.push(PGSQL_AF_INET6); // ip_family + buf.push(net.prefix()); // ip_bits + buf.push(0); // is_cidr + buf.push(16); // nb (number of bytes) + buf.extend_from_slice(&net.ip().octets()); // address + } + } + + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + match self { + IpNetwork::V4(_) => 8, + IpNetwork::V6(_) => 20, + } + } +} + +impl Decode<'_, Postgres> for IpNetwork { + fn decode(value: PgValueRef<'_>) -> Result { + let bytes = match value.format() { + PgValueFormat::Binary => value.as_bytes()?, + PgValueFormat::Text => { + return Ok(value.as_str()?.parse()?); + } + }; + + if bytes.len() >= 8 { + let family = bytes[0]; + let prefix = bytes[1]; + let _is_cidr = bytes[2] != 0; + let len = bytes[3]; + + match family { + PGSQL_AF_INET => { + if bytes.len() == 8 && len == 4 { + let inet = Ipv4Network::new( + Ipv4Addr::new(bytes[4], bytes[5], bytes[6], bytes[7]), + prefix, + )?; + + return Ok(IpNetwork::V4(inet)); + } + } + + PGSQL_AF_INET6 => { + if bytes.len() == 20 && len == 16 { + let inet = Ipv6Network::new( + Ipv6Addr::from([ + bytes[4], bytes[5], bytes[6], bytes[7], bytes[8], bytes[9], + bytes[10], bytes[11], bytes[12], bytes[13], bytes[14], bytes[15], + bytes[16], bytes[17], bytes[18], bytes[19], + ]), + prefix, + )?; + + return Ok(IpNetwork::V6(inet)); + } + } + + _ => { + return Err(format!("unknown ip family {family}").into()); + } + } + } + + Err("invalid data received when expecting an INET".into()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/mod.rs new file mode 100644 index 00000000..de40244c --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ipnetwork/mod.rs @@ -0,0 +1,5 @@ +mod ipaddr; + +// Parent module is named after the `ipnetwork` crate, this is named after the `IpNetwork` type. +#[allow(clippy::module_inception)] +mod ipnetwork; diff --git a/src-tauri/vendor/sqlx-postgres/src/types/json.rs b/src-tauri/vendor/sqlx-postgres/src/types/json.rs new file mode 100644 index 00000000..567e4801 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/json.rs @@ -0,0 +1,99 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::array_compatible; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use serde::{Deserialize, Serialize}; +use serde_json::value::RawValue as JsonRawValue; +use serde_json::Value as JsonValue; +pub(crate) use sqlx_core::types::{Json, Type}; + +// + +// In general, most applications should prefer to store JSON data as jsonb, +// unless there are quite specialized needs, such as legacy assumptions +// about ordering of object keys. + +impl Type for Json { + fn type_info() -> PgTypeInfo { + PgTypeInfo::JSONB + } + + fn compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::JSON || *ty == PgTypeInfo::JSONB + } +} + +impl PgHasArrayType for Json { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::JSONB_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + array_compatible::>(ty) + } +} + +impl PgHasArrayType for JsonValue { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::JSONB_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + array_compatible::(ty) + } +} + +impl PgHasArrayType for JsonRawValue { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::JSONB_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + array_compatible::(ty) + } +} + +impl<'q, T> Encode<'q, Postgres> for Json +where + T: Serialize, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // we have a tiny amount of dynamic behavior depending if we are resolved to be JSON + // instead of JSONB + buf.patch(|buf, ty: &PgTypeInfo| { + if *ty == PgTypeInfo::JSON || *ty == PgTypeInfo::JSON_ARRAY { + buf[0] = b' '; + } + }); + + // JSONB version (as of 2020-03-20) + buf.push(1); + + // the JSON data written to the buffer is the same regardless of parameter type + serde_json::to_writer(&mut **buf, &self.0)?; + + Ok(IsNull::No) + } +} + +impl<'r, T: 'r> Decode<'r, Postgres> for Json +where + T: Deserialize<'r>, +{ + fn decode(value: PgValueRef<'r>) -> Result { + let mut buf = value.as_bytes()?; + + if value.format() == PgValueFormat::Binary && value.type_info == PgTypeInfo::JSONB { + assert_eq!( + buf[0], 1, + "unsupported JSONB format version {}; please open an issue", + buf[0] + ); + + buf = &buf[1..]; + } + + serde_json::from_slice(buf).map(Json).map_err(Into::into) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/lquery.rs b/src-tauri/vendor/sqlx-postgres/src/types/lquery.rs new file mode 100644 index 00000000..a20fef54 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/lquery.rs @@ -0,0 +1,341 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use bitflags::bitflags; +use std::fmt::{self, Display, Formatter}; +use std::io::Write; +use std::ops::Deref; +use std::str::FromStr; + +use crate::types::ltree::{PgLTreeLabel, PgLTreeParseError}; + +/// Represents lquery specific errors +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum PgLQueryParseError { + #[error("lquery cannot be empty")] + EmptyString, + #[error("unexpected character in lquery")] + UnexpectedCharacter, + #[error("error parsing integer: {0}")] + ParseIntError(#[from] std::num::ParseIntError), + #[error("error parsing integer: {0}")] + LTreeParrseError(#[from] PgLTreeParseError), + /// LQuery version not supported + #[error("lquery version not supported")] + InvalidLqueryVersion, +} + +/// Container for a Label Tree Query (`lquery`) in Postgres. +/// +/// See +/// +/// ### Note: Requires Postgres 13+ +/// +/// This integration requires that the `lquery` type support the binary format in the Postgres +/// wire protocol, which only became available in Postgres 13. +/// ([Postgres 13.0 Release Notes, Additional Modules](https://www.postgresql.org/docs/13/release-13.html#id-1.11.6.11.5.14)) +/// +/// Ideally, SQLx's Postgres driver should support falling back to text format for types +/// which don't have `typsend` and `typrecv` entries in `pg_type`, but that work still needs +/// to be done. +/// +/// ### Note: Extension Required +/// The `ltree` extension is not enabled by default in Postgres. You will need to do so explicitly: +/// +/// ```ignore +/// CREATE EXTENSION IF NOT EXISTS "ltree"; +/// ``` +#[derive(Clone, Debug, Default, PartialEq)] +pub struct PgLQuery { + levels: Vec, +} + +// TODO: maybe a QueryBuilder pattern would be nice here +impl PgLQuery { + /// creates default/empty lquery + pub fn new() -> Self { + Self::default() + } + + pub fn from(levels: Vec) -> Self { + Self { levels } + } + + /// push a query level + pub fn push(&mut self, level: PgLQueryLevel) { + self.levels.push(level); + } + + /// pop a query level + pub fn pop(&mut self) -> Option { + self.levels.pop() + } + + /// creates lquery from an iterator with checking labels + // TODO: this should just be removed but I didn't want to bury it in a massive diff + #[deprecated = "renamed to `try_from_iter()`"] + #[allow(clippy::should_implement_trait)] + pub fn from_iter(levels: I) -> Result + where + S: Into, + I: IntoIterator, + { + let mut lquery = Self::default(); + for level in levels { + lquery.push(PgLQueryLevel::from_str(&level.into())?); + } + Ok(lquery) + } + + /// Create an `LQUERY` from an iterator of label strings. + /// + /// Returns an error if any label fails to parse according to [`PgLQueryLevel::from_str()`]. + pub fn try_from_iter(levels: I) -> Result + where + S: AsRef, + I: IntoIterator, + { + levels + .into_iter() + .map(|level| level.as_ref().parse::()) + .collect() + } +} + +impl FromIterator for PgLQuery { + fn from_iter>(iter: T) -> Self { + Self::from(iter.into_iter().collect()) + } +} + +impl IntoIterator for PgLQuery { + type Item = PgLQueryLevel; + type IntoIter = std::vec::IntoIter; + + fn into_iter(self) -> Self::IntoIter { + self.levels.into_iter() + } +} + +impl FromStr for PgLQuery { + type Err = PgLQueryParseError; + + fn from_str(s: &str) -> Result { + Ok(Self { + levels: s + .split('.') + .map(PgLQueryLevel::from_str) + .collect::>()?, + }) + } +} + +impl Display for PgLQuery { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + let mut iter = self.levels.iter(); + if let Some(label) = iter.next() { + write!(f, "{label}")?; + for label in iter { + write!(f, ".{label}")?; + } + } + Ok(()) + } +} + +impl Deref for PgLQuery { + type Target = [PgLQueryLevel]; + + fn deref(&self) -> &Self::Target { + &self.levels + } +} + +impl Type for PgLQuery { + fn type_info() -> PgTypeInfo { + // Since `ltree` is enabled by an extension, it does not have a stable OID. + PgTypeInfo::with_name("lquery") + } +} + +impl PgHasArrayType for PgLQuery { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_lquery") + } +} + +impl Encode<'_, Postgres> for PgLQuery { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(1i8.to_le_bytes()); + write!(buf, "{self}")?; + + Ok(IsNull::No) + } +} + +impl<'r> Decode<'r, Postgres> for PgLQuery { + fn decode(value: PgValueRef<'r>) -> Result { + match value.format() { + PgValueFormat::Binary => { + let bytes = value.as_bytes()?; + let version = i8::from_le_bytes([bytes[0]; 1]); + if version != 1 { + return Err(Box::new(PgLQueryParseError::InvalidLqueryVersion)); + } + Ok(Self::from_str(std::str::from_utf8(&bytes[1..])?)?) + } + PgValueFormat::Text => Ok(Self::from_str(value.as_str()?)?), + } + } +} + +bitflags! { + /// Modifiers that can be set to non-star labels + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + pub struct PgLQueryVariantFlag: u16 { + /// * - Match any label with this prefix, for example foo* matches foobar + const ANY_END = 0x01; + /// @ - Match case-insensitively, for example a@ matches A + const IN_CASE = 0x02; + /// % - Match initial underscore-separated words + const SUBLEXEME = 0x04; + } +} + +impl Display for PgLQueryVariantFlag { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + if self.contains(PgLQueryVariantFlag::ANY_END) { + write!(f, "*")?; + } + if self.contains(PgLQueryVariantFlag::IN_CASE) { + write!(f, "@")?; + } + if self.contains(PgLQueryVariantFlag::SUBLEXEME) { + write!(f, "%")?; + } + + Ok(()) + } +} + +#[derive(Clone, Debug, PartialEq)] +pub struct PgLQueryVariant { + label: PgLTreeLabel, + modifiers: PgLQueryVariantFlag, +} + +impl Display for PgLQueryVariant { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + write!(f, "{}{}", self.label, self.modifiers) + } +} + +#[derive(Clone, Debug, PartialEq)] +pub enum PgLQueryLevel { + /// match any label (*) with optional at least / at most numbers + Star(Option, Option), + /// match any of specified labels with optional flags + NonStar(Vec), + /// match none of specified labels with optional flags + NotNonStar(Vec), +} + +impl FromStr for PgLQueryLevel { + type Err = PgLQueryParseError; + + fn from_str(s: &str) -> Result { + let bytes = s.as_bytes(); + if bytes.is_empty() { + Err(PgLQueryParseError::EmptyString) + } else { + match bytes[0] { + b'*' => { + if bytes.len() > 1 { + let parts = s[2..s.len() - 1].split(',').collect::>(); + match parts.len() { + 1 => { + let number = parts[0].parse()?; + Ok(PgLQueryLevel::Star(Some(number), Some(number))) + } + 2 => Ok(PgLQueryLevel::Star( + Some(parts[0].parse()?), + Some(parts[1].parse()?), + )), + _ => Err(PgLQueryParseError::UnexpectedCharacter), + } + } else { + Ok(PgLQueryLevel::Star(None, None)) + } + } + b'!' => Ok(PgLQueryLevel::NotNonStar( + s[1..] + .split('|') + .map(PgLQueryVariant::from_str) + .collect::, PgLQueryParseError>>()?, + )), + _ => Ok(PgLQueryLevel::NonStar( + s.split('|') + .map(PgLQueryVariant::from_str) + .collect::, PgLQueryParseError>>()?, + )), + } + } + } +} + +impl FromStr for PgLQueryVariant { + type Err = PgLQueryParseError; + + fn from_str(s: &str) -> Result { + let mut label_length = s.len(); + let mut modifiers = PgLQueryVariantFlag::empty(); + + for b in s.bytes().rev() { + match b { + b'@' => modifiers.insert(PgLQueryVariantFlag::IN_CASE), + b'*' => modifiers.insert(PgLQueryVariantFlag::ANY_END), + b'%' => modifiers.insert(PgLQueryVariantFlag::SUBLEXEME), + _ => break, + } + label_length -= 1; + } + + Ok(PgLQueryVariant { + label: PgLTreeLabel::new(&s[0..label_length])?, + modifiers, + }) + } +} + +fn write_variants(f: &mut Formatter<'_>, variants: &[PgLQueryVariant], not: bool) -> fmt::Result { + let mut iter = variants.iter(); + if let Some(variant) = iter.next() { + write!(f, "{}{}", if not { "!" } else { "" }, variant)?; + for variant in iter { + write!(f, ".{variant}")?; + } + } + Ok(()) +} + +impl Display for PgLQueryLevel { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + match self { + PgLQueryLevel::Star(Some(at_least), Some(at_most)) => { + if at_least == at_most { + write!(f, "*{{{at_least}}}") + } else { + write!(f, "*{{{at_least},{at_most}}}") + } + } + PgLQueryLevel::Star(Some(at_least), _) => write!(f, "*{{{at_least},}}"), + PgLQueryLevel::Star(_, Some(at_most)) => write!(f, "*{{,{at_most}}}"), + PgLQueryLevel::Star(_, _) => write!(f, "*"), + PgLQueryLevel::NonStar(variants) => write_variants(f, variants, false), + PgLQueryLevel::NotNonStar(variants) => write_variants(f, variants, true), + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/ltree.rs b/src-tauri/vendor/sqlx-postgres/src/types/ltree.rs new file mode 100644 index 00000000..531f5065 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/ltree.rs @@ -0,0 +1,228 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use std::fmt::{self, Display, Formatter}; +use std::io::Write; +use std::ops::Deref; +use std::str::FromStr; + +/// Represents ltree specific errors +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum PgLTreeParseError { + /// LTree labels can only contain [A-Za-z0-9_] + #[error("ltree label contains invalid characters")] + InvalidLtreeLabel, + + /// LTree version not supported + #[error("ltree version not supported")] + InvalidLtreeVersion, +} + +#[derive(Clone, Debug, Default, PartialEq)] +pub struct PgLTreeLabel(String); + +impl PgLTreeLabel { + pub fn new(label: S) -> Result + where + S: Into, + { + let label = label.into(); + if label.len() <= 256 + && label + .bytes() + .all(|c| c.is_ascii_alphabetic() || c.is_ascii_digit() || c == b'_') + { + Ok(Self(label)) + } else { + Err(PgLTreeParseError::InvalidLtreeLabel) + } + } +} + +impl Deref for PgLTreeLabel { + type Target = str; + + fn deref(&self) -> &Self::Target { + self.0.as_str() + } +} + +impl FromStr for PgLTreeLabel { + type Err = PgLTreeParseError; + + fn from_str(s: &str) -> Result { + PgLTreeLabel::new(s) + } +} + +impl Display for PgLTreeLabel { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +/// Container for a Label Tree (`ltree`) in Postgres. +/// +/// See +/// +/// ### Note: Requires Postgres 13+ +/// +/// This integration requires that the `ltree` type support the binary format in the Postgres +/// wire protocol, which only became available in Postgres 13. +/// ([Postgres 13.0 Release Notes, Additional Modules](https://www.postgresql.org/docs/13/release-13.html#id-1.11.6.11.5.14)) +/// +/// Ideally, SQLx's Postgres driver should support falling back to text format for types +/// which don't have `typsend` and `typrecv` entries in `pg_type`, but that work still needs +/// to be done. +/// +/// ### Note: Extension Required +/// The `ltree` extension is not enabled by default in Postgres. You will need to do so explicitly: +/// +/// ```ignore +/// CREATE EXTENSION IF NOT EXISTS "ltree"; +/// ``` +#[derive(Clone, Debug, Default, PartialEq)] +pub struct PgLTree { + labels: Vec, +} + +impl PgLTree { + /// creates default/empty ltree + pub fn new() -> Self { + Self::default() + } + + /// creates ltree from a [`Vec`] + pub fn from(labels: Vec) -> Self { + Self { labels } + } + + /// creates ltree from an iterator with checking labels + // TODO: this should just be removed but I didn't want to bury it in a massive diff + #[deprecated = "renamed to `try_from_iter()`"] + #[allow(clippy::should_implement_trait)] + pub fn from_iter(labels: I) -> Result + where + String: From, + I: IntoIterator, + { + let mut ltree = Self::default(); + for label in labels { + ltree.push(PgLTreeLabel::new(label)?); + } + Ok(ltree) + } + + /// Create an `LTREE` from an iterator of label strings. + /// + /// Returns an error if any label fails to parse according to [`PgLTreeLabel::new()`]. + pub fn try_from_iter(labels: I) -> Result + where + S: Into, + I: IntoIterator, + { + labels.into_iter().map(PgLTreeLabel::new).collect() + } + + /// push a label to ltree + pub fn push(&mut self, label: PgLTreeLabel) { + self.labels.push(label); + } + + /// pop a label from ltree + pub fn pop(&mut self) -> Option { + self.labels.pop() + } +} + +impl FromIterator for PgLTree { + fn from_iter>(iter: T) -> Self { + Self { + labels: iter.into_iter().collect(), + } + } +} + +impl IntoIterator for PgLTree { + type Item = PgLTreeLabel; + type IntoIter = std::vec::IntoIter; + + fn into_iter(self) -> Self::IntoIter { + self.labels.into_iter() + } +} + +impl FromStr for PgLTree { + type Err = PgLTreeParseError; + + fn from_str(s: &str) -> Result { + Ok(Self { + labels: s + .split('.') + .map(PgLTreeLabel::new) + .collect::, Self::Err>>()?, + }) + } +} + +impl Display for PgLTree { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + let mut iter = self.labels.iter(); + if let Some(label) = iter.next() { + write!(f, "{label}")?; + for label in iter { + write!(f, ".{label}")?; + } + } + Ok(()) + } +} + +impl Deref for PgLTree { + type Target = [PgLTreeLabel]; + + fn deref(&self) -> &Self::Target { + &self.labels + } +} + +impl Type for PgLTree { + fn type_info() -> PgTypeInfo { + // Since `ltree` is enabled by an extension, it does not have a stable OID. + PgTypeInfo::with_name("ltree") + } +} + +impl PgHasArrayType for PgLTree { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::with_name("_ltree") + } +} + +impl Encode<'_, Postgres> for PgLTree { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(1i8.to_le_bytes()); + write!(buf, "{self}")?; + + Ok(IsNull::No) + } +} + +impl<'r> Decode<'r, Postgres> for PgLTree { + fn decode(value: PgValueRef<'r>) -> Result { + match value.format() { + PgValueFormat::Binary => { + let bytes = value.as_bytes()?; + let version = i8::from_le_bytes([bytes[0]; 1]); + if version != 1 { + return Err(Box::new(PgLTreeParseError::InvalidLtreeVersion)); + } + Ok(Self::from_str(std::str::from_utf8(&bytes[1..])?)?) + } + PgValueFormat::Text => Ok(Self::from_str(value.as_str()?)?), + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/mac_address.rs b/src-tauri/vendor/sqlx-postgres/src/types/mac_address.rs new file mode 100644 index 00000000..23766e70 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/mac_address.rs @@ -0,0 +1,51 @@ +use mac_address::MacAddress; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +impl Type for MacAddress { + fn type_info() -> PgTypeInfo { + PgTypeInfo::MACADDR + } + + fn compatible(ty: &PgTypeInfo) -> bool { + *ty == PgTypeInfo::MACADDR + } +} + +impl PgHasArrayType for MacAddress { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::MACADDR_ARRAY + } +} + +impl Encode<'_, Postgres> for MacAddress { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend_from_slice(&self.bytes()); // write just the address + Ok(IsNull::No) + } + + fn size_hint(&self) -> usize { + 6 + } +} + +impl Decode<'_, Postgres> for MacAddress { + fn decode(value: PgValueRef<'_>) -> Result { + let bytes = match value.format() { + PgValueFormat::Binary => value.as_bytes()?, + PgValueFormat::Text => { + return Ok(value.as_str()?.parse()?); + } + }; + + if bytes.len() == 6 { + return Ok(MacAddress::new(bytes.try_into().unwrap())); + } + + Err("invalid data received when expecting an MACADDR".into()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/mod.rs new file mode 100644 index 00000000..0faefbb4 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/mod.rs @@ -0,0 +1,314 @@ +//! Conversions between Rust and **Postgres** types. +//! +//! # Types +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `bool` | BOOL | +//! | `i8` | "CHAR" | +//! | `i16` | SMALLINT, SMALLSERIAL, INT2 | +//! | `i32` | INT, SERIAL, INT4 | +//! | `i64` | BIGINT, BIGSERIAL, INT8 | +//! | `f32` | REAL, FLOAT4 | +//! | `f64` | DOUBLE PRECISION, FLOAT8 | +//! | `&str`, [`String`] | VARCHAR, CHAR(N), TEXT, NAME, CITEXT | +//! | `&[u8]`, `Vec` | BYTEA | +//! | `()` | VOID | +//! | [`PgInterval`] | INTERVAL | +//! | [`PgRange`](PgRange) | INT8RANGE, INT4RANGE, TSRANGE, TSTZRANGE, DATERANGE, NUMRANGE | +//! | [`PgMoney`] | MONEY | +//! | [`PgLTree`] | LTREE | +//! | [`PgLQuery`] | LQUERY | +//! | [`PgCiText`] | CITEXT1 | +//! | [`PgCube`] | CUBE | +//! | [`PgPoint`] | POINT | +//! | [`PgLine`] | LINE | +//! | [`PgLSeg`] | LSEG | +//! | [`PgBox`] | BOX | +//! | [`PgPath`] | PATH | +//! | [`PgPolygon`] | POLYGON | +//! | [`PgCircle`] | CIRCLE | +//! | [`PgHstore`] | HSTORE | +//! +//! 1 SQLx generally considers `CITEXT` to be compatible with `String`, `&str`, etc., +//! but this wrapper type is available for edge cases, such as `CITEXT[]` which Postgres +//! does not consider to be compatible with `TEXT[]`. +//! +//! ### [`bigdecimal`](https://crates.io/crates/bigdecimal) +//! Requires the `bigdecimal` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `bigdecimal::BigDecimal` | NUMERIC | +//! +#![doc=include_str!("bigdecimal-range.md")] +//! +//! ### [`rust_decimal`](https://crates.io/crates/rust_decimal) +//! Requires the `rust_decimal` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `rust_decimal::Decimal` | NUMERIC | +//! +#![doc=include_str!("rust_decimal-range.md")] +//! +//! ### [`chrono`](https://crates.io/crates/chrono) +//! +//! Requires the `chrono` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `chrono::DateTime` | TIMESTAMPTZ | +//! | `chrono::DateTime` | TIMESTAMPTZ | +//! | `chrono::NaiveDateTime` | TIMESTAMP | +//! | `chrono::NaiveDate` | DATE | +//! | `chrono::NaiveTime` | TIME | +//! | [`PgTimeTz`] | TIMETZ | +//! +//! ### [`time`](https://crates.io/crates/time) +//! +//! Requires the `time` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `time::PrimitiveDateTime` | TIMESTAMP | +//! | `time::OffsetDateTime` | TIMESTAMPTZ | +//! | `time::Date` | DATE | +//! | `time::Time` | TIME | +//! | [`PgTimeTz`] | TIMETZ | +//! +//! ### [`uuid`](https://crates.io/crates/uuid) +//! +//! Requires the `uuid` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `uuid::Uuid` | UUID | +//! +//! ### [`ipnetwork`](https://crates.io/crates/ipnetwork) +//! +//! Requires the `ipnetwork` Cargo feature flag (takes precedence over `ipnet` if both are used). +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `ipnetwork::IpNetwork` | INET, CIDR | +//! | `std::net::IpAddr` | INET, CIDR | +//! +//! Note that because `IpAddr` does not support network prefixes, it is an error to attempt to decode +//! an `IpAddr` from a `INET` or `CIDR` value with a network prefix smaller than the address' full width: +//! `/32` for IPv4 addresses and `/128` for IPv6 addresses. +//! +//! `IpNetwork` does not have this limitation. +//! +//! ### [`ipnet`](https://crates.io/crates/ipnet) +//! +//! Requires the `ipnet` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `ipnet::IpNet` | INET, CIDR | +//! | `std::net::IpAddr` | INET, CIDR | +//! +//! The same `IpAddr` limitation for smaller network prefixes applies as with `ipnet`. +//! +//! ### [`mac_address`](https://crates.io/crates/mac_address) +//! +//! Requires the `mac_address` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `mac_address::MacAddress` | MACADDR | +//! +//! ### [`bit-vec`](https://crates.io/crates/bit-vec) +//! +//! Requires the `bit-vec` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | `bit_vec::BitVec` | BIT, VARBIT | +//! +//! ### [`json`](https://crates.io/crates/serde_json) +//! +//! Requires the `json` Cargo feature flag. +//! +//! | Rust type | Postgres type(s) | +//! |---------------------------------------|------------------------------------------------------| +//! | [`Json`] | JSON, JSONB | +//! | `serde_json::Value` | JSON, JSONB | +//! | `&serde_json::value::RawValue` | JSON, JSONB | +//! +//! `Value` and `RawValue` from `serde_json` can be used for unstructured JSON data with +//! Postgres. +//! +//! [`Json`](crate::types::Json) can be used for structured JSON data with Postgres. +//! +//! # [Composite types](https://www.postgresql.org/docs/current/rowtypes.html) +//! +//! User-defined composite types are supported through a derive for `Type`. +//! +//! ```text +//! CREATE TYPE inventory_item AS ( +//! name text, +//! supplier_id integer, +//! price numeric +//! ); +//! ``` +//! +//! ```rust,ignore +//! #[derive(sqlx::Type)] +//! #[sqlx(type_name = "inventory_item")] +//! struct InventoryItem { +//! name: String, +//! supplier_id: i32, +//! price: BigDecimal, +//! } +//! ``` +//! +//! Anonymous composite types are represented as tuples. Note that anonymous composites may only +//! be returned and not sent to Postgres (this is a limitation of postgres). +//! +//! # Arrays +//! +//! One-dimensional arrays are supported as `Vec` or `&[T]` where `T` implements `Type`. +//! +//! # [Enumerations](https://www.postgresql.org/docs/current/datatype-enum.html) +//! +//! User-defined enumerations are supported through a derive for `Type`. +//! +//! ```text +//! CREATE TYPE mood AS ENUM ('sad', 'ok', 'happy'); +//! ``` +//! +//! ```rust,ignore +//! #[derive(sqlx::Type)] +//! #[sqlx(type_name = "mood", rename_all = "lowercase")] +//! enum Mood { Sad, Ok, Happy } +//! ``` +//! +//! Rust enumerations may also be defined to be represented as an integer using `repr`. +//! The following type expects a SQL type of `INTEGER` or `INT4` and will convert to/from the +//! Rust enumeration. +//! +//! ```rust,ignore +//! #[derive(sqlx::Type)] +//! #[repr(i32)] +//! enum Mood { Sad = 0, Ok = 1, Happy = 2 } +//! ``` +//! +//! Rust enumerations may also be defined to be represented as a string using `type_name = "text"`. +//! The following type expects a SQL type of `TEXT` and will convert to/from the Rust enumeration. +//! +//! ```rust,ignore +//! #[derive(sqlx::Type)] +//! #[sqlx(type_name = "text")] +//! enum Mood { Sad, Ok, Happy } +//! ``` +//! +//! Note that an error can occur if you attempt to decode a value not contained within the enum +//! definition. +//! + +use crate::type_info::PgTypeKind; +use crate::{PgTypeInfo, Postgres}; + +pub(crate) use sqlx_core::types::{Json, Type}; + +mod array; +mod bool; +mod bytes; +mod citext; +mod float; +mod hstore; +mod int; +mod interval; +mod lquery; +mod ltree; +// Not behind a Cargo feature because we require JSON in the driver implementation. +mod json; +mod money; +mod oid; +mod range; +mod record; +mod str; +mod text; +mod tuple; +mod void; + +#[cfg(any(feature = "chrono", feature = "time"))] +mod time_tz; + +#[cfg(feature = "bigdecimal")] +mod bigdecimal; + +mod cube; + +mod geometry; + +#[cfg(any(feature = "bigdecimal", feature = "rust_decimal"))] +mod numeric; + +#[cfg(feature = "rust_decimal")] +mod rust_decimal; + +#[cfg(feature = "chrono")] +mod chrono; + +#[cfg(feature = "time")] +mod time; + +#[cfg(feature = "uuid")] +mod uuid; + +#[cfg(feature = "ipnet")] +mod ipnet; + +#[cfg(feature = "ipnetwork")] +mod ipnetwork; + +#[cfg(feature = "mac_address")] +mod mac_address; + +#[cfg(feature = "bit-vec")] +mod bit_vec; + +pub use array::PgHasArrayType; +pub use citext::PgCiText; +pub use cube::PgCube; +pub use geometry::circle::PgCircle; +pub use geometry::line::PgLine; +pub use geometry::line_segment::PgLSeg; +pub use geometry::path::PgPath; +pub use geometry::point::PgPoint; +pub use geometry::polygon::PgPolygon; +pub use geometry::r#box::PgBox; +pub use hstore::PgHstore; +pub use interval::PgInterval; +pub use lquery::PgLQuery; +pub use lquery::PgLQueryLevel; +pub use lquery::PgLQueryVariant; +pub use lquery::PgLQueryVariantFlag; +pub use ltree::PgLTree; +pub use ltree::PgLTreeLabel; +pub use ltree::PgLTreeParseError; +pub use money::PgMoney; +pub use oid::Oid; +pub use range::PgRange; + +#[cfg(any(feature = "chrono", feature = "time"))] +pub use time_tz::PgTimeTz; + +// used in derive(Type) for `struct` +// but the interface is not considered part of the public API +#[doc(hidden)] +pub use record::{PgRecordDecoder, PgRecordEncoder}; + +// Type::compatible impl appropriate for arrays +fn array_compatible + ?Sized>(ty: &PgTypeInfo) -> bool { + // we require the declared type to be an _array_ with an + // element type that is acceptable + if let PgTypeKind::Array(element) = &ty.kind() { + return E::compatible(element); + } + + false +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/money.rs b/src-tauri/vendor/sqlx-postgres/src/types/money.rs new file mode 100644 index 00000000..52fc6879 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/money.rs @@ -0,0 +1,365 @@ +use crate::{ + decode::Decode, + encode::{Encode, IsNull}, + error::BoxDynError, + types::Type, + {PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}, +}; +use byteorder::{BigEndian, ByteOrder}; +use std::{ + io, + ops::{Add, AddAssign, Sub, SubAssign}, +}; + +/// The PostgreSQL [`MONEY`] type stores a currency amount with a fixed fractional +/// precision. The fractional precision is determined by the database's +/// `lc_monetary` setting. +/// +/// Data is read and written as 64-bit signed integers, and conversion into a +/// decimal should be done using the right precision. +/// +/// Reading `MONEY` value in text format is not supported and will cause an error. +/// +/// ### `locale_frac_digits` +/// This parameter corresponds to the number of digits after the decimal separator. +/// +/// This value must match what Postgres is expecting for the locale set in the database +/// or else the decimal value you see on the client side will not match the `money` value +/// on the server side. +/// +/// **For _most_ locales, this value is `2`.** +/// +/// If you're not sure what locale your database is set to or how many decimal digits it specifies, +/// you can execute `SHOW lc_monetary;` to get the locale name, and then look it up in this list +/// (you can ignore the `.utf8` prefix): +/// +/// +/// If that link is dead and you're on a POSIX-compliant system (Unix, FreeBSD) you can also execute: +/// +/// ```sh +/// $ LC_MONETARY= locale -k frac_digits +/// ``` +/// +/// And the value you want is `N` in `frac_digits=N`. If you have shell access to the database +/// server you should execute it there as available locales may differ between machines. +/// +/// Note that if `frac_digits` for the locale is outside the range `[0, 10]`, Postgres assumes +/// it's a sentinel value and defaults to 2: +/// +/// +/// [`MONEY`]: https://www.postgresql.org/docs/current/datatype-money.html +#[derive(Debug, PartialEq, Eq, Clone, Copy, Default)] +pub struct PgMoney( + /// The raw integer value sent over the wire; for locales with `frac_digits=2` (i.e. most + /// of them), this will be the value in whole cents. + /// + /// E.g. for `select '$123.45'::money` with a locale of `en_US` (`frac_digits=2`), + /// this will be `12345`. + /// + /// If the currency of your locale does not have fractional units, e.g. Yen, then this will + /// just be the units of the currency. + /// + /// See the type-level docs for an explanation of `locale_frac_units`. + pub i64, +); + +impl PgMoney { + /// Convert the money value into a [`BigDecimal`] using `locale_frac_digits`. + /// + /// See the type-level docs for an explanation of `locale_frac_digits`. + /// + /// [`BigDecimal`]: bigdecimal::BigDecimal + #[cfg(feature = "bigdecimal")] + pub fn to_bigdecimal(self, locale_frac_digits: i64) -> bigdecimal::BigDecimal { + let digits = num_bigint::BigInt::from(self.0); + + bigdecimal::BigDecimal::new(digits, locale_frac_digits) + } + + /// Convert the money value into a [`Decimal`] using `locale_frac_digits`. + /// + /// See the type-level docs for an explanation of `locale_frac_digits`. + /// + /// [`Decimal`]: rust_decimal::Decimal + #[cfg(feature = "rust_decimal")] + pub fn to_decimal(self, locale_frac_digits: u32) -> rust_decimal::Decimal { + rust_decimal::Decimal::new(self.0, locale_frac_digits) + } + + /// Convert a [`Decimal`] value into money using `locale_frac_digits`. + /// + /// See the type-level docs for an explanation of `locale_frac_digits`. + /// + /// Note that `Decimal` has 96 bits of precision, but `PgMoney` only has 63 plus the sign bit. + /// If the value is larger than 63 bits it will be truncated. + /// + /// [`Decimal`]: rust_decimal::Decimal + #[cfg(feature = "rust_decimal")] + pub fn from_decimal(mut decimal: rust_decimal::Decimal, locale_frac_digits: u32) -> Self { + // this is all we need to convert to our expected locale's `frac_digits` + decimal.rescale(locale_frac_digits); + + /// a mask to bitwise-AND with an `i64` to zero the sign bit + const SIGN_MASK: i64 = i64::MAX; + + let is_negative = decimal.is_sign_negative(); + let serialized = decimal.serialize(); + + // interpret bytes `4..12` as an i64, ignoring the sign bit + // this is where truncation occurs + let value = i64::from_le_bytes( + *<&[u8; 8]>::try_from(&serialized[4..12]) + .expect("BUG: slice of serialized should be 8 bytes"), + ) & SIGN_MASK; // zero out the sign bit + + // negate if necessary + Self(if is_negative { -value } else { value }) + } + + /// Convert a [`BigDecimal`](bigdecimal::BigDecimal) value into money using the correct precision + /// defined in the PostgreSQL settings. The default precision is two. + #[cfg(feature = "bigdecimal")] + pub fn from_bigdecimal( + decimal: bigdecimal::BigDecimal, + locale_frac_digits: u32, + ) -> Result { + use bigdecimal::ToPrimitive; + + let multiplier = bigdecimal::BigDecimal::new( + num_bigint::BigInt::from(10i128.pow(locale_frac_digits)), + 0, + ); + + let cents = decimal * multiplier; + + let money = cents.to_i64().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "Provided BigDecimal could not convert to i64: overflow.", + ) + })?; + + Ok(Self(money)) + } +} + +impl Type for PgMoney { + fn type_info() -> PgTypeInfo { + PgTypeInfo::MONEY + } +} + +impl PgHasArrayType for PgMoney { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::MONEY_ARRAY + } +} + +impl From for PgMoney +where + T: Into, +{ + fn from(num: T) -> Self { + Self(num.into()) + } +} + +impl Encode<'_, Postgres> for PgMoney { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.0.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for PgMoney { + fn decode(value: PgValueRef<'_>) -> Result { + match value.format() { + PgValueFormat::Binary => { + let cents = BigEndian::read_i64(value.as_bytes()?); + + Ok(PgMoney(cents)) + } + PgValueFormat::Text => { + let error = io::Error::new( + io::ErrorKind::InvalidData, + "Reading a `MONEY` value in text format is not supported.", + ); + + Err(Box::new(error)) + } + } + } +} + +impl Add for PgMoney { + type Output = PgMoney; + + /// Adds two monetary values. + /// + /// # Panics + /// Panics if overflowing the `i64::MAX`. + fn add(self, rhs: PgMoney) -> Self::Output { + self.0 + .checked_add(rhs.0) + .map(PgMoney) + .expect("overflow adding money amounts") + } +} + +impl AddAssign for PgMoney { + /// An assigning add for two monetary values. + /// + /// # Panics + /// Panics if overflowing the `i64::MAX`. + fn add_assign(&mut self, rhs: PgMoney) { + self.0 = self + .0 + .checked_add(rhs.0) + .expect("overflow adding money amounts") + } +} + +impl Sub for PgMoney { + type Output = PgMoney; + + /// Subtracts two monetary values. + /// + /// # Panics + /// Panics if underflowing the `i64::MIN`. + fn sub(self, rhs: PgMoney) -> Self::Output { + self.0 + .checked_sub(rhs.0) + .map(PgMoney) + .expect("overflow subtracting money amounts") + } +} + +impl SubAssign for PgMoney { + /// An assigning subtract for two monetary values. + /// + /// # Panics + /// Panics if underflowing the `i64::MIN`. + fn sub_assign(&mut self, rhs: PgMoney) { + self.0 = self + .0 + .checked_sub(rhs.0) + .expect("overflow subtracting money amounts") + } +} + +#[cfg(test)] +mod tests { + use super::PgMoney; + + #[test] + fn adding_works() { + assert_eq!(PgMoney(3), PgMoney(1) + PgMoney(2)) + } + + #[test] + fn add_assign_works() { + let mut money = PgMoney(1); + money += PgMoney(2); + + assert_eq!(PgMoney(3), money); + } + + #[test] + fn subtracting_works() { + assert_eq!(PgMoney(4), PgMoney(5) - PgMoney(1)) + } + + #[test] + fn sub_assign_works() { + let mut money = PgMoney(1); + money -= PgMoney(2); + + assert_eq!(PgMoney(-1), money); + } + + #[test] + fn default_value() { + let money = PgMoney::default(); + + assert_eq!(money, PgMoney(0)); + } + + #[test] + #[should_panic] + fn add_overflow_panics() { + let _ = PgMoney(i64::MAX) + PgMoney(1); + } + + #[test] + #[should_panic] + fn add_assign_overflow_panics() { + let mut money = PgMoney(i64::MAX); + money += PgMoney(1); + } + + #[test] + #[should_panic] + fn sub_overflow_panics() { + let _ = PgMoney(i64::MIN) - PgMoney(1); + } + + #[test] + #[should_panic] + fn sub_assign_overflow_panics() { + let mut money = PgMoney(i64::MIN); + money -= PgMoney(1); + } + + #[test] + #[cfg(feature = "bigdecimal")] + fn conversion_to_bigdecimal_works() { + let money = PgMoney(12345); + + assert_eq!( + bigdecimal::BigDecimal::new(num_bigint::BigInt::from(12345), 2), + money.to_bigdecimal(2) + ); + } + + #[test] + #[cfg(feature = "rust_decimal")] + fn conversion_to_decimal_works() { + assert_eq!( + rust_decimal::Decimal::new(12345, 2), + PgMoney(12345).to_decimal(2) + ); + } + + #[test] + #[cfg(feature = "rust_decimal")] + fn conversion_from_decimal_works() { + assert_eq!( + PgMoney(12345), + PgMoney::from_decimal(rust_decimal::Decimal::new(12345, 2), 2) + ); + + assert_eq!( + PgMoney(12345), + PgMoney::from_decimal(rust_decimal::Decimal::new(123450, 3), 2) + ); + + assert_eq!( + PgMoney(-12345), + PgMoney::from_decimal(rust_decimal::Decimal::new(-123450, 3), 2) + ); + + assert_eq!( + PgMoney(-12300), + PgMoney::from_decimal(rust_decimal::Decimal::new(-123, 0), 2) + ); + } + + #[test] + #[cfg(feature = "bigdecimal")] + fn conversion_from_bigdecimal_works() { + let dec = bigdecimal::BigDecimal::new(num_bigint::BigInt::from(12345), 2); + + assert_eq!(PgMoney(12345), PgMoney::from_bigdecimal(dec, 2).unwrap()); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/numeric.rs b/src-tauri/vendor/sqlx-postgres/src/types/numeric.rs new file mode 100644 index 00000000..67713d76 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/numeric.rs @@ -0,0 +1,172 @@ +use sqlx_core::bytes::Buf; +use std::num::Saturating; + +use crate::error::BoxDynError; +use crate::PgArgumentBuffer; + +/// Represents a `NUMERIC` value in the **Postgres** wire protocol. +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum PgNumeric { + /// Equivalent to the `'NaN'` value in Postgres. The result of, e.g. `1 / 0`. + NotANumber, + + /// A populated `NUMERIC` value. + /// + /// A description of these fields can be found here (although the type being described is the + /// version for in-memory calculations, the field names are the same): + /// https://github.com/postgres/postgres/blob/bcd1c3630095e48bc3b1eb0fc8e8c8a7c851eba1/src/backend/utils/adt/numeric.c#L224-L269 + Number { + /// The sign of the value: positive (also set for 0 and -0), or negative. + sign: PgNumericSign, + + /// The digits of the number in base-10000 with the most significant digit first + /// (big-endian). + /// + /// The length of this vector must not overflow `i16` for the binary protocol. + /// + /// *Note*: the `Encode` implementation will panic if any digit is `>= 10000`. + digits: Vec, + + /// The scaling factor of the number, such that the value will be interpreted as + /// + /// ```text + /// digits[0] * 10,000 ^ weight + /// + digits[1] * 10,000 ^ (weight - 1) + /// ... + /// + digits[N] * 10,000 ^ (weight - N) where N = digits.len() - 1 + /// ``` + /// May be negative. + weight: i16, + + /// How many _decimal_ (base-10) digits following the decimal point to consider in + /// arithmetic regardless of how many actually follow the decimal point as determined by + /// `weight`--the comment in the Postgres code linked above recommends using this only for + /// ignoring unnecessary trailing zeroes (as trimming nonzero digits means reducing the + /// precision of the value). + /// + /// Must be `>= 0`. + scale: i16, + }, +} + +// https://github.com/postgres/postgres/blob/bcd1c3630095e48bc3b1eb0fc8e8c8a7c851eba1/src/backend/utils/adt/numeric.c#L167-L170 +const SIGN_POS: u16 = 0x0000; +const SIGN_NEG: u16 = 0x4000; +const SIGN_NAN: u16 = 0xC000; // overflows i16 (C equivalent truncates from integer literal) + +/// Possible sign values for [PgNumeric]. +#[derive(Copy, Clone, Debug, PartialEq, Eq)] +#[repr(u16)] +pub(crate) enum PgNumericSign { + Positive = SIGN_POS, + Negative = SIGN_NEG, +} + +impl PgNumericSign { + fn try_from_u16(val: u16) -> Result { + match val { + SIGN_POS => Ok(PgNumericSign::Positive), + SIGN_NEG => Ok(PgNumericSign::Negative), + + SIGN_NAN => unreachable!("sign value for NaN passed to PgNumericSign"), + + _ => Err(format!("invalid value for PgNumericSign: {val:#04X}").into()), + } + } +} + +impl PgNumeric { + /// Equivalent value of `0::numeric`. + pub const ZERO: Self = PgNumeric::Number { + sign: PgNumericSign::Positive, + digits: vec![], + weight: 0, + scale: 0, + }; + + pub(crate) fn is_valid_digit(digit: i16) -> bool { + (0..10_000).contains(&digit) + } + + pub(crate) fn size_hint(decimal_digits: u64) -> usize { + let mut size_hint = Saturating(decimal_digits); + + // BigDecimal::digits() gives us base-10 digits, so we divide by 4 to get base-10000 digits + // and since this is just a hint we just always round up + size_hint /= 4; + size_hint += 1; + + // Times two bytes for each base-10000 digit + size_hint *= 2; + + // Plus `weight` and `scale` + size_hint += 8; + + usize::try_from(size_hint.0).unwrap_or(usize::MAX) + } + + pub(crate) fn decode(mut buf: &[u8]) -> Result { + // https://github.com/postgres/postgres/blob/bcd1c3630095e48bc3b1eb0fc8e8c8a7c851eba1/src/backend/utils/adt/numeric.c#L874 + let num_digits = buf.get_u16(); + let weight = buf.get_i16(); + let sign = buf.get_u16(); + let scale = buf.get_i16(); + + if sign == SIGN_NAN { + Ok(PgNumeric::NotANumber) + } else { + let digits: Vec<_> = (0..num_digits).map(|_| buf.get_i16()).collect::<_>(); + + Ok(PgNumeric::Number { + sign: PgNumericSign::try_from_u16(sign)?, + scale, + weight, + digits, + }) + } + } + + /// ### Errors + /// + /// * If `digits.len()` overflows `i16` + /// * If any element in `digits` is greater than or equal to 10000 + pub(crate) fn encode(&self, buf: &mut PgArgumentBuffer) -> Result<(), String> { + match *self { + PgNumeric::Number { + ref digits, + sign, + scale, + weight, + } => { + let digits_len = i16::try_from(digits.len()).map_err(|_| { + format!( + "PgNumeric digits.len() ({}) should not overflow i16", + digits.len() + ) + })?; + + buf.extend(&digits_len.to_be_bytes()); + buf.extend(&weight.to_be_bytes()); + buf.extend(&(sign as i16).to_be_bytes()); + buf.extend(&scale.to_be_bytes()); + + for (i, &digit) in digits.iter().enumerate() { + if !Self::is_valid_digit(digit) { + return Err(format!("{i}th PgNumeric digit out of range: {digit}")); + } + + buf.extend(&digit.to_be_bytes()); + } + } + + PgNumeric::NotANumber => { + buf.extend(&0_i16.to_be_bytes()); + buf.extend(&0_i16.to_be_bytes()); + buf.extend(&SIGN_NAN.to_be_bytes()); + buf.extend(&0_i16.to_be_bytes()); + } + } + + Ok(()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/oid.rs b/src-tauri/vendor/sqlx-postgres/src/types/oid.rs new file mode 100644 index 00000000..04c5ef83 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/oid.rs @@ -0,0 +1,65 @@ +use byteorder::{BigEndian, ByteOrder}; +use serde::{de::Deserializer, ser::Serializer, Deserialize, Serialize}; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +/// The PostgreSQL [`OID`] type stores an object identifier, +/// used internally by PostgreSQL as primary keys for various system tables. +/// +/// [`OID`]: https://www.postgresql.org/docs/current/datatype-oid.html +#[derive(Debug, Copy, Clone, Hash, PartialEq, Eq, Default)] +pub struct Oid( + /// The raw unsigned integer value sent over the wire + pub u32, +); + +impl Type for Oid { + fn type_info() -> PgTypeInfo { + PgTypeInfo::OID + } +} + +impl PgHasArrayType for Oid { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::OID_ARRAY + } +} + +impl Encode<'_, Postgres> for Oid { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(&self.0.to_be_bytes()); + + Ok(IsNull::No) + } +} + +impl Decode<'_, Postgres> for Oid { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(Self(match value.format() { + PgValueFormat::Binary => BigEndian::read_u32(value.as_bytes()?), + PgValueFormat::Text => value.as_str()?.parse()?, + })) + } +} + +impl Serialize for Oid { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + self.0.serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for Oid { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + u32::deserialize(deserializer).map(Self) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/range.rs b/src-tauri/vendor/sqlx-postgres/src/types/range.rs new file mode 100644 index 00000000..0d9c14bd --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/range.rs @@ -0,0 +1,522 @@ +use std::fmt::{self, Debug, Display, Formatter}; +use std::ops::{Bound, Range, RangeBounds, RangeFrom, RangeInclusive, RangeTo, RangeToInclusive}; + +use bitflags::bitflags; +use sqlx_core::bytes::Buf; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::type_info::PgTypeKind; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +// https://github.com/postgres/postgres/blob/2f48ede080f42b97b594fb14102c82ca1001b80c/src/include/utils/rangetypes.h#L35-L44 +bitflags! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + struct RangeFlags: u8 { + const EMPTY = 0x01; + const LB_INC = 0x02; + const UB_INC = 0x04; + const LB_INF = 0x08; + const UB_INF = 0x10; + const LB_NULL = 0x20; // not used + const UB_NULL = 0x40; // not used + const CONTAIN_EMPTY = 0x80; // internal + } +} + +#[derive(Debug, PartialEq, Eq, Clone, Copy)] +pub struct PgRange { + pub start: Bound, + pub end: Bound, +} + +impl From<[Bound; 2]> for PgRange { + fn from(v: [Bound; 2]) -> Self { + let [start, end] = v; + Self { start, end } + } +} + +impl From<(Bound, Bound)> for PgRange { + fn from(v: (Bound, Bound)) -> Self { + Self { + start: v.0, + end: v.1, + } + } +} + +impl From> for PgRange { + fn from(v: Range) -> Self { + Self { + start: Bound::Included(v.start), + end: Bound::Excluded(v.end), + } + } +} + +impl From> for PgRange { + fn from(v: RangeFrom) -> Self { + Self { + start: Bound::Included(v.start), + end: Bound::Unbounded, + } + } +} + +impl From> for PgRange { + fn from(v: RangeInclusive) -> Self { + let (start, end) = v.into_inner(); + Self { + start: Bound::Included(start), + end: Bound::Included(end), + } + } +} + +impl From> for PgRange { + fn from(v: RangeTo) -> Self { + Self { + start: Bound::Unbounded, + end: Bound::Excluded(v.end), + } + } +} + +impl From> for PgRange { + fn from(v: RangeToInclusive) -> Self { + Self { + start: Bound::Unbounded, + end: Bound::Included(v.end), + } + } +} + +impl RangeBounds for PgRange { + fn start_bound(&self) -> Bound<&T> { + match self.start { + Bound::Included(ref start) => Bound::Included(start), + Bound::Excluded(ref start) => Bound::Excluded(start), + Bound::Unbounded => Bound::Unbounded, + } + } + + fn end_bound(&self) -> Bound<&T> { + match self.end { + Bound::Included(ref end) => Bound::Included(end), + Bound::Excluded(ref end) => Bound::Excluded(end), + Bound::Unbounded => Bound::Unbounded, + } + } +} + +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INT4_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::INT8_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "bigdecimal")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::NUM_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "rust_decimal")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::NUM_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "chrono")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::DATE_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "chrono")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TS_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "chrono")] +impl Type for PgRange> { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TSTZ_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::>(ty) + } +} + +#[cfg(feature = "time")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::DATE_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "time")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TS_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +#[cfg(feature = "time")] +impl Type for PgRange { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TSTZ_RANGE + } + + fn compatible(ty: &PgTypeInfo) -> bool { + range_compatible::(ty) + } +} + +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT4_RANGE_ARRAY + } +} + +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::INT8_RANGE_ARRAY + } +} + +#[cfg(feature = "bigdecimal")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::NUM_RANGE_ARRAY + } +} + +#[cfg(feature = "rust_decimal")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::NUM_RANGE_ARRAY + } +} + +#[cfg(feature = "chrono")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::DATE_RANGE_ARRAY + } +} + +#[cfg(feature = "chrono")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TS_RANGE_ARRAY + } +} + +#[cfg(feature = "chrono")] +impl PgHasArrayType for PgRange> { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TSTZ_RANGE_ARRAY + } +} + +#[cfg(feature = "time")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::DATE_RANGE_ARRAY + } +} + +#[cfg(feature = "time")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TS_RANGE_ARRAY + } +} + +#[cfg(feature = "time")] +impl PgHasArrayType for PgRange { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TSTZ_RANGE_ARRAY + } +} + +impl<'q, T> Encode<'q, Postgres> for PgRange +where + T: Encode<'q, Postgres>, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // https://github.com/postgres/postgres/blob/2f48ede080f42b97b594fb14102c82ca1001b80c/src/backend/utils/adt/rangetypes.c#L245 + + let mut flags = RangeFlags::empty(); + + flags |= match self.start { + Bound::Included(_) => RangeFlags::LB_INC, + Bound::Unbounded => RangeFlags::LB_INF, + Bound::Excluded(_) => RangeFlags::empty(), + }; + + flags |= match self.end { + Bound::Included(_) => RangeFlags::UB_INC, + Bound::Unbounded => RangeFlags::UB_INF, + Bound::Excluded(_) => RangeFlags::empty(), + }; + + buf.push(flags.bits()); + + if let Bound::Included(v) | Bound::Excluded(v) = &self.start { + buf.encode(v)?; + } + + if let Bound::Included(v) | Bound::Excluded(v) = &self.end { + buf.encode(v)?; + } + + // ranges are themselves never null + Ok(IsNull::No) + } +} + +impl<'r, T> Decode<'r, Postgres> for PgRange +where + T: Type + for<'a> Decode<'a, Postgres>, +{ + fn decode(value: PgValueRef<'r>) -> Result { + match value.format { + PgValueFormat::Binary => { + let element_ty = if let PgTypeKind::Range(element) = &value.type_info.0.kind() { + element + } else { + return Err(format!("unexpected non-range type {}", value.type_info).into()); + }; + + let mut buf = value.as_bytes()?; + + let mut start = Bound::Unbounded; + let mut end = Bound::Unbounded; + + let flags = RangeFlags::from_bits_truncate(buf.get_u8()); + + if flags.contains(RangeFlags::EMPTY) { + return Ok(PgRange { start, end }); + } + + if !flags.contains(RangeFlags::LB_INF) { + let value = + T::decode(PgValueRef::get(&mut buf, value.format, element_ty.clone())?)?; + + start = if flags.contains(RangeFlags::LB_INC) { + Bound::Included(value) + } else { + Bound::Excluded(value) + }; + } + + if !flags.contains(RangeFlags::UB_INF) { + let value = + T::decode(PgValueRef::get(&mut buf, value.format, element_ty.clone())?)?; + + end = if flags.contains(RangeFlags::UB_INC) { + Bound::Included(value) + } else { + Bound::Excluded(value) + }; + } + + Ok(PgRange { start, end }) + } + + PgValueFormat::Text => { + // https://github.com/postgres/postgres/blob/2f48ede080f42b97b594fb14102c82ca1001b80c/src/backend/utils/adt/rangetypes.c#L2046 + + let mut start = None; + let mut end = None; + + let s = value.as_str()?; + + // remember the bounds + let sb = s.as_bytes(); + let lower = sb[0] as char; + let upper = sb[sb.len() - 1] as char; + + // trim the wrapping braces/brackets + let s = &s[1..(s.len() - 1)]; + + let mut chars = s.chars(); + + let mut element = String::new(); + let mut done = false; + let mut quoted = false; + let mut in_quotes = false; + let mut in_escape = false; + let mut prev_ch = '\0'; + let mut count = 0; + + while !done { + element.clear(); + + loop { + match chars.next() { + Some(ch) => { + match ch { + _ if in_escape => { + element.push(ch); + in_escape = false; + } + + '"' if in_quotes => { + in_quotes = false; + } + + '"' => { + in_quotes = true; + quoted = true; + + if prev_ch == '"' { + element.push('"') + } + } + + '\\' if !in_escape => { + in_escape = true; + } + + ',' if !in_quotes => break, + + _ => { + element.push(ch); + } + } + prev_ch = ch; + } + + None => { + done = true; + break; + } + } + } + + count += 1; + if !element.is_empty() || quoted { + let value = Some(T::decode(PgValueRef { + type_info: T::type_info(), + format: PgValueFormat::Text, + value: Some(element.as_bytes()), + row: None, + })?); + + if count == 1 { + start = value; + } else if count == 2 { + end = value; + } else { + return Err("more than 2 elements found in a range".into()); + } + } + } + + let start = parse_bound(lower, start)?; + let end = parse_bound(upper, end)?; + + Ok(PgRange { start, end }) + } + } + } +} + +fn parse_bound(ch: char, value: Option) -> Result, BoxDynError> { + Ok(if let Some(value) = value { + match ch { + '(' | ')' => Bound::Excluded(value), + '[' | ']' => Bound::Included(value), + + _ => { + return Err(format!( + "expected `(`, ')', '[', or `]` but found `{ch}` for range literal" + ) + .into()); + } + } + } else { + Bound::Unbounded + }) +} + +impl Display for PgRange +where + T: Display, +{ + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + match &self.start { + Bound::Unbounded => f.write_str("(,")?, + Bound::Excluded(v) => write!(f, "({v},")?, + Bound::Included(v) => write!(f, "[{v},")?, + } + + match &self.end { + Bound::Unbounded => f.write_str(")")?, + Bound::Excluded(v) => write!(f, "{v})")?, + Bound::Included(v) => write!(f, "{v}]")?, + } + + Ok(()) + } +} + +fn range_compatible>(ty: &PgTypeInfo) -> bool { + // we require the declared type to be a _range_ with an + // element type that is acceptable + if let PgTypeKind::Range(element) = &ty.kind() { + return E::compatible(element); + } + + false +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/record.rs b/src-tauri/vendor/sqlx-postgres/src/types/record.rs new file mode 100644 index 00000000..6e37182c --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/record.rs @@ -0,0 +1,205 @@ +use sqlx_core::bytes::Buf; + +use crate::decode::Decode; +use crate::encode::Encode; +use crate::error::{mismatched_types, BoxDynError}; +use crate::type_info::TypeInfo; +use crate::type_info::{PgType, PgTypeKind}; +use crate::types::Oid; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +#[doc(hidden)] +pub struct PgRecordEncoder<'a> { + buf: &'a mut PgArgumentBuffer, + off: usize, + num: u32, +} + +impl<'a> PgRecordEncoder<'a> { + #[doc(hidden)] + pub fn new(buf: &'a mut PgArgumentBuffer) -> Self { + let off = buf.len(); + + // reserve space for a field count + buf.extend(&(0_u32).to_be_bytes()); + + Self { buf, off, num: 0 } + } + + #[doc(hidden)] + pub fn finish(&mut self) { + // fill in the record length + self.buf[self.off..(self.off + 4)].copy_from_slice(&self.num.to_be_bytes()); + } + + #[doc(hidden)] + pub fn encode<'q, T>(&mut self, value: T) -> Result<&mut Self, BoxDynError> + where + 'a: 'q, + T: Encode<'q, Postgres> + Type, + { + let ty = value.produces().unwrap_or_else(T::type_info); + + match ty.0 { + // push a hole for this type ID + // to be filled in on query execution + PgType::DeclareWithName(name) => self.buf.patch_type_by_name(&name), + PgType::DeclareArrayOf(array) => self.buf.patch_array_type(array), + // write type id + pg_type => self.buf.extend(&pg_type.oid().0.to_be_bytes()), + } + + self.buf.encode(value)?; + self.num += 1; + + Ok(self) + } +} + +#[doc(hidden)] +pub struct PgRecordDecoder<'r> { + buf: &'r [u8], + typ: PgTypeInfo, + fmt: PgValueFormat, + ind: usize, +} + +impl<'r> PgRecordDecoder<'r> { + #[doc(hidden)] + pub fn new(value: PgValueRef<'r>) -> Result { + let fmt = value.format(); + let mut buf = value.as_bytes()?; + let typ = value.type_info; + + match fmt { + PgValueFormat::Binary => { + let _len = buf.get_u32(); + } + + PgValueFormat::Text => { + // remove the enclosing `(` .. `)` + buf = &buf[1..(buf.len() - 1)]; + } + } + + Ok(Self { + buf, + fmt, + typ, + ind: 0, + }) + } + + #[doc(hidden)] + pub fn try_decode(&mut self) -> Result + where + T: for<'a> Decode<'a, Postgres> + Type, + { + if self.buf.is_empty() { + return Err(format!("no field `{0}` found on record", self.ind).into()); + } + + match self.fmt { + PgValueFormat::Binary => { + let element_type_oid = Oid(self.buf.get_u32()); + let element_type_opt = match self.typ.0.kind() { + PgTypeKind::Simple if self.typ.0 == PgType::Record => { + PgTypeInfo::try_from_oid(element_type_oid) + } + + PgTypeKind::Composite(fields) => { + let ty = fields[self.ind].1.clone(); + if ty.0.oid() != element_type_oid { + return Err("unexpected mismatch of composite type information".into()); + } + + Some(ty) + } + + _ => { + return Err( + "unexpected non-composite type being decoded as a composite type" + .into(), + ); + } + }; + + if let Some(ty) = &element_type_opt { + if !ty.is_null() && !T::compatible(ty) { + return Err(mismatched_types::(ty)); + } + } + + let element_type = + element_type_opt + .ok_or_else(|| BoxDynError::from(format!("custom types in records are not fully supported yet: failed to retrieve type info for field {} with type oid {}", self.ind, element_type_oid.0)))?; + + self.ind += 1; + + T::decode(PgValueRef::get(&mut self.buf, self.fmt, element_type)?) + } + + PgValueFormat::Text => { + let mut element = String::new(); + let mut quoted = false; + let mut in_quotes = false; + let mut in_escape = false; + let mut prev_ch = '\0'; + + while !self.buf.is_empty() { + let ch = self.buf.get_u8() as char; + match ch { + _ if in_escape => { + element.push(ch); + in_escape = false; + } + + '"' if in_quotes => { + in_quotes = false; + } + + '"' => { + in_quotes = true; + quoted = true; + + if prev_ch == '"' { + element.push('"') + } + } + + '\\' if !in_escape => { + in_escape = true; + } + + ',' if !in_quotes => break, + + _ => { + element.push(ch); + } + } + prev_ch = ch; + } + + let buf = if element.is_empty() && !quoted { + // completely empty input means NULL + None + } else { + Some(element.as_bytes()) + }; + + // NOTE: we do not call [`accepts`] or give a chance to from a user as + // TEXT sequences are not strongly typed + + T::decode(PgValueRef { + // NOTE: We pass `0` as the type ID because we don't have a reasonable value + // we could use. + type_info: PgTypeInfo::with_oid(Oid(0)), + format: self.fmt, + value: buf, + row: None, + }) + } + } + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal-range.md b/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal-range.md new file mode 100644 index 00000000..f986d616 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal-range.md @@ -0,0 +1,10 @@ +#### Note: `rust_decimal::Decimal` Has a Smaller Range than `NUMERIC` +`NUMERIC` is can have up to 131,072 digits before the decimal point, and 16,384 digits after it. +See [Section 8.1, Numeric Types] of the Postgres manual for details. + +However, `rust_decimal::Decimal` is limited to a maximum absolute magnitude of 296 - 1, +a number with 67 decimal digits, and a minimum absolute magnitude of 10-28, a number with, unsurprisingly, +28 decimal digits. + +Thus, in contrast with `BigDecimal`, `NUMERIC` can actually represent every possible value of `rust_decimal::Decimal`, +but not the other way around. This means that encoding should never fail, but decoding can. diff --git a/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal.rs b/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal.rs new file mode 100644 index 00000000..8321e828 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/rust_decimal.rs @@ -0,0 +1,493 @@ +use rust_decimal::Decimal; + +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::numeric::{PgNumeric, PgNumericSign}; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; + +use rust_decimal::MathematicalOps; + +impl Type for Decimal { + fn type_info() -> PgTypeInfo { + PgTypeInfo::NUMERIC + } +} + +impl PgHasArrayType for Decimal { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::NUMERIC_ARRAY + } +} + +impl TryFrom for Decimal { + type Error = BoxDynError; + + fn try_from(numeric: PgNumeric) -> Result { + Decimal::try_from(&numeric) + } +} + +impl TryFrom<&'_ PgNumeric> for Decimal { + type Error = BoxDynError; + + fn try_from(numeric: &'_ PgNumeric) -> Result { + let (digits, sign, mut weight, scale) = match *numeric { + PgNumeric::Number { + ref digits, + sign, + weight, + scale, + } => (digits, sign, weight, scale), + + PgNumeric::NotANumber => { + return Err("Decimal does not support NaN values".into()); + } + }; + + if digits.is_empty() { + // Postgres returns an empty digit array for 0 + return Ok(Decimal::ZERO); + } + + let scale = u32::try_from(scale) + .map_err(|_| format!("invalid scale value for Pg NUMERIC: {scale}"))?; + + let mut value = Decimal::ZERO; + + // Sum over `digits`, multiply each by its weight and add it to `value`. + for &digit in digits { + let mul = Decimal::from(10_000i16) + .checked_powi(weight as i64) + .ok_or("value not representable as rust_decimal::Decimal")?; + + let part = Decimal::from(digit) * mul; + + value = value + .checked_add(part) + .ok_or("value not representable as rust_decimal::Decimal")?; + + weight = weight.checked_sub(1).ok_or("weight underflowed")?; + } + + match sign { + PgNumericSign::Positive => value.set_sign_positive(true), + PgNumericSign::Negative => value.set_sign_negative(true), + } + + value.rescale(scale); + + Ok(value) + } +} + +impl From for PgNumeric { + fn from(value: Decimal) -> Self { + PgNumeric::from(&value) + } +} + +// This impl is effectively infallible because `NUMERIC` has a greater range than `Decimal`. +impl From<&'_ Decimal> for PgNumeric { + // Impl has been manually validated. + #[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)] + fn from(decimal: &Decimal) -> Self { + if Decimal::is_zero(decimal) { + return PgNumeric::ZERO; + } + + assert!( + (0u32..=28).contains(&decimal.scale()), + "decimal scale out of range {:?}", + decimal.unpack(), + ); + + // Cannot overflow: always in the range [0, 28] + let scale = decimal.scale() as u16; + + let mut mantissa = decimal.mantissa().unsigned_abs(); + + // If our scale is not a multiple of 4, we need to go to the next multiple. + let groups_diff = scale % 4; + if groups_diff > 0 { + let remainder = 4 - groups_diff as u32; + let power = 10u32.pow(remainder) as u128; + + // Impossible to overflow; 0 <= mantissa <= 2^96, + // and we're multiplying by at most 1,000 (giving us a result < 2^106) + mantissa *= power; + } + + // Array to store max mantissa of Decimal in Postgres decimal format. + let mut digits = Vec::with_capacity(8); + + // Convert to base-10000. + while mantissa != 0 { + // Cannot overflow or wrap because of the modulus + digits.push((mantissa % 10_000) as i16); + mantissa /= 10_000; + } + + // We started with the low digits first, but they should actually be at the end. + digits.reverse(); + + // Cannot overflow: strictly smaller than `scale`. + let digits_after_decimal = scale.div_ceil(4) as i16; + + // `mantissa` contains at most 29 decimal digits (log10(2^96)), + // split into at most 8 4-digit segments. + assert!( + digits.len() <= 8, + "digits.len() out of range: {}; unpacked: {:?}", + digits.len(), + decimal.unpack() + ); + + // Cannot overflow; at most 8 + let num_digits = digits.len() as i16; + + // Find how many 4-digit segments should go before the decimal point. + // `weight = 0` puts just `digit[0]` before the decimal point, and the rest after. + let weight = num_digits - digits_after_decimal - 1; + + // Remove non-significant zeroes. + while let Some(&0) = digits.last() { + digits.pop(); + } + + PgNumeric::Number { + sign: match decimal.is_sign_negative() { + false => PgNumericSign::Positive, + true => PgNumericSign::Negative, + }, + // Cannot overflow; between 0 and 28 + scale: scale as i16, + weight, + digits, + } + } +} + +impl Encode<'_, Postgres> for Decimal { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + PgNumeric::from(self).encode(buf)?; + + Ok(IsNull::No) + } +} + +#[doc=include_str!("rust_decimal-range.md")] +impl Decode<'_, Postgres> for Decimal { + fn decode(value: PgValueRef<'_>) -> Result { + match value.format() { + PgValueFormat::Binary => PgNumeric::decode(value.as_bytes()?)?.try_into(), + PgValueFormat::Text => Ok(value.as_str()?.parse::()?), + } + } +} + +#[cfg(test)] +mod tests { + use super::{Decimal, PgNumeric, PgNumericSign}; + use std::convert::TryFrom; + + #[test] + fn zero() { + let zero: Decimal = "0".parse().unwrap(); + + assert_eq!(PgNumeric::from(&zero), PgNumeric::ZERO,); + + assert_eq!(Decimal::try_from(&PgNumeric::ZERO).unwrap(), Decimal::ZERO); + } + + #[test] + fn one() { + let one: Decimal = "1".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![1] + } + ); + } + + #[test] + fn ten() { + let ten: Decimal = "10".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&ten).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![10] + } + ); + } + + #[test] + fn one_hundred() { + let one_hundred: Decimal = "100".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_hundred).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![100] + } + ); + } + + #[test] + fn ten_thousand() { + // Decimal doesn't normalize here + let ten_thousand: Decimal = "10000".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&ten_thousand).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1] + } + ); + } + + #[test] + fn two_digits() { + let two_digits: Decimal = "12345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&two_digits).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1, 2345] + } + ); + } + + #[test] + fn one_tenth() { + let one_tenth: Decimal = "0.1".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&one_tenth).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 1, + weight: -1, + digits: vec![1000] + } + ); + } + + #[test] + fn decimal_1() { + let decimal: Decimal = "1.2345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 4, + weight: 0, + digits: vec![1, 2345] + } + ); + } + + #[test] + fn decimal_2() { + let decimal: Decimal = "0.12345".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: -1, + digits: vec![1234, 5000] + } + ); + } + + #[test] + fn decimal_3() { + let decimal: Decimal = "0.01234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&decimal).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: -1, + digits: vec![0123, 4000] + } + ); + } + + #[test] + fn decimal_4() { + let decimal: Decimal = "12345.67890".parse().unwrap(); + let expected_numeric = PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 5, + weight: 1, + digits: vec![1, 2345, 6789], + }; + assert_eq!(PgNumeric::try_from(&decimal).unwrap(), expected_numeric); + + let actual_decimal = Decimal::try_from(expected_numeric).unwrap(); + assert_eq!(actual_decimal, decimal); + assert_eq!(actual_decimal.mantissa(), 1234567890); + assert_eq!(actual_decimal.scale(), 5); + } + + #[test] + fn one_digit_decimal() { + let one_digit_decimal: Decimal = "0.00001234".parse().unwrap(); + let expected_numeric = PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 8, + weight: -2, + digits: vec![1234], + }; + assert_eq!( + PgNumeric::try_from(&one_digit_decimal).unwrap(), + expected_numeric + ); + + let actual_decimal = Decimal::try_from(expected_numeric).unwrap(); + assert_eq!(actual_decimal, one_digit_decimal); + assert_eq!(actual_decimal.mantissa(), 1234); + assert_eq!(actual_decimal.scale(), 8); + } + + #[test] + fn max_value() { + let expected_numeric = PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 7, + digits: vec![7, 9228, 1625, 1426, 4337, 5935, 4395, 0335], + }; + assert_eq!( + PgNumeric::try_from(&Decimal::MAX).unwrap(), + expected_numeric + ); + + let actual_decimal = Decimal::try_from(expected_numeric).unwrap(); + assert_eq!(actual_decimal, Decimal::MAX); + // Value split by 10,000's to match the expected digits[] + assert_eq!( + actual_decimal.mantissa(), + 7_9228_1625_1426_4337_5935_4395_0335 + ); + assert_eq!(actual_decimal.scale(), 0); + } + + #[test] + fn max_value_max_scale() { + let mut max_value_max_scale = Decimal::MAX; + max_value_max_scale.set_scale(28).unwrap(); + + let expected_numeric = PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 28, + weight: 0, + digits: vec![7, 9228, 1625, 1426, 4337, 5935, 4395, 0335], + }; + assert_eq!( + PgNumeric::try_from(&max_value_max_scale).unwrap(), + expected_numeric + ); + + let actual_decimal = Decimal::try_from(expected_numeric).unwrap(); + assert_eq!(actual_decimal, max_value_max_scale); + assert_eq!( + actual_decimal.mantissa(), + 79_228_162_514_264_337_593_543_950_335 + ); + assert_eq!(actual_decimal.scale(), 28); + } + + #[test] + fn issue_423_four_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let four_digit: Decimal = "1234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&four_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 0, + digits: vec![1234] + } + ); + } + + #[test] + fn issue_423_negative_four_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let negative_four_digit: Decimal = "-1234".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&negative_four_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Negative, + scale: 0, + weight: 0, + digits: vec![1234] + } + ); + } + + #[test] + fn issue_423_eight_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let eight_digit: Decimal = "12345678".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&eight_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 0, + weight: 1, + digits: vec![1234, 5678] + } + ); + } + + #[test] + fn issue_423_negative_eight_digit() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/423 + let negative_eight_digit: Decimal = "-12345678".parse().unwrap(); + assert_eq!( + PgNumeric::try_from(&negative_eight_digit).unwrap(), + PgNumeric::Number { + sign: PgNumericSign::Negative, + scale: 0, + weight: 1, + digits: vec![1234, 5678] + } + ); + } + + #[test] + fn issue_2247_trailing_zeros() { + // This is a regression test for https://github.com/launchbadge/sqlx/issues/2247 + let one_hundred: Decimal = "100.00".parse().unwrap(); + let expected_numeric = PgNumeric::Number { + sign: PgNumericSign::Positive, + scale: 2, + weight: 0, + digits: vec![100], + }; + assert_eq!(PgNumeric::try_from(&one_hundred).unwrap(), expected_numeric); + + let actual_decimal = Decimal::try_from(expected_numeric).unwrap(); + assert_eq!(actual_decimal, one_hundred); + assert_eq!(actual_decimal.mantissa(), 10000); + assert_eq!(actual_decimal.scale(), 2); + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/str.rs b/src-tauri/vendor/sqlx-postgres/src/types/str.rs new file mode 100644 index 00000000..ca7e20a5 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/str.rs @@ -0,0 +1,148 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::array_compatible; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueRef, Postgres}; +use std::borrow::Cow; + +impl Type for str { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TEXT + } + + fn compatible(ty: &PgTypeInfo) -> bool { + [ + PgTypeInfo::TEXT, + PgTypeInfo::NAME, + PgTypeInfo::BPCHAR, + PgTypeInfo::VARCHAR, + PgTypeInfo::UNKNOWN, + PgTypeInfo::with_name("citext"), + ] + .contains(ty) + } +} + +impl Type for Cow<'_, str> { + fn type_info() -> PgTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Type for Box { + fn type_info() -> PgTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl Type for String { + fn type_info() -> PgTypeInfo { + <&str as Type>::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + <&str as Type>::compatible(ty) + } +} + +impl PgHasArrayType for &'_ str { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TEXT_ARRAY + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + array_compatible::<&str>(ty) + } +} + +impl PgHasArrayType for Cow<'_, str> { + fn array_type_info() -> PgTypeInfo { + <&str as PgHasArrayType>::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + <&str as PgHasArrayType>::array_compatible(ty) + } +} + +impl PgHasArrayType for Box { + fn array_type_info() -> PgTypeInfo { + <&str as PgHasArrayType>::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + <&str as PgHasArrayType>::array_compatible(ty) + } +} + +impl PgHasArrayType for String { + fn array_type_info() -> PgTypeInfo { + <&str as PgHasArrayType>::array_type_info() + } + + fn array_compatible(ty: &PgTypeInfo) -> bool { + <&str as PgHasArrayType>::array_compatible(ty) + } +} + +impl Encode<'_, Postgres> for &'_ str { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + buf.extend(self.as_bytes()); + + Ok(IsNull::No) + } +} + +impl Encode<'_, Postgres> for Cow<'_, str> { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + match self { + Cow::Borrowed(str) => <&str as Encode>::encode(*str, buf), + Cow::Owned(str) => <&str as Encode>::encode(&**str, buf), + } + } +} + +impl Encode<'_, Postgres> for Box { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&str as Encode>::encode(&**self, buf) + } +} + +impl Encode<'_, Postgres> for String { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + <&str as Encode>::encode(&**self, buf) + } +} + +impl<'r> Decode<'r, Postgres> for &'r str { + fn decode(value: PgValueRef<'r>) -> Result { + value.as_str() + } +} + +impl<'r> Decode<'r, Postgres> for Cow<'r, str> { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(Cow::Borrowed(value.as_str()?)) + } +} + +impl<'r> Decode<'r, Postgres> for Box { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(Box::from(value.as_str()?)) + } +} + +impl Decode<'_, Postgres> for String { + fn decode(value: PgValueRef<'_>) -> Result { + Ok(value.as_str()?.to_owned()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/text.rs b/src-tauri/vendor/sqlx-postgres/src/types/text.rs new file mode 100644 index 00000000..b5b0a5ed --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/text.rs @@ -0,0 +1,40 @@ +use crate::{PgArgumentBuffer, PgTypeInfo, PgValueRef, Postgres}; +use sqlx_core::decode::Decode; +use sqlx_core::encode::{Encode, IsNull}; +use sqlx_core::error::BoxDynError; +use sqlx_core::types::{Text, Type}; +use std::fmt::Display; +use std::str::FromStr; + +use std::io::Write; + +impl Type for Text { + fn type_info() -> PgTypeInfo { + >::type_info() + } + + fn compatible(ty: &PgTypeInfo) -> bool { + >::compatible(ty) + } +} + +impl<'q, T> Encode<'q, Postgres> for Text +where + T: Display, +{ + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + write!(**buf, "{}", self.0)?; + Ok(IsNull::No) + } +} + +impl<'r, T> Decode<'r, Postgres> for Text +where + T: FromStr, + BoxDynError: From<::Err>, +{ + fn decode(value: PgValueRef<'r>) -> Result { + let s: &str = Decode::::decode(value)?; + Ok(Self(s.parse()?)) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/time/date.rs b/src-tauri/vendor/sqlx-postgres/src/types/time/date.rs new file mode 100644 index 00000000..2afa57ee --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/time/date.rs @@ -0,0 +1,52 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::time::PG_EPOCH; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use std::mem; +use time::macros::format_description; +use time::{Date, Duration}; + +impl Type for Date { + fn type_info() -> PgTypeInfo { + PgTypeInfo::DATE + } +} + +impl PgHasArrayType for Date { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::DATE_ARRAY + } +} + +impl Encode<'_, Postgres> for Date { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // DATE is encoded as number of days since epoch (2000-01-01) + let days: i32 = (*self - PG_EPOCH).whole_days().try_into().map_err(|_| { + format!("value {self:?} would overflow binary encoding for Postgres DATE") + })?; + Encode::::encode(days, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for Date { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // DATE is encoded as the days since epoch + let days: i32 = Decode::::decode(value)?; + PG_EPOCH + Duration::days(days.into()) + } + + PgValueFormat::Text => Date::parse( + value.as_str()?, + &format_description!("[year]-[month]-[day]"), + )?, + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/time/datetime.rs b/src-tauri/vendor/sqlx-postgres/src/types/time/datetime.rs new file mode 100644 index 00000000..3484116b --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/time/datetime.rs @@ -0,0 +1,108 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::time::PG_EPOCH; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use std::borrow::Cow; +use std::mem; +use time::macros::format_description; +use time::macros::offset; +use time::{Duration, OffsetDateTime, PrimitiveDateTime}; + +impl Type for PrimitiveDateTime { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMP + } +} + +impl Type for OffsetDateTime { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMPTZ + } +} + +impl PgHasArrayType for PrimitiveDateTime { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMP_ARRAY + } +} + +impl PgHasArrayType for OffsetDateTime { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIMESTAMPTZ_ARRAY + } +} + +impl Encode<'_, Postgres> for PrimitiveDateTime { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // TIMESTAMP is encoded as the microseconds since the epoch + let micros: i64 = (*self - PG_EPOCH.midnight()) + .whole_microseconds() + .try_into() + .map_err(|_| { + format!("value {self:?} would overflow binary encoding for Postgres TIME") + })?; + Encode::::encode(micros, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for PrimitiveDateTime { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // TIMESTAMP is encoded as the microseconds since the epoch + let us = Decode::::decode(value)?; + PG_EPOCH.midnight() + Duration::microseconds(us) + } + + PgValueFormat::Text => { + let s = value.as_str()?; + + // If there is no decimal point we need to add one. + let s = if s.contains('.') { + Cow::Borrowed(s) + } else { + Cow::Owned(format!("{s}.0")) + }; + + // Contains a time-zone specifier + // This is given for timestamptz for some reason + // Postgres already guarantees this to always be UTC + if s.contains('+') { + PrimitiveDateTime::parse(&s, &format_description!("[year]-[month]-[day] [hour]:[minute]:[second].[subsecond][offset_hour]"))? + } else { + PrimitiveDateTime::parse( + &s, + &format_description!( + "[year]-[month]-[day] [hour]:[minute]:[second].[subsecond]" + ), + )? + } + } + }) + } +} + +impl Encode<'_, Postgres> for OffsetDateTime { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + let utc = self.to_offset(offset!(UTC)); + let primitive = PrimitiveDateTime::new(utc.date(), utc.time()); + + Encode::::encode(primitive, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for OffsetDateTime { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(>::decode(value)?.assume_utc()) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/time/mod.rs b/src-tauri/vendor/sqlx-postgres/src/types/time/mod.rs new file mode 100644 index 00000000..9a45ba83 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/time/mod.rs @@ -0,0 +1,9 @@ +mod date; +mod datetime; + +// Parent module is named after the `time` crate, this module is named after the `TIME` SQL type. +#[allow(clippy::module_inception)] +mod time; + +#[rustfmt::skip] +const PG_EPOCH: ::time::Date = ::time::macros::date!(2000-1-1); diff --git a/src-tauri/vendor/sqlx-postgres/src/types/time/time.rs b/src-tauri/vendor/sqlx-postgres/src/types/time/time.rs new file mode 100644 index 00000000..635170d1 --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/time/time.rs @@ -0,0 +1,53 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use std::mem; +use time::macros::format_description; +use time::{Duration, Time}; + +impl Type for Time { + fn type_info() -> PgTypeInfo { + PgTypeInfo::TIME + } +} + +impl PgHasArrayType for Time { + fn array_type_info() -> PgTypeInfo { + PgTypeInfo::TIME_ARRAY + } +} + +impl Encode<'_, Postgres> for Time { + fn encode_by_ref(&self, buf: &mut PgArgumentBuffer) -> Result { + // TIME is encoded as the microseconds since midnight. + // + // A truncating cast is fine because `self - Time::MIDNIGHT` cannot exceed a span of 24 hours. + #[allow(clippy::cast_possible_truncation)] + let micros: i64 = (*self - Time::MIDNIGHT).whole_microseconds() as i64; + Encode::::encode(micros, buf) + } + + fn size_hint(&self) -> usize { + mem::size_of::() + } +} + +impl<'r> Decode<'r, Postgres> for Time { + fn decode(value: PgValueRef<'r>) -> Result { + Ok(match value.format() { + PgValueFormat::Binary => { + // TIME is encoded as the microseconds since midnight + let us = Decode::::decode(value)?; + Time::MIDNIGHT + Duration::microseconds(us) + } + + PgValueFormat::Text => Time::parse( + value.as_str()?, + // Postgres will not include the subsecond part if it's zero. + &format_description!("[hour]:[minute]:[second][optional [.[subsecond]]]"), + )?, + }) + } +} diff --git a/src-tauri/vendor/sqlx-postgres/src/types/time_tz.rs b/src-tauri/vendor/sqlx-postgres/src/types/time_tz.rs new file mode 100644 index 00000000..e3de79ea --- /dev/null +++ b/src-tauri/vendor/sqlx-postgres/src/types/time_tz.rs @@ -0,0 +1,176 @@ +use crate::decode::Decode; +use crate::encode::{Encode, IsNull}; +use crate::error::BoxDynError; +use crate::types::Type; +use crate::{PgArgumentBuffer, PgHasArrayType, PgTypeInfo, PgValueFormat, PgValueRef, Postgres}; +use byteorder::{BigEndian, ReadBytesExt}; +use std::io::Cursor; +use std::mem; + +#[cfg(feature = "time")] +type DefaultTime = ::time::Time; + +#[cfg(all(not(feature = "time"), feature = "chrono"))] +type DefaultTime = ::chrono::NaiveTime; + +#[cfg(feature = "time")] +type DefaultOffset = ::time::UtcOffset; + +#[cfg(all(not(feature = "time"), feature = "chrono"))] +type DefaultOffset = ::chrono::FixedOffset; + +/// Represents a moment of time, in a specified timezone. +/// +/// # Warning +/// +/// `PgTimeTz` provides `TIMETZ` and is supported only for reading from legacy databases. +/// [PostgreSQL recommends] to use `TIMESTAMPTZ` instead. +/// +/// [PostgreSQL recommends]: https://wiki.postgresql.org/wiki/Don't_Do_This#Don.27t_use_timetz +#[derive(Debug, PartialEq, Clone, Copy)] +pub struct PgTimeTz