use bytes::Bytes;
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct ColumnInfo {
pub name_to_index: HashMap<String, usize>,
pub oids: Vec<u32>,
pub formats: Vec<i16>,
}
impl ColumnInfo {
pub fn from_fields(fields: &[crate::protocol::FieldDescription]) -> Self {
let mut name_to_index = HashMap::with_capacity(fields.len());
let mut oids = Vec::with_capacity(fields.len());
let mut formats = Vec::with_capacity(fields.len());
for (i, field) in fields.iter().enumerate() {
name_to_index.entry(field.name.clone()).or_insert(i);
oids.push(field.type_oid);
formats.push(field.format);
}
Self {
name_to_index,
oids,
formats,
}
}
}
pub struct PgRow {
pub columns: Vec<Option<Vec<u8>>>,
pub column_info: Option<Arc<ColumnInfo>>,
}
#[derive(Debug, Clone, Default)]
pub struct PgBytesRow {
pub(crate) payload: Bytes,
pub(crate) spans: Vec<Option<(usize, usize)>>,
pub column_info: Option<Arc<ColumnInfo>>,
}
#[derive(Debug)]
pub enum PgError {
Connection(String),
Protocol(String),
Auth(String),
Query(String),
QueryServer(PgServerError),
NoRows,
Io(std::io::Error),
Encode(String),
Timeout(String),
PoolExhausted {
max: usize,
},
PoolClosed,
}
pub(crate) const TLS_UNSUPPORTED_BY_SERVER: &str = "Server does not support TLS";
impl PgError {
pub(crate) fn tls_unsupported_by_server() -> Self {
PgError::Connection(TLS_UNSUPPORTED_BY_SERVER.to_string())
}
pub(crate) fn is_tls_unsupported_by_server(&self) -> bool {
matches!(self, PgError::Connection(msg) if msg == TLS_UNSUPPORTED_BY_SERVER)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PgServerError {
pub severity: String,
pub code: String,
pub message: String,
pub detail: Option<String>,
pub hint: Option<String>,
}
impl From<crate::protocol::ErrorFields> for PgServerError {
fn from(value: crate::protocol::ErrorFields) -> Self {
Self {
severity: value.severity,
code: value.code,
message: value.message,
detail: value.detail,
hint: value.hint,
}
}
}
impl std::fmt::Display for PgError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PgError::Connection(e) => write!(f, "Connection error: {}", e),
PgError::Protocol(e) => write!(f, "Protocol error: {}", e),
PgError::Auth(e) => write!(f, "Auth error: {}", e),
PgError::Query(e) => write!(f, "Query error: {}", e),
PgError::QueryServer(e) => write!(f, "Query error [{}]: {}", e.code, e.message),
PgError::NoRows => write!(f, "No rows returned"),
PgError::Io(e) => write!(f, "I/O error: {}", e),
PgError::Encode(e) => write!(f, "Encode error: {}", e),
PgError::Timeout(ctx) => write!(f, "Timeout: {}", ctx),
PgError::PoolExhausted { max } => write!(f, "Pool exhausted ({} max connections)", max),
PgError::PoolClosed => write!(f, "Connection pool is closed"),
}
}
}
impl std::error::Error for PgError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
PgError::Io(e) => Some(e),
_ => None,
}
}
}
impl From<std::io::Error> for PgError {
fn from(e: std::io::Error) -> Self {
PgError::Io(e)
}
}
impl From<crate::protocol::EncodeError> for PgError {
fn from(e: crate::protocol::EncodeError) -> Self {
PgError::Encode(e.to_string())
}
}
impl PgError {
pub fn server_error(&self) -> Option<&PgServerError> {
match self {
PgError::QueryServer(err) => Some(err),
_ => None,
}
}
pub fn sqlstate(&self) -> Option<&str> {
self.server_error().map(|e| e.code.as_str())
}
pub fn is_prepared_statement_retryable(&self) -> bool {
let Some(err) = self.server_error() else {
return false;
};
let code = err.code.as_str();
let message = err.message.to_ascii_lowercase();
if code.eq_ignore_ascii_case("26000")
&& message.contains("prepared statement")
&& message.contains("does not exist")
{
return true;
}
if code.eq_ignore_ascii_case("0A000") && message.contains("cached plan must be replanned") {
return true;
}
message.contains("cached plan must be replanned")
}
pub fn is_prepared_statement_already_exists(&self) -> bool {
let Some(err) = self.server_error() else {
return false;
};
if !err.code.eq_ignore_ascii_case("42P05") {
return false;
}
let message = err.message.to_ascii_lowercase();
message.contains("prepared statement") && message.contains("already exists")
}
pub fn is_transient_server_error(&self) -> bool {
match self {
PgError::Timeout(_) => return true,
PgError::Io(io) => {
return matches!(
io.kind(),
std::io::ErrorKind::TimedOut
| std::io::ErrorKind::ConnectionRefused
| std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::BrokenPipe
| std::io::ErrorKind::Interrupted
);
}
PgError::Connection(_) => return true,
_ => {}
}
if self.is_prepared_statement_retryable() {
return true;
}
let Some(code) = self.sqlstate() else {
return false;
};
matches!(
code,
"40001"
| "40P01"
| "57P03"
| "57P01"
| "57P02"
) || code.starts_with("08") }
}
#[cfg(test)]
mod tests {
use super::{ColumnInfo, PgError, TLS_UNSUPPORTED_BY_SERVER};
use crate::protocol::FieldDescription;
#[test]
fn tls_sentinel_matches_only_the_exact_message() {
assert!(PgError::tls_unsupported_by_server().is_tls_unsupported_by_server());
let prefixed = PgError::Connection(format!("connect failed: {TLS_UNSUPPORTED_BY_SERVER}"));
let suffixed = PgError::Connection(format!("{TLS_UNSUPPORTED_BY_SERVER}: retrying"));
let handshake = PgError::Connection("TLS handshake failed: bad cert".to_string());
assert!(!prefixed.is_tls_unsupported_by_server());
assert!(!suffixed.is_tls_unsupported_by_server());
assert!(!handshake.is_tls_unsupported_by_server());
assert!(
!PgError::Protocol(TLS_UNSUPPORTED_BY_SERVER.to_string())
.is_tls_unsupported_by_server()
);
}
fn field(name: &str, type_oid: u32) -> FieldDescription {
FieldDescription {
name: name.to_string(),
table_oid: 0,
column_attr: 0,
type_oid,
type_size: -1,
type_modifier: -1,
format: 0,
}
}
#[test]
fn column_info_preserves_first_duplicate_column_name() {
let info = ColumnInfo::from_fields(&[field("id", 23), field("id", 25)]);
assert_eq!(info.name_to_index.get("id").copied(), Some(0));
assert_eq!(info.oids, vec![23, 25]);
}
}
pub type PgResult<T> = Result<T, PgError>;
#[inline]
pub(crate) fn is_ignorable_session_message(msg: &crate::protocol::BackendMessage) -> bool {
matches!(
msg,
crate::protocol::BackendMessage::NoticeResponse(_)
| crate::protocol::BackendMessage::ParameterStatus { .. }
)
}
#[inline]
pub(crate) fn unexpected_backend_message(
phase: &str,
msg: &crate::protocol::BackendMessage,
) -> PgError {
PgError::Protocol(format!(
"Unexpected backend message during {} phase: {:?}",
phase, msg
))
}
#[inline]
pub(crate) fn is_ignorable_session_msg_type(msg_type: u8) -> bool {
matches!(msg_type, b'N' | b'S')
}
#[inline]
pub(crate) fn unexpected_backend_msg_type(phase: &str, msg_type: u8) -> PgError {
let printable = if msg_type.is_ascii_graphic() {
msg_type as char
} else {
'?'
};
PgError::Protocol(format!(
"Unexpected backend message type during {} phase: byte={} char={}",
phase, msg_type, printable
))
}
#[derive(Debug, Clone)]
pub struct QueryResult {
pub columns: Vec<String>,
pub rows: Vec<Vec<Option<String>>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ResultFormat {
#[default]
Text,
Binary,
}
impl ResultFormat {
#[inline]
pub(crate) fn as_wire_code(self) -> i16 {
match self {
ResultFormat::Text => crate::protocol::PgEncoder::FORMAT_TEXT,
ResultFormat::Binary => crate::protocol::PgEncoder::FORMAT_BINARY,
}
}
}