use pgwire::api::Type;
use pgwire::api::portal::Format;
use pgwire::api::results::{FieldFormat, FieldInfo};
use crate::control::server::response_shape::types::DdlColType;
pub(super) fn pg_type_to_ddl_col_type(t: &Type) -> DdlColType {
if *t == Type::INT8 {
DdlColType::Int8
} else if *t == Type::INT4 {
DdlColType::Int4
} else if *t == Type::INT2 {
DdlColType::Int2
} else if *t == Type::FLOAT8 {
DdlColType::Float8
} else if *t == Type::FLOAT4 {
DdlColType::Float4
} else if *t == Type::BOOL {
DdlColType::Bool
} else if *t == Type::BYTEA {
DdlColType::Bytea
} else if *t == Type::VARCHAR {
DdlColType::Varchar
} else if *t == Type::JSON {
DdlColType::Json
} else if *t == Type::JSONB {
DdlColType::Jsonb
} else if *t == Type::TIMESTAMP {
DdlColType::Timestamp
} else if *t == Type::TIMESTAMPTZ {
DdlColType::Timestamptz
} else if *t == Type::FLOAT4_ARRAY {
DdlColType::Float4Array
} else if *t == Type::FLOAT8_ARRAY {
DdlColType::Float8Array
} else {
DdlColType::Text
}
}
pub(super) fn binary_supported(ct: DdlColType) -> bool {
matches!(
ct,
DdlColType::Int8
| DdlColType::Int4
| DdlColType::Int2
| DdlColType::Float8
| DdlColType::Float4
| DdlColType::Bool
| DdlColType::Text
| DdlColType::Varchar
)
}
fn requested_format(fmt: &Format, idx: usize) -> FieldFormat {
match fmt {
Format::Individual(codes) => codes
.get(idx)
.map(|c| FieldFormat::from(*c))
.unwrap_or(FieldFormat::Text),
_ => fmt.format_for(idx),
}
}
pub(super) fn resolve_result_formats(fields: &[FieldInfo], fmt: &Format) -> Vec<FieldFormat> {
fields
.iter()
.enumerate()
.map(|(i, f)| {
let ct = pg_type_to_ddl_col_type(f.datatype());
if requested_format(fmt, i) == FieldFormat::Binary && binary_supported(ct) {
FieldFormat::Binary
} else {
FieldFormat::Text
}
})
.collect()
}
pub(super) fn stamp_formats(fields: &[FieldInfo], formats: &[FieldFormat]) -> Vec<FieldInfo> {
fields
.iter()
.enumerate()
.map(|(i, f)| {
let format = formats.get(i).copied().unwrap_or(FieldFormat::Text);
FieldInfo::new(
f.name().to_owned(),
None,
None,
f.datatype().clone(),
format,
)
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn maps_scalar_pg_types() {
assert_eq!(pg_type_to_ddl_col_type(&Type::INT8), DdlColType::Int8);
assert_eq!(pg_type_to_ddl_col_type(&Type::FLOAT8), DdlColType::Float8);
assert_eq!(pg_type_to_ddl_col_type(&Type::BOOL), DdlColType::Bool);
assert_eq!(pg_type_to_ddl_col_type(&Type::BYTEA), DdlColType::Bytea);
assert_eq!(
pg_type_to_ddl_col_type(&Type::TIMESTAMP),
DdlColType::Timestamp
);
assert_eq!(pg_type_to_ddl_col_type(&Type::UUID), DdlColType::Text);
}
#[test]
fn binary_supported_excludes_feature_blocked() {
assert!(binary_supported(DdlColType::Int8));
assert!(binary_supported(DdlColType::Bool));
assert!(binary_supported(DdlColType::Text));
assert!(!binary_supported(DdlColType::Bytea));
assert!(!binary_supported(DdlColType::Timestamp));
assert!(!binary_supported(DdlColType::Json));
assert!(!binary_supported(DdlColType::Float8Array));
}
#[test]
fn unified_binary_downgrades_blocked_types() {
let fields = vec![
FieldInfo::new("a".into(), None, None, Type::INT8, FieldFormat::Text),
FieldInfo::new("b".into(), None, None, Type::TIMESTAMP, FieldFormat::Text),
];
let formats = resolve_result_formats(&fields, &Format::UnifiedBinary);
assert_eq!(formats[0], FieldFormat::Binary);
assert_eq!(formats[1], FieldFormat::Text);
}
#[test]
fn individual_shorter_than_columns_defaults_text() {
let fields = vec![
FieldInfo::new("a".into(), None, None, Type::INT8, FieldFormat::Text),
FieldInfo::new("b".into(), None, None, Type::INT8, FieldFormat::Text),
];
let formats = resolve_result_formats(&fields, &Format::Individual(vec![1]));
assert_eq!(formats[0], FieldFormat::Binary);
assert_eq!(formats[1], FieldFormat::Text);
}
#[test]
fn unified_text_is_all_text() {
let fields = vec![FieldInfo::new(
"a".into(),
None,
None,
Type::INT8,
FieldFormat::Text,
)];
let formats = resolve_result_formats(&fields, &Format::UnifiedText);
assert_eq!(formats[0], FieldFormat::Text);
}
}