ironwork_rt/sql/postgres/
mod.rs1mod 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: 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 #[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 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 pub fn source(&self) -> &str {
46 &self.source
47 }
48
49 fn run(&mut self, call: &Call, max_rows: i32) -> Answer {
51 self.guarded(call, |pg, sql| pg.statement(sql, call, max_rows))
52 }
53
54 fn guarded(&mut self, call: &Call, work: impl FnOnce(&mut Self, &str) -> Result<Outcome, Failure>) -> Answer {
56 let sql = dialect::rewrite(call.text, call.cursor);
57 if self.conn.status == b'I' {
58 self.conn.simple("BEGIN").map_err(abandon)?;
59 }
60 self.conn.simple("SAVEPOINT ironwork").map_err(abandon)?;
61 match work(self, &sql) {
62 Ok(outcome) => {
63 self.conn.simple("RELEASE SAVEPOINT ironwork").map_err(abandon)?;
64 Ok(outcome)
65 }
66 Err(Failure::Refused { state, message }) => {
67 let db2 = dialect::db2_error(&state);
68 let undo = if db2.is_some_and(|(code, _)| code == -911) { "ROLLBACK" } else { "ROLLBACK TO SAVEPOINT ironwork; RELEASE SAVEPOINT ironwork" };
70 self.conn.simple(undo).map_err(abandon)?;
71 match db2 {
72 Some((code, state)) => Ok(Outcome::error(code, state)),
73 None => Err(Abandoned {
74 code: "SQL",
75 message: format!("PostgreSQL refused {} {sql} with SQLSTATE {state}, which ironwork's table gives no Db2 SQLCODE: {message}", call.verb),
76 }),
77 }
78 }
79 Err(broken) => Err(abandon(broken)),
80 }
81 }
82
83 fn parsed(&mut self, sql: &str) -> Result<(String, Described), Failure> {
85 if let Some(p) = self.prepared.get(sql) {
86 return Ok(p.clone());
87 }
88 let name = format!("ironwork{}", self.prepared.len() + 1);
89 let described = self.conn.prepare(&name, sql)?;
90 self.prepared.insert(sql.to_owned(), (name.clone(), described.clone()));
91 Ok((name, described))
92 }
93
94 fn statement(&mut self, sql: &str, call: &Call, max_rows: i32) -> Result<Outcome, Failure> {
95 let (name, described) = self.parsed(sql)?;
96 if described.parameters.len() != call.inputs.len() {
97 return Err(Failure::Broken(format!("PostgreSQL reads {} parameters in {sql}, and the program sends {}", described.parameters.len(), call.inputs.len())));
98 }
99 let parameters: Vec<Option<String>> = call.inputs.iter().zip(&described.parameters).map(|(v, &oid)| dialect::text(v, oid)).collect();
100 let executed = self.conn.execute(&name, ¶meters, max_rows)?;
101 let mut rows = Vec::new();
102 for row in &executed.rows {
103 let values = row.iter().zip(&described.columns).map(|(column, &oid)| column.as_deref().map_or(Ok(Value::Null), |text| dialect::value(oid, text)));
104 rows.push(values.collect::<Result<Vec<_>, _>>().map_err(Failure::Broken)?);
105 }
106 let changed = matches!(executed.tag.split(' ').next(), Some("INSERT" | "UPDATE" | "DELETE" | "MERGE"));
107 let affected = if changed { executed.tag.rsplit(' ').next().and_then(|n| n.parse().ok()).unwrap_or(0) } else { 0 };
108 Ok(Outcome { affected, rows, ..Outcome::ok() })
109 }
110
111 fn end(&mut self, verb: &str) -> Answer {
112 if self.conn.status == b'I' {
113 return Ok(Outcome::ok());
114 }
115 match self.conn.simple(verb) {
116 Ok(()) => Ok(Outcome::ok()),
117 Err(Failure::Refused { state, message }) => match dialect::db2_error(&state) {
118 Some((code, state)) => Ok(Outcome::error(code, state)),
119 None => Err(Abandoned { code: "SQL", message: format!("PostgreSQL refused {verb} with SQLSTATE {state}: {message}") }),
120 },
121 Err(broken) => Err(abandon(broken)),
122 }
123 }
124}
125
126impl Database for Postgres {
127 fn execute(&mut self, call: &Call) -> Answer {
129 self.run(call, 2)
130 }
131 fn prepare(&mut self, call: &Call) -> Answer {
133 self.guarded(call, |pg, sql| pg.parsed(sql).map(|_| Outcome::ok()))
134 }
135 fn open(&mut self, call: &Call) -> Answer {
136 self.run(call, 0)
137 }
138 fn fetch(&mut self, call: &Call) -> Answer {
139 self.run(call, 0)
140 }
141 fn close(&mut self, call: &Call) -> Answer {
142 self.run(call, 0)
143 }
144 fn commit(&mut self, _: &Call) -> Answer {
145 self.end("COMMIT")
146 }
147 fn rollback(&mut self, _: &Call) -> Answer {
148 self.end("ROLLBACK")
149 }
150 fn close_all(&mut self) -> Result<(), Abandoned> {
151 self.conn.simple("CLOSE ALL").map_err(abandon)
152 }
153}
154
155#[cfg(test)]
158mod tests {
159 use super::*;
160
161 fn url() -> Option<String> {
162 std::env::var("IRONWORK_PG_URL").ok()
163 }
164
165 #[test]
166 fn a_wrong_password_is_refused() {
167 let Some(url) = url() else { return };
168 let Some((head, tail)) = url.split_once('@') else { return };
169 let wrong = format!("{}:wrong@{tail}", head.rsplit_once(':').map_or(head, |(user, _)| user));
170 let refused = Postgres::connect(&wrong, None).err().expect("refused");
171 assert!(refused.contains("28P01"), "{refused}");
172 }
173}