use std::io::{Read, Write};
use crate::error::{Error, Result};
pub const PROTOCOL_MAGIC: u32 = 0x4E55_5341;
pub const PROTOCOL_MAJOR: u16 = 1;
pub const PROTOCOL_MINOR: u16 = 1;
pub const HEADER_LEN: usize = 5;
pub const MAX_FRAME_LEN: u32 = 256 * 1024 * 1024;
pub const SCRAM_MECHANISM: &str = "SCRAM-SHA-256";
#[allow(
dead_code,
reason = "complete protocol taxonomy; COPY bytes are reserved for COPY support"
)]
pub mod backend {
pub const AUTH: u8 = b'R';
pub const BACKEND_KEY: u8 = b'K';
pub const READY: u8 = b'Z';
pub const COMMAND_COMPLETE: u8 = b'C';
pub const ERROR: u8 = b'E';
pub const ROW_DESCRIPTION: u8 = b'T';
pub const ROW_DESCRIPTION_TYPED: u8 = b'y';
pub const DATA_ROW: u8 = b'D';
pub const PARSE_COMPLETE: u8 = b'1';
pub const BIND_COMPLETE: u8 = b'2';
pub const CLOSE_COMPLETE: u8 = b'3';
pub const PARAMETER_DESCRIPTION: u8 = b't';
pub const NO_DATA: u8 = b'n';
pub const PORTAL_SUSPENDED: u8 = b'z';
pub const COPY_IN: u8 = b'G';
pub const COPY_OUT: u8 = b'H';
pub const COPY_DATA: u8 = b'd';
pub const COPY_DONE: u8 = b'c';
pub const PARAMETER_STATUS: u8 = b'S';
pub const NOTIFICATION_RESPONSE: u8 = b'A';
}
#[allow(
dead_code,
reason = "complete protocol taxonomy; CLOSE is reserved for explicit statement close"
)]
pub mod frontend {
pub const STARTUP: u8 = b'S';
pub const QUERY: u8 = b'Q';
pub const PARSE: u8 = b'P';
pub const BIND: u8 = b'B';
pub const DESCRIBE: u8 = b'D';
pub const EXECUTE: u8 = b'E';
pub const SYNC: u8 = b'Y';
pub const CLOSE: u8 = b'C';
pub const SASL_INITIAL: u8 = b'p';
pub const SASL_RESPONSE: u8 = b'r';
pub const CANCEL: u8 = b'K';
pub const TERMINATE: u8 = b'X';
pub const COPY_DATA: u8 = b'd';
pub const COPY_DONE: u8 = b'c';
pub const COPY_FAIL: u8 = b'f';
}
pub mod auth {
pub const OK: u32 = 0;
pub const SASL: u32 = 10;
pub const SASL_CONTINUE: u32 = 11;
pub const SASL_FINAL: u32 = 12;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum TypeTag {
Unknown,
Bool,
Int,
Float,
Numeric,
Text,
Bytes,
Date,
Time,
TimeTz,
Timestamp,
TimestampTz,
Interval,
Uuid,
Json,
Array,
Vector,
}
impl TypeTag {
#[must_use]
pub const fn from_byte(tag: u8) -> Self {
match tag {
0x01 => Self::Bool,
0x02 => Self::Int,
0x03 => Self::Float,
0x04 => Self::Numeric,
0x05 => Self::Text,
0x06 => Self::Bytes,
0x07 => Self::Date,
0x08 => Self::Time,
0x09 => Self::TimeTz,
0x0A => Self::Timestamp,
0x0B => Self::TimestampTz,
0x0C => Self::Interval,
0x0D => Self::Uuid,
0x0E => Self::Json,
0x0F => Self::Array,
0x10 => Self::Vector,
_ => Self::Unknown,
}
}
#[must_use]
pub const fn name(self) -> &'static str {
match self {
Self::Unknown => "UNKNOWN",
Self::Bool => "BOOL",
Self::Int => "INT",
Self::Float => "FLOAT",
Self::Numeric => "NUMERIC",
Self::Text => "TEXT",
Self::Bytes => "BYTES",
Self::Date => "DATE",
Self::Time => "TIME",
Self::TimeTz => "TIMETZ",
Self::Timestamp => "TIMESTAMP",
Self::TimestampTz => "TIMESTAMPTZ",
Self::Interval => "INTERVAL",
Self::Uuid => "UUID",
Self::Json => "JSON",
Self::Array => "ARRAY",
Self::Vector => "VECTOR",
}
}
}
#[derive(Debug, Clone)]
pub struct Frame {
pub kind: u8,
pub payload: Vec<u8>,
}
#[derive(Default)]
pub struct Writer {
buf: Vec<u8>,
}
impl Writer {
#[must_use]
pub fn new() -> Self {
Self {
buf: Vec::with_capacity(32),
}
}
pub fn u8(&mut self, v: u8) -> &mut Self {
self.buf.push(v);
self
}
pub fn u16(&mut self, v: u16) -> &mut Self {
self.buf.extend_from_slice(&v.to_be_bytes());
self
}
pub fn u32(&mut self, v: u32) -> &mut Self {
self.buf.extend_from_slice(&v.to_be_bytes());
self
}
pub fn str(&mut self, s: &str) -> &mut Self {
self.u32(s.len() as u32);
self.buf.extend_from_slice(s.as_bytes());
self
}
pub fn raw(&mut self, bytes: &[u8]) -> &mut Self {
self.buf.extend_from_slice(bytes);
self
}
pub fn fields(&mut self, values: &[Option<Vec<u8>>]) -> &mut Self {
self.u16(values.len() as u16);
for v in values {
match v {
None => {
self.u8(0);
}
Some(bytes) => {
self.u8(1);
self.u32(bytes.len() as u32);
self.buf.extend_from_slice(bytes);
}
}
}
self
}
#[must_use]
pub fn frame(&self, kind: u8) -> Vec<u8> {
let total = (self.buf.len() + HEADER_LEN) as u32;
let mut out = Vec::with_capacity(self.buf.len() + HEADER_LEN);
out.push(kind);
out.extend_from_slice(&total.to_be_bytes());
out.extend_from_slice(&self.buf);
out
}
}
pub struct Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Reader<'a> {
#[must_use]
pub fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
#[must_use]
pub fn remaining(&self) -> usize {
self.buf.len().saturating_sub(self.pos)
}
fn take(&mut self, n: usize) -> Result<&'a [u8]> {
if self.pos + n > self.buf.len() {
return Err(Error::Protocol(format!(
"truncated payload: need {n} bytes, {} remain",
self.remaining()
)));
}
let slice = &self.buf[self.pos..self.pos + n];
self.pos += n;
Ok(slice)
}
pub fn u8(&mut self) -> Result<u8> {
Ok(self.take(1)?[0])
}
pub fn u16(&mut self) -> Result<u16> {
let b = self.take(2)?;
Ok(u16::from_be_bytes([b[0], b[1]]))
}
pub fn u32(&mut self) -> Result<u32> {
let b = self.take(4)?;
Ok(u32::from_be_bytes([b[0], b[1], b[2], b[3]]))
}
pub fn str(&mut self) -> Result<String> {
let n = self.u32()? as usize;
let bytes = self.take(n)?;
String::from_utf8(bytes.to_vec())
.map_err(|_| Error::Protocol("invalid UTF-8 in string field".to_owned()))
}
pub fn rest(&mut self) -> Result<&'a [u8]> {
let n = self.remaining();
self.take(n)
}
pub fn fields(&mut self) -> Result<Vec<Option<Vec<u8>>>> {
let count = self.u16()? as usize;
if count > self.remaining() + 1 {
return Err(Error::Protocol(format!(
"field count {count} exceeds remaining payload"
)));
}
let mut out = Vec::with_capacity(count);
for _ in 0..count {
if self.u8()? == 0 {
out.push(None);
} else {
let n = self.u32()? as usize;
out.push(Some(self.take(n)?.to_vec()));
}
}
Ok(out)
}
}
pub fn read_frame<R: Read>(r: &mut R) -> Result<Frame> {
let mut header = [0u8; HEADER_LEN];
r.read_exact(&mut header)?;
let kind = header[0];
let total = u32::from_be_bytes([header[1], header[2], header[3], header[4]]);
if total < HEADER_LEN as u32 {
return Err(Error::Protocol(format!("malformed frame: len {total} < 5")));
}
if total > MAX_FRAME_LEN {
return Err(Error::Protocol(format!(
"frame too large: {total} > {MAX_FRAME_LEN}"
)));
}
let payload_len = (total - HEADER_LEN as u32) as usize;
let mut payload = vec![0u8; payload_len];
r.read_exact(&mut payload)?;
Ok(Frame { kind, payload })
}
pub fn write_frame<W: Write>(w: &mut W, frame: &[u8]) -> Result<()> {
w.write_all(frame)?;
w.flush()?;
Ok(())
}