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, Column, 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
20/// The statements a connection keeps prepared. Past it they are all deallocated and each prepared
21/// again when next run, so dynamic statements that carry their values in their text do not grow
22/// the server's memory without end.
23const PREPARED_LIMIT: usize = 1024;
24
25fn abandon(f: Failure) -> Abandoned {
26    let message = match f {
27        Failure::Refused { state, message } => format!("PostgreSQL refused a request of ironwork's own ({state}): {message}"),
28        Failure::Broken(m) => m,
29    };
30    Abandoned { code: "SQL", message }
31}
32
33impl Postgres {
34    /// Runs `script` outside any unit of work, to prepare a server for a test.
35    #[doc(hidden)]
36    pub fn load_script(&mut self, script: &str) -> Result<(), String> {
37        self.conn.simple(script).map_err(|f| abandon(f).message)
38    }
39
40    /// `tls` is None in ironwork's own build, which then connects without TLS.
41    pub fn connect(url: &str, tls: Option<&dyn Tls>) -> Result<Self, String> {
42        let target = Target::parse(url)?;
43        let conn = Connection::open(&target, tls)?;
44        let over = if conn.encrypted { " over TLS" } else { "" };
45        let source = format!("PostgreSQL {} at {}:{}/{}{over}", conn.server_version, target.host, target.port, target.database);
46        Ok(Self { conn, prepared: HashMap::new(), source })
47    }
48
49    /// The server and database, as a recording's header names them.
50    pub fn source(&self) -> &str {
51        &self.source
52    }
53
54    /// Runs a call as one statement, returning at most `max_rows` rows (0 for all).
55    fn run(&mut self, call: &Call, max_rows: i32) -> Answer {
56        self.guarded(call, |pg, sql| pg.statement(sql, call, max_rows))
57    }
58
59    /// `work` on the call's PostgreSQL text under a savepoint, its refusal mapped to a Db2 SQLCODE.
60    fn guarded(&mut self, call: &Call, work: impl FnOnce(&mut Self, &str) -> Result<Outcome, Failure>) -> Answer {
61        let sql = dialect::rewrite(call.text, call.cursor);
62        if self.conn.status == b'I' {
63            self.conn.simple("BEGIN").map_err(abandon)?;
64        }
65        self.conn.simple("SAVEPOINT ironwork").map_err(abandon)?;
66        match work(self, &sql) {
67            Ok(outcome) => {
68                self.conn.simple("RELEASE SAVEPOINT ironwork").map_err(abandon)?;
69                Ok(outcome)
70            }
71            Err(Failure::Refused { state, message }) => {
72                let db2 = dialect::db2_error(&state);
73                // -911 rolls back the whole unit of work, as Db2 does after a deadlock or timeout.
74                let undo = if db2.is_some_and(|(code, _)| code == -911) { "ROLLBACK" } else { "ROLLBACK TO SAVEPOINT ironwork; RELEASE SAVEPOINT ironwork" };
75                self.conn.simple(undo).map_err(abandon)?;
76                match db2 {
77                    Some((code, state)) => Ok(Outcome::error(code, state)),
78                    None => Err(Abandoned {
79                        code: "SQL",
80                        message: format!("PostgreSQL refused {} {sql} with SQLSTATE {state}, which ironwork's table gives no Db2 SQLCODE: {message}", call.verb),
81                    }),
82                }
83            }
84            Err(broken) => Err(abandon(broken)),
85        }
86    }
87
88    /// The PostgreSQL statement for `sql`, prepared on first use and kept while fewer than
89    /// [`PREPARED_LIMIT`] are.
90    fn parsed(&mut self, sql: &str) -> Result<(String, Described), Failure> {
91        if let Some(p) = self.prepared.get(sql) {
92            return Ok(p.clone());
93        }
94        if self.prepared.len() >= PREPARED_LIMIT {
95            self.conn.simple("DEALLOCATE ALL")?;
96            self.prepared.clear();
97        }
98        let name = format!("ironwork{}", self.prepared.len() + 1);
99        let described = self.conn.prepare(&name, sql)?;
100        self.prepared.insert(sql.to_owned(), (name.clone(), described.clone()));
101        Ok((name, described))
102    }
103
104    fn statement(&mut self, sql: &str, call: &Call, max_rows: i32) -> Result<Outcome, Failure> {
105        let (name, described) = self.parsed(sql)?;
106        if described.parameters.len() != call.inputs.len() {
107            return Err(Failure::Broken(format!("PostgreSQL reads {} parameters in {sql}, and the program sends {}", described.parameters.len(), call.inputs.len())));
108        }
109        let parameters: Vec<Option<String>> = call.inputs.iter().zip(&described.parameters).map(|(v, &oid)| dialect::text(v, oid)).collect();
110        let executed = self.conn.execute(&name, &parameters, max_rows)?;
111        let mut rows = Vec::new();
112        for row in &executed.rows {
113            let values = row.iter().zip(&described.columns).map(|(column, field)| column.as_deref().map_or(Ok(Value::Null), |text| dialect::value(field.oid, text)));
114            rows.push(values.collect::<Result<Vec<_>, _>>().map_err(Failure::Broken)?);
115        }
116        let changed = matches!(executed.tag.split(' ').next(), Some("INSERT" | "UPDATE" | "DELETE" | "MERGE"));
117        let affected = if changed { executed.tag.rsplit(' ').next().and_then(|n| n.parse().ok()).unwrap_or(0) } else { 0 };
118        Ok(Outcome { affected, rows, ..Outcome::ok() })
119    }
120
121    /// Each result column with its Db2 type, its name upper-cased and NULL allowed unless it is a
122    /// table's column declared NOT NULL (assumption C403).
123    fn columns(&mut self, described: &Described) -> Result<Vec<Column>, Failure> {
124        let from_tables: Vec<String> = described.columns.iter().filter(|f| f.table != 0).map(|f| format!("({}, {})", f.table, f.attnum)).collect();
125        let mut not_null = Vec::new();
126        if !from_tables.is_empty() {
127            let sql = format!("SELECT attrelid, attnum FROM pg_attribute WHERE attnotnull AND (attrelid, attnum) IN ({})", from_tables.join(", "));
128            let (name, _) = self.parsed(&sql)?;
129            for row in self.conn.execute(&name, &[], 0)?.rows {
130                if let [Some(table), Some(attnum)] = row.as_slice() {
131                    not_null.push((table.clone(), attnum.clone()));
132                }
133            }
134        }
135        Ok(described
136            .columns
137            .iter()
138            .map(|f| Column {
139                name: f.name.to_uppercase(),
140                ty: dialect::column_type(f.oid, f.typmod),
141                nullable: !not_null.contains(&(f.table.to_string(), f.attnum.to_string())),
142            })
143            .collect())
144    }
145
146    fn end(&mut self, verb: &str) -> Answer {
147        if self.conn.status == b'I' {
148            return Ok(Outcome::ok());
149        }
150        match self.conn.simple(verb) {
151            Ok(()) => Ok(Outcome::ok()),
152            Err(Failure::Refused { state, message }) => match dialect::db2_error(&state) {
153                Some((code, state)) => Ok(Outcome::error(code, state)),
154                None => Err(Abandoned { code: "SQL", message: format!("PostgreSQL refused {verb} with SQLSTATE {state}: {message}") }),
155            },
156            Err(broken) => Err(abandon(broken)),
157        }
158    }
159}
160
161impl Database for Postgres {
162    /// Two rows are enough to tell a SELECT INTO's one row from its too many.
163    fn execute(&mut self, call: &Call) -> Answer {
164        self.run(call, 2)
165    }
166    /// PostgreSQL parses the statement string as Db2's PREPARE does, so its errors come at PREPARE,
167    /// and describes its result columns.
168    fn prepare(&mut self, call: &Call) -> Answer {
169        self.guarded(call, |pg, sql| {
170            let (_, described) = pg.parsed(sql)?;
171            Ok(Outcome { columns: pg.columns(&described)?, ..Outcome::ok() })
172        })
173    }
174    fn open(&mut self, call: &Call) -> Answer {
175        self.run(call, 0)
176    }
177    fn fetch(&mut self, call: &Call) -> Answer {
178        self.run(call, 0)
179    }
180    fn fetch_rows(&mut self, call: &Call, rows: u32) -> Answer {
181        let text = format!("FETCH FORWARD {rows} FROM {}", call.cursor.unwrap_or_default());
182        self.run(&Call { text: &text, inputs: &[], ..*call }, 0)
183    }
184    /// ATOMIC inserts every row under the call's one savepoint, so a failure undoes them all; NOT
185    /// ATOMIC gives each row its own, as Db2 keeps the rows that went in
186    /// ([`numeric::assumptions::NOT_ATOMIC_SUMMARY`]).
187    fn insert_rows(&mut self, call: &Call, rows: &[Vec<Value>], atomic: bool) -> Answer {
188        if atomic {
189            return self.guarded(call, |pg, sql| {
190                let mut affected = 0;
191                for row in rows {
192                    affected += pg.statement(sql, &Call { inputs: row, ..*call }, 0)?.affected;
193                }
194                Ok(Outcome { affected, ..Outcome::ok() })
195            });
196        }
197        let (mut inserted, mut failed) = (0, 0);
198        for row in rows {
199            let one = Call { inputs: row, ..*call };
200            let outcome = self.guarded(&one, |pg, sql| pg.statement(sql, &one, 0))?;
201            if outcome.sqlcode < 0 {
202                failed += 1;
203            } else {
204                inserted += outcome.affected;
205            }
206        }
207        Ok(match (failed, inserted) {
208            (0, _) => Outcome { affected: inserted, ..Outcome::ok() },
209            (_, 0) => Outcome::error(-254, "22530"),
210            _ => Outcome { affected: inserted, ..Outcome::error(-253, "22529") },
211        })
212    }
213    fn close(&mut self, call: &Call) -> Answer {
214        self.run(call, 0)
215    }
216    fn commit(&mut self, _: &Call) -> Answer {
217        self.end("COMMIT")
218    }
219    fn rollback(&mut self, _: &Call) -> Answer {
220        self.end("ROLLBACK")
221    }
222    fn close_all(&mut self) -> Result<(), Abandoned> {
223        self.conn.simple("CLOSE ALL").map_err(abandon)
224    }
225}
226
227/// Against a live server named by IRONWORK_PG_URL, which `tools/pg-test.sh` starts in a container;
228/// without it this test passes without running. The recorded run of a whole program is in exec.
229#[cfg(test)]
230mod tests {
231    use super::*;
232
233    fn url() -> Option<String> {
234        std::env::var("IRONWORK_PG_URL").ok()
235    }
236
237    #[test]
238    fn a_wrong_password_is_refused() {
239        let Some(url) = url() else { return };
240        let Some((head, tail)) = url.split_once('@') else { return };
241        let wrong = format!("{}:wrong@{tail}", head.rsplit_once(':').map_or(head, |(user, _)| user));
242        let refused = Postgres::connect(&wrong, None).err().expect("refused");
243        assert!(refused.contains("28P01"), "{refused}");
244    }
245
246    #[test]
247    fn distinct_statements_do_not_pile_up_on_the_server() {
248        let Some(url) = url() else { return };
249        let mut pg = Postgres::connect(&url, None).expect("connects");
250        fn call(text: &str) -> Call<'_> {
251            Call { program: "P", ordinal: 1, verb: "SELECT", cursor: None, text, inputs: &[] }
252        }
253        for n in 0..PREPARED_LIMIT as i64 + 8 {
254            let text = format!("VALUES {n}");
255            assert_eq!(pg.execute(&call(&text)).expect("answers").rows, [[Value::Int(n)]]);
256        }
257        let rows = pg.execute(&call("SELECT COUNT(*) FROM PG_PREPARED_STATEMENTS")).expect("answers").rows;
258        assert!(matches!(rows.as_slice(), [row] if matches!(row.as_slice(), [Value::Int(k)] if *k as usize <= PREPARED_LIMIT)), "{rows:?}");
259        assert!(pg.prepared.len() <= PREPARED_LIMIT);
260    }
261}