Skip to main content

ironwork_rt/sql/postgres/
wire.rs

1//! PostgreSQL's frontend/backend protocol, version 3: startup and authentication, the simple query
2//! cycle, and the extended one (Parse, Describe, Bind, Execute, Sync) with text values.
3
4use super::scram::{Scram, nonce};
5use std::io::{BufReader, Read, Write};
6use std::net::TcpStream;
7use std::path::{Path, PathBuf};
8
9/// A duplex byte stream: a socket, or TLS over one.
10pub trait Stream: Read + Write + Send {}
11
12impl<T: Read + Write + Send> Stream for T {}
13
14/// Wraps a connected socket in TLS, verifying the server's certificate chain against `roots` (a PEM
15/// file) or built-in roots, and its name as `host`. ironwork's own build has no implementation, so
16/// that it keeps no dependencies; the build in `tls/` supplies one.
17pub trait Tls: Send + Sync {
18    fn wrap(&self, socket: TcpStream, host: &str, roots: Option<&Path>) -> std::io::Result<Box<dyn Stream>>;
19}
20
21/// The URL's `sslmode`: the two of libpq's modes that ironwork offers.
22#[derive(Clone, Copy, Debug, PartialEq, Eq)]
23pub enum SslMode {
24    Disable,
25    /// TLS, with the server's certificate chain and name verified.
26    VerifyFull,
27}
28
29/// Where and as whom to connect.
30#[derive(Clone, Debug, PartialEq, Eq)]
31pub struct Target {
32    pub host: String,
33    pub port: u16,
34    pub user: String,
35    pub password: Option<String>,
36    pub database: String,
37    /// None where the URL does not say.
38    pub ssl: Option<SslMode>,
39    pub root_cert: Option<PathBuf>,
40}
41
42impl Target {
43    /// `postgres://user[:password]@host[:port]/database[?option=value&...]`, with the options
44    /// `host` (a directory, for a Unix socket), `sslmode` and `sslrootcert`. The password may come
45    /// from PGPASSWORD instead.
46    pub fn parse(url: &str) -> Result<Self, String> {
47        let rest = url.strip_prefix("postgres://").or_else(|| url.strip_prefix("postgresql://")).ok_or("a database URL starts postgres://")?;
48        let (rest, query) = rest.split_once('?').unwrap_or((rest, ""));
49        let (authority, database) = rest.split_once('/').unwrap_or((rest, ""));
50        let (userinfo, hostport) = authority.rsplit_once('@').map_or((None, authority), |(u, h)| (Some(u), h));
51        let (user, password) = match userinfo.map(|u| u.split_once(':').map_or((u, None), |(u, p)| (u, Some(p)))) {
52            Some((u, p)) => (Some(decode(u)?), p.map(decode).transpose()?),
53            None => (None, None),
54        };
55        let (mut host, port) = match hostport.rsplit_once(':') {
56            Some((h, p)) => (h.to_owned(), p.parse().map_err(|_| format!("{p} is not a port"))?),
57            None => (hostport.to_owned(), 5432),
58        };
59        let (mut ssl, mut root_cert) = (None, None);
60        for pair in query.split('&').filter(|p| !p.is_empty()) {
61            match pair.split_once('=') {
62                Some(("host", h)) => host = decode(h)?,
63                Some(("sslmode", "disable")) => ssl = Some(SslMode::Disable),
64                Some(("sslmode", "verify-full")) => ssl = Some(SslMode::VerifyFull),
65                Some(("sslmode", other)) => {
66                    return Err(format!("sslmode={other} is not offered: disable, or verify-full, which checks the server's certificate and name"));
67                }
68                Some(("sslrootcert", path)) => root_cert = Some(PathBuf::from(decode(path)?)),
69                _ => return Err(format!("{pair} is not a URL option ironwork reads")),
70            }
71        }
72        let user = match user {
73            Some(u) => u,
74            None => std::env::var("USER").map_err(|_| "the URL names no user, and USER is not set")?,
75        };
76        Ok(Self {
77            host: if host.is_empty() { "localhost".into() } else { decode(&host)? },
78            port,
79            database: if database.is_empty() { user.clone() } else { decode(database)? },
80            password: password.or_else(|| std::env::var("PGPASSWORD").ok()),
81            user,
82            ssl,
83            root_cert,
84        })
85    }
86}
87
88fn decode(text: &str) -> Result<String, String> {
89    let bytes = text.as_bytes();
90    let (mut out, mut i) = (Vec::new(), 0);
91    while i < bytes.len() {
92        if bytes[i] == b'%' {
93            let hex = text.get(i + 1..i + 3).and_then(|h| u8::from_str_radix(h, 16).ok()).ok_or("the URL has a bad % escape")?;
94            out.push(hex);
95            i += 3;
96        } else {
97            out.push(bytes[i]);
98            i += 1;
99        }
100    }
101    String::from_utf8(out).map_err(|_| "the URL is not UTF-8 once decoded".to_owned())
102}
103
104/// Why a request failed: the server refused it, naming an SQLSTATE, or the connection broke.
105#[derive(Debug, PartialEq, Eq)]
106pub enum Failure {
107    Refused { state: String, message: String },
108    Broken(String),
109}
110
111impl From<std::io::Error> for Failure {
112    fn from(e: std::io::Error) -> Self {
113        Failure::Broken(format!("the connection to PostgreSQL failed: {e}"))
114    }
115}
116
117/// What Execute returned: each row's columns as text, and the command tag.
118#[derive(Debug, Default)]
119pub struct Executed {
120    pub rows: Vec<Vec<Option<String>>>,
121    pub tag: String,
122}
123
124/// A prepared statement's parameter and column types.
125#[derive(Clone, Debug, Default)]
126pub struct Described {
127    pub parameters: Vec<u32>,
128    pub columns: Vec<Field>,
129}
130
131/// A result column as RowDescription gives it: its table and column number where it is a table's
132/// column, else 0, and its type and type modifier.
133#[derive(Clone, Debug, Default)]
134pub struct Field {
135    pub name: String,
136    pub table: u32,
137    pub attnum: i16,
138    pub oid: u32,
139    pub typmod: i32,
140}
141
142pub struct Connection {
143    stream: BufReader<Box<dyn Stream>>,
144    /// Messages waiting for the next flush.
145    out: Vec<u8>,
146    pub server_version: String,
147    pub encrypted: bool,
148    /// ReadyForQuery's transaction status: `I` idle, `T` in a transaction, `E` in a failed one.
149    pub status: u8,
150}
151
152struct Body<'a> {
153    bytes: &'a [u8],
154    at: usize,
155}
156
157impl<'a> Body<'a> {
158    fn take(&mut self, n: usize) -> Result<&'a [u8], Failure> {
159        let got = self.bytes.get(self.at..self.at + n).ok_or_else(|| Failure::Broken("PostgreSQL sent a short message".into()))?;
160        self.at += n;
161        Ok(got)
162    }
163    fn i16(&mut self) -> Result<i16, Failure> {
164        self.take(2).map(|b| i16::from_be_bytes([b[0], b[1]]))
165    }
166    fn i32(&mut self) -> Result<i32, Failure> {
167        self.take(4).map(|b| i32::from_be_bytes([b[0], b[1], b[2], b[3]]))
168    }
169    fn cstr(&mut self) -> Result<String, Failure> {
170        let end = self.bytes[self.at..].iter().position(|&b| b == 0).ok_or_else(|| Failure::Broken("PostgreSQL sent an unterminated string".into()))?;
171        let s = String::from_utf8_lossy(&self.bytes[self.at..self.at + end]).into_owned();
172        self.at += end + 1;
173        Ok(s)
174    }
175    fn rest(&mut self) -> &'a [u8] {
176        let r = &self.bytes[self.at..];
177        self.at = self.bytes.len();
178        r
179    }
180}
181
182fn cstr(out: &mut Vec<u8>, s: &str) {
183    out.extend_from_slice(s.as_bytes());
184    out.push(0);
185}
186
187/// An ErrorResponse's SQLSTATE and message.
188fn refusal(body: &[u8]) -> Failure {
189    let mut b = Body { bytes: body, at: 0 };
190    let (mut state, mut message) = (String::new(), String::new());
191    while let Ok(field) = b.take(1) {
192        if field[0] == 0 {
193            break;
194        }
195        let Ok(value) = b.cstr() else { break };
196        match field[0] {
197            b'C' => state = value,
198            b'M' => message = value,
199            _ => {}
200        }
201    }
202    Failure::Refused { state, message }
203}
204
205impl Connection {
206    /// Connects, over TLS where `sslmode` asks for it. Without an sslmode, a TCP connection uses TLS
207    /// when this build has it, and a Unix socket never does.
208    pub fn open(target: &Target, tls: Option<&dyn Tls>) -> Result<Self, String> {
209        let at = format!("PostgreSQL at {}:{}", target.host, target.port);
210        let broken = |e: std::io::Error| format!("cannot reach {at}: {e}");
211        let socket = target.host.starts_with('/');
212        let mode = match (target.ssl, tls) {
213            (Some(mode), _) => mode,
214            (None, Some(_)) if !socket => SslMode::VerifyFull,
215            (None, _) => SslMode::Disable,
216        };
217        let stream: Box<dyn Stream> = match (mode, socket, tls) {
218            (SslMode::Disable, true, _) => unix_socket(target).map_err(broken)?,
219            (SslMode::VerifyFull, true, _) => return Err("sslmode=verify-full is for TCP; a Unix socket needs no TLS".into()),
220            (SslMode::VerifyFull, false, None) => {
221                return Err("sslmode=verify-full needs TLS, which this build of ironwork leaves out to keep it free of dependencies; build tls/ for it".into());
222            }
223            (SslMode::Disable, false, _) => Box::new(tcp(target).map_err(broken)?),
224            (SslMode::VerifyFull, false, Some(tls)) => {
225                let mut socket = tcp(target).map_err(broken)?;
226                socket.write_all(&[0, 0, 0, 8, 0x04, 0xD2, 0x16, 0x2F]).map_err(broken)?;
227                let mut answer = [0u8];
228                socket.read_exact(&mut answer).map_err(broken)?;
229                if answer[0] != b'S' {
230                    return Err(format!("{at} does not offer TLS; give sslmode=disable to connect without it"));
231                }
232                tls.wrap(socket, &target.host, target.root_cert.as_deref()).map_err(|e| format!("TLS with {at} failed: {e}"))?
233            }
234        };
235        let mut conn = Self { stream: BufReader::new(stream), out: Vec::new(), server_version: String::new(), encrypted: mode == SslMode::VerifyFull, status: b'I' };
236        conn.start(target).map_err(|f| match f {
237            Failure::Refused { state, message } => format!("PostgreSQL refused the connection ({state}): {message}"),
238            Failure::Broken(m) => m,
239        })?;
240        Ok(conn)
241    }
242
243    fn start(&mut self, target: &Target) -> Result<(), Failure> {
244        let mut body = 196_608i32.to_be_bytes().to_vec();
245        let params = [
246            ("user", target.user.as_str()),
247            ("database", &target.database),
248            ("client_encoding", "UTF8"),
249            ("DateStyle", "ISO"),
250            ("TimeZone", "UTC"),
251            ("extra_float_digits", "3"),
252            ("application_name", "ironwork"),
253        ];
254        for (name, value) in params {
255            cstr(&mut body, name);
256            cstr(&mut body, value);
257        }
258        body.push(0);
259        self.out.extend_from_slice(&(body.len() as i32 + 4).to_be_bytes());
260        self.out.extend_from_slice(&body);
261        self.flush()?;
262        let password = || target.password.clone().ok_or_else(|| Failure::Broken("PostgreSQL asks for a password: give one in the URL or PGPASSWORD".into()));
263        let mut scram = None;
264        loop {
265            let (kind, body) = self.receive()?;
266            let mut b = Body { bytes: &body, at: 0 };
267            match kind {
268                b'R' => match b.i32()? {
269                    0 => {}
270                    3 => {
271                        let mut reply = Vec::new();
272                        cstr(&mut reply, &password()?);
273                        self.send(b'p', &reply)?;
274                        self.flush()?;
275                    }
276                    10 => {
277                        let mechanisms: Vec<String> = std::iter::from_fn(|| b.cstr().ok().filter(|m| !m.is_empty())).collect();
278                        if !mechanisms.iter().any(|m| m == "SCRAM-SHA-256") {
279                            return Err(Failure::Broken(format!("PostgreSQL offers SASL {mechanisms:?}, and ironwork speaks SCRAM-SHA-256")));
280                        }
281                        let mut exchange = Scram::new(&password()?, nonce());
282                        let first = exchange.client_first("");
283                        let mut reply = Vec::new();
284                        cstr(&mut reply, "SCRAM-SHA-256");
285                        reply.extend_from_slice(&(first.len() as i32).to_be_bytes());
286                        reply.extend_from_slice(first.as_bytes());
287                        self.send(b'p', &reply)?;
288                        self.flush()?;
289                        scram = Some(exchange);
290                    }
291                    11 => {
292                        let exchange = scram.as_mut().ok_or_else(|| Failure::Broken("PostgreSQL continued a SASL exchange that had not begun".into()))?;
293                        let reply = exchange.client_final(&String::from_utf8_lossy(b.rest())).map_err(Failure::Broken)?;
294                        self.send(b'p', reply.as_bytes())?;
295                        self.flush()?;
296                    }
297                    12 => {
298                        let exchange = scram.as_ref().ok_or_else(|| Failure::Broken("PostgreSQL ended a SASL exchange that had not begun".into()))?;
299                        exchange.verify(&String::from_utf8_lossy(b.rest())).map_err(Failure::Broken)?;
300                    }
301                    5 => return Err(Failure::Broken("PostgreSQL asks for MD5 authentication; ironwork speaks SCRAM-SHA-256".into())),
302                    other => return Err(Failure::Broken(format!("PostgreSQL asks for authentication method {other}, which ironwork does not speak"))),
303                },
304                b'Z' => {
305                    self.status = body.first().copied().unwrap_or(b'I');
306                    return Ok(());
307                }
308                b'E' => return Err(refusal(&body)),
309                _ => self.note(kind, &body),
310            }
311        }
312    }
313
314    fn send(&mut self, kind: u8, body: &[u8]) -> Result<(), Failure> {
315        self.out.push(kind);
316        self.out.extend_from_slice(&(body.len() as i32 + 4).to_be_bytes());
317        self.out.extend_from_slice(body);
318        Ok(())
319    }
320
321    fn flush(&mut self) -> Result<(), Failure> {
322        let stream = self.stream.get_mut();
323        stream.write_all(&self.out)?;
324        stream.flush()?;
325        self.out.clear();
326        Ok(())
327    }
328
329    fn receive(&mut self) -> Result<(u8, Vec<u8>), Failure> {
330        let mut head = [0u8; 5];
331        self.stream.read_exact(&mut head)?;
332        let len = i32::from_be_bytes([head[1], head[2], head[3], head[4]]);
333        let mut body = vec![0u8; (len.max(4) - 4) as usize];
334        self.stream.read_exact(&mut body)?;
335        Ok((head[0], body))
336    }
337
338    /// Messages that may arrive at any time: ParameterStatus, notices and notifications.
339    fn note(&mut self, kind: u8, body: &[u8]) {
340        if kind == b'S' {
341            let mut b = Body { bytes: body, at: 0 };
342            if let (Ok(name), Ok(value)) = (b.cstr(), b.cstr())
343                && name == "server_version"
344            {
345                self.server_version = value;
346            }
347        }
348    }
349
350    /// Reads to ReadyForQuery, handing each message to `each`, and fails with the first refusal.
351    fn until_ready(&mut self, mut each: impl FnMut(u8, &[u8]) -> Result<(), Failure>) -> Result<(), Failure> {
352        let mut refused = None;
353        loop {
354            let (kind, body) = self.receive()?;
355            match kind {
356                b'Z' => {
357                    self.status = body.first().copied().unwrap_or(b'I');
358                    return refused.map_or(Ok(()), Err);
359                }
360                b'E' => refused = refused.or(Some(refusal(&body))),
361                b'S' | b'N' | b'A' => self.note(kind, &body),
362                _ if refused.is_none() => each(kind, &body)?,
363                _ => {}
364            }
365        }
366    }
367
368    pub fn simple(&mut self, sql: &str) -> Result<(), Failure> {
369        let mut body = Vec::new();
370        cstr(&mut body, sql);
371        self.send(b'Q', &body)?;
372        self.flush()?;
373        self.until_ready(|_, _| Ok(()))
374    }
375
376    pub fn prepare(&mut self, name: &str, sql: &str) -> Result<Described, Failure> {
377        let mut parse = Vec::new();
378        cstr(&mut parse, name);
379        cstr(&mut parse, sql);
380        parse.extend_from_slice(&0i16.to_be_bytes());
381        self.send(b'P', &parse)?;
382        let mut describe = vec![b'S'];
383        cstr(&mut describe, name);
384        self.send(b'D', &describe)?;
385        self.send(b'S', &[])?;
386        self.flush()?;
387        let mut described = Described::default();
388        self.until_ready(|kind, body| {
389            let mut b = Body { bytes: body, at: 0 };
390            match kind {
391                b't' => {
392                    let n = b.i16()?;
393                    described.parameters = (0..n).map(|_| b.i32().map(|o| o as u32)).collect::<Result<_, _>>()?;
394                }
395                b'T' => {
396                    let n = b.i16()?;
397                    for _ in 0..n {
398                        let name = b.cstr()?;
399                        let (table, attnum, oid) = (b.i32()? as u32, b.i16()?, b.i32()? as u32);
400                        b.take(2)?;
401                        let typmod = b.i32()?;
402                        b.take(2)?;
403                        described.columns.push(Field { name, table, attnum, oid, typmod });
404                    }
405                }
406                _ => {}
407            }
408            Ok(())
409        })?;
410        Ok(described)
411    }
412
413    /// Binds text parameters (None is NULL) to a prepared statement and executes it, returning at
414    /// most `max_rows` rows (0 for all).
415    pub fn execute(&mut self, name: &str, parameters: &[Option<String>], max_rows: i32) -> Result<Executed, Failure> {
416        let mut bind = Vec::new();
417        cstr(&mut bind, "");
418        cstr(&mut bind, name);
419        bind.extend_from_slice(&0i16.to_be_bytes());
420        bind.extend_from_slice(&(parameters.len() as i16).to_be_bytes());
421        for p in parameters {
422            match p {
423                None => bind.extend_from_slice(&(-1i32).to_be_bytes()),
424                Some(text) => {
425                    bind.extend_from_slice(&(text.len() as i32).to_be_bytes());
426                    bind.extend_from_slice(text.as_bytes());
427                }
428            }
429        }
430        bind.extend_from_slice(&0i16.to_be_bytes());
431        self.send(b'B', &bind)?;
432        let mut execute = Vec::new();
433        cstr(&mut execute, "");
434        execute.extend_from_slice(&max_rows.to_be_bytes());
435        self.send(b'E', &execute)?;
436        self.send(b'S', &[])?;
437        self.flush()?;
438        let mut executed = Executed::default();
439        self.until_ready(|kind, body| {
440            let mut b = Body { bytes: body, at: 0 };
441            match kind {
442                b'D' => {
443                    let n = b.i16()?;
444                    let mut row = Vec::new();
445                    for _ in 0..n {
446                        let len = b.i32()?;
447                        row.push(if len < 0 { None } else { Some(String::from_utf8_lossy(b.take(len as usize)?).into_owned()) });
448                    }
449                    executed.rows.push(row);
450                }
451                b'C' => executed.tag = b.cstr()?,
452                _ => {}
453            }
454            Ok(())
455        })?;
456        Ok(executed)
457    }
458}
459
460fn tcp(target: &Target) -> std::io::Result<TcpStream> {
461    let socket = TcpStream::connect((target.host.as_str(), target.port))?;
462    socket.set_nodelay(true)?;
463    Ok(socket)
464}
465
466#[cfg(unix)]
467fn unix_socket(target: &Target) -> std::io::Result<Box<dyn Stream>> {
468    Ok(Box::new(std::os::unix::net::UnixStream::connect(format!("{}/.s.PGSQL.{}", target.host, target.port))?))
469}
470
471#[cfg(not(unix))]
472fn unix_socket(_: &Target) -> std::io::Result<Box<dyn Stream>> {
473    Err(std::io::Error::new(std::io::ErrorKind::Unsupported, "Unix sockets need a Unix system"))
474}
475
476#[cfg(test)]
477mod tests {
478    use super::*;
479
480    #[test]
481    fn urls() {
482        let t = Target::parse("postgres://ironwork:p%40ss@db.example:6543/payroll").unwrap();
483        let expected = Target { host: "db.example".into(), port: 6543, user: "ironwork".into(), password: Some("p@ss".into()), database: "payroll".into(), ssl: None, root_cert: None };
484        assert_eq!(t, expected);
485        let t = Target::parse("postgres://me@db/payroll?sslmode=verify-full&sslrootcert=/etc/ca.pem").unwrap();
486        assert_eq!((t.ssl, t.root_cert.as_deref()), (Some(SslMode::VerifyFull), Some(Path::new("/etc/ca.pem"))));
487        assert!(Target::parse("postgres://me@db/payroll?sslmode=require").unwrap_err().contains("verify-full"));
488        let t = Target::parse("postgres://me@/payroll?host=/var/run/postgresql").unwrap();
489        assert_eq!((t.host.as_str(), t.port, t.database.as_str()), ("/var/run/postgresql", 5432, "payroll"));
490        assert!(Target::parse("mysql://x").is_err());
491        assert!(Target::parse("postgres://me@h:port/db").is_err());
492    }
493}