use pylon_value::DecodedValue;
use crate::numeric;
pub use crate::error::Error;
pub type Result<T> = crate::Result<T>;
const OID_BOOL: u32 = 16;
const OID_BYTEA: u32 = 17;
const OID_INT8: u32 = 20;
const OID_INT2: u32 = 21;
const OID_INT4: u32 = 23;
const OID_TEXT: u32 = 25;
const OID_JSONB: u32 = 3802;
const OID_FLOAT4: u32 = 700;
const OID_FLOAT8: u32 = 701;
const OID_BPCHAR: u32 = 1042;
const OID_VARCHAR: u32 = 1043;
const OID_NUMERIC: u32 = 1700;
const OID_DATE: u32 = 1082;
const OID_TIME: u32 = 1083;
const OID_TIMESTAMP: u32 = 1114;
const OID_TIMESTAMPTZ: u32 = 1184;
const OID_INTERVAL: u32 = 1186;
const OID_UUID: u32 = 2950;
const OID_RECORD: u32 = 2249;
const OID_RECORD_ARRAY: u32 = 2287;
const OID_UNKNOWN: u32 = 705;
const OID_NAME: u32 = 19;
const OID_BOOL_ARRAY: u32 = 1000;
const OID_BYTEA_ARRAY: u32 = 1001;
const OID_INT2_ARRAY: u32 = 1005;
const OID_INT4_ARRAY: u32 = 1007;
const OID_TEXT_ARRAY: u32 = 1009;
const OID_BPCHAR_ARRAY: u32 = 1014;
const OID_VARCHAR_ARRAY: u32 = 1015;
const OID_INT8_ARRAY: u32 = 1016;
const OID_FLOAT4_ARRAY: u32 = 1021;
const OID_FLOAT8_ARRAY: u32 = 1022;
const OID_NUMERIC_ARRAY: u32 = 1231;
const OID_UUID_ARRAY: u32 = 2951;
const OID_JSONB_ARRAY: u32 = 3807;
const OID_INT4RANGE: u32 = 3904;
const OID_INT8RANGE: u32 = 3926;
const OID_NUMRANGE: u32 = 3906;
const OID_TSRANGE: u32 = 3908;
const OID_TSTZRANGE: u32 = 3910;
const OID_DATERANGE: u32 = 3912;
const OID_INT4MULTIRANGE: u32 = 4451;
const OID_INT8MULTIRANGE: u32 = 4536;
const OID_NUMMULTIRANGE: u32 = 4532;
const OID_TSMULTIRANGE: u32 = 4533;
const OID_TSTZMULTIRANGE: u32 = 4534;
const OID_DATEMULTIRANGE: u32 = 4535;
fn range_element_oid(oid: u32) -> Option<u32> {
match oid {
OID_INT4RANGE | OID_INT4MULTIRANGE => Some(OID_INT4),
OID_INT8RANGE | OID_INT8MULTIRANGE => Some(OID_INT8),
OID_NUMRANGE | OID_NUMMULTIRANGE => Some(OID_NUMERIC),
OID_TSRANGE | OID_TSMULTIRANGE => Some(OID_TIMESTAMP),
OID_TSTZRANGE | OID_TSTZMULTIRANGE => Some(OID_TIMESTAMPTZ),
OID_DATERANGE | OID_DATEMULTIRANGE => Some(OID_DATE),
_ => None,
}
}
#[derive(Debug, Clone, Default)]
pub struct ExtensionOids {
pub vector: Option<u32>,
pub enums: std::collections::HashSet<u32>,
pub domains: std::collections::HashMap<u32, u32>,
pub arrays: std::collections::HashSet<u32>,
}
pub(crate) const TYPE_DISCOVERY_SQL: &str = "\
SELECT t.oid::int8, t.typtype::text, COALESCE(b.oid, 0)::int8, t.typname::text \
FROM pg_type t \
LEFT JOIN pg_type b ON b.oid = t.typbasetype \
WHERE t.typtype IN ('e', 'd') OR t.typname = 'vector' \
UNION ALL \
SELECT a.oid::int8, 'A', e.oid::int8, a.typname::text \
FROM pg_type a \
JOIN pg_type e ON e.oid = a.typelem \
WHERE a.typcategory = 'A' AND e.typtype IN ('e', 'd')";
impl ExtensionOids {
pub(crate) fn from_discovery_rows(rows: impl IntoIterator<Item = (u32, String, u32, String)>) -> Self {
let mut out = Self::default();
for (oid, typtype, base_oid, typname) in rows {
match typtype.as_str() {
"e" => {
out.enums.insert(oid);
}
"d" if base_oid != 0 => {
out.domains.insert(oid, base_oid);
}
"A" => {
out.arrays.insert(oid);
}
_ => {}
}
if typname == "vector" {
out.vector = Some(oid);
}
}
out
}
}
pub fn decode_value(oid: u32, data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
if let Some(vector_oid) = ext.vector
&& oid == vector_oid
{
return Ok(DecodedValue::Array(decode_vector(data)?));
}
match oid {
OID_BOOL => Ok(DecodedValue::Bool(data.first().copied().unwrap_or(0) != 0)),
OID_INT2 => Ok(DecodedValue::I64(i16::from_be_bytes(data.try_into()?) as i64)),
OID_INT4 => Ok(DecodedValue::I64(i32::from_be_bytes(data.try_into()?) as i64)),
OID_INT8 => Ok(DecodedValue::I64(i64::from_be_bytes(data.try_into()?))),
OID_FLOAT4 => Ok(DecodedValue::F64(f32::from_be_bytes(data.try_into()?) as f64)),
OID_FLOAT8 => Ok(DecodedValue::F64(f64::from_be_bytes(data.try_into()?))),
OID_TEXT | OID_VARCHAR | OID_BPCHAR | OID_UNKNOWN | OID_NAME => {
Ok(DecodedValue::Str(std::str::from_utf8(data)?.to_string()))
}
OID_UUID => {
let mut bytes = [0u8; 16];
bytes.copy_from_slice(data);
Ok(DecodedValue::Uuid(bytes))
}
OID_BYTEA => Ok(DecodedValue::Bytes(data.to_vec())),
OID_NUMERIC => decode_numeric(data),
OID_INTERVAL => decode_interval(data),
OID_DATE => Ok(DecodedValue::Date(i32::from_be_bytes(data.try_into()?))),
OID_TIME => Ok(DecodedValue::Time(i64::from_be_bytes(data.try_into()?))),
OID_TIMESTAMP => Ok(DecodedValue::Timestamp(i64::from_be_bytes(data.try_into()?))),
OID_TIMESTAMPTZ => Ok(DecodedValue::Timestamptz(i64::from_be_bytes(data.try_into()?))),
OID_JSONB => decode_jsonb(data),
OID_RECORD => decode_record(data, ext),
OID_RECORD_ARRAY => decode_array(data, ext),
OID_BOOL_ARRAY | OID_BYTEA_ARRAY | OID_INT2_ARRAY | OID_INT4_ARRAY | OID_INT8_ARRAY | OID_TEXT_ARRAY
| OID_BPCHAR_ARRAY | OID_VARCHAR_ARRAY | OID_FLOAT4_ARRAY | OID_FLOAT8_ARRAY | OID_NUMERIC_ARRAY
| OID_UUID_ARRAY | OID_JSONB_ARRAY => decode_array(data, ext),
OID_INT4RANGE | OID_INT8RANGE | OID_NUMRANGE | OID_TSRANGE | OID_TSTZRANGE | OID_DATERANGE => {
decode_range(data, range_element_oid(oid).expect("range OID"), ext)
}
OID_INT4MULTIRANGE | OID_INT8MULTIRANGE | OID_NUMMULTIRANGE | OID_TSMULTIRANGE | OID_TSTZMULTIRANGE
| OID_DATEMULTIRANGE => decode_multirange(data, range_element_oid(oid).expect("multirange OID"), ext),
_ if ext.enums.contains(&oid) => {
Ok(DecodedValue::Str(std::str::from_utf8(data)?.to_string()))
}
_ if ext.arrays.contains(&oid) => decode_array(data, ext),
_ => match ext.domains.get(&oid) {
Some(&base_oid) => decode_value(base_oid, data, ext),
None => Err(Error::UnknownTypeOid { oid }),
},
}
}
fn decode_numeric(data: &[u8]) -> Result<DecodedValue> {
Ok(DecodedValue::Decimal(numeric::decode(data)?))
}
fn decode_interval(data: &[u8]) -> Result<DecodedValue> {
if data.len() != 16 {
return Err(Error::message(format!(
"malformed interval: expected 16 bytes, got {}",
data.len()
)));
}
let microseconds = i64::from_be_bytes(data[0..8].try_into()?);
let days = i32::from_be_bytes(data[8..12].try_into()?);
let months = i32::from_be_bytes(data[12..16].try_into()?);
Ok(DecodedValue::Interval {
months,
days,
microseconds,
})
}
const RANGE_EMPTY: u8 = 0x01;
const RANGE_LB_INC: u8 = 0x02;
const RANGE_UB_INC: u8 = 0x04;
const RANGE_LB_INF: u8 = 0x08;
const RANGE_UB_INF: u8 = 0x10;
fn decode_range(data: &[u8], element_oid: u32, ext: &ExtensionOids) -> Result<DecodedValue> {
let flags = data[0];
let mut offset = 1usize;
if flags & RANGE_EMPTY != 0 {
return Ok(DecodedValue::Range {
lower: None,
upper: None,
inc_lower: false,
inc_upper: false,
empty: true,
});
}
let lower = if flags & RANGE_LB_INF != 0 {
None
} else {
let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
let value = decode_value(element_oid, &data[offset..offset + len], ext)?;
offset += len;
Some(Box::new(value))
};
let upper = if flags & RANGE_UB_INF != 0 {
None
} else {
let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
Some(Box::new(decode_value(element_oid, &data[offset..offset + len], ext)?))
};
Ok(DecodedValue::Range {
lower,
upper,
inc_lower: flags & RANGE_LB_INC != 0,
inc_upper: flags & RANGE_UB_INC != 0,
empty: false,
})
}
fn decode_multirange(data: &[u8], element_oid: u32, ext: &ExtensionOids) -> Result<DecodedValue> {
let mut offset = 0usize;
let count = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
let mut ranges = Vec::with_capacity(count);
for _ in 0..count {
let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
ranges.push(decode_range(&data[offset..offset + len], element_oid, ext)?);
offset += len;
}
Ok(DecodedValue::Array(ranges))
}
fn decode_jsonb(data: &[u8]) -> Result<DecodedValue> {
let text = std::str::from_utf8(&data[1..])?;
let value: serde_json::Value = serde_json::from_str(text)?;
Ok(json_to_cached(value))
}
fn json_to_cached(value: serde_json::Value) -> DecodedValue {
match value {
serde_json::Value::Null => DecodedValue::Null,
serde_json::Value::Bool(b) => DecodedValue::Bool(b),
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
DecodedValue::I64(i)
} else {
DecodedValue::F64(n.as_f64().unwrap_or(f64::NAN))
}
}
serde_json::Value::String(s) => DecodedValue::Str(s),
serde_json::Value::Array(items) => DecodedValue::Array(items.into_iter().map(json_to_cached).collect()),
serde_json::Value::Object(map) => {
DecodedValue::Object(map.into_iter().map(|(k, v)| (k, json_to_cached(v))).collect())
}
}
}
fn decode_vector(data: &[u8]) -> Result<Vec<DecodedValue>> {
let ndim = u16::from_be_bytes(data[0..2].try_into()?) as usize;
let mut values = Vec::with_capacity(ndim);
for i in 0..ndim {
let start = 4 + i * 4;
let f = f32::from_be_bytes(data[start..start + 4].try_into()?);
values.push(DecodedValue::F64(f as f64));
}
Ok(values)
}
fn encode_vector(items: &[DecodedValue], out: &mut bytes::BytesMut) -> Result<()> {
let ndim: u16 = items
.len()
.try_into()
.map_err(|_| Error::message("vector has too many dimensions to encode"))?;
out.put_u16(ndim);
out.put_u16(0); for item in items {
let f = match item {
DecodedValue::F64(f) => *f as f32,
DecodedValue::I64(i) => *i as f32,
other => return Err(Error::message(format!("cannot encode {other:?} as a vector element"))),
};
out.put_f32(f);
}
Ok(())
}
fn decode_record(data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
let mut offset = 0usize;
let nfields = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
let mut fields = Vec::with_capacity(nfields);
for _ in 0..nfields {
let type_oid = u32::from_be_bytes(data[offset..offset + 4].try_into()?);
offset += 4;
let field_len = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
offset += 4;
if field_len == -1 {
fields.push(DecodedValue::Null);
} else {
let len = field_len as usize;
fields.push(decode_value(type_oid, &data[offset..offset + len], ext)?);
offset += len;
}
}
Ok(DecodedValue::Composite(fields))
}
fn decode_array(data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
let mut offset = 0usize;
let ndim = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
offset += 4;
offset += 4; let element_oid = u32::from_be_bytes(data[offset..offset + 4].try_into()?);
offset += 4;
if ndim == 0 {
return Ok(DecodedValue::Array(vec![]));
}
let dim_size = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
offset += 4;
offset += 4;
let mut items = Vec::with_capacity(dim_size);
for _ in 0..dim_size {
let elem_len = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
offset += 4;
if elem_len == -1 {
items.push(DecodedValue::Null);
} else {
let len = elem_len as usize;
items.push(decode_value(element_oid, &data[offset..offset + len], ext)?);
offset += len;
}
}
Ok(DecodedValue::Array(items))
}
use bytes::BufMut;
use postgres_types::{IsNull, Kind, Type};
fn accepts_text_bytes(ty: &Type) -> bool {
match ty.kind() {
Kind::Enum(_) => true,
Kind::Domain(base) => accepts_text_bytes(base),
_ => matches!(
*ty,
Type::TEXT | Type::VARCHAR | Type::BPCHAR | Type::NAME | Type::JSON | Type::BYTEA | Type::UNKNOWN
),
}
}
pub fn encode_value(value: &DecodedValue, ty: &Type, out: &mut bytes::BytesMut) -> Result<IsNull> {
let DecodedValue::Null = value else {
return encode_non_null(value, ty, out);
};
Ok(IsNull::Yes)
}
fn encode_non_null(value: &DecodedValue, ty: &Type, out: &mut bytes::BytesMut) -> Result<IsNull> {
if *ty == Type::NUMERIC {
let text = match value {
DecodedValue::Decimal(s) | DecodedValue::Str(s) => s.clone(),
DecodedValue::I64(i) => i.to_string(),
DecodedValue::F64(f) => format!("{f}"),
_ => return Err(Error::message("cannot bind this value as a numeric parameter")),
};
numeric::encode(&text, out)?;
return Ok(IsNull::No);
}
match value {
DecodedValue::Null => unreachable!("caller already handled NULL"),
DecodedValue::Bool(b) => out.put_u8(*b as u8),
DecodedValue::I64(i) => {
if *ty == Type::INT2 {
out.put_i16(*i as i16);
} else if *ty == Type::INT4 {
out.put_i32(*i as i32);
} else {
out.put_i64(*i);
}
}
DecodedValue::F64(f) => {
if *ty == Type::FLOAT4 {
out.put_f32(*f as f32);
} else {
out.put_f64(*f);
}
}
DecodedValue::Str(s) => {
if *ty == Type::UUID {
out.put_slice(&parse_uuid_str(s)?);
} else if *ty == Type::JSONB {
out.put_u8(1);
out.put_slice(s.as_bytes());
} else if accepts_text_bytes(ty) {
out.put_slice(s.as_bytes());
} else {
return Err(Error::message(format!(
"cannot bind a string as a parameter of type {:?}",
ty.name()
)));
}
}
DecodedValue::Bytes(b) => out.put_slice(b),
DecodedValue::Uuid(bytes) => out.put_slice(bytes),
DecodedValue::Decimal(s) => numeric::encode(s, out)?,
DecodedValue::Array(items) if ty.name() == "vector" => {
encode_vector(items, out)?;
}
DecodedValue::Array(items) => {
let element_ty = match ty.kind() {
Kind::Array(inner) => inner.clone(),
_ => Type::TEXT,
};
encode_array(items, &element_ty, out)?;
}
DecodedValue::Composite(_) => {
return Err(Error::message("cannot bind a composite value as a query parameter"));
}
DecodedValue::Object(fields) => {
let json = cached_object_to_json(fields);
out.put_u8(1); out.put_slice(json.to_string().as_bytes());
}
DecodedValue::Interval {
months,
days,
microseconds,
} => {
out.put_i64(*microseconds);
out.put_i32(*days);
out.put_i32(*months);
}
DecodedValue::Date(days) => out.put_i32(*days),
DecodedValue::Time(us) => out.put_i64(*us),
DecodedValue::Timestamp(us) => out.put_i64(*us),
DecodedValue::Timestamptz(us) => out.put_i64(*us),
DecodedValue::Range {
lower,
upper,
inc_lower,
inc_upper,
empty,
} => {
if *empty {
out.put_u8(RANGE_EMPTY);
return Ok(IsNull::No);
}
let element_ty = match ty.kind() {
Kind::Range(inner) => inner.clone(),
_ => Type::TEXT,
};
let mut flags = 0u8;
if *inc_lower {
flags |= RANGE_LB_INC;
}
if *inc_upper {
flags |= RANGE_UB_INC;
}
if lower.is_none() {
flags |= RANGE_LB_INF;
}
if upper.is_none() {
flags |= RANGE_UB_INF;
}
out.put_u8(flags);
for bound in [lower, upper].into_iter().flatten() {
let mut buf = bytes::BytesMut::new();
encode_value(bound, &element_ty, &mut buf)?;
out.put_i32(buf.len() as i32);
out.put_slice(&buf);
}
}
}
Ok(IsNull::No)
}
fn cached_object_to_json(fields: &[(String, DecodedValue)]) -> serde_json::Value {
serde_json::Value::Object(fields.iter().map(|(k, v)| (k.clone(), cached_to_json(v))).collect())
}
fn cached_to_json(value: &DecodedValue) -> serde_json::Value {
match value {
DecodedValue::Null => serde_json::Value::Null,
DecodedValue::Bool(b) => serde_json::Value::Bool(*b),
DecodedValue::I64(i) => serde_json::Value::Number((*i).into()),
DecodedValue::F64(f) => serde_json::Number::from_f64(*f)
.map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null),
DecodedValue::Str(s) => serde_json::Value::String(s.clone()),
DecodedValue::Bytes(b) => serde_json::Value::String(hex::encode(b)),
DecodedValue::Uuid(bytes) => serde_json::Value::String(format_uuid(bytes)),
DecodedValue::Decimal(s) => serde_json::Value::String(s.clone()),
DecodedValue::Array(items) | DecodedValue::Composite(items) => {
serde_json::Value::Array(items.iter().map(cached_to_json).collect())
}
DecodedValue::Object(fields) => cached_object_to_json(fields),
DecodedValue::Interval {
months,
days,
microseconds,
} => serde_json::json!({
"months": months, "days": days, "microseconds": microseconds,
}),
DecodedValue::Date(days) => serde_json::json!({ "days_since_2000_01_01": days }),
DecodedValue::Time(us) => serde_json::json!({ "microseconds_since_midnight": us }),
DecodedValue::Timestamp(us) => serde_json::json!({ "microseconds_since_2000_01_01": us }),
DecodedValue::Timestamptz(us) => serde_json::json!({ "microseconds_since_2000_01_01_utc": us }),
DecodedValue::Range {
lower,
upper,
inc_lower,
inc_upper,
empty,
} => serde_json::json!({
"lower": lower.as_deref().map(cached_to_json),
"upper": upper.as_deref().map(cached_to_json),
"inc_lower": inc_lower,
"inc_upper": inc_upper,
"empty": empty,
}),
}
}
fn format_uuid(bytes: &[u8; 16]) -> String {
let hex = hex::encode(bytes);
format!(
"{}-{}-{}-{}-{}",
&hex[0..8],
&hex[8..12],
&hex[12..16],
&hex[16..20],
&hex[20..32]
)
}
fn parse_uuid_str(s: &str) -> Result<[u8; 16]> {
let hex_only: String = s.chars().filter(|c| *c != '-').collect();
let bytes = hex::decode(&hex_only).map_err(|_| Error::message(format!("invalid UUID string: {s:?}")))?;
bytes
.try_into()
.map_err(|_: Vec<u8>| Error::message(format!("invalid UUID string: {s:?}")))
}
fn encode_array(items: &[DecodedValue], element_ty: &Type, out: &mut bytes::BytesMut) -> Result<()> {
if items.is_empty() {
out.put_i32(0); out.put_i32(0); out.put_u32(element_ty.oid());
return Ok(());
}
let has_null = items.iter().any(|v| matches!(v, DecodedValue::Null));
out.put_i32(1); out.put_i32(has_null as i32);
out.put_u32(element_ty.oid());
out.put_i32(items.len() as i32); out.put_i32(1);
for item in items {
if matches!(item, DecodedValue::Null) {
out.put_i32(-1);
continue;
}
let start = out.len();
out.put_i32(0); let is_null = encode_value(item, element_ty, out)?;
let len = (out.len() - start - 4) as i32;
let len = if matches!(is_null, IsNull::Yes) { -1 } else { len };
out[start..start + 4].copy_from_slice(&len.to_be_bytes());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn no_ext() -> ExtensionOids {
ExtensionOids::default()
}
#[test]
fn decodes_bool() {
assert_eq!(
decode_value(OID_BOOL, &[1], &no_ext()).unwrap(),
DecodedValue::Bool(true)
);
assert_eq!(
decode_value(OID_BOOL, &[0], &no_ext()).unwrap(),
DecodedValue::Bool(false)
);
}
#[test]
fn decodes_integers() {
assert_eq!(
decode_value(OID_INT2, &7i16.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::I64(7)
);
assert_eq!(
decode_value(OID_INT4, &(-42i32).to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::I64(-42)
);
assert_eq!(
decode_value(OID_INT8, &9_223_372_036_854_775_807i64.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::I64(9_223_372_036_854_775_807)
);
}
#[test]
fn decodes_floats() {
assert_eq!(
decode_value(OID_FLOAT4, &1.5f32.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::F64(1.5)
);
assert_eq!(
decode_value(OID_FLOAT8, &2.25f64.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::F64(2.25)
);
}
#[test]
fn decodes_text_varchar_bpchar() {
for oid in [OID_TEXT, OID_VARCHAR, OID_BPCHAR] {
assert_eq!(
decode_value(oid, "hello".as_bytes(), &no_ext()).unwrap(),
DecodedValue::Str("hello".to_string())
);
}
}
#[test]
fn decodes_unicode_text() {
assert_eq!(
decode_value(OID_TEXT, "héllo wörld 🎉".as_bytes(), &no_ext()).unwrap(),
DecodedValue::Str("héllo wörld 🎉".to_string())
);
}
#[test]
fn decodes_uuid() {
let bytes: [u8; 16] = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16];
assert_eq!(
decode_value(OID_UUID, &bytes, &no_ext()).unwrap(),
DecodedValue::Uuid(bytes)
);
}
#[test]
fn decodes_bytea() {
assert_eq!(
decode_value(OID_BYTEA, &[1, 2, 3, 255], &no_ext()).unwrap(),
DecodedValue::Bytes(vec![1, 2, 3, 255])
);
}
#[test]
fn decodes_interval() {
let mut data = Vec::new();
data.extend_from_slice(&3_600_000_000i64.to_be_bytes()); data.extend_from_slice(&2i32.to_be_bytes()); data.extend_from_slice(&1i32.to_be_bytes()); assert_eq!(
decode_value(OID_INTERVAL, &data, &no_ext()).unwrap(),
DecodedValue::Interval {
months: 1,
days: 2,
microseconds: 3_600_000_000
}
);
}
#[test]
fn encodes_interval() {
let value = DecodedValue::Interval {
months: 1,
days: 2,
microseconds: 3_600_000_000,
};
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::INTERVAL, &mut out).unwrap();
assert_eq!(decode_value(OID_INTERVAL, &out, &no_ext()).unwrap(), value);
}
#[test]
fn decodes_date_time_timestamp_timestamptz() {
assert_eq!(
decode_value(OID_DATE, &9525i32.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::Date(9525)
);
assert_eq!(
decode_value(OID_TIME, &3_600_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::Time(3_600_000_000)
);
assert_eq!(
decode_value(OID_TIMESTAMP, &1_000_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::Timestamp(1_000_000_000)
);
assert_eq!(
decode_value(OID_TIMESTAMPTZ, &1_000_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
DecodedValue::Timestamptz(1_000_000_000)
);
}
#[test]
fn encodes_date_time_timestamp_timestamptz() {
for (value, ty) in [
(DecodedValue::Date(9525), postgres_types::Type::DATE),
(DecodedValue::Time(3_600_000_000), postgres_types::Type::TIME),
(DecodedValue::Timestamp(1_000_000_000), postgres_types::Type::TIMESTAMP),
(
DecodedValue::Timestamptz(1_000_000_000),
postgres_types::Type::TIMESTAMPTZ,
),
] {
let mut out = bytes::BytesMut::new();
encode_value(&value, &ty, &mut out).unwrap();
let oid = match &value {
DecodedValue::Date(_) => OID_DATE,
DecodedValue::Time(_) => OID_TIME,
DecodedValue::Timestamp(_) => OID_TIMESTAMP,
DecodedValue::Timestamptz(_) => OID_TIMESTAMPTZ,
_ => unreachable!(),
};
assert_eq!(decode_value(oid, &out, &no_ext()).unwrap(), value);
}
}
#[test]
fn decodes_a_bounded_int8range() {
let mut data = vec![RANGE_LB_INC];
data.extend_from_slice(&8i32.to_be_bytes());
data.extend_from_slice(&1i64.to_be_bytes());
data.extend_from_slice(&8i32.to_be_bytes());
data.extend_from_slice(&10i64.to_be_bytes());
assert_eq!(
decode_value(OID_INT8RANGE, &data, &no_ext()).unwrap(),
DecodedValue::Range {
lower: Some(Box::new(DecodedValue::I64(1))),
upper: Some(Box::new(DecodedValue::I64(10))),
inc_lower: true,
inc_upper: false,
empty: false,
}
);
}
#[test]
fn decodes_an_empty_range() {
assert_eq!(
decode_value(OID_INT8RANGE, &[RANGE_EMPTY], &no_ext()).unwrap(),
DecodedValue::Range {
lower: None,
upper: None,
inc_lower: false,
inc_upper: false,
empty: true
}
);
}
#[test]
fn decodes_an_unbounded_range() {
let data = [RANGE_LB_INF | RANGE_UB_INF];
assert_eq!(
decode_value(OID_INT8RANGE, &data, &no_ext()).unwrap(),
DecodedValue::Range {
lower: None,
upper: None,
inc_lower: false,
inc_upper: false,
empty: false
}
);
}
#[test]
fn encodes_and_round_trips_an_int8range() {
let value = DecodedValue::Range {
lower: Some(Box::new(DecodedValue::I64(1))),
upper: Some(Box::new(DecodedValue::I64(10))),
inc_lower: true,
inc_upper: false,
empty: false,
};
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::INT8_RANGE, &mut out).unwrap();
assert_eq!(decode_value(OID_INT8RANGE, &out, &no_ext()).unwrap(), value);
}
#[test]
fn decodes_a_multirange_of_int8ranges() {
let mut range1 = vec![RANGE_LB_INC];
range1.extend_from_slice(&8i32.to_be_bytes());
range1.extend_from_slice(&1i64.to_be_bytes());
range1.extend_from_slice(&8i32.to_be_bytes());
range1.extend_from_slice(&3i64.to_be_bytes());
let mut range2 = vec![RANGE_LB_INC];
range2.extend_from_slice(&8i32.to_be_bytes());
range2.extend_from_slice(&5i64.to_be_bytes());
range2.extend_from_slice(&8i32.to_be_bytes());
range2.extend_from_slice(&7i64.to_be_bytes());
let mut data = 2i32.to_be_bytes().to_vec();
data.extend_from_slice(&(range1.len() as i32).to_be_bytes());
data.extend_from_slice(&range1);
data.extend_from_slice(&(range2.len() as i32).to_be_bytes());
data.extend_from_slice(&range2);
let decoded = decode_value(OID_INT8MULTIRANGE, &data, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Array(vec![
DecodedValue::Range {
lower: Some(Box::new(DecodedValue::I64(1))),
upper: Some(Box::new(DecodedValue::I64(3))),
inc_lower: true,
inc_upper: false,
empty: false,
},
DecodedValue::Range {
lower: Some(Box::new(DecodedValue::I64(5))),
upper: Some(Box::new(DecodedValue::I64(7))),
inc_lower: true,
inc_upper: false,
empty: false,
},
])
);
}
#[test]
fn errors_on_an_oid_no_rule_or_discovery_covers() {
let err = decode_value(999_999, &[0xff, 0xfe], &no_ext()).unwrap_err();
assert!(matches!(err, Error::UnknownTypeOid { oid: 999_999 }), "got {err:?}");
}
#[test]
fn decodes_a_discovered_enum_oid_as_its_label_text() {
let ext = ExtensionOids {
enums: std::collections::HashSet::from([50_001]),
..Default::default()
};
assert_eq!(
decode_value(50_001, "Active".as_bytes(), &ext).unwrap(),
DecodedValue::Str("Active".to_string())
);
}
#[test]
fn decodes_a_discovered_domain_through_its_base_type() {
let ext = ExtensionOids {
domains: std::collections::HashMap::from([(50_002, OID_INT8)]),
..Default::default()
};
assert_eq!(
decode_value(50_002, &7i64.to_be_bytes(), &ext).unwrap(),
DecodedValue::I64(7)
);
}
#[test]
fn discovery_rows_populate_vector_enums_and_domains() {
let ext = ExtensionOids::from_discovery_rows([
(50_000, "b".to_string(), 0, "vector".to_string()),
(50_001, "e".to_string(), 0, "status".to_string()),
(50_002, "d".to_string(), OID_INT8, "positive_int".to_string()),
(50_003, "d".to_string(), 0, "broken".to_string()),
]);
assert_eq!(ext.vector, Some(50_000));
assert!(ext.enums.contains(&50_001));
assert_eq!(ext.domains.get(&50_002), Some(&OID_INT8));
assert!(!ext.domains.contains_key(&50_003));
}
#[test]
fn discovery_rows_record_array_types() {
let ext = ExtensionOids::from_discovery_rows([
(50_001, "e".to_string(), 0, "status".to_string()),
(50_010, "A".to_string(), 50_001, "_status".to_string()),
]);
assert!(ext.enums.contains(&50_001));
assert!(ext.arrays.contains(&50_010));
}
#[test]
fn decodes_an_array_of_a_discovered_enum() {
let ext = ExtensionOids {
enums: std::collections::HashSet::from([50_001]),
arrays: std::collections::HashSet::from([50_010]),
..Default::default()
};
let encoded = encode_array(50_001, &[Some(b"Password"), Some(b"Passkey")]);
assert_eq!(
decode_value(50_010, &encoded, &ext).unwrap(),
DecodedValue::Array(vec![
DecodedValue::Str("Password".to_string()),
DecodedValue::Str("Passkey".to_string()),
])
);
}
#[test]
fn an_array_of_an_undiscovered_type_is_still_an_error() {
let ext = ExtensionOids::default();
let encoded = encode_array(50_001, &[Some(b"Password")]);
let err = decode_value(50_010, &encoded, &ext).unwrap_err();
assert!(matches!(err, Error::UnknownTypeOid { oid: 50_010 }), "got {err:?}");
}
#[test]
fn a_vector_inside_a_record_decodes_when_discovery_ran() {
let mut vec_bytes = 2u16.to_be_bytes().to_vec();
vec_bytes.extend_from_slice(&0u16.to_be_bytes());
vec_bytes.extend_from_slice(&1.5f32.to_be_bytes());
vec_bytes.extend_from_slice(&2.5f32.to_be_bytes());
let rec = encode_record(&[(OID_TEXT, Some(b"doc")), (50_000, Some(&vec_bytes))]);
let ext = ExtensionOids {
vector: Some(50_000),
..Default::default()
};
let decoded = decode_value(OID_RECORD, &rec, &ext).unwrap();
let DecodedValue::Composite(fields) = decoded else {
panic!("expected Composite, got {decoded:?}")
};
assert_eq!(fields[0], DecodedValue::Str("doc".to_string()));
assert_eq!(
fields[1],
DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
);
let err = decode_value(OID_RECORD, &rec, &no_ext()).unwrap_err();
assert!(matches!(err, Error::UnknownTypeOid { oid: 50_000 }), "got {err:?}");
}
fn encode_numeric(sign: u16, weight: i16, dscale: i16, digits: &[u16]) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(&(digits.len() as u16).to_be_bytes());
buf.extend_from_slice(&weight.to_be_bytes());
buf.extend_from_slice(&sign.to_be_bytes());
buf.extend_from_slice(&dscale.to_be_bytes());
for d in digits {
buf.extend_from_slice(&d.to_be_bytes());
}
buf
}
#[test]
fn decodes_numeric_integer() {
let data = encode_numeric(0x0000, 1, 0, &[1, 2345]);
assert_eq!(
decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
DecodedValue::Decimal("12345".to_string())
);
}
#[test]
fn decodes_numeric_with_fraction() {
let data = encode_numeric(0x0000, 0, 2, &[12, 5000]);
assert_eq!(
decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
DecodedValue::Decimal("12.50".to_string())
);
}
#[test]
fn decodes_negative_numeric() {
let data = encode_numeric(0x4000, 0, 2, &[12, 5000]);
assert_eq!(
decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
DecodedValue::Decimal("-12.50".to_string())
);
}
#[test]
fn refuses_a_string_for_a_parameter_whose_binary_form_is_not_text() {
for ty in [
Type::INT8,
Type::INT4,
Type::BOOL,
Type::INTERVAL,
Type::TIMESTAMPTZ,
Type::DATE,
] {
let mut buffer = bytes::BytesMut::new();
let Err(err) = encode_value(&DecodedValue::Str("25 days".into()), &ty, &mut buffer) else {
panic!("{} must refuse a string", ty.name());
};
assert!(
err.to_string().contains(ty.name()),
"the message must name the type that was wanted: {err}"
);
assert!(buffer.is_empty(), "nothing may be written for a refused parameter");
}
}
#[test]
fn a_string_still_reaches_the_types_it_is_the_wire_form_of() {
for ty in [
Type::TEXT,
Type::VARCHAR,
Type::BPCHAR,
Type::NAME,
Type::JSON,
Type::BYTEA,
Type::UNKNOWN,
] {
let mut buffer = bytes::BytesMut::new();
encode_value(&DecodedValue::Str("hello".into()), &ty, &mut buffer)
.unwrap_or_else(|e| panic!("{} must take a string: {e}", ty.name()));
assert_eq!(&buffer[..], b"hello", "{}", ty.name());
}
}
#[test]
fn a_string_still_converts_for_uuid_jsonb_and_numeric() {
let mut buffer = bytes::BytesMut::new();
encode_value(
&DecodedValue::Str("00000000-0000-0000-0000-000000000001".into()),
&Type::UUID,
&mut buffer,
)
.expect("uuid takes a string");
assert_eq!(buffer.len(), 16, "a uuid is converted, not copied");
let mut buffer = bytes::BytesMut::new();
encode_value(&DecodedValue::Str(r#"{"a":1}"#.into()), &Type::JSONB, &mut buffer).expect("jsonb takes a string");
assert_eq!(buffer[0], 1, "jsonb needs its version byte");
let mut buffer = bytes::BytesMut::new();
encode_value(&DecodedValue::Str("12.50".into()), &Type::NUMERIC, &mut buffer).expect("numeric takes a string");
assert_eq!(numeric::decode(&buffer).unwrap(), "12.50");
}
#[test]
fn decodes_jsonb_object() {
let mut data = vec![1u8]; data.extend_from_slice(br#"{"a":1,"b":"two","c":[1,2,3],"d":null}"#);
let decoded = decode_value(OID_JSONB, &data, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Object(vec![
("a".into(), DecodedValue::I64(1)),
("b".into(), DecodedValue::Str("two".into())),
(
"c".into(),
DecodedValue::Array(vec![DecodedValue::I64(1), DecodedValue::I64(2), DecodedValue::I64(3)])
),
("d".into(), DecodedValue::Null),
])
);
}
#[test]
fn decodes_jsonb_scalar_and_array() {
let mut data = vec![1u8];
data.extend_from_slice(b"42");
assert_eq!(
decode_value(OID_JSONB, &data, &no_ext()).unwrap(),
DecodedValue::I64(42)
);
let mut data2 = vec![1u8];
data2.extend_from_slice(b"[1.5, 2.5]");
assert_eq!(
decode_value(OID_JSONB, &data2, &no_ext()).unwrap(),
DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
);
}
fn encode_record(fields: &[(u32, Option<&[u8]>)]) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(&(fields.len() as i32).to_be_bytes());
for (oid, data) in fields {
buf.extend_from_slice(&oid.to_be_bytes());
match data {
None => buf.extend_from_slice(&(-1i32).to_be_bytes()),
Some(bytes) => {
buf.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
buf.extend_from_slice(bytes);
}
}
}
buf
}
#[test]
fn decodes_flat_record() {
let data = encode_record(&[
(OID_INT8, Some(&42i64.to_be_bytes())),
(OID_TEXT, Some(b"alice")),
(OID_BOOL, None),
]);
let decoded = decode_value(OID_RECORD, &data, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Composite(vec![
DecodedValue::I64(42),
DecodedValue::Str("alice".into()),
DecodedValue::Null
])
);
}
#[test]
fn decodes_nested_record() {
let inner = encode_record(&[(OID_INT8, Some(&1i64.to_be_bytes()))]);
let outer = encode_record(&[(OID_RECORD, Some(&inner)), (OID_TEXT, Some(b"outer"))]);
let decoded = decode_value(OID_RECORD, &outer, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Composite(vec![
DecodedValue::Composite(vec![DecodedValue::I64(1)]),
DecodedValue::Str("outer".into()),
])
);
}
fn encode_array(element_oid: u32, elements: &[Option<&[u8]>]) -> Vec<u8> {
if elements.is_empty() {
let mut buf = Vec::new();
buf.extend_from_slice(&0i32.to_be_bytes());
buf.extend_from_slice(&0i32.to_be_bytes());
buf.extend_from_slice(&element_oid.to_be_bytes());
return buf;
}
let mut buf = Vec::new();
buf.extend_from_slice(&1i32.to_be_bytes());
buf.extend_from_slice(&0i32.to_be_bytes());
buf.extend_from_slice(&element_oid.to_be_bytes());
buf.extend_from_slice(&(elements.len() as i32).to_be_bytes());
buf.extend_from_slice(&1i32.to_be_bytes());
for data in elements {
match data {
None => buf.extend_from_slice(&(-1i32).to_be_bytes()),
Some(bytes) => {
buf.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
buf.extend_from_slice(bytes);
}
}
}
buf
}
#[test]
fn decodes_array_of_scalars() {
let data = encode_array(OID_TEXT, &[Some(b"a"), Some(b"b"), None]);
let decoded = decode_value(OID_TEXT_ARRAY, &data, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Array(vec![
DecodedValue::Str("a".into()),
DecodedValue::Str("b".into()),
DecodedValue::Null
])
);
}
#[test]
fn decodes_empty_array() {
let data = encode_array(OID_TEXT, &[]);
assert_eq!(
decode_value(OID_TEXT_ARRAY, &data, &no_ext()).unwrap(),
DecodedValue::Array(vec![])
);
}
#[test]
fn decodes_array_of_records() {
let rec1 = encode_record(&[(OID_INT8, Some(&1i64.to_be_bytes()))]);
let rec2 = encode_record(&[(OID_INT8, Some(&2i64.to_be_bytes()))]);
let data = encode_array(OID_RECORD, &[Some(&rec1), Some(&rec2)]);
let decoded = decode_value(OID_RECORD_ARRAY, &data, &no_ext()).unwrap();
assert_eq!(
decoded,
DecodedValue::Array(vec![
DecodedValue::Composite(vec![DecodedValue::I64(1)]),
DecodedValue::Composite(vec![DecodedValue::I64(2)]),
])
);
}
#[test]
fn decodes_vector_when_extension_oid_known() {
let mut data = 2u16.to_be_bytes().to_vec(); data.extend_from_slice(&0u16.to_be_bytes()); data.extend_from_slice(&1.5f32.to_be_bytes());
data.extend_from_slice(&2.5f32.to_be_bytes());
let ext = ExtensionOids {
vector: Some(50_000),
..Default::default()
};
let decoded = decode_value(50_000, &data, &ext).unwrap();
assert_eq!(
decoded,
DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
);
}
#[test]
fn unknown_oid_without_vector_extension_is_an_error_not_a_guess() {
let err = decode_value(50_000, "some-domain-value".as_bytes(), &no_ext()).unwrap_err();
assert!(matches!(err, Error::UnknownTypeOid { oid: 50_000 }), "got {err:?}");
}
#[test]
fn encodes_a_str_value_as_uuid_binary_when_the_target_type_is_uuid() {
let value = DecodedValue::Str("11111111-2222-3333-4444-555555555555".to_string());
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::UUID, &mut out).unwrap();
assert_eq!(
out.as_ref(),
&[
0x11, 0x11, 0x11, 0x11, 0x22, 0x22, 0x33, 0x33, 0x44, 0x44, 0x55, 0x55, 0x55, 0x55, 0x55, 0x55
]
);
}
#[test]
fn a_str_value_still_encodes_as_plain_text_for_a_text_target() {
let value = DecodedValue::Str("11111111-2222-3333-4444-555555555555".to_string());
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::TEXT, &mut out).unwrap();
assert_eq!(out.as_ref(), "11111111-2222-3333-4444-555555555555".as_bytes());
}
#[test]
fn rejects_a_malformed_uuid_string_instead_of_sending_garbage_bytes() {
let value = DecodedValue::Str("not-a-uuid".to_string());
let mut out = bytes::BytesMut::new();
assert!(encode_value(&value, &postgres_types::Type::UUID, &mut out).is_err());
}
fn vector_type() -> Type {
Type::new(
"vector".to_string(),
50_000,
postgres_types::Kind::Simple,
"public".to_string(),
)
}
#[test]
fn encodes_an_array_value_as_pgvector_binary_when_the_target_type_is_vector() {
let value = DecodedValue::Array(vec![
DecodedValue::F64(1.5),
DecodedValue::F64(-2.25),
DecodedValue::F64(0.0),
]);
let mut out = bytes::BytesMut::new();
encode_value(&value, &vector_type(), &mut out).unwrap();
let mut expected = vec![0u8, 3, 0, 0];
expected.extend_from_slice(&1.5f32.to_be_bytes());
expected.extend_from_slice(&(-2.25f32).to_be_bytes());
expected.extend_from_slice(&0.0f32.to_be_bytes());
assert_eq!(out.as_ref(), expected.as_slice());
}
#[test]
fn a_vector_encoded_value_round_trips_through_decode_vector() {
let value = DecodedValue::Array(vec![
DecodedValue::F64(1.0),
DecodedValue::F64(2.0),
DecodedValue::F64(3.0),
]);
let mut out = bytes::BytesMut::new();
encode_value(&value, &vector_type(), &mut out).unwrap();
let decoded = decode_vector(out.as_ref()).unwrap();
assert_eq!(
decoded,
vec![DecodedValue::F64(1.0), DecodedValue::F64(2.0), DecodedValue::F64(3.0)]
);
}
#[test]
fn an_array_value_still_encodes_as_a_plain_postgres_array_for_a_non_vector_target() {
let value = DecodedValue::Array(vec![DecodedValue::F64(1.0), DecodedValue::F64(2.0)]);
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::FLOAT8_ARRAY, &mut out).unwrap();
assert_eq!(&out.as_ref()[0..4], &1i32.to_be_bytes());
}
#[test]
fn encodes_a_str_value_as_jsonb_binary_when_the_target_type_is_jsonb() {
let value = DecodedValue::Str(r#"{"a":1}"#.to_string());
let mut out = bytes::BytesMut::new();
encode_value(&value, &postgres_types::Type::JSONB, &mut out).unwrap();
assert_eq!(out.as_ref(), [&[1u8][..], br#"{"a":1}"#].concat());
assert_eq!(
decode_value(OID_JSONB, &out, &no_ext()).unwrap(),
DecodedValue::Object(vec![("a".into(), DecodedValue::I64(1))])
);
}
}