use std::collections::HashMap;
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,
}
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 };
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> {
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, _ => {}
}
}
Ok(PgQueryResult {
column_names,
rows,
notices,
row_count,
})
}
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 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()
}