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 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 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, ¶meters, 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 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#[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}