use std::collections::{HashMap, HashSet};
use std::io::{Read, Write};
use std::net::TcpStream;
pub enum Value {
String(String),
Null,
Bool(bool),
Integer(i64),
Float(f64),
Bytes(Vec<u8>),
}
pub struct Notice {
pub fields: HashMap<String, String>,
}
pub struct PgQueryResult {
pub column_names: Vec<String>,
pub rows: Vec<HashMap<String, Value>>,
pub notices: Vec<Notice>,
pub row_count: usize,
}
pub struct PgwireLite {
stream: TcpStream,
seen_notice_sigs: HashSet<String>,
}
impl PgwireLite {
pub fn new(host: &str, port: u16, _ssl: bool, _verbosity: &str) -> Result<Self, String> {
let addr = format!("{}:{}", host, port);
let stream = TcpStream::connect(&addr)
.map_err(|e| format!("Connection to {} failed: {}", addr, e))?;
let mut client = PgwireLite {
stream,
seen_notice_sigs: HashSet::new(),
};
client.startup()?;
Ok(client)
}
pub fn libpq_version(&self) -> String {
"pure-rust-pgwire-client".to_string()
}
fn startup(&mut self) -> Result<(), String> {
const PROTOCOL_V3: i32 = 196608;
let params = b"user\0stackql\0database\0stackql\0\0";
let total_len = 4 + 4 + params.len();
let mut msg = Vec::with_capacity(total_len);
msg.extend_from_slice(&(total_len as i32).to_be_bytes());
msg.extend_from_slice(&PROTOCOL_V3.to_be_bytes());
msg.extend_from_slice(params);
self.stream
.write_all(&msg)
.map_err(|e| format!("Startup write error: {}", e))?;
loop {
let msg_type = self.read_byte()?;
let payload_len = self.read_i32()? as usize;
let data = self.read_bytes(payload_len.saturating_sub(4))?;
match msg_type {
b'R' => {
let auth_type =
i32::from_be_bytes(data[..4].try_into().map_err(|_| "Bad auth")?);
if auth_type != 0 {
return Err(format!(
"Unsupported authentication type {} from server",
auth_type
));
}
}
b'K' => {} b'S' => {} b'Z' => break, b'E' => return Err(parse_error_fields(&data)),
b'N' => {} _ => {} }
}
Ok(())
}
pub fn query(&mut self, sql: &str) -> Result<PgQueryResult, String> {
self.drain_pending();
let sql_bytes = sql.as_bytes();
let payload_len = 4 + sql_bytes.len() + 1;
let mut msg = Vec::with_capacity(1 + payload_len);
msg.push(b'Q');
msg.extend_from_slice(&(payload_len as i32).to_be_bytes());
msg.extend_from_slice(sql_bytes);
msg.push(0u8);
self.stream
.write_all(&msg)
.map_err(|e| format!("Query write error: {}", e))?;
let mut column_names: Vec<String> = Vec::new();
let mut rows: Vec<HashMap<String, Value>> = Vec::new();
let mut notices: Vec<Notice> = Vec::new();
let mut row_count: usize = 0;
loop {
let msg_type = self.read_byte()?;
let payload_len = self.read_i32()? as usize;
let data = self.read_bytes(payload_len.saturating_sub(4))?;
match msg_type {
b'T' => {
column_names = parse_row_description(&data);
}
b'D' => {
let row = parse_data_row(&data, &column_names);
rows.push(row);
}
b'C' => {
let tag = std::str::from_utf8(data.strip_suffix(b"\0").unwrap_or(&data))
.unwrap_or("")
.to_string();
if let Some(n) = tag.split_whitespace().last().and_then(|s| s.parse().ok()) {
row_count = n;
}
}
b'N' => {
notices.push(parse_notice_fields(&data));
}
b'E' => {
let err_msg = parse_error_fields(&data);
loop {
let drain_type = self.read_byte()?;
let drain_len = self.read_i32()? as usize;
let _drain_data = self.read_bytes(drain_len.saturating_sub(4))?;
if drain_type == b'Z' {
break;
}
}
return Err(err_msg);
}
b'I' => {} b'Z' => break, _ => {}
}
}
let kept = filter_stale_notices(notices, &mut self.seen_notice_sigs);
Ok(PgQueryResult {
column_names,
rows,
notices: kept,
row_count,
})
}
fn drain_pending(&mut self) {
if self.stream.set_nonblocking(true).is_err() {
return;
}
let mut buf = [0u8; 4096];
loop {
match self.stream.read(&mut buf) {
Ok(0) => break, Ok(_) => continue, Err(_) => break, }
}
let _ = self.stream.set_nonblocking(false);
}
fn read_byte(&mut self) -> Result<u8, String> {
let mut buf = [0u8; 1];
self.stream
.read_exact(&mut buf)
.map_err(|e| format!("Read error: {}", e))?;
Ok(buf[0])
}
fn read_i32(&mut self) -> Result<i32, String> {
let mut buf = [0u8; 4];
self.stream
.read_exact(&mut buf)
.map_err(|e| format!("Read error: {}", e))?;
Ok(i32::from_be_bytes(buf))
}
fn read_bytes(&mut self, n: usize) -> Result<Vec<u8>, String> {
let mut buf = vec![0u8; n];
self.stream
.read_exact(&mut buf)
.map_err(|e| format!("Read error: {}", e))?;
Ok(buf)
}
}
fn canonical_line(line: &str) -> String {
line.trim().to_string()
}
fn filter_stale_notices(notices: Vec<Notice>, seen: &mut HashSet<String>) -> Vec<Notice> {
let mut kept = Vec::with_capacity(notices.len());
for n in notices {
let message_stale = n
.fields
.get("message")
.map(|m| seen.contains(&format!("M|{}", canonical_line(m))))
.unwrap_or(false);
let filtered_detail = n.fields.get("detail").map(|detail| {
detail
.lines()
.filter(|line| {
let c = canonical_line(line);
c.is_empty() || !seen.contains(&format!("D|{}", c))
})
.collect::<Vec<_>>()
.join("\n")
});
let has_detail_content = filtered_detail
.as_deref()
.map(|d| !d.trim().is_empty())
.unwrap_or(false);
if message_stale && !has_detail_content {
continue;
}
if let Some(m) = n.fields.get("message") {
let c = canonical_line(m);
if !c.is_empty() {
seen.insert(format!("M|{}", c));
}
}
if let Some(d) = filtered_detail.as_deref() {
for line in d.lines() {
let c = canonical_line(line);
if !c.is_empty() {
seen.insert(format!("D|{}", c));
}
}
}
let mut new_fields = n.fields.clone();
if let Some(d) = filtered_detail {
if d.trim().is_empty() {
new_fields.remove("detail");
} else {
new_fields.insert("detail".to_string(), d);
}
}
kept.push(Notice { fields: new_fields });
}
kept
}
fn parse_row_description(data: &[u8]) -> Vec<String> {
let mut names = Vec::new();
if data.len() < 2 {
return names;
}
let num_fields = u16::from_be_bytes([data[0], data[1]]) as usize;
let mut pos = 2;
for _ in 0..num_fields {
let Some(null_off) = data[pos..].iter().position(|&b| b == 0) else {
break;
};
let name = String::from_utf8_lossy(&data[pos..pos + null_off]).into_owned();
names.push(name);
pos += null_off + 1 + 18;
}
names
}
fn parse_data_row(data: &[u8], columns: &[String]) -> HashMap<String, Value> {
let mut row = HashMap::new();
if data.len() < 2 {
return row;
}
let num_cols = u16::from_be_bytes([data[0], data[1]]) as usize;
let mut pos = 2;
for col_name in columns.iter().take(num_cols) {
if pos + 4 > data.len() {
break;
}
let col_len = i32::from_be_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]);
pos += 4;
let value = if col_len < 0 {
Value::Null
} else {
let len = col_len as usize;
if pos + len > data.len() {
break;
}
let s = String::from_utf8_lossy(&data[pos..pos + len]).into_owned();
pos += len;
Value::String(s)
};
row.insert(col_name.clone(), value);
}
row
}
fn parse_notice_fields(data: &[u8]) -> Notice {
let mut fields = HashMap::new();
let mut pos = 0;
while pos < data.len() {
let field_code = data[pos];
pos += 1;
if field_code == 0 {
break;
}
let Some(null_off) = data[pos..].iter().position(|&b| b == 0) else {
break;
};
let value = String::from_utf8_lossy(&data[pos..pos + null_off]).into_owned();
pos += null_off + 1;
let key = match field_code {
b'S' => "severity",
b'M' => "message",
b'D' => "detail",
b'H' => "hint",
b'C' => "code",
b'P' => "position",
b'W' => "where",
_ => continue,
};
fields.insert(key.to_string(), value);
}
Notice { fields }
}
fn parse_error_fields(data: &[u8]) -> String {
let mut pos = 0;
while pos < data.len() {
let field_code = data[pos];
pos += 1;
if field_code == 0 {
break;
}
let Some(null_off) = data[pos..].iter().position(|&b| b == 0) else {
break;
};
let value = String::from_utf8_lossy(&data[pos..pos + null_off]).into_owned();
pos += null_off + 1;
if field_code == b'M' {
return value;
}
}
"Unknown server error".to_string()
}