fix(postgresql): fall back to text for custom types

This commit is contained in:
lewis 2026-07-23 22:42:43 +08:00 committed by GitHub
parent 094d1fefde
commit e0ad5dcf86
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 426 additions and 127 deletions

View File

@ -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<String>, unsupported_type: String },
}
struct PreparedSelectMetadata {
columns: Vec<String>,
column_types: Vec<String>,
column_classes: Vec<PgColType>,
unsupported_type: Option<String>,
}
fn prepared_select_metadata(stmt: &tokio_postgres::Statement) -> PreparedSelectMetadata {
let columns: Vec<String> = stmt.columns().iter().map(|c| c.name().to_string()).collect();
let column_types: Vec<String> = 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<QueryResult, tokio_postgres::Error> {
) -> Result<PreparedSelectOutcome, tokio_postgres::Error> {
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<String> = stmt.columns().iter().map(|c| c.name().to_string()).collect();
let column_types: Vec<String> = 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<String>>) -> Vec<String> {
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<Vec<String>>,
) -> Result<QueryResult, String> {
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<String> = Vec::new();
let mut result_rows: Vec<Vec<serde_json::Value>> = 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<QueryResult, String> {
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<QueryResult, String> {
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<String>, 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<usize>,
on_item: &mut impl FnMut(PostgresQueryStreamItem) -> Result<(), String>,
) -> Result<u64, PostgresQueryStreamError> {
let stmt =
client.prepare_cached(sql).await.map_err(|err| PostgresQueryStreamError::Postgres { err, emitted: false })?;
let columns: Vec<String> = stmt.columns().iter().map(|c| c.name().to_string()).collect();
let column_types: Vec<String> = 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<usize>,
prepared_column_types: Option<Vec<String>>,
on_item: &mut impl FnMut(PostgresQueryStreamItem) -> Result<(), String>,
) -> Result<u64, String> {
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<usize>,
@ -943,6 +1035,13 @@ async fn stream_select_query_inner(
) -> Result<u64, String> {
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<u64, String> {
let stmt = client.prepare_cached(sql).await.map_err(pg_error_to_string)?;
let column_types: Vec<String> = 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<u64, String> {
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::<u32>()
.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<Vec<String>> {
columns.iter().find(|column| column.name == "state").and_then(|column| column.enum_values.clone())
}

View File

@ -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<String> =
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<db::QueryResult, String> {
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<String> = stmt.columns().iter().map(|c| c.name().to_string()).collect();
let column_types: Vec<String> = 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<serde_json::Value>> = 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<serde_json::Value> = (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 {