#![cfg_attr(
not(test),
deny(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::unreachable,
clippy::todo,
clippy::unimplemented,
clippy::indexing_slicing,
clippy::string_slice,
clippy::arithmetic_side_effects,
)
)]
use bytes::{Bytes, BytesMut};
pub const MAX_FRAME_ACCUM: usize = 8 * 1024 * 1024;
pub const TAG_UNTAGGED: u8 = 0;
pub const TAG_READY_FOR_QUERY: u8 = b'Z';
pub const TAG_ERROR_RESPONSE: u8 = b'E';
pub const TAG_PARAMETER_STATUS: u8 = b'S';
pub const TAG_BACKEND_KEY_DATA: u8 = b'K';
#[allow(
dead_code,
reason = "tag-table completeness: the recorder matches RowDescription frames by \
position inside an exchange rather than by tag, so only the builder and \
this module's tests name it"
)]
pub const TAG_ROW_DESCRIPTION: u8 = b'T';
pub const TAG_DATA_ROW: u8 = b'D';
pub const TAG_COMMAND_COMPLETE: u8 = b'C';
pub const TAG_COPY_IN_RESPONSE: u8 = b'G';
pub const TAG_COPY_OUT_RESPONSE: u8 = b'H';
pub const TAG_COPY_BOTH_RESPONSE: u8 = b'W';
pub const SSL_REQUEST_CODE: u32 = 80_877_103;
pub const GSSENC_REQUEST_CODE: u32 = 80_877_104;
pub const CANCEL_REQUEST_CODE: u32 = 80_877_102;
#[allow(
dead_code,
reason = "startup-code table completeness: the recorder only needs to tell a startup \
packet from the SSL/GSS/cancel codes, so the version itself is named by \
this module's tests"
)]
pub const PROTOCOL_VERSION_3: u32 = 196_608;
pub const MARKER_GUC: &str = "autumn.capsule_request";
const MARKER_GUC_UPPER: &str = "AUTUMN.CAPSULE_REQUEST";
const MIN_TAGGED_LEN: usize = 4;
const MIN_UNTAGGED_LEN: usize = 8;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Direction {
Frontend,
Backend,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Frame {
pub tag: u8,
pub bytes: Bytes,
}
impl Frame {
pub fn payload(&self) -> &[u8] {
let skip = if self.tag == TAG_UNTAGGED { 4 } else { 5 };
self.bytes.get(skip..).unwrap_or(&[])
}
pub const fn is_untagged(&self) -> bool {
self.tag == TAG_UNTAGGED
}
#[allow(
dead_code,
reason = "the splitter reacts to an SSL answer where it is read (an `S` marks the \
connection unrecordable) rather than through this accessor, which \
documents the frame shape and is exercised by this module's tests"
)]
pub fn ssl_answer(&self) -> Option<u8> {
if self.is_untagged() && self.bytes.len() == 1 {
self.bytes.first().copied()
} else {
None
}
}
pub fn startup_code(&self) -> Option<u32> {
if !self.is_untagged() {
return None;
}
be_u32(self.payload())
}
}
#[derive(Debug)]
pub struct FrameSplitter {
direction: Direction,
phase: Phase,
buf: BytesMut,
unrecordable: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Phase {
Untagged,
BackendFirstByte,
Tagged,
}
enum Step {
Frame(Frame),
Need,
Fail,
}
impl FrameSplitter {
pub fn new_frontend() -> Self {
Self {
direction: Direction::Frontend,
phase: Phase::Untagged,
buf: BytesMut::new(),
unrecordable: false,
}
}
pub fn new_backend() -> Self {
Self {
direction: Direction::Backend,
phase: Phase::BackendFirstByte,
buf: BytesMut::new(),
unrecordable: false,
}
}
#[allow(
dead_code,
reason = "the recorder owns one splitter per direction and knows which is which \
from the field it reads it out of; the accessor keeps the type \
self-describing and is exercised by this module's tests"
)]
pub const fn direction(&self) -> Direction {
self.direction
}
pub const fn is_unrecordable(&self) -> bool {
self.unrecordable
}
#[allow(
dead_code,
reason = "the accumulator bound is enforced inside `push`, so production code never \
asks; the partial-frame tests assert on it"
)]
pub fn buffered(&self) -> usize {
self.buf.len()
}
pub fn mark_unrecordable(&mut self) {
self.unrecordable = true;
self.buf = BytesMut::new();
}
pub fn push(&mut self, bytes: &[u8]) -> Vec<Frame> {
let mut frames = Vec::new();
if self.unrecordable {
return frames;
}
self.buf.extend_from_slice(bytes);
while !self.unrecordable {
match self.next_frame() {
Step::Frame(frame) => frames.push(frame),
Step::Need => break,
Step::Fail => self.mark_unrecordable(),
}
}
if self.buf.len() > MAX_FRAME_ACCUM {
self.mark_unrecordable();
}
frames
}
fn next_frame(&mut self) -> Step {
match self.phase {
Phase::Untagged => self.next_untagged(),
Phase::BackendFirstByte => self.next_backend_first_byte(),
Phase::Tagged => self.next_tagged(),
}
}
fn next_untagged(&mut self) -> Step {
let Some(len) = be_i32(&self.buf) else {
return Step::Need;
};
let Ok(total) = usize::try_from(len) else {
return Step::Fail;
};
if !(MIN_UNTAGGED_LEN..=MAX_FRAME_ACCUM).contains(&total) {
return Step::Fail;
}
if self.buf.len() < total {
return Step::Need;
}
let frame = Frame {
tag: TAG_UNTAGGED,
bytes: self.buf.split_to(total).freeze(),
};
let still_untagged = matches!(
frame.startup_code(),
Some(SSL_REQUEST_CODE | GSSENC_REQUEST_CODE | CANCEL_REQUEST_CODE)
);
if !still_untagged {
self.phase = Phase::Tagged;
}
Step::Frame(frame)
}
fn next_backend_first_byte(&mut self) -> Step {
let Some(&first) = self.buf.first() else {
return Step::Need;
};
if first == b'S' || first == b'N' {
let plausible_len = be_i32(self.buf.get(1..).unwrap_or_default()).is_some_and(|len| {
usize::try_from(len)
.is_ok_and(|len| (MIN_TAGGED_LEN..=MAX_FRAME_ACCUM).contains(&len))
});
let decidable = self.buf.len() >= 5 || self.buf.len() == 1;
if !decidable {
return Step::Need;
}
if !plausible_len {
let frame = Frame {
tag: TAG_UNTAGGED,
bytes: self.buf.split_to(1).freeze(),
};
self.phase = Phase::Tagged;
if first == b'S' {
self.mark_unrecordable();
}
return Step::Frame(frame);
}
}
self.phase = Phase::Tagged;
self.next_tagged()
}
fn next_tagged(&mut self) -> Step {
let Some(&tag) = self.buf.first() else {
return Step::Need;
};
let Some(len) = be_i32(self.buf.get(1..).unwrap_or_default()) else {
return Step::Need;
};
let Ok(len) = usize::try_from(len) else {
return Step::Fail;
};
if !(MIN_TAGGED_LEN..=MAX_FRAME_ACCUM).contains(&len) {
return Step::Fail;
}
let Some(total) = len.checked_add(1) else {
return Step::Fail;
};
if self.buf.len() < total {
return Step::Need;
}
let frame = Frame {
tag,
bytes: self.buf.split_to(total).freeze(),
};
if self.direction == Direction::Backend && is_copy_start(tag) {
self.mark_unrecordable();
}
Step::Frame(frame)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum FrontendMessage {
Startup,
SslRequest,
Parse {
name: String,
sql: String,
param_oids: Vec<u32>,
},
Bind {
portal: String,
statement: String,
params: Vec<Option<Vec<u8>>>,
},
Describe {
kind: u8,
name: String,
},
Execute,
Query(String),
Sync,
Flush,
Close {
kind: u8,
name: String,
},
Terminate,
Other(u8),
}
pub fn parse_frontend(frame: &Frame) -> FrontendMessage {
if frame.is_untagged() {
return match frame.startup_code() {
Some(SSL_REQUEST_CODE) => FrontendMessage::SslRequest,
_ => FrontendMessage::Startup,
};
}
let mut reader = Reader::new(frame.payload());
let parsed = match frame.tag {
b'P' => parse_parse(&mut reader),
b'B' => parse_bind(&mut reader),
b'D' => parse_describe(&mut reader),
b'Q' => reader.cstr().map(FrontendMessage::Query),
b'E' => Some(FrontendMessage::Execute),
b'S' => Some(FrontendMessage::Sync),
b'H' => Some(FrontendMessage::Flush),
b'C' => parse_close(&mut reader),
b'X' => Some(FrontendMessage::Terminate),
_ => None,
};
parsed.unwrap_or(FrontendMessage::Other(frame.tag))
}
fn parse_parse(reader: &mut Reader<'_>) -> Option<FrontendMessage> {
let name = reader.cstr()?;
let sql = reader.cstr()?;
let count = reader.count()?;
let mut param_oids = Vec::with_capacity(count.min(SANE_COUNT));
for _ in 0..count {
param_oids.push(reader.u32()?);
}
Some(FrontendMessage::Parse {
name,
sql,
param_oids,
})
}
fn parse_bind(reader: &mut Reader<'_>) -> Option<FrontendMessage> {
let portal = reader.cstr()?;
let statement = reader.cstr()?;
let formats = reader.count()?;
for _ in 0..formats {
reader.i16()?;
}
let count = reader.count()?;
let mut params = Vec::with_capacity(count.min(SANE_COUNT));
for _ in 0..count {
let len = reader.i32()?;
if len < 0 {
params.push(None);
} else {
let len = usize::try_from(len).ok()?;
params.push(Some(reader.take(len)?.to_vec()));
}
}
Some(FrontendMessage::Bind {
portal,
statement,
params,
})
}
fn parse_describe(reader: &mut Reader<'_>) -> Option<FrontendMessage> {
let kind = reader.u8()?;
let name = reader.cstr()?;
Some(FrontendMessage::Describe { kind, name })
}
fn parse_close(reader: &mut Reader<'_>) -> Option<FrontendMessage> {
let kind = reader.u8()?;
let name = reader.cstr()?;
Some(FrontendMessage::Close { kind, name })
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum MarkerId {
Set(String),
Clear,
Invalid,
}
pub fn marker_request_id(sql: &str) -> Option<MarkerId> {
let mut found = None;
for statement in split_statements(sql) {
if let Some(marker) = parse_marker_statement(statement) {
found = Some(marker);
}
}
found
}
const MAX_MARKER_ID_LEN: usize = 64;
pub fn is_valid_marker_id(id: &str) -> bool {
!id.is_empty()
&& id.len() <= MAX_MARKER_ID_LEN
&& id
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'-')
}
pub fn marker_set_sql(id: &str) -> Option<String> {
if id.is_empty() || is_valid_marker_id(id) {
Some(format!("SET {MARKER_GUC} = '{id}'"))
} else {
None
}
}
pub fn is_session_housekeeping(sql: &str) -> bool {
let mut saw_statement = false;
for statement in split_statements(sql) {
let statement = statement.trim();
if statement.is_empty() {
continue;
}
saw_statement = true;
if !is_housekeeping_statement(statement) {
return false;
}
}
saw_statement
}
fn is_housekeeping_statement(statement: &str) -> bool {
let statement = statement.trim().to_ascii_uppercase();
let Some(setting) = strip_keyword(&statement, "SET") else {
return false;
};
let setting = setting.trim_start();
if strip_keyword(setting, "LOCAL").is_some() {
return false;
}
[
"TIME ZONE",
"CLIENT_ENCODING",
"STATEMENT_TIMEOUT",
MARKER_GUC_UPPER,
]
.iter()
.any(|name| {
setting
.strip_prefix(name)
.is_some_and(|rest| !rest.chars().next().is_some_and(is_setting_name_char))
})
}
const fn is_setting_name_char(c: char) -> bool {
c.is_ascii_alphanumeric() || matches!(c, '_' | '$' | '.')
}
pub fn split_statements(sql: &str) -> Vec<&str> {
let mut out = Vec::new();
let mut start = 0usize;
let mut in_single = false;
let mut in_double = false;
for (index, byte) in sql.bytes().enumerate() {
match byte {
b'\'' if !in_double => in_single = !in_single,
b'"' if !in_single => in_double = !in_double,
b';' if !in_single && !in_double => {
if let Some(statement) = sql.get(start..index) {
out.push(statement);
}
start = index.saturating_add(1);
}
_ => {}
}
}
if let Some(statement) = sql.get(start..) {
out.push(statement);
}
out
}
fn parse_marker_statement(statement: &str) -> Option<MarkerId> {
let rest = strip_keyword(statement.trim_start(), "set")?;
let rest = strip_keyword(rest, "session")
.or_else(|| strip_keyword(rest, "local"))
.unwrap_or(rest);
let rest = strip_keyword_exact(rest, MARKER_GUC)?;
let rest = match rest.strip_prefix('=') {
Some(rest) => rest,
None => strip_keyword(rest, "to")?,
};
let Some(value) = single_quoted_literal(rest.trim_start()) else {
return Some(MarkerId::Invalid);
};
if value.is_empty() {
Some(MarkerId::Clear)
} else if is_valid_marker_id(&value) {
Some(MarkerId::Set(value))
} else {
Some(MarkerId::Invalid)
}
}
fn strip_keyword<'a>(input: &'a str, keyword: &str) -> Option<&'a str> {
let rest = strip_prefix_ignore_ascii_case(input, keyword)?;
if !rest.starts_with(char::is_whitespace) {
return None;
}
Some(rest.trim_start())
}
fn strip_keyword_exact<'a>(input: &'a str, keyword: &str) -> Option<&'a str> {
let rest = strip_prefix_ignore_ascii_case(input, keyword)?;
match rest.chars().next() {
None => Some(rest),
Some(next) if next.is_whitespace() || next == '=' => Some(rest.trim_start()),
Some(_) => None,
}
}
fn strip_prefix_ignore_ascii_case<'a>(input: &'a str, prefix: &str) -> Option<&'a str> {
if !input.get(..prefix.len())?.eq_ignore_ascii_case(prefix) {
return None;
}
input.get(prefix.len()..)
}
fn single_quoted_literal(input: &str) -> Option<String> {
let inner = input.strip_prefix('\'')?;
let mut value = String::new();
let mut chars = inner.char_indices();
while let Some((index, ch)) = chars.next() {
if ch != '\'' {
value.push(ch);
continue;
}
let after = inner.get(index.saturating_add(1)..).unwrap_or_default();
if after.starts_with('\'') {
value.push('\'');
chars.next();
continue;
}
return after.trim().is_empty().then_some(value);
}
None
}
pub fn is_catalog_sql(sql: &str) -> bool {
let trimmed = sql.trim();
DRIVER_TYPEINFO_QUERIES
.iter()
.any(|query| query.trim() == trimmed)
}
const DRIVER_TYPEINFO_QUERIES: &[&str] = &[
"SELECT t.typname, t.typtype, t.typelem, r.rngsubtype, t.typbasetype, n.nspname, t.typrelid
FROM pg_catalog.pg_type t
LEFT OUTER JOIN pg_catalog.pg_range r ON r.rngtypid = t.oid
INNER JOIN pg_catalog.pg_namespace n ON t.typnamespace = n.oid
WHERE t.oid = $1
",
"SELECT t.typname, t.typtype, t.typelem, NULL::OID, t.typbasetype, n.nspname, t.typrelid
FROM pg_catalog.pg_type t
INNER JOIN pg_catalog.pg_namespace n ON t.typnamespace = n.oid
WHERE t.oid = $1
",
"SELECT enumlabel
FROM pg_catalog.pg_enum
WHERE enumtypid = $1
ORDER BY enumsortorder
",
"SELECT enumlabel
FROM pg_catalog.pg_enum
WHERE enumtypid = $1
ORDER BY oid
",
"SELECT attname, atttypid
FROM pg_catalog.pg_attribute
WHERE attrelid = $1
AND NOT attisdropped
AND attnum > 0
ORDER BY attnum
",
];
pub const fn terminates_exchange(frame: &Frame) -> bool {
frame.tag == TAG_READY_FOR_QUERY
}
#[allow(
dead_code,
reason = "recording and replay both treat ReadyForQuery as an opaque terminator and \
replay its recorded bytes verbatim, so the status byte is only read by this \
module's tests; keeping it documents what the terminator carries"
)]
pub fn ready_for_query_state(frame: &Frame) -> Option<u8> {
if frame.tag != TAG_READY_FOR_QUERY {
return None;
}
frame.payload().first().copied()
}
pub const fn is_copy_start(tag: u8) -> bool {
matches!(
tag,
TAG_COPY_IN_RESPONSE | TAG_COPY_OUT_RESPONSE | TAG_COPY_BOTH_RESPONSE
)
}
#[allow(
dead_code,
reason = "the stub server replays a recorded handshake verbatim and otherwise sends a \
canned parameter set, so it never has to take one apart; the accessor is \
exercised by this module's tests and is the tool for reading one back"
)]
pub fn parameter_status_pair(frame: &Frame) -> Option<(String, String)> {
if frame.tag != TAG_PARAMETER_STATUS {
return None;
}
let mut reader = Reader::new(frame.payload());
Some((reader.cstr()?, reader.cstr()?))
}
pub fn error_response_fields(frame: &Frame) -> Option<(String, String)> {
if frame.tag != TAG_ERROR_RESPONSE {
return None;
}
let mut reader = Reader::new(frame.payload());
let mut code = None;
let mut message = None;
loop {
let field = reader.u8()?;
if field == 0 {
break;
}
let value = reader.cstr()?;
match field {
b'C' => code = Some(value),
b'M' => message = Some(value),
_ => {}
}
}
Some((code?, message?))
}
pub mod build {
fn frame(tag: u8, payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(payload.len().saturating_add(5));
out.push(tag);
let len = i32::try_from(payload.len().saturating_add(4)).unwrap_or(i32::MAX);
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(payload);
out
}
fn push_cstr(out: &mut Vec<u8>, text: &str) {
out.extend(text.bytes().filter(|byte| *byte != 0));
out.push(0);
}
fn push_count(out: &mut Vec<u8>, count: usize) {
let count = i16::try_from(count).unwrap_or(i16::MAX);
out.extend_from_slice(&count.to_be_bytes());
}
pub fn authentication_ok() -> Vec<u8> {
frame(b'R', &0i32.to_be_bytes())
}
pub fn parameter_status(key: &str, value: &str) -> Vec<u8> {
let mut payload = Vec::new();
push_cstr(&mut payload, key);
push_cstr(&mut payload, value);
frame(super::TAG_PARAMETER_STATUS, &payload)
}
pub fn backend_key_data(pid: i32, secret: i32) -> Vec<u8> {
let mut payload = pid.to_be_bytes().to_vec();
payload.extend_from_slice(&secret.to_be_bytes());
frame(super::TAG_BACKEND_KEY_DATA, &payload)
}
pub fn ready_for_query(state: u8) -> Vec<u8> {
frame(super::TAG_READY_FOR_QUERY, &[state])
}
pub fn command_complete(tag: &str) -> Vec<u8> {
let mut payload = Vec::new();
push_cstr(&mut payload, tag);
frame(super::TAG_COMMAND_COMPLETE, &payload)
}
pub fn parse_complete() -> Vec<u8> {
frame(b'1', &[])
}
pub fn bind_complete() -> Vec<u8> {
frame(b'2', &[])
}
#[allow(
dead_code,
reason = "the stub server answers from recorded response bytes rather than \
synthesising result frames, so the row builders serve this module's \
tests and any future synthetic tape"
)]
pub fn row_description(cols: &[(String, u32)]) -> Vec<u8> {
let mut payload = Vec::new();
push_count(&mut payload, cols.len());
for (name, type_oid) in cols {
push_cstr(&mut payload, name);
payload.extend_from_slice(&0i32.to_be_bytes()); payload.extend_from_slice(&0i16.to_be_bytes()); payload.extend_from_slice(&type_oid.to_be_bytes());
payload.extend_from_slice(&(-1i16).to_be_bytes()); payload.extend_from_slice(&(-1i32).to_be_bytes()); payload.extend_from_slice(&0i16.to_be_bytes()); }
frame(super::TAG_ROW_DESCRIPTION, &payload)
}
#[allow(
dead_code,
reason = "companion to `row_description`: recorded tapes carry their own result \
bytes, so this builds rows for tests and synthetic tapes"
)]
pub fn data_row(fields: &[Option<Vec<u8>>]) -> Vec<u8> {
let mut payload = Vec::new();
push_count(&mut payload, fields.len());
for field in fields {
match field {
Some(value) => {
let len = i32::try_from(value.len()).unwrap_or(i32::MAX);
payload.extend_from_slice(&len.to_be_bytes());
payload.extend_from_slice(value);
}
None => payload.extend_from_slice(&(-1i32).to_be_bytes()),
}
}
frame(super::TAG_DATA_ROW, &payload)
}
pub fn parameter_description(oids: &[u32]) -> Vec<u8> {
let mut payload = Vec::new();
push_count(&mut payload, oids.len());
for oid in oids {
payload.extend_from_slice(&oid.to_be_bytes());
}
frame(b't', &payload)
}
pub fn no_data() -> Vec<u8> {
frame(b'n', &[])
}
#[allow(
dead_code,
reason = "an empty query is recorded like any other exchange and replayed from its \
own bytes; the builder completes the backend vocabulary and is exercised \
by this module's tests"
)]
pub fn empty_query_response() -> Vec<u8> {
frame(b'I', &[])
}
pub fn error_response(code: &str, message: &str) -> Vec<u8> {
let mut payload = Vec::new();
payload.push(b'S');
push_cstr(&mut payload, "ERROR");
payload.push(b'V');
push_cstr(&mut payload, "ERROR");
payload.push(b'C');
push_cstr(&mut payload, code);
payload.push(b'M');
push_cstr(&mut payload, message);
payload.push(0);
frame(super::TAG_ERROR_RESPONSE, &payload)
}
}
const SANE_COUNT: usize = 64;
struct Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Reader<'a> {
const fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
fn take(&mut self, len: usize) -> Option<&'a [u8]> {
let (head, _) = self.buf.get(self.pos..)?.split_at_checked(len)?;
self.pos = self.pos.checked_add(len)?;
Some(head)
}
fn u8(&mut self) -> Option<u8> {
self.take(1)?.first().copied()
}
fn i16(&mut self) -> Option<i16> {
let head: [u8; 2] = self.take(2)?.try_into().ok()?;
Some(i16::from_be_bytes(head))
}
fn count(&mut self) -> Option<usize> {
usize::try_from(self.i16()?).ok()
}
fn i32(&mut self) -> Option<i32> {
let head: [u8; 4] = self.take(4)?.try_into().ok()?;
Some(i32::from_be_bytes(head))
}
fn u32(&mut self) -> Option<u32> {
let head: [u8; 4] = self.take(4)?.try_into().ok()?;
Some(u32::from_be_bytes(head))
}
fn cstr(&mut self) -> Option<String> {
let rest = self.buf.get(self.pos..)?;
let nul = rest.iter().position(|byte| *byte == 0)?;
let (text, _) = rest.split_at_checked(nul)?;
self.pos = self.pos.checked_add(nul)?.checked_add(1)?;
Some(String::from_utf8_lossy(text).into_owned())
}
}
fn be_u32(bytes: &[u8]) -> Option<u32> {
let head: [u8; 4] = bytes.get(..4)?.try_into().ok()?;
Some(u32::from_be_bytes(head))
}
fn be_i32(bytes: &[u8]) -> Option<i32> {
let head: [u8; 4] = bytes.get(..4)?.try_into().ok()?;
Some(i32::from_be_bytes(head))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_custom_guc_sharing_a_housekeeping_prefix_is_not_housekeeping() {
assert!(
!is_session_housekeeping("SET autumn.capsule_request_mode = 'audit'"),
"a longer setting name sharing the marker's prefix is the app's own"
);
assert!(
!is_session_housekeeping("SET statement_timeout_policy = 'strict'"),
"a longer setting name sharing statement_timeout's prefix is the app's own"
);
assert!(is_session_housekeeping("SET TIME ZONE 'UTC'"));
assert!(is_session_housekeeping("SET client_encoding TO 'UTF8'"));
assert!(is_session_housekeeping("SET statement_timeout = 5000"));
assert!(is_session_housekeeping(
"SET autumn.capsule_request = 'req-1'"
));
assert!(is_session_housekeeping("SET autumn.capsule_request = ''"));
}
fn tagged(tag: u8, payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(payload.len() + 5);
out.push(tag);
let len = i32::try_from(payload.len() + 4).unwrap();
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(payload);
out
}
fn cstr(s: &str) -> Vec<u8> {
let mut out = s.as_bytes().to_vec();
out.push(0);
out
}
fn untagged(payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(payload.len() + 4);
let len = i32::try_from(payload.len() + 4).unwrap();
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(payload);
out
}
fn startup_packet() -> Vec<u8> {
let mut payload = PROTOCOL_VERSION_3.to_be_bytes().to_vec();
payload.extend_from_slice(&cstr("user"));
payload.extend_from_slice(&cstr("postgres"));
payload.extend_from_slice(&cstr("database"));
payload.extend_from_slice(&cstr("autumn"));
payload.push(0);
untagged(&payload)
}
fn ssl_request() -> Vec<u8> {
untagged(&SSL_REQUEST_CODE.to_be_bytes())
}
fn parse_msg(name: &str, sql: &str, oids: &[u32]) -> Vec<u8> {
let mut p = cstr(name);
p.extend_from_slice(&cstr(sql));
p.extend_from_slice(&i16::try_from(oids.len()).unwrap().to_be_bytes());
for oid in oids {
p.extend_from_slice(&oid.to_be_bytes());
}
tagged(b'P', &p)
}
fn bind_msg(portal: &str, statement: &str, params: &[Option<&[u8]>]) -> Vec<u8> {
let mut p = cstr(portal);
p.extend_from_slice(&cstr(statement));
p.extend_from_slice(&1i16.to_be_bytes());
p.extend_from_slice(&1i16.to_be_bytes());
p.extend_from_slice(&i16::try_from(params.len()).unwrap().to_be_bytes());
for param in params {
match param {
Some(value) => {
p.extend_from_slice(&i32::try_from(value.len()).unwrap().to_be_bytes());
p.extend_from_slice(value);
}
None => p.extend_from_slice(&(-1i32).to_be_bytes()),
}
}
p.extend_from_slice(&1i16.to_be_bytes());
p.extend_from_slice(&0i16.to_be_bytes());
tagged(b'B', &p)
}
fn query_msg(sql: &str) -> Vec<u8> {
tagged(b'Q', &cstr(sql))
}
fn describe_msg(kind: u8, name: &str) -> Vec<u8> {
let mut p = vec![kind];
p.extend_from_slice(&cstr(name));
tagged(b'D', &p)
}
fn execute_msg(portal: &str) -> Vec<u8> {
let mut p = cstr(portal);
p.extend_from_slice(&0i32.to_be_bytes());
tagged(b'E', &p)
}
fn copy_response(tag: u8) -> Vec<u8> {
let mut p = vec![0u8];
p.extend_from_slice(&1i16.to_be_bytes());
p.extend_from_slice(&0i16.to_be_bytes());
tagged(tag, &p)
}
fn backend_fixture() -> Vec<u8> {
let mut s = Vec::new();
s.extend_from_slice(&build::authentication_ok());
s.extend_from_slice(&build::parameter_status("server_version", "16.2"));
s.extend_from_slice(&build::parameter_status("client_encoding", "UTF8"));
s.extend_from_slice(&build::backend_key_data(4242, 99));
s.extend_from_slice(&build::ready_for_query(b'I'));
s.extend_from_slice(&build::parse_complete());
s.extend_from_slice(&build::parameter_description(&[23]));
s.extend_from_slice(&build::row_description(&[
("id".to_owned(), 23),
("name".to_owned(), 25),
]));
s.extend_from_slice(&build::bind_complete());
s.extend_from_slice(&build::data_row(&[Some(b"1".to_vec()), None]));
s.extend_from_slice(&build::data_row(&[
Some(b"2".to_vec()),
Some("héllo".as_bytes().to_vec()),
]));
s.extend_from_slice(&build::command_complete("SELECT 2"));
s.extend_from_slice(&build::no_data());
s.extend_from_slice(&build::empty_query_response());
s.extend_from_slice(&build::error_response("58000", "autumn replay divergence"));
s.extend_from_slice(&build::ready_for_query(b'E'));
s
}
fn frontend_fixture() -> Vec<u8> {
let mut s = startup_packet();
s.extend_from_slice(&parse_msg(
"s1",
"SELECT id, name FROM users WHERE id = $1",
&[23],
));
s.extend_from_slice(&describe_msg(b'S', "s1"));
s.extend_from_slice(&bind_msg("", "s1", &[Some(&[0, 0, 0, 7]), None]));
s.extend_from_slice(&describe_msg(b'P', ""));
s.extend_from_slice(&execute_msg(""));
s.extend_from_slice(&tagged(b'S', &[]));
s.extend_from_slice(&query_msg("BEGIN"));
s.extend_from_slice(&tagged(b'X', &[]));
s
}
fn tags(frames: &[Frame]) -> Vec<u8> {
frames.iter().map(|f| f.tag).collect()
}
#[test]
fn splits_backend_frames_on_tag_and_length() {
let mut stream = build::authentication_ok();
stream.extend_from_slice(&build::parameter_status("client_encoding", "UTF8"));
stream.extend_from_slice(&build::backend_key_data(4242, 99));
stream.extend_from_slice(&build::ready_for_query(b'I'));
let mut splitter = FrameSplitter::new_backend();
let frames = splitter.push(&stream);
assert_eq!(tags(&frames), vec![b'R', b'S', b'K', b'Z']);
assert!(!splitter.is_unrecordable());
assert_eq!(splitter.buffered(), 0);
let rejoined: Vec<u8> = frames.iter().flat_map(|f| f.bytes.to_vec()).collect();
assert_eq!(rejoined, stream);
assert_eq!(
parameter_status_pair(frames.get(1).unwrap()),
Some(("client_encoding".to_owned(), "UTF8".to_owned()))
);
}
#[test]
fn partial_frame_is_buffered_until_complete() {
let mut stream = build::row_description(&[("id".to_owned(), 23)]);
stream.extend_from_slice(&build::data_row(&[Some(b"1".to_vec())]));
let mut splitter = FrameSplitter::new_backend();
assert!(splitter.push(stream.get(..1).unwrap()).is_empty());
assert!(splitter.push(stream.get(1..4).unwrap()).is_empty());
assert!(splitter.push(stream.get(4..8).unwrap()).is_empty());
assert!(!splitter.is_unrecordable());
assert!(splitter.buffered() > 0);
let frames = splitter.push(stream.get(8..).unwrap());
assert_eq!(tags(&frames), vec![TAG_ROW_DESCRIPTION, TAG_DATA_ROW]);
assert_eq!(splitter.buffered(), 0);
}
#[test]
fn frame_split_at_every_byte_boundary_yields_same_frames() {
for (label, stream, backend) in [
("backend", backend_fixture(), true),
("frontend", frontend_fixture(), false),
] {
let mut whole = if backend {
FrameSplitter::new_backend()
} else {
FrameSplitter::new_frontend()
};
let expected = whole.push(&stream);
assert!(
!expected.is_empty(),
"{label}: fixture produced no frames at all"
);
assert!(!whole.is_unrecordable(), "{label}: fixture unrecordable");
assert_eq!(whole.buffered(), 0, "{label}: fixture left a partial frame");
for split in 0..=stream.len() {
let mut splitter = if backend {
FrameSplitter::new_backend()
} else {
FrameSplitter::new_frontend()
};
let mut got = splitter.push(stream.get(..split).unwrap());
got.extend(splitter.push(stream.get(split..).unwrap()));
assert_eq!(got, expected, "{label}: mismatch splitting at {split}");
assert_eq!(
splitter.buffered(),
0,
"{label}: leftover bytes splitting at {split}"
);
}
let mut splitter = if backend {
FrameSplitter::new_backend()
} else {
FrameSplitter::new_frontend()
};
let mut got = Vec::new();
for byte in &stream {
got.extend(splitter.push(&[*byte]));
}
assert_eq!(
got, expected,
"{label}: mismatch feeding one byte at a time"
);
}
}
#[test]
fn parses_parse_message_name_sql_and_param_oids() {
let mut splitter = FrameSplitter::new_frontend();
let _ = splitter.push(&startup_packet());
let frames = splitter.push(&parse_msg("s3", "SELECT $1::int4, $2::text", &[23, 25]));
let frame = frames.first().expect("one Parse frame");
assert_eq!(
parse_frontend(frame),
FrontendMessage::Parse {
name: "s3".to_owned(),
sql: "SELECT $1::int4, $2::text".to_owned(),
param_oids: vec![23, 25],
}
);
}
#[test]
fn parses_bind_statement_name_and_parameter_values() {
let mut splitter = FrameSplitter::new_frontend();
let _ = splitter.push(&startup_packet());
let frames = splitter.push(&bind_msg(
"",
"s3",
&[Some(&[0, 0, 0, 42]), None, Some(b"")],
));
let frame = frames.first().expect("one Bind frame");
assert_eq!(
parse_frontend(frame),
FrontendMessage::Bind {
portal: String::new(),
statement: "s3".to_owned(),
params: vec![Some(vec![0, 0, 0, 42]), None, Some(Vec::new())],
}
);
}
#[test]
fn parses_simple_query_text() {
let mut splitter = FrameSplitter::new_frontend();
let _ = splitter.push(&startup_packet());
let frames = splitter.push(&query_msg("BEGIN; SELECT 1; COMMIT"));
let frame = frames.first().expect("one Query frame");
assert_eq!(
parse_frontend(frame),
FrontendMessage::Query("BEGIN; SELECT 1; COMMIT".to_owned())
);
}
#[test]
fn recognizes_ready_for_query_as_exchange_terminator() {
let mut splitter = FrameSplitter::new_backend();
let mut stream = build::command_complete("SELECT 1");
stream.extend_from_slice(&build::ready_for_query(b'T'));
let frames = splitter.push(&stream);
let complete = frames.first().expect("CommandComplete");
let ready = frames.get(1).expect("ReadyForQuery");
assert!(!terminates_exchange(complete));
assert!(terminates_exchange(ready));
assert_eq!(ready.tag, TAG_READY_FOR_QUERY);
assert_eq!(ready_for_query_state(ready), Some(b'T'));
assert_eq!(ready_for_query_state(complete), None);
}
#[test]
fn recognizes_capsule_marker_and_extracts_request_id() {
assert_eq!(
marker_request_id("SET autumn.capsule_request = 'req-4f2a_01'"),
Some(MarkerId::Set("req-4f2a_01".to_owned()))
);
assert_eq!(
marker_request_id(
"SET statement_timeout = 5000; SET autumn.capsule_request = 'abc123'"
),
Some(MarkerId::Set("abc123".to_owned()))
);
assert_eq!(
marker_request_id("set AUTUMN.CAPSULE_REQUEST to 'Zz9'"),
Some(MarkerId::Set("Zz9".to_owned()))
);
let mut splitter = FrameSplitter::new_frontend();
let _ = splitter.push(&startup_packet());
let frames = splitter.push(&query_msg(
"SET statement_timeout = 5000; SET autumn.capsule_request = 'r-1'",
));
let FrontendMessage::Query(sql) = parse_frontend(frames.first().expect("Query frame"))
else {
panic!("expected a simple Query message");
};
assert_eq!(
marker_request_id(&sql),
Some(MarkerId::Set("r-1".to_owned()))
);
assert_eq!(marker_request_id("SELECT * FROM users"), None);
assert_eq!(marker_request_id("SET statement_timeout = 5000"), None);
}
#[test]
fn empty_marker_clears_binding() {
assert_eq!(
marker_request_id("SET autumn.capsule_request = ''"),
Some(MarkerId::Clear)
);
assert_eq!(
marker_request_id("SET statement_timeout = 5000; SET autumn.capsule_request = ''"),
Some(MarkerId::Clear)
);
assert_eq!(
marker_request_id("SET autumn.capsule_request = 'a1'; SET autumn.capsule_request = ''"),
Some(MarkerId::Clear)
);
}
#[test]
fn rejects_unsafe_marker_id() {
assert_eq!(
marker_request_id("SET autumn.capsule_request = 'a''; DROP TABLE users; --'"),
Some(MarkerId::Invalid)
);
assert_eq!(
marker_request_id("SET autumn.capsule_request = 'has space'"),
Some(MarkerId::Invalid)
);
let long = "x".repeat(65);
assert_eq!(
marker_request_id(&format!("SET autumn.capsule_request = '{long}'")),
Some(MarkerId::Invalid)
);
assert!(is_valid_marker_id("abc-123_XYZ"));
assert!(is_valid_marker_id(&"x".repeat(64)));
assert!(!is_valid_marker_id(""));
assert!(!is_valid_marker_id(&"x".repeat(65)));
assert!(!is_valid_marker_id("a'b"));
assert!(!is_valid_marker_id("a b"));
assert!(!is_valid_marker_id("naïve"));
assert_eq!(
marker_set_sql("r-1").as_deref(),
Some("SET autumn.capsule_request = 'r-1'")
);
assert_eq!(
marker_set_sql("").as_deref(),
Some("SET autumn.capsule_request = ''")
);
assert_eq!(marker_set_sql("a'; DROP TABLE users; --"), None);
let sql = marker_set_sql("r-1").expect("valid id");
assert_eq!(
marker_request_id(&sql),
Some(MarkerId::Set("r-1".to_owned()))
);
}
#[test]
fn oversized_accumulator_marks_connection_unrecordable() {
let mut splitter = FrameSplitter::new_backend();
let mut header = vec![TAG_DATA_ROW];
header.extend_from_slice(&i32::try_from(MAX_FRAME_ACCUM + 1).unwrap().to_be_bytes());
let frames = splitter.push(&header);
assert!(frames.is_empty());
assert!(splitter.is_unrecordable());
assert_eq!(splitter.buffered(), 0);
assert!(splitter.push(&build::ready_for_query(b'I')).is_empty());
assert!(splitter.is_unrecordable());
let mut fe = FrameSplitter::new_frontend();
assert!(
fe.push(&i32::try_from(MAX_FRAME_ACCUM + 1).unwrap().to_be_bytes())
.is_empty()
);
assert!(fe.is_unrecordable());
}
#[test]
fn copy_in_response_marks_connection_unrecordable() {
for tag in [
TAG_COPY_IN_RESPONSE,
TAG_COPY_OUT_RESPONSE,
TAG_COPY_BOTH_RESPONSE,
] {
assert!(is_copy_start(tag), "tag {tag} should start a copy");
let mut splitter = FrameSplitter::new_backend();
let _ = splitter.push(&build::authentication_ok());
assert!(!splitter.is_unrecordable());
let frames = splitter.push(©_response(tag));
assert_eq!(tags(&frames), vec![tag]);
assert!(
splitter.is_unrecordable(),
"backend tag {tag} must mark the connection unrecordable"
);
}
let mut fe = FrameSplitter::new_frontend();
let _ = fe.push(&startup_packet());
let frames = fe.push(&tagged(b'H', &[]));
assert_eq!(tags(&frames), vec![b'H']);
assert_eq!(
parse_frontend(frames.first().unwrap()),
FrontendMessage::Flush
);
assert!(!fe.is_unrecordable());
}
#[test]
fn frontend_startup_phase_is_untagged_then_tagged_forever() {
let mut splitter = FrameSplitter::new_frontend();
let frames = splitter.push(&ssl_request());
let ssl = frames.first().expect("SSLRequest frame");
assert_eq!(ssl.tag, TAG_UNTAGGED);
assert_eq!(ssl.bytes.len(), 8);
assert_eq!(ssl.startup_code(), Some(SSL_REQUEST_CODE));
assert_eq!(parse_frontend(ssl), FrontendMessage::SslRequest);
let frames = splitter.push(&startup_packet());
let startup = frames.first().expect("startup frame");
assert_eq!(startup.tag, TAG_UNTAGGED);
assert_eq!(startup.startup_code(), Some(PROTOCOL_VERSION_3));
assert_eq!(parse_frontend(startup), FrontendMessage::Startup);
let frames = splitter.push(&parse_msg("", "SELECT 1", &[]));
assert_eq!(
parse_frontend(frames.first().expect("Parse frame")),
FrontendMessage::Parse {
name: String::new(),
sql: "SELECT 1".to_owned(),
param_oids: Vec::new(),
}
);
}
#[test]
fn bare_ssl_refusal_byte_is_its_own_backend_frame() {
let mut splitter = FrameSplitter::new_backend();
let frames = splitter.push(b"N");
let refusal = frames.first().expect("SSL refusal frame");
assert_eq!(refusal.tag, TAG_UNTAGGED);
assert_eq!(refusal.ssl_answer(), Some(b'N'));
assert!(!splitter.is_unrecordable());
let frames = splitter.push(&build::authentication_ok());
assert_eq!(tags(&frames), vec![b'R']);
let mut coalesced = b"N".to_vec();
coalesced.extend_from_slice(&build::authentication_ok());
coalesced.extend_from_slice(&build::ready_for_query(b'I'));
let mut splitter = FrameSplitter::new_backend();
let frames = splitter.push(&coalesced);
assert_eq!(tags(&frames), vec![TAG_UNTAGGED, b'R', b'Z']);
let mut splitter = FrameSplitter::new_backend();
let notice = tagged(b'N', &[b'S', b'W', b'A', b'R', b'N', 0, 0]);
let frames = splitter.push(¬ice);
let frame = frames.first().expect("NoticeResponse frame");
assert_eq!(frame.tag, b'N');
assert_eq!(frame.bytes.len(), notice.len());
let mut splitter = FrameSplitter::new_backend();
let _ = splitter.push(b"S");
assert!(splitter.is_unrecordable());
}
#[test]
fn builders_round_trip_through_the_splitter() {
let mut splitter = FrameSplitter::new_backend();
let frames = splitter.push(&backend_fixture());
assert_eq!(
tags(&frames),
vec![
b'R', b'S', b'S', b'K', b'Z', b'1', b't', b'T', b'2', b'D', b'D', b'C', b'n', b'I',
b'E', b'Z'
]
);
let auth = frames.first().unwrap();
assert_eq!(auth.bytes.as_ref(), &[b'R', 0, 0, 0, 8, 0, 0, 0, 0]);
let key = frames.get(3).unwrap();
assert_eq!(key.payload().len(), 8);
assert_eq!(be_u32(key.payload()), Some(4242));
assert_eq!(ready_for_query_state(frames.get(4).unwrap()), Some(b'I'));
let params = frames.get(6).unwrap();
assert_eq!(params.payload(), &[0, 1, 0, 0, 0, 23]);
let row_desc = frames.get(7).unwrap();
let payload = row_desc.payload();
assert_eq!(payload.get(..2), Some([0u8, 2].as_slice()));
assert_eq!(payload.get(2..5), Some(b"id\0".as_slice()));
assert_eq!(
payload.get(5..23),
Some(
[
0, 0, 0, 0, 0, 0, 0, 0, 0, 23, 255, 255, 255, 255, 255, 255, 0, 0
]
.as_slice()
)
);
let null_row = frames.get(9).unwrap();
assert_eq!(
null_row.payload(),
&[0, 2, 0, 0, 0, 1, b'1', 255, 255, 255, 255]
);
assert_eq!(frames.get(11).unwrap().payload(), b"SELECT 2\0");
assert!(frames.get(12).unwrap().payload().is_empty());
assert!(frames.get(13).unwrap().payload().is_empty());
assert_eq!(
error_response_fields(frames.get(14).unwrap()),
Some(("58000".to_owned(), "autumn replay divergence".to_owned()))
);
assert_eq!(error_response_fields(frames.get(11).unwrap()), None);
assert!(!splitter.is_unrecordable());
assert_eq!(splitter.buffered(), 0);
}
#[test]
fn driver_typeinfo_probes_are_catalog_sql() {
for query in DRIVER_TYPEINFO_QUERIES {
assert!(is_catalog_sql(query), "driver probe must match: {query:?}");
assert!(is_catalog_sql(query.trim()));
}
}
#[test]
fn application_catalog_reads_are_not_catalog_probes() {
assert!(!is_catalog_sql("SELECT * FROM pg_settings"));
assert!(!is_catalog_sql(
"SELECT column_name FROM information_schema.columns"
));
assert!(!is_catalog_sql(
"SELECT t.oid, t.typname FROM pg_catalog.pg_type t WHERE t.oid = $1"
));
assert!(!is_catalog_sql("SELECT id, name FROM users WHERE id = $1"));
}
#[test]
fn malformed_payloads_degrade_to_other_instead_of_panicking() {
let truncated = Frame {
tag: b'P',
bytes: Bytes::from_static(&[b'P', 0, 0, 0, 6, b'x']),
};
assert_eq!(parse_frontend(&truncated), FrontendMessage::Other(b'P'));
let empty_bind = Frame {
tag: b'B',
bytes: Bytes::from_static(&[b'B', 0, 0, 0, 4]),
};
assert_eq!(parse_frontend(&empty_bind), FrontendMessage::Other(b'B'));
let unknown = Frame {
tag: b'@',
bytes: Bytes::from_static(&[b'@', 0, 0, 0, 4]),
};
assert_eq!(parse_frontend(&unknown), FrontendMessage::Other(b'@'));
}
#[test]
fn short_length_field_marks_connection_unrecordable() {
let mut splitter = FrameSplitter::new_backend();
let frames = splitter.push(&[b'D', 0, 0, 0, 3, 0, 0, 0]);
assert!(frames.is_empty());
assert!(splitter.is_unrecordable());
}
#[test]
fn frontend_control_messages_are_recognized() {
let mut splitter = FrameSplitter::new_frontend();
let _ = splitter.push(&startup_packet());
let mut stream = describe_msg(b'S', "s1");
stream.extend_from_slice(&execute_msg("p1"));
stream.extend_from_slice(&tagged(b'S', &[]));
stream.extend_from_slice(&tagged(b'C', &[b'S', 0]));
stream.extend_from_slice(&tagged(b'X', &[]));
let frames = splitter.push(&stream);
let parsed: Vec<FrontendMessage> = frames.iter().map(parse_frontend).collect();
assert_eq!(
parsed,
vec![
FrontendMessage::Describe {
kind: b'S',
name: "s1".to_owned()
},
FrontendMessage::Execute,
FrontendMessage::Sync,
FrontendMessage::Close {
kind: b'S',
name: String::new()
},
FrontendMessage::Terminate,
]
);
}
}