diff --git a/crates/dbx-core/src/db/postgres.rs b/crates/dbx-core/src/db/postgres.rs index 1e1b189f3..9fdcc5c12 100644 --- a/crates/dbx-core/src/db/postgres.rs +++ b/crates/dbx-core/src/db/postgres.rs @@ -16,7 +16,7 @@ use std::sync::atomic::AtomicBool; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio_postgres::config::SslMode; -use tokio_postgres::types::{FromSql, Type}; +use tokio_postgres::types::{FromSql, Kind, Type}; use tokio_postgres::{NoTls, Row, SimpleQueryMessage}; use tokio_util::sync::CancellationToken; @@ -338,6 +338,20 @@ pub(crate) enum PgColType { Other, } +const POSTGRES_FIRST_NORMAL_OBJECT_ID: u32 = 16_384; + +fn pg_scalar_type_requires_text_protocol(oid: u32, col_type: PgColType) -> bool { + oid >= POSTGRES_FIRST_NORMAL_OBJECT_ID && !matches!(col_type, PgColType::Vector | PgColType::Geometry) +} + +fn pg_type_requires_text_protocol(pg_type: &Type, col_type: PgColType) -> bool { + match pg_type.kind() { + Kind::Array(element_type) => element_type.oid() >= POSTGRES_FIRST_NORMAL_OBJECT_ID, + Kind::Simple => pg_scalar_type_requires_text_protocol(pg_type.oid(), col_type), + _ => pg_type.oid() >= POSTGRES_FIRST_NORMAL_OBJECT_ID, + } +} + pub(crate) fn classify_pg_type(type_name: &str) -> PgColType { let upper = type_name.to_uppercase(); @@ -687,22 +701,59 @@ async fn postgres_query_one_cached( } } +enum PreparedSelectOutcome { + Complete(QueryResult), + TextFallback { column_types: Vec, unsupported_type: String }, +} + +struct PreparedSelectMetadata { + columns: Vec, + column_types: Vec, + column_classes: Vec, + unsupported_type: Option, +} + +fn prepared_select_metadata(stmt: &tokio_postgres::Statement) -> PreparedSelectMetadata { + let columns: Vec = stmt.columns().iter().map(|c| c.name().to_string()).collect(); + let column_types: Vec = stmt.columns().iter().map(|c| c.type_().name().to_string()).collect(); + let column_classes = classify_pg_column_types(&column_types); + let unsupported_type = stmt.columns().iter().zip(&column_classes).find_map(|(column, col_type)| { + let pg_type = column.type_(); + pg_type_requires_text_protocol(pg_type, *col_type).then(|| pg_type.name().to_string()) + }); + PreparedSelectMetadata { columns, column_types, column_classes, unsupported_type } +} + +async fn prepare_select_with_metadata( + client: &deadpool_postgres::Client, + sql: &str, +) -> Result<(tokio_postgres::Statement, PreparedSelectMetadata), tokio_postgres::Error> { + let mut stmt = client.prepare_cached(sql).await?; + let mut metadata = prepared_select_metadata(&stmt); + if metadata.unsupported_type.is_some() { + stmt = client.prepare(sql).await?; + metadata = prepared_select_metadata(&stmt); + } + Ok((stmt, metadata)) +} + async fn execute_select_prepared( client: &deadpool_postgres::Client, sql: &str, start: Instant, row_limit: usize, -) -> Result { +) -> Result { let prepared_start = Instant::now(); - let stmt = client.prepare_cached(sql).await?; + let (stmt, metadata) = prepare_select_with_metadata(client, sql).await?; log::info!( "[postgres][select:prepare_cached:done] elapsed_ms={} total_ms={}", prepared_start.elapsed().as_millis(), start.elapsed().as_millis() ); - let columns: Vec = stmt.columns().iter().map(|c| c.name().to_string()).collect(); - let column_types: Vec = stmt.columns().iter().map(|c| c.type_().name().to_string()).collect(); - let column_classes = classify_pg_column_types(&column_types); + let PreparedSelectMetadata { columns, column_types, column_classes, unsupported_type } = metadata; + if let Some(unsupported_type) = unsupported_type { + return Ok(PreparedSelectOutcome::TextFallback { column_types, unsupported_type }); + } let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = Vec::new(); let query_start = Instant::now(); @@ -740,7 +791,7 @@ async fn execute_select_prepared( truncated ); - Ok(QueryResult { + Ok(PreparedSelectOutcome::Complete(QueryResult { columns, column_types, column_sortables: Vec::new(), @@ -750,7 +801,11 @@ async fn execute_select_prepared( truncated, session_id: None, has_more: false, - }) + })) +} + +fn matching_pg_text_column_types(columns: &[String], prepared: Option>) -> Vec { + prepared.filter(|types| types.len() == columns.len()).unwrap_or_default() } async fn execute_select_text( @@ -758,24 +813,26 @@ async fn execute_select_text( sql: &str, start: Instant, row_limit: usize, + prepared_column_types: Option>, ) -> Result { - let messages = client.simple_query(sql).await.map_err(pg_error_to_string)?; + let stream = client.simple_query_raw(sql).await.map_err(pg_error_to_string)?; + tokio::pin!(stream); let mut columns: Vec = Vec::new(); let mut result_rows: Vec> = Vec::new(); let mut truncated = false; - for message in messages { + while let Some(message) = stream.next().await { match message { - SimpleQueryMessage::RowDescription(cols) => { + Ok(SimpleQueryMessage::RowDescription(cols)) => { columns = cols.iter().map(|c| c.name().to_string()).collect(); } - SimpleQueryMessage::Row(row) => { + Ok(SimpleQueryMessage::Row(row)) => { if columns.is_empty() { columns = row.columns().iter().map(|c| c.name().to_string()).collect(); } if result_rows.len() >= row_limit { truncated = true; - continue; + break; } let mut values = Vec::with_capacity(row.len()); for i in 0..row.len() { @@ -786,14 +843,19 @@ async fn execute_select_text( } result_rows.push(values); } - SimpleQueryMessage::CommandComplete(_) => {} - _ => {} + Err(_) if result_rows.len() >= row_limit => { + truncated = true; + break; + } + Err(err) => return Err(pg_error_to_string(err)), + Ok(SimpleQueryMessage::CommandComplete(_)) => {} + Ok(_) => {} } } Ok(QueryResult { + column_types: matching_pg_text_column_types(&columns, prepared_column_types), columns, - column_types: Vec::new(), column_sortables: Vec::new(), rows: result_rows, affected_rows: 0, @@ -804,14 +866,33 @@ async fn execute_select_text( }) } -async fn execute_select_query( +async fn finish_prepared_select( + client: &deadpool_postgres::Client, + sql: &str, + start: Instant, + row_limit: usize, + outcome: PreparedSelectOutcome, +) -> Result { + match outcome { + PreparedSelectOutcome::Complete(result) => Ok(result), + PreparedSelectOutcome::TextFallback { column_types, unsupported_type } => { + log::info!( + "[postgres][select:text_fallback] unsupported_type={} switching_to=simple_query", + unsupported_type + ); + execute_select_text(client, sql, start, row_limit, Some(column_types)).await + } + } +} + +pub(crate) async fn execute_select_query( client: &deadpool_postgres::Client, sql: &str, start: Instant, row_limit: usize, ) -> Result { match execute_select_prepared(client, sql, start, row_limit).await { - Ok(result) => Ok(result), + Ok(outcome) => finish_prepared_select(client, sql, start, row_limit, outcome).await, Err(err) if should_retry_postgres_stale_cache(&err) => { // The cached prepared statement is stale (e.g. the view or table // schema changed since the statement was prepared). Evict the @@ -819,14 +900,16 @@ async fn execute_select_query( log::warn!("[postgres][select:stale_cache] evicting cached statement: {}", pg_error_to_string(err)); client.statement_cache.remove(sql, &[]); match execute_select_prepared(client, sql, start, row_limit).await { - Ok(result) => Ok(result), + Ok(outcome) => finish_prepared_select(client, sql, start, row_limit, outcome).await, Err(err) if should_retry_postgres_text_query(&err) => { - execute_select_text(client, sql, start, row_limit).await + execute_select_text(client, sql, start, row_limit, None).await } Err(err) => Err(pg_error_to_string(err)), } } - Err(err) if should_retry_postgres_text_query(&err) => execute_select_text(client, sql, start, row_limit).await, + Err(err) if should_retry_postgres_text_query(&err) => { + execute_select_text(client, sql, start, row_limit, None).await + } Err(err) => Err(pg_error_to_string(err)), } } @@ -838,6 +921,7 @@ pub enum PostgresQueryStreamItem { enum PostgresQueryStreamError { Postgres { err: tokio_postgres::Error, emitted: bool }, + TextFallback { column_types: Vec, unsupported_type: String }, Export(String), } @@ -845,6 +929,9 @@ impl PostgresQueryStreamError { fn into_string(self) -> String { match self { Self::Postgres { err, .. } => pg_error_to_string(err), + Self::TextFallback { unsupported_type, .. } => { + format!("PostgreSQL type {unsupported_type} requires text protocol") + } Self::Export(err) => err, } } @@ -856,11 +943,13 @@ async fn stream_select_query_prepared( row_limit: Option, on_item: &mut impl FnMut(PostgresQueryStreamItem) -> Result<(), String>, ) -> Result { - let stmt = - client.prepare_cached(sql).await.map_err(|err| PostgresQueryStreamError::Postgres { err, emitted: false })?; - let columns: Vec = stmt.columns().iter().map(|c| c.name().to_string()).collect(); - let column_types: Vec = stmt.columns().iter().map(|c| c.type_().name().to_string()).collect(); - let column_classes = classify_pg_column_types(&column_types); + let (stmt, metadata) = prepare_select_with_metadata(client, sql) + .await + .map_err(|err| PostgresQueryStreamError::Postgres { err, emitted: false })?; + let PreparedSelectMetadata { columns, column_types, column_classes, unsupported_type } = metadata; + if let Some(unsupported_type) = unsupported_type { + return Err(PostgresQueryStreamError::TextFallback { column_types, unsupported_type }); + } let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = Vec::new(); let stream = client @@ -898,6 +987,7 @@ async fn stream_select_query_text( client: &deadpool_postgres::Client, sql: &str, row_limit: Option, + prepared_column_types: Option>, on_item: &mut impl FnMut(PostgresQueryStreamItem) -> Result<(), String>, ) -> Result { let stream = client.simple_query_raw(sql).await.map_err(pg_error_to_string)?; @@ -908,7 +998,8 @@ async fn stream_select_query_text( match message.map_err(pg_error_to_string)? { SimpleQueryMessage::RowDescription(cols) => { columns = cols.iter().map(|c| c.name().to_string()).collect(); - on_item(PostgresQueryStreamItem::Columns { columns: columns.clone(), column_types: Vec::new() })?; + let column_types = matching_pg_text_column_types(&columns, prepared_column_types.clone()); + on_item(PostgresQueryStreamItem::Columns { columns: columns.clone(), column_types })?; } SimpleQueryMessage::Row(row) => { if row_limit.is_some_and(|limit| rows_streamed as usize >= limit) { @@ -916,7 +1007,8 @@ async fn stream_select_query_text( } if columns.is_empty() { columns = row.columns().iter().map(|c| c.name().to_string()).collect(); - on_item(PostgresQueryStreamItem::Columns { columns: columns.clone(), column_types: Vec::new() })?; + let column_types = matching_pg_text_column_types(&columns, prepared_column_types.clone()); + on_item(PostgresQueryStreamItem::Columns { columns: columns.clone(), column_types })?; } let mut values = Vec::with_capacity(row.len()); for i in 0..row.len() { @@ -935,7 +1027,7 @@ async fn stream_select_query_text( Ok(rows_streamed) } -async fn stream_select_query_inner( +pub(crate) async fn stream_select_query_inner( client: &deadpool_postgres::Client, sql: &str, row_limit: Option, @@ -943,6 +1035,13 @@ async fn stream_select_query_inner( ) -> Result { match stream_select_query_prepared(client, sql, row_limit, on_item).await { Ok(rows) => Ok(rows), + Err(PostgresQueryStreamError::TextFallback { column_types, unsupported_type }) => { + log::info!( + "[postgres][stream:text_fallback] unsupported_type={} switching_to=simple_query", + unsupported_type + ); + stream_select_query_text(client, sql, row_limit, Some(column_types), on_item).await + } Err(PostgresQueryStreamError::Postgres { err, emitted: false }) if should_retry_postgres_stale_cache(&err) => { // The cached prepared statement can become stale after schema changes. // Evict and retry once, matching the normal query execution path. @@ -953,13 +1052,20 @@ async fn stream_select_query_inner( Err(PostgresQueryStreamError::Postgres { err, emitted: false }) if should_retry_postgres_text_query(&err) => { - stream_select_query_text(client, sql, row_limit, on_item).await + stream_select_query_text(client, sql, row_limit, None, on_item).await + } + Err(PostgresQueryStreamError::TextFallback { column_types, unsupported_type }) => { + log::info!( + "[postgres][stream:text_fallback] unsupported_type={} switching_to=simple_query", + unsupported_type + ); + stream_select_query_text(client, sql, row_limit, Some(column_types), on_item).await } Err(err) => Err(err.into_string()), } } Err(PostgresQueryStreamError::Postgres { err, emitted: false }) if should_retry_postgres_text_query(&err) => { - stream_select_query_text(client, sql, row_limit, on_item).await + stream_select_query_text(client, sql, row_limit, None, on_item).await } Err(err) => Err(err.into_string()), } @@ -989,9 +1095,15 @@ async fn stream_query_rows_on_client( cancelled: &AtomicBool, on_row: &mut impl FnMut(&[serde_json::Value]) -> Result<(), String>, ) -> Result { - let stmt = client.prepare_cached(sql).await.map_err(pg_error_to_string)?; - let column_types: Vec = stmt.columns().iter().map(|c| c.type_().name().to_string()).collect(); - let column_classes = classify_pg_column_types(&column_types); + let (stmt, metadata) = prepare_select_with_metadata(client, sql).await.map_err(pg_error_to_string)?; + let PreparedSelectMetadata { column_classes, unsupported_type, .. } = metadata; + if let Some(unsupported_type) = unsupported_type { + log::info!( + "[postgres][row_stream:text_fallback] unsupported_type={} switching_to=simple_query", + unsupported_type + ); + return stream_query_rows_text_on_client(client, sql, max_rows, cancelled, on_row).await; + } let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = Vec::new(); let stream = client.query_raw(&stmt, params).await.map_err(pg_error_to_string)?; tokio::pin!(stream); @@ -1023,17 +1135,19 @@ async fn stream_query_rows_text_on_client( cancelled: &AtomicBool, on_row: &mut impl FnMut(&[serde_json::Value]) -> Result<(), String>, ) -> Result { - let messages = client.simple_query(sql).await.map_err(pg_error_to_string)?; + let stream = client.simple_query_raw(sql).await.map_err(pg_error_to_string)?; + tokio::pin!(stream); let row_limit = max_rows.unwrap_or(usize::MAX); let mut rows_exported = 0_u64; - for message in messages { + while let Some(message) = stream.next().await { if cancelled.load(std::sync::atomic::Ordering::SeqCst) { return Err(crate::query::canceled_error()); } if rows_exported as usize >= row_limit { break; } + let message = message.map_err(pg_error_to_string)?; if let SimpleQueryMessage::Row(row) = message { let mut values = Vec::with_capacity(row.len()); for i in 0..row.len() { @@ -3630,6 +3744,45 @@ mod tests { use std::time::Instant; use tokio_postgres::types::FromSql; + #[test] + fn postgres_custom_other_type_requires_text_protocol() { + assert!(pg_scalar_type_requires_text_protocol(POSTGRES_FIRST_NORMAL_OBJECT_ID, PgColType::Other)); + assert!(pg_scalar_type_requires_text_protocol(98_765, PgColType::Other)); + assert!(pg_scalar_type_requires_text_protocol(98_765, PgColType::GenericArray)); + } + + #[test] + fn postgres_builtin_or_supported_type_keeps_binary_protocol() { + assert!(!pg_scalar_type_requires_text_protocol(POSTGRES_FIRST_NORMAL_OBJECT_ID - 1, PgColType::Other)); + assert!(!pg_type_requires_text_protocol(&Type::INT4, PgColType::Other)); + assert!(!pg_type_requires_text_protocol(&Type::VARCHAR, PgColType::Other)); + assert!(!pg_type_requires_text_protocol(&Type::INT4_ARRAY, PgColType::GenericArray)); + assert!(!pg_scalar_type_requires_text_protocol(98_765, PgColType::Vector)); + assert!(!pg_scalar_type_requires_text_protocol(98_765, PgColType::Geometry)); + } + + #[test] + fn postgres_query_uses_text_when_any_output_type_is_unsupported() { + let columns = + [(Type::INT4.oid(), PgColType::Other), (98_765, PgColType::Other), (Type::TEXT.oid(), PgColType::Other)]; + assert!(columns.into_iter().any(|(oid, col_type)| pg_scalar_type_requires_text_protocol(oid, col_type))); + } + + #[test] + fn postgres_text_fallback_keeps_matching_prepared_column_types() { + let columns = vec!["payload".to_string(), "id".to_string()]; + let types = vec!["payload_type".to_string(), "int4".to_string()]; + assert_eq!(matching_pg_text_column_types(&columns, Some(types.clone())), types); + } + + #[test] + fn postgres_text_fallback_discards_misaligned_column_types() { + let columns = vec!["payload".to_string(), "id".to_string()]; + let types = vec!["payload_type".to_string()]; + assert!(matching_pg_text_column_types(&columns, Some(types)).is_empty()); + assert!(matching_pg_text_column_types(&columns, None).is_empty()); + } + #[test] fn postgres_query_search_path_preserves_public_after_catalog() { assert_eq!( @@ -3787,6 +3940,226 @@ mod tests { } } + async fn assert_postgres_18(pool: &Pool) { + let version = execute_query(pool, "SHOW server_version_num").await.expect("query PostgreSQL version"); + let version_num = version.rows[0][0] + .as_str() + .expect("server_version_num should be text") + .parse::() + .expect("server_version_num should be numeric"); + assert!((180_000..190_000).contains(&version_num), "expected PostgreSQL 18, got {version_num}"); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a writable PostgreSQL 18 database"] + async fn postgres_custom_composite_result_uses_server_text_output() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool = connect(&url, Duration::from_secs(5)).await.expect("connect postgres"); + assert_postgres_18(&pool).await; + let schema = format!("dbx_custom_text_{}", uuid::Uuid::new_v4().simple()); + let schema_ident = pg_quote_ident(&schema); + let payload_type = format!("{schema_ident}.payload"); + execute_query(&pool, &format!("CREATE SCHEMA {schema_ident}")).await.expect("create schema"); + let exercise = async { + execute_query(&pool, &format!("CREATE TYPE {payload_type} AS (id integer, label text)")).await?; + let custom = + execute_query(&pool, &format!("SELECT ROW(7, 'alpha')::{payload_type} AS payload, 42::int4 AS id")) + .await?; + let builtin = execute_query(&pool, "SELECT 42::int4 AS id").await?; + Ok::<_, String>((custom, builtin)) + } + .await; + + let cleanup = execute_query(&pool, &format!("DROP SCHEMA {schema_ident} CASCADE")).await; + cleanup.expect("drop schema"); + let (custom, builtin) = exercise.expect("exercise custom composite fallback"); + + assert_eq!(custom.columns, vec!["payload", "id"]); + assert_eq!(custom.column_types, vec!["payload", "int4"]); + assert_eq!(custom.rows[0][0], serde_json::Value::String("(7,alpha)".to_string())); + assert_eq!(custom.rows[0][1], serde_json::Value::String("42".to_string())); + assert!(!custom.rows[0][0].as_str().unwrap().chars().any(char::is_control)); + assert_eq!(builtin.column_types, vec!["int4"]); + assert_eq!(builtin.rows[0][0], serde_json::Value::Number(42.into())); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a writable PostgreSQL 18 database"] + async fn postgres_custom_type_arrays_and_exports_use_server_text_output() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool = connect(&url, Duration::from_secs(5)).await.expect("connect postgres"); + assert_postgres_18(&pool).await; + let schema = format!("dbx_custom_array_{}", uuid::Uuid::new_v4().simple()); + let schema_ident = pg_quote_ident(&schema); + let payload_type = format!("{schema_ident}.payload"); + let mood_type = format!("{schema_ident}.mood"); + let score_type = format!("{schema_ident}.positive_int"); + let underscore_scalar_type = format!("{schema_ident}._hidden"); + let vector_named_enum_type = format!("{schema_ident}.vector"); + let table = format!("{schema_ident}.custom_arrays"); + let select_sql = format!("SELECT payloads, moods, scores FROM {table}"); + execute_query(&pool, &format!("CREATE SCHEMA {schema_ident}")).await.expect("create schema"); + + let exercise = async { + execute_query(&pool, &format!("CREATE TYPE {payload_type} AS (id integer, label text)")).await?; + execute_query(&pool, &format!("CREATE TYPE {mood_type} AS ENUM ('ready', 'done')")).await?; + execute_query(&pool, &format!("CREATE DOMAIN {score_type} AS integer CHECK (VALUE > 0)")).await?; + execute_query(&pool, &format!("CREATE TYPE {underscore_scalar_type} AS ENUM ('secret')")).await?; + execute_query(&pool, &format!("CREATE TYPE {vector_named_enum_type} AS ENUM ('label')")).await?; + execute_query( + &pool, + &format!( + "CREATE TABLE {table} (payloads {payload_type}[], moods {mood_type}[], scores {score_type}[])" + ), + ) + .await?; + execute_query( + &pool, + &format!( + "INSERT INTO {table} VALUES \ + (ARRAY[ROW(7, 'alpha')::{payload_type}], ARRAY['ready'::{mood_type}], ARRAY[7::{score_type}])" + ), + ) + .await?; + + let query = execute_query(&pool, &select_sql).await?; + let underscore_scalar = + execute_query(&pool, &format!("SELECT 'secret'::{underscore_scalar_type} AS hidden")).await?; + let vector_named_enum = + execute_query(&pool, &format!("SELECT 'label'::{vector_named_enum_type} AS label")).await?; + let client = checkout_postgres_client(&pool, None, Duration::from_secs(5)).await?; + + let mut query_export_rows = Vec::new(); + stream_select_query_inner(&client, &select_sql, None, &mut |item| { + if let PostgresQueryStreamItem::Row(row) = item { + query_export_rows.push(row); + } + Ok(()) + }) + .await?; + + let cancelled = AtomicBool::new(false); + let mut table_export_rows = Vec::new(); + stream_query_rows_on_client(&client, &select_sql, None, &cancelled, &mut |row| { + table_export_rows.push(row.to_vec()); + Ok(()) + }) + .await?; + drop(client); + + Ok::<_, String>((query, underscore_scalar, vector_named_enum, query_export_rows, table_export_rows)) + } + .await; + + let cleanup = execute_query(&pool, &format!("DROP SCHEMA {schema_ident} CASCADE")).await; + cleanup.expect("drop schema"); + let (query, underscore_scalar, vector_named_enum, query_export_rows, table_export_rows) = + exercise.expect("exercise custom array fallbacks"); + let expected = vec![ + serde_json::Value::String(r#"{"(7,alpha)"}"#.to_string()), + serde_json::Value::String("{ready}".to_string()), + serde_json::Value::String("{7}".to_string()), + ]; + + assert_eq!(query.rows, vec![expected.clone()]); + assert_eq!(underscore_scalar.rows, vec![vec![serde_json::Value::String("secret".to_string())]]); + assert_eq!(vector_named_enum.rows, vec![vec![serde_json::Value::String("label".to_string())]]); + assert_eq!(query_export_rows, vec![expected.clone()]); + assert_eq!(table_export_rows, vec![expected]); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a writable PostgreSQL 18 database"] + async fn postgres_custom_type_fallback_refreshes_stale_cached_metadata() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool_a = connect(&url, Duration::from_secs(5)).await.expect("connect postgres pool A"); + let pool_b = connect(&url, Duration::from_secs(5)).await.expect("connect postgres pool B"); + assert_postgres_18(&pool_a).await; + let schema = format!("dbx_custom_stale_{}", uuid::Uuid::new_v4().simple()); + let schema_ident = pg_quote_ident(&schema); + let payload_type = format!("{schema_ident}.payload"); + let view = format!("{schema_ident}.cached_payload"); + let view_sql = format!("SELECT payload FROM {view}"); + execute_query(&pool_a, &format!("CREATE SCHEMA {schema_ident}")).await.expect("create schema"); + + let exercise = async { + execute_query(&pool_a, &format!("CREATE TYPE {payload_type} AS (id integer, label text)")).await?; + execute_query(&pool_a, &format!("CREATE VIEW {view} AS SELECT ROW(7, 'alpha')::{payload_type} AS payload")) + .await?; + let custom = execute_query(&pool_a, &view_sql).await?; + + execute_query(&pool_b, &format!("DROP VIEW {view}")).await?; + execute_query(&pool_b, &format!("CREATE VIEW {view} AS SELECT 42::int4 AS payload")).await?; + let builtin = execute_query(&pool_a, &view_sql).await?; + Ok::<_, String>((custom, builtin)) + } + .await; + + let cleanup = execute_query(&pool_a, &format!("DROP SCHEMA {schema_ident} CASCADE")).await; + cleanup.expect("drop schema"); + let (custom, builtin) = exercise.expect("exercise stale cached custom metadata"); + assert_eq!(custom.column_types, vec!["payload"]); + assert_eq!(custom.rows[0][0], serde_json::Value::String("(7,alpha)".to_string())); + assert_eq!(builtin.column_types, vec!["int4"]); + assert_eq!(builtin.rows[0][0], serde_json::Value::Number(42.into())); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a writable PostgreSQL 18 database"] + async fn postgres_text_fallback_stops_before_late_row_error_at_limit() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool = connect(&url, Duration::from_secs(5)).await.expect("connect postgres"); + assert_postgres_18(&pool).await; + let schema = format!("dbx_custom_limit_{}", uuid::Uuid::new_v4().simple()); + let schema_ident = pg_quote_ident(&schema); + let payload_type = format!("{schema_ident}.payload"); + let fail_after_two = format!("{schema_ident}.fail_after_two"); + execute_query(&pool, &format!("CREATE SCHEMA {schema_ident}")).await.expect("create schema"); + + let exercise = async { + execute_query(&pool, &format!("CREATE TYPE {payload_type} AS (id integer)")).await?; + execute_query( + &pool, + &format!( + "CREATE FUNCTION {fail_after_two}(i integer) RETURNS integer LANGUAGE plpgsql AS $$ \ + BEGIN IF i >= 2 THEN RAISE EXCEPTION 'late row failure'; END IF; RETURN i; END $$" + ), + ) + .await?; + let client = checkout_postgres_client(&pool, None, Duration::from_secs(5)).await?; + let custom_sql = format!( + "SELECT ROW({fail_after_two}(i))::{payload_type} AS payload \ + FROM generate_series(1, 2) AS series(i)" + ); + let limited = execute_select_query(&client, &custom_sql, Instant::now(), 1).await; + let cancelled = AtomicBool::new(false); + let mut exported_rows = Vec::new(); + let exported = stream_query_rows_on_client(&client, &custom_sql, Some(1), &cancelled, &mut |row| { + exported_rows.push(row.to_vec()); + Ok(()) + }) + .await; + let recovery = execute_select_query(&client, "SELECT 1::int4 AS value", Instant::now(), 1).await; + drop(client); + Ok::<_, String>((limited, exported, exported_rows, recovery)) + } + .await; + + let cleanup = execute_query(&pool, &format!("DROP SCHEMA {schema_ident} CASCADE")).await; + cleanup.expect("drop schema"); + let (limited, exported, exported_rows, recovery) = exercise.expect("set up late row error query"); + let limited = limited.expect("query should stop before late row error"); + let exported = exported.expect("streamed export should stop before late row error"); + let recovery = recovery.expect("connection should remain reusable"); + assert_eq!(limited.column_types, vec!["payload"]); + assert_eq!(limited.rows, vec![vec![serde_json::Value::String("(1)".to_string())]]); + assert!(limited.truncated); + assert_eq!(exported, 1); + assert_eq!(exported_rows, vec![vec![serde_json::Value::String("(1)".to_string())]]); + assert_eq!(recovery.column_types, vec!["int4"]); + assert_eq!(recovery.rows[0][0], serde_json::Value::Number(1.into())); + } + fn state_enum_values(columns: &[ColumnInfo]) -> Option> { columns.iter().find(|column| column.name == "state").and_then(|column| column.enum_values.clone()) } diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index c7ac0088a..6f523c5a1 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -3291,61 +3291,24 @@ where let mut conn = connection.lock().await; let stream_result = match &mut *conn { TxnConnection::Postgres(conn) => { - let stmt = conn.prepare_cached(sql).await.map_err(|e| format!("Prepare failed: {e}")); - match stmt { - Ok(stmt) => { - let column_types: Vec = - stmt.columns().iter().map(|column| column.type_().name().to_string()).collect(); - let column_classes = db::postgres::classify_pg_column_types(&column_types); - let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = Vec::new(); - match conn.query_raw(&stmt, params).await { - Ok(stream) => { - tokio::pin!(stream); - let mut batch = Vec::with_capacity(batch_size); - let mut total_rows = 0_u64; - let mut error = None; - while let Some(row_result) = stream.next().await { - match row_result { - Ok(row) => { - let values = (0..row.columns().len()) - .map(|index| { - db::postgres::pg_value_to_json_classified( - &row, - index, - column_classes - .get(index) - .copied() - .unwrap_or(db::postgres::PgColType::Other), - ) - }) - .collect(); - batch.push(values); - total_rows += 1; - if batch.len() >= batch_size { - if let Err(err) = on_batch(std::mem::take(&mut batch)) { - error = Some(err); - break; - } - batch = Vec::with_capacity(batch_size); - } - } - Err(err) => { - error = Some(format!("Query failed: {err}")); - break; - } - } - } - if error.is_none() && !batch.is_empty() { - if let Err(err) = on_batch(batch) { - error = Some(err); - } - } - error.map_or(Ok(total_rows), Err) - } - Err(err) => Err(format!("Query failed: {err}")), + let mut batch = Vec::with_capacity(batch_size); + let mut total_rows = 0_u64; + let result = db::postgres::stream_select_query_inner(conn, sql, None, &mut |item| { + if let db::postgres::PostgresQueryStreamItem::Row(row) = item { + batch.push(row); + total_rows += 1; + if batch.len() >= batch_size { + on_batch(std::mem::take(&mut batch))?; + batch = Vec::with_capacity(batch_size); } } - Err(err) => Err(err), + Ok(()) + }) + .await; + match result { + Ok(_) if !batch.is_empty() => on_batch(batch).map(|_| total_rows), + Ok(_) => Ok(total_rows), + Err(error) => Err(error), } } TxnConnection::Mysql(conn) => match conn.query_iter(sql).await { @@ -3468,44 +3431,7 @@ async fn execute_manual_txn_postgres_statement( row_limit: usize, ) -> Result { if starts_with_executable_sql_keyword(sql, &["SELECT", "SHOW", "EXPLAIN", "WITH", "TABLE"]) { - let start = std::time::Instant::now(); - let stmt = conn.prepare_cached(sql).await.map_err(|e| format!("Prepare failed: {e}"))?; - let columns: Vec = stmt.columns().iter().map(|c| c.name().to_string()).collect(); - let column_types: Vec = stmt.columns().iter().map(|c| c.type_().name().to_string()).collect(); - let column_classes = db::postgres::classify_pg_column_types(&column_types); - let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = Vec::new(); - let stream = conn.query_raw(&stmt, params).await.map_err(|e| format!("Query failed: {e}"))?; - tokio::pin!(stream); - let mut data: Vec> = Vec::with_capacity(row_limit.min(1024)); - let mut truncated = false; - while let Some(row_result) = stream.next().await { - if data.len() >= row_limit { - truncated = true; - break; - } - let row = row_result.map_err(|e| format!("Query failed: {e}"))?; - let values: Vec = (0..row.columns().len()) - .map(|i| { - db::postgres::pg_value_to_json_classified( - &row, - i, - column_classes.get(i).copied().unwrap_or(db::postgres::PgColType::Other), - ) - }) - .collect(); - data.push(values); - } - Ok(db::QueryResult { - columns, - column_types, - column_sortables: vec![], - rows: data, - affected_rows: 0, - execution_time_ms: start.elapsed().as_millis(), - truncated, - session_id: None, - has_more: false, - }) + db::postgres::execute_select_query(conn, sql, std::time::Instant::now(), row_limit).await } else { let affected = conn.execute(sql, &[]).await.map_err(|e| format!("Query failed: {e}"))?; Ok(db::QueryResult {