fix(postgresql): fall back to text for custom types
This commit is contained in:
parent
094d1fefde
commit
e0ad5dcf86
|
|
@ -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())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
Loading…
Reference in New Issue