use super::error::{Error, Result};
use crate::types::Oid;
use std::borrow::Cow;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ParamFormat {
Text,
#[default]
Binary,
}
impl ParamFormat {
pub const BINARY_CODE: i16 = 1;
pub const TEXT_CODE: i16 = 0;
#[must_use]
pub fn to_code(self) -> i16 {
match self {
ParamFormat::Text => Self::TEXT_CODE,
ParamFormat::Binary => Self::BINARY_CODE,
}
}
#[must_use]
pub fn is_binary(self) -> bool {
matches!(self, ParamFormat::Binary)
}
}
const BROADCAST_BINARY: &[i16] = &[ParamFormat::BINARY_CODE];
pub(crate) fn bind_format_codes(
formats: &[ParamFormat],
param_count: usize,
) -> Result<Cow<'static, [i16]>> {
if !formats.is_empty() && formats.len() != param_count {
return Err(Error::protocol(format!(
"parameter format count ({}) does not match parameter count ({param_count})",
formats.len()
)));
}
if param_count == 0 {
return Ok(Cow::Borrowed(&[]));
}
if formats.is_empty() || formats.iter().copied().all(ParamFormat::is_binary) {
return Ok(Cow::Borrowed(BROADCAST_BINARY));
}
Ok(Cow::Owned(
formats.iter().copied().map(ParamFormat::to_code).collect(),
))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ColumnFormat {
#[default]
Text,
Binary,
HyperBinary,
}
impl ColumnFormat {
#[must_use]
pub fn from_code(code: i16) -> Self {
match code {
0 => ColumnFormat::Text,
1 => ColumnFormat::Binary,
2 => ColumnFormat::HyperBinary,
_ => ColumnFormat::Text, }
}
#[must_use]
pub fn to_code(self) -> i16 {
match self {
ColumnFormat::Text => 0,
ColumnFormat::Binary => 1,
ColumnFormat::HyperBinary => 2,
}
}
#[must_use]
pub fn is_binary(self) -> bool {
matches!(self, ColumnFormat::Binary | ColumnFormat::HyperBinary)
}
}
#[derive(Debug, Clone)]
pub struct Column {
pub(crate) name: String,
pub(crate) type_oid: Oid,
pub(crate) type_modifier: i32,
pub(crate) format: ColumnFormat,
}
impl Column {
#[inline]
pub(crate) fn new(
name: String,
type_oid: Oid,
type_modifier: i32,
format: ColumnFormat,
) -> Self {
Column {
name,
type_oid,
type_modifier,
format,
}
}
#[inline]
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[inline]
#[must_use]
pub fn type_oid(&self) -> Oid {
self.type_oid
}
#[inline]
#[must_use]
pub fn type_modifier(&self) -> i32 {
self.type_modifier
}
#[inline]
#[must_use]
pub fn format(&self) -> ColumnFormat {
self.format
}
}
#[cfg(test)]
mod tests {
use super::{ParamFormat, bind_format_codes};
#[test]
fn all_binary_collapses_to_one_broadcast_code() {
for n in 1..=8 {
let formats = vec![ParamFormat::Binary; n];
assert_eq!(
&*bind_format_codes(&formats, n).unwrap(),
&[1_i16],
"{n} binary params should ship one broadcast code"
);
}
}
#[test]
fn empty_formats_mean_all_binary_not_all_text() {
for n in 1..=8 {
assert_eq!(
&*bind_format_codes(&[], n).unwrap(),
&[1_i16],
"empty formats with {n} params must broadcast binary"
);
}
}
#[test]
fn zero_parameters_send_no_format_codes() {
assert!(bind_format_codes(&[], 0).unwrap().is_empty());
}
#[test]
fn format_count_must_match_parameter_count() {
for (formats, param_count) in [
(vec![ParamFormat::Text], 3),
(vec![ParamFormat::Text, ParamFormat::Binary], 3),
(vec![ParamFormat::Binary; 3], 2),
(vec![ParamFormat::Binary], 0),
] {
let err = bind_format_codes(&formats, param_count)
.expect_err("length mismatch must be rejected");
let msg = err.to_string();
assert!(
msg.contains("does not match parameter count"),
"unexpected error for {}/{param_count}: {msg}",
formats.len()
);
}
}
#[test]
fn mixed_formats_expand_to_one_code_per_parameter() {
assert_eq!(
&*bind_format_codes(&[ParamFormat::Binary, ParamFormat::Text], 2).unwrap(),
&[1_i16, 0]
);
assert_eq!(
&*bind_format_codes(&[ParamFormat::Text, ParamFormat::Binary], 2).unwrap(),
&[0_i16, 1]
);
assert_eq!(
&*bind_format_codes(&[ParamFormat::Text, ParamFormat::Text], 2).unwrap(),
&[0_i16, 0]
);
}
#[test]
fn param_format_codes_match_the_protocol() {
assert_eq!(ParamFormat::Text.to_code(), 0);
assert_eq!(ParamFormat::Binary.to_code(), 1);
assert!(ParamFormat::Binary.is_binary());
assert!(!ParamFormat::Text.is_binary());
assert_eq!(ParamFormat::default(), ParamFormat::Binary);
}
}