1use super::scram::{Scram, nonce};
5use std::io::{BufReader, Read, Write};
6use std::net::TcpStream;
7use std::path::{Path, PathBuf};
8
9pub trait Stream: Read + Write + Send {}
11
12impl<T: Read + Write + Send> Stream for T {}
13
14pub trait Tls: Send + Sync {
18 fn wrap(&self, socket: TcpStream, host: &str, roots: Option<&Path>) -> std::io::Result<Box<dyn Stream>>;
19}
20
21#[derive(Clone, Copy, Debug, PartialEq, Eq)]
23pub enum SslMode {
24 Disable,
25 VerifyFull,
27}
28
29#[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 pub ssl: Option<SslMode>,
39 pub root_cert: Option<PathBuf>,
40}
41
42impl Target {
43 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#[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#[derive(Debug, Default)]
119pub struct Executed {
120 pub rows: Vec<Vec<Option<String>>>,
121 pub tag: String,
122}
123
124#[derive(Clone, Debug, Default)]
126pub struct Described {
127 pub parameters: Vec<u32>,
128 pub columns: Vec<Field>,
129}
130
131#[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 out: Vec<u8>,
146 pub server_version: String,
147 pub encrypted: bool,
148 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
187fn 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
205const LONGEST_MESSAGE: i32 = 64 << 20;
208
209fn body_len(len: i32) -> Result<usize, Failure> {
211 if len > LONGEST_MESSAGE {
212 return Err(Failure::Broken(format!("PostgreSQL sent a message of {len} bytes, more than the {LONGEST_MESSAGE} ironwork reads")));
213 }
214 Ok((len.max(4) - 4) as usize)
215}
216
217impl Connection {
218 pub fn open(target: &Target, tls: Option<&dyn Tls>) -> Result<Self, String> {
221 let at = format!("PostgreSQL at {}:{}", target.host, target.port);
222 let broken = |e: std::io::Error| format!("cannot reach {at}: {e}");
223 let socket = target.host.starts_with('/');
224 let mode = match (target.ssl, tls) {
225 (Some(mode), _) => mode,
226 (None, Some(_)) if !socket => SslMode::VerifyFull,
227 (None, _) => SslMode::Disable,
228 };
229 let stream: Box<dyn Stream> = match (mode, socket, tls) {
230 (SslMode::Disable, true, _) => unix_socket(target).map_err(broken)?,
231 (SslMode::VerifyFull, true, _) => return Err("sslmode=verify-full is for TCP; a Unix socket needs no TLS".into()),
232 (SslMode::VerifyFull, false, None) => {
233 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());
234 }
235 (SslMode::Disable, false, _) => Box::new(tcp(target).map_err(broken)?),
236 (SslMode::VerifyFull, false, Some(tls)) => {
237 let mut socket = tcp(target).map_err(broken)?;
238 socket.write_all(&[0, 0, 0, 8, 0x04, 0xD2, 0x16, 0x2F]).map_err(broken)?;
239 let mut answer = [0u8];
240 socket.read_exact(&mut answer).map_err(broken)?;
241 if answer[0] != b'S' {
242 return Err(format!("{at} does not offer TLS; give sslmode=disable to connect without it"));
243 }
244 tls.wrap(socket, &target.host, target.root_cert.as_deref()).map_err(|e| format!("TLS with {at} failed: {e}"))?
245 }
246 };
247 let mut conn = Self { stream: BufReader::new(stream), out: Vec::new(), server_version: String::new(), encrypted: mode == SslMode::VerifyFull, status: b'I' };
248 conn.start(target).map_err(|f| match f {
249 Failure::Refused { state, message } => format!("PostgreSQL refused the connection ({state}): {message}"),
250 Failure::Broken(m) => m,
251 })?;
252 Ok(conn)
253 }
254
255 fn start(&mut self, target: &Target) -> Result<(), Failure> {
256 let mut body = 196_608i32.to_be_bytes().to_vec();
257 let params = [
258 ("user", target.user.as_str()),
259 ("database", &target.database),
260 ("client_encoding", "UTF8"),
261 ("DateStyle", "ISO"),
262 ("TimeZone", "UTC"),
263 ("extra_float_digits", "3"),
264 ("application_name", "ironwork"),
265 ];
266 for (name, value) in params {
267 cstr(&mut body, name);
268 cstr(&mut body, value);
269 }
270 body.push(0);
271 self.out.extend_from_slice(&(body.len() as i32 + 4).to_be_bytes());
272 self.out.extend_from_slice(&body);
273 self.flush()?;
274 let password = || target.password.clone().ok_or_else(|| Failure::Broken("PostgreSQL asks for a password: give one in the URL or PGPASSWORD".into()));
275 let mut scram = None;
276 loop {
277 let (kind, body) = self.receive()?;
278 let mut b = Body { bytes: &body, at: 0 };
279 match kind {
280 b'R' => match b.i32()? {
281 0 => {}
282 3 => {
283 let mut reply = Vec::new();
284 cstr(&mut reply, &password()?);
285 self.send(b'p', &reply)?;
286 self.flush()?;
287 }
288 10 => {
289 let mechanisms: Vec<String> = std::iter::from_fn(|| b.cstr().ok().filter(|m| !m.is_empty())).collect();
290 if !mechanisms.iter().any(|m| m == "SCRAM-SHA-256") {
291 return Err(Failure::Broken(format!("PostgreSQL offers SASL {mechanisms:?}, and ironwork speaks SCRAM-SHA-256")));
292 }
293 let mut exchange = Scram::new(&password()?, nonce());
294 let first = exchange.client_first("");
295 let mut reply = Vec::new();
296 cstr(&mut reply, "SCRAM-SHA-256");
297 reply.extend_from_slice(&(first.len() as i32).to_be_bytes());
298 reply.extend_from_slice(first.as_bytes());
299 self.send(b'p', &reply)?;
300 self.flush()?;
301 scram = Some(exchange);
302 }
303 11 => {
304 let exchange = scram.as_mut().ok_or_else(|| Failure::Broken("PostgreSQL continued a SASL exchange that had not begun".into()))?;
305 let reply = exchange.client_final(&String::from_utf8_lossy(b.rest())).map_err(Failure::Broken)?;
306 self.send(b'p', reply.as_bytes())?;
307 self.flush()?;
308 }
309 12 => {
310 let exchange = scram.as_ref().ok_or_else(|| Failure::Broken("PostgreSQL ended a SASL exchange that had not begun".into()))?;
311 exchange.verify(&String::from_utf8_lossy(b.rest())).map_err(Failure::Broken)?;
312 }
313 5 => return Err(Failure::Broken("PostgreSQL asks for MD5 authentication; ironwork speaks SCRAM-SHA-256".into())),
314 other => return Err(Failure::Broken(format!("PostgreSQL asks for authentication method {other}, which ironwork does not speak"))),
315 },
316 b'Z' => {
317 self.status = body.first().copied().unwrap_or(b'I');
318 return Ok(());
319 }
320 b'E' => return Err(refusal(&body)),
321 _ => self.note(kind, &body),
322 }
323 }
324 }
325
326 fn send(&mut self, kind: u8, body: &[u8]) -> Result<(), Failure> {
327 self.out.push(kind);
328 self.out.extend_from_slice(&(body.len() as i32 + 4).to_be_bytes());
329 self.out.extend_from_slice(body);
330 Ok(())
331 }
332
333 fn flush(&mut self) -> Result<(), Failure> {
334 let stream = self.stream.get_mut();
335 stream.write_all(&self.out)?;
336 stream.flush()?;
337 self.out.clear();
338 Ok(())
339 }
340
341 fn receive(&mut self) -> Result<(u8, Vec<u8>), Failure> {
342 let mut head = [0u8; 5];
343 self.stream.read_exact(&mut head)?;
344 let mut body = vec![0u8; body_len(i32::from_be_bytes([head[1], head[2], head[3], head[4]]))?];
345 self.stream.read_exact(&mut body)?;
346 Ok((head[0], body))
347 }
348
349 fn note(&mut self, kind: u8, body: &[u8]) {
351 if kind == b'S' {
352 let mut b = Body { bytes: body, at: 0 };
353 if let (Ok(name), Ok(value)) = (b.cstr(), b.cstr())
354 && name == "server_version"
355 {
356 self.server_version = value;
357 }
358 }
359 }
360
361 fn until_ready(&mut self, mut each: impl FnMut(u8, &[u8]) -> Result<(), Failure>) -> Result<(), Failure> {
363 let mut refused = None;
364 loop {
365 let (kind, body) = self.receive()?;
366 match kind {
367 b'Z' => {
368 self.status = body.first().copied().unwrap_or(b'I');
369 return refused.map_or(Ok(()), Err);
370 }
371 b'E' => refused = refused.or(Some(refusal(&body))),
372 b'S' | b'N' | b'A' => self.note(kind, &body),
373 _ if refused.is_none() => each(kind, &body)?,
374 _ => {}
375 }
376 }
377 }
378
379 pub fn simple(&mut self, sql: &str) -> Result<(), Failure> {
380 let mut body = Vec::new();
381 cstr(&mut body, sql);
382 self.send(b'Q', &body)?;
383 self.flush()?;
384 self.until_ready(|_, _| Ok(()))
385 }
386
387 pub fn prepare(&mut self, name: &str, sql: &str) -> Result<Described, Failure> {
388 let mut parse = Vec::new();
389 cstr(&mut parse, name);
390 cstr(&mut parse, sql);
391 parse.extend_from_slice(&0i16.to_be_bytes());
392 self.send(b'P', &parse)?;
393 let mut describe = vec![b'S'];
394 cstr(&mut describe, name);
395 self.send(b'D', &describe)?;
396 self.send(b'S', &[])?;
397 self.flush()?;
398 let mut described = Described::default();
399 self.until_ready(|kind, body| {
400 let mut b = Body { bytes: body, at: 0 };
401 match kind {
402 b't' => {
403 let n = b.i16()?;
404 described.parameters = (0..n).map(|_| b.i32().map(|o| o as u32)).collect::<Result<_, _>>()?;
405 }
406 b'T' => {
407 let n = b.i16()?;
408 for _ in 0..n {
409 let name = b.cstr()?;
410 let (table, attnum, oid) = (b.i32()? as u32, b.i16()?, b.i32()? as u32);
411 b.take(2)?;
412 let typmod = b.i32()?;
413 b.take(2)?;
414 described.columns.push(Field { name, table, attnum, oid, typmod });
415 }
416 }
417 _ => {}
418 }
419 Ok(())
420 })?;
421 Ok(described)
422 }
423
424 pub fn execute(&mut self, name: &str, parameters: &[Option<String>], max_rows: i32) -> Result<Executed, Failure> {
427 let mut bind = Vec::new();
428 cstr(&mut bind, "");
429 cstr(&mut bind, name);
430 bind.extend_from_slice(&0i16.to_be_bytes());
431 bind.extend_from_slice(&(parameters.len() as i16).to_be_bytes());
432 for p in parameters {
433 match p {
434 None => bind.extend_from_slice(&(-1i32).to_be_bytes()),
435 Some(text) => {
436 bind.extend_from_slice(&(text.len() as i32).to_be_bytes());
437 bind.extend_from_slice(text.as_bytes());
438 }
439 }
440 }
441 bind.extend_from_slice(&0i16.to_be_bytes());
442 self.send(b'B', &bind)?;
443 let mut execute = Vec::new();
444 cstr(&mut execute, "");
445 execute.extend_from_slice(&max_rows.to_be_bytes());
446 self.send(b'E', &execute)?;
447 self.send(b'S', &[])?;
448 self.flush()?;
449 let mut executed = Executed::default();
450 self.until_ready(|kind, body| {
451 let mut b = Body { bytes: body, at: 0 };
452 match kind {
453 b'D' => {
454 let n = b.i16()?;
455 let mut row = Vec::new();
456 for _ in 0..n {
457 let len = b.i32()?;
458 row.push(if len < 0 { None } else { Some(String::from_utf8_lossy(b.take(len as usize)?).into_owned()) });
459 }
460 executed.rows.push(row);
461 }
462 b'C' => executed.tag = b.cstr()?,
463 _ => {}
464 }
465 Ok(())
466 })?;
467 Ok(executed)
468 }
469}
470
471fn tcp(target: &Target) -> std::io::Result<TcpStream> {
472 let socket = TcpStream::connect((target.host.as_str(), target.port))?;
473 socket.set_nodelay(true)?;
474 Ok(socket)
475}
476
477#[cfg(unix)]
478fn unix_socket(target: &Target) -> std::io::Result<Box<dyn Stream>> {
479 Ok(Box::new(std::os::unix::net::UnixStream::connect(format!("{}/.s.PGSQL.{}", target.host, target.port))?))
480}
481
482#[cfg(not(unix))]
483fn unix_socket(_: &Target) -> std::io::Result<Box<dyn Stream>> {
484 Err(std::io::Error::new(std::io::ErrorKind::Unsupported, "Unix sockets need a Unix system"))
485}
486
487#[cfg(test)]
488mod tests {
489 use super::*;
490
491 #[test]
492 fn a_message_longer_than_ironwork_reads_is_refused() {
493 assert_eq!(body_len(4).ok(), Some(0));
494 assert_eq!(body_len(LONGEST_MESSAGE).ok(), Some(LONGEST_MESSAGE as usize - 4));
495 assert!(matches!(body_len(i32::MAX), Err(Failure::Broken(m)) if m.contains("more than")));
496 }
497
498 #[test]
499 fn urls() {
500 let t = Target::parse("postgres://ironwork:p%40ss@db.example:6543/payroll").unwrap();
501 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 };
502 assert_eq!(t, expected);
503 let t = Target::parse("postgres://me@db/payroll?sslmode=verify-full&sslrootcert=/etc/ca.pem").unwrap();
504 assert_eq!((t.ssl, t.root_cert.as_deref()), (Some(SslMode::VerifyFull), Some(Path::new("/etc/ca.pem"))));
505 assert!(Target::parse("postgres://me@db/payroll?sslmode=require").unwrap_err().contains("verify-full"));
506 let t = Target::parse("postgres://me@/payroll?host=/var/run/postgresql").unwrap();
507 assert_eq!((t.host.as_str(), t.port, t.database.as_str()), ("/var/run/postgresql", 5432, "payroll"));
508 assert!(Target::parse("mysql://x").is_err());
509 assert!(Target::parse("postgres://me@h:port/db").is_err());
510 }
511}