Skip to main content

ironwork_rt/sql/postgres/
mod.rs

1//! The PostgreSQL backend. Each call runs under a savepoint inside the unit of work, because Db2
2//! undoes a failed statement and keeps the rest of the unit, where PostgreSQL would abort it all.
3
4mod dialect;
5mod scram;
6mod wire;
7
8use super::{Abandoned, Answer, Call, Database, Outcome, Value};
9use std::collections::HashMap;
10pub use wire::{Stream, Tls};
11use wire::{Connection, Described, Failure, Target};
12
13pub struct Postgres {
14    conn: Connection,
15    /// Prepared statements by their PostgreSQL text: the statement's name and its types.
16    prepared: HashMap<String, (String, Described)>,
17    source: String,
18}
19
20fn abandon(f: Failure) -> Abandoned {
21    let message = match f {
22        Failure::Refused { state, message } => format!("PostgreSQL refused a request of ironwork's own ({state}): {message}"),
23        Failure::Broken(m) => m,
24    };
25    Abandoned { code: "SQL", message }
26}
27
28impl Postgres {
29    /// Runs `script` outside any unit of work, to prepare a server for a test.
30    #[doc(hidden)]
31    pub fn load_script(&mut self, script: &str) -> Result<(), String> {
32        self.conn.simple(script).map_err(|f| abandon(f).message)
33    }
34
35    /// `tls` is None in ironwork's own build, which then connects without TLS.
36    pub fn connect(url: &str, tls: Option<&dyn Tls>) -> Result<Self, String> {
37        let target = Target::parse(url)?;
38        let conn = Connection::open(&target, tls)?;
39        let over = if conn.encrypted { " over TLS" } else { "" };
40        let source = format!("PostgreSQL {} at {}:{}/{}{over}", conn.server_version, target.host, target.port, target.database);
41        Ok(Self { conn, prepared: HashMap::new(), source })
42    }
43
44    /// The server and database, as a recording's header names them.
45    pub fn source(&self) -> &str {
46        &self.source
47    }
48
49    /// Runs a call as one statement, returning at most `max_rows` rows (0 for all).
50    fn run(&mut self, call: &Call, max_rows: i32) -> Answer {
51        let sql = dialect::rewrite(call.text, call.cursor);
52        if self.conn.status == b'I' {
53            self.conn.simple("BEGIN").map_err(abandon)?;
54        }
55        self.conn.simple("SAVEPOINT ironwork").map_err(abandon)?;
56        match self.statement(&sql, call, max_rows) {
57            Ok(outcome) => {
58                self.conn.simple("RELEASE SAVEPOINT ironwork").map_err(abandon)?;
59                Ok(outcome)
60            }
61            Err(Failure::Refused { state, message }) => {
62                let db2 = dialect::db2_error(&state);
63                // -911 rolls back the whole unit of work, as Db2 does after a deadlock or timeout.
64                let undo = if db2.is_some_and(|(code, _)| code == -911) { "ROLLBACK" } else { "ROLLBACK TO SAVEPOINT ironwork; RELEASE SAVEPOINT ironwork" };
65                self.conn.simple(undo).map_err(abandon)?;
66                match db2 {
67                    Some((code, state)) => Ok(Outcome::error(code, state)),
68                    None => Err(Abandoned {
69                        code: "SQL",
70                        message: format!("PostgreSQL refused {} {sql} with SQLSTATE {state}, which ironwork's table gives no Db2 SQLCODE: {message}", call.verb),
71                    }),
72                }
73            }
74            Err(broken) => Err(abandon(broken)),
75        }
76    }
77
78    fn statement(&mut self, sql: &str, call: &Call, max_rows: i32) -> Result<Outcome, Failure> {
79        let (name, described) = match self.prepared.get(sql) {
80            Some(p) => p.clone(),
81            None => {
82                let name = format!("ironwork{}", self.prepared.len() + 1);
83                let described = self.conn.prepare(&name, sql)?;
84                self.prepared.insert(sql.to_owned(), (name.clone(), described.clone()));
85                (name, described)
86            }
87        };
88        if described.parameters.len() != call.inputs.len() {
89            return Err(Failure::Broken(format!("PostgreSQL reads {} parameters in {sql}, and the program sends {}", described.parameters.len(), call.inputs.len())));
90        }
91        let parameters: Vec<Option<String>> = call.inputs.iter().zip(&described.parameters).map(|(v, &oid)| dialect::text(v, oid)).collect();
92        let executed = self.conn.execute(&name, &parameters, max_rows)?;
93        let mut rows = Vec::new();
94        for row in &executed.rows {
95            let values = row.iter().zip(&described.columns).map(|(column, &oid)| column.as_deref().map_or(Ok(Value::Null), |text| dialect::value(oid, text)));
96            rows.push(values.collect::<Result<Vec<_>, _>>().map_err(Failure::Broken)?);
97        }
98        let changed = matches!(executed.tag.split(' ').next(), Some("INSERT" | "UPDATE" | "DELETE" | "MERGE"));
99        let affected = if changed { executed.tag.rsplit(' ').next().and_then(|n| n.parse().ok()).unwrap_or(0) } else { 0 };
100        Ok(Outcome { affected, rows, ..Outcome::ok() })
101    }
102
103    fn end(&mut self, verb: &str) -> Answer {
104        if self.conn.status == b'I' {
105            return Ok(Outcome::ok());
106        }
107        match self.conn.simple(verb) {
108            Ok(()) => Ok(Outcome::ok()),
109            Err(Failure::Refused { state, message }) => match dialect::db2_error(&state) {
110                Some((code, state)) => Ok(Outcome::error(code, state)),
111                None => Err(Abandoned { code: "SQL", message: format!("PostgreSQL refused {verb} with SQLSTATE {state}: {message}") }),
112            },
113            Err(broken) => Err(abandon(broken)),
114        }
115    }
116}
117
118impl Database for Postgres {
119    /// Two rows are enough to tell a SELECT INTO's one row from its too many.
120    fn execute(&mut self, call: &Call) -> Answer {
121        self.run(call, 2)
122    }
123    fn open(&mut self, call: &Call) -> Answer {
124        self.run(call, 0)
125    }
126    fn fetch(&mut self, call: &Call) -> Answer {
127        self.run(call, 0)
128    }
129    fn close(&mut self, call: &Call) -> Answer {
130        self.run(call, 0)
131    }
132    fn commit(&mut self, _: &Call) -> Answer {
133        self.end("COMMIT")
134    }
135    fn rollback(&mut self, _: &Call) -> Answer {
136        self.end("ROLLBACK")
137    }
138    fn close_all(&mut self) -> Result<(), Abandoned> {
139        self.conn.simple("CLOSE ALL").map_err(abandon)
140    }
141}
142
143/// Against a live server named by IRONWORK_PG_URL, which `tools/pg-test.sh` starts in a container;
144/// without it this test passes without running. The recorded run of a whole program is in exec.
145#[cfg(test)]
146mod tests {
147    use super::*;
148
149    fn url() -> Option<String> {
150        std::env::var("IRONWORK_PG_URL").ok()
151    }
152
153    #[test]
154    fn a_wrong_password_is_refused() {
155        let Some(url) = url() else { return };
156        let Some((head, tail)) = url.split_once('@') else { return };
157        let wrong = format!("{}:wrong@{tail}", head.rsplit_once(':').map_or(head, |(user, _)| user));
158        let refused = Postgres::connect(&wrong, None).err().expect("refused");
159        assert!(refused.contains("28P01"), "{refused}");
160    }
161}